[NVIDIA][CuTe,Fwd/Bwd,sm120] Attention dropout (#2436), forward + backward - #2657
Open
johnnynunez wants to merge 8 commits into
Open
johnnynunez wants to merge 8 commits into
johnnynunez wants to merge 8 commits into
Conversation
Contributor
Author
On SM120 (Blackwell GeForce / RTX PRO 6000 / DGX Spark) the forward kernel set
`use_tma_O = self.arch >= Arch.sm_90`, enabling the TMA-based O-store epilogue.
But SM120 does not build the TMA store atom (tma_atom_O is None), so any forward
call crashes in cpasync.tma_partition with:
AttributeError: 'NoneType' object has no attribute '_trait'
This makes the CuTe-DSL forward unusable on every SM120 GPU.
Restrict the TMA O-store to sm_90..sm_119, which is where the WGMMA-era epilogue
path is actually available:
self.use_tma_O = Arch.sm_90 <= self.arch < Arch.sm_120
SM120 falls back to the non-TMA register->gmem O store (already used for the
SM80 path), which is correct and what the CpAsync SM120 kernel expects.
Verified on RTX PRO 6000 Blackwell (sm_120, cc 12.0), torch 2.12.0+cu130,
nvidia-cutlass-dsl 4.5.2: forward now runs and matches PyTorch SDPA reference
for hdim 64/96/128, causal and non-causal (max abs err <= 8e-3 in bf16). Before
this fix every SM120 forward call raised the AttributeError above.
…ed softmax_scale The CuTe-DSL backward did not run on SM120 at all; two bugs: 1. interface._flash_attn_bwd: the SM120 branch never defined `dQ_single_wg`, but the shared compile_key references it -> UnboundLocalError on the first bwd call. Define dQ_single_wg = False for SM120 (single MMA warp-group dQ). 2. flash_bwd.py __call__ used utils.compute_softmax_scale_log2(), which returns softmax_scale = None when score_mod is None. But the kernel arg is typed cutlass.Float32 and the backward multiplies acc_dK by softmax_scale (dK rescale when qhead_per_kvhead == 1), so launch crashed with "None to Float conversion is not supported". Compute softmax_scale_log2 inline (fold LOG2_E) and keep softmax_scale a real Float32, matching the SM100 backward. Also make utils.atomic_add_fp32 robust to the nvvm.atomicrmw binding signature, which differs between nvidia-cutlass-dsl CUDA-toolkit builds (cu12.x uses res=/op=/ptr=/a= kwargs; cu13 uses positional op, ptr, a with inferred result). Try the cu13 form, fall back to the cu12.x form. Verified on RTX PRO 6000 Blackwell (sm_120): full fwd+bwd now runs and gradients match a PyTorch fp32 reference for hdim 64/96/128, causal and non-causal (max |dq/dk/dv - ref| <= 1.3e-2 in bf16). Before this, every SM120 backward call raised UnboundLocalError / None-to-Float. Stacked on the use_tma_O fix (Dao-AILab#2649).
On SM120 (Blackwell GeForce / RTX PRO 6000 / DGX Spark) the forward kernel set
`use_tma_O = self.arch >= Arch.sm_90`, enabling the TMA-based O-store epilogue.
But SM120 does not build the TMA store atom (tma_atom_O is None), so any forward
call crashes in cpasync.tma_partition with:
AttributeError: 'NoneType' object has no attribute '_trait'
This makes the CuTe-DSL forward unusable on every SM120 GPU.
Restrict the TMA O-store to sm_90..sm_119, which is where the WGMMA-era epilogue
path is actually available:
self.use_tma_O = Arch.sm_90 <= self.arch < Arch.sm_120
SM120 falls back to the non-TMA register->gmem O store (already used for the
SM80 path), which is correct and what the CpAsync SM120 kernel expects.
Verified on RTX PRO 6000 Blackwell (sm_120, cc 12.0), torch 2.12.0+cu130,
nvidia-cutlass-dsl 4.5.2: forward now runs and matches PyTorch SDPA reference
for hdim 64/96/128, causal and non-causal (max abs err <= 8e-3 in bf16). Before
this fix every SM120 forward call raised the AttributeError above.
…softmax_scale The CuTe-DSL backward did not run on SM120 at all; two bugs: 1. interface._flash_attn_bwd: the SM120 branch never defined `dQ_single_wg`, but the shared compile_key references it -> UnboundLocalError on the first bwd call. Define dQ_single_wg = False for SM120 (single MMA warp-group dQ). 2. flash_bwd.py __call__ used utils.compute_softmax_scale_log2(), which returns softmax_scale = None when score_mod is None. But the kernel arg is typed cutlass.Float32 and the backward multiplies acc_dK by softmax_scale (dK rescale when qhead_per_kvhead == 1), so launch crashed with "None to Float conversion is not supported". Compute softmax_scale_log2 inline (fold LOG2_E) and keep softmax_scale as a real Float32, matching the SM100 bwd. Also make utils.atomic_add_fp32 robust to the nvvm.atomicrmw binding signature, which differs between nvidia-cutlass-dsl CUDA-toolkit builds (cu12.x uses res=/op=/ptr=/a= kwargs; cu13 uses positional op, ptr, a with inferred result). Try the cu13 form, fall back to the cu12.x form. Selected once at trace time. Verified on RTX PRO 6000 Blackwell (sm_120): full fwd+bwd now runs and gradients match a PyTorch fp32 reference for hdim 64/96/128, causal and non-causal (max |dq/dk/dv - ref| <= 1.3e-2 in bf16). Before this, every SM120 backward call raised UnboundLocalError / None-to-Float.
…ackward Implements attention dropout (fwd + bwd) for the SM80/SM120 CuTe-DSL kernels, self-contained on top of the SM120 fwd (Dao-AILab#2649) and bwd enablement fixes. - philox.py: Philox4x32 counter-based RNG + apply_dropout(), keyed by GLOBAL (batch, head, q_idx, kv_idx) via the identity coordinate tensor, so the forward and the backward recompute draw the IDENTICAL keep-mask regardless of fragment layout (the backward S tile is transposed under SdP_swapAB; handled). - Forward: apply inverted dropout (zero dropped, scale kept by 1/(1-p)) to P right after online_softmax/rescale_O. One RNG seed is drawn per call in autograd and reused in the backward via ctx. - Backward: regenerate the same keep-mask on the recomputed P; zeroing dropped P automatically zeroes the dV and dS contributions, matching the forward. - interface: dropout_p/seed/offset threaded through _flash_attn_fwd/_flash_attn_bwd (compile keys + arg lists, SM120 only), FlashAttnFunc fwd/bwd, and flash_attn_func. Verified on RTX PRO 6000 Blackwell (sm_120): p=0 matches PyTorch SDPA (err<=9e-3) with finite grads; dropout changes output and keeps grads finite; fwd+bwd are deterministic for a fixed seed; forward keep-fraction 0.702 vs 0.70 target with exact inverted-dropout scaling (1/S/(1-p)). dropout_p default 0.0 = no-op.
johnnynunez
force-pushed
the
feat/cute-dropout-2436-full
branch
from
June 28, 2026 18:08
1bb944f to
5a54626
Compare
The functional use_tma_O guard already landed via Dao-AILab#2656; this branch now carries only the explanatory comment for the sm_120 TMA-O restriction (issue Dao-AILab#2649).
…0-backward-enable Upstream Dao-AILab#2671 already landed equivalent fixes for the two bugs this PR addressed (dQ_single_wg defined in the SM120 branch, softmax_scale kept as a real Float32 via utils.compute_softmax_scale_log2), so this merge adopts the upstream implementations. Review feedback from v0i0 incorporated: - flash_bwd.py: drop the inline LOG2_E computation in favor of upstream's utils.compute_softmax_scale_log2 (the inline version is no longer needed). - utils.py: drop the try/except dual-signature nvvm.atomicrmw shim; the repo pins nvidia-cutlass-dsl==4.6.0.dev0 which supports a single call form.
…te-dropout-2436-full Brings in upstream main (through Dao-AILab#2671/Dao-AILab#2656) plus the resolved SM120 fix branches, leaving this PR dropout-only as promised: - flash_bwd.py: adopt upstream's utils.compute_softmax_scale_log2 and the local-attention AttentionMask/window_size wiring + the m_block guard indentation from Dao-AILab#2671; re-attach dropout_fn wiring inside the guarded mainloop. - flash_fwd.py: keep the philox import alongside upstream's use_tma_O guard. - utils.py: take upstream's single-form nvvm.atomicrmw (per v0i0's review on Dao-AILab#2658; the pinned nvidia-cutlass-dsl==4.6.0.dev0 supports one form). - interface.py: upstream's dQ_single_wg definition supersedes ours.
Contributor
Author
Contributor
|
Hi @johnnynunez, can you please add tests? Benchmarks would also be helpful for the PR. |
Contributor
|
Validated this on DGX Spark (GB10, sm_121a): dropout_p=0 is bit-exact against the no-dropout baseline for both forward and gradients; same-seed runs are bit-identical and different seeds diverge as expected; dropout backward runs with finite gradients; and the output mean scaling is consistent with the 1/(1-p) rescale (ratio 1.38 at p=0.5). Looks good from the GB10 side. I am closing my #2439 in favor of this, since it covers the same #2436 ask fwd+bwd on current main. |
12 tasks
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Implements attention dropout (#2436) for the CuTe-DSL forward AND backward on SM120 (Blackwell GeForce / RTX PRO 6000 / DGX Spark), self-contained on upstream
main.This PR is complete and standalone: it also carries the two SM120 prerequisites the dropout backward needs to run at all (so it works on a clean
maincheckout without depending on any other open PR). Theuse_tma_Ofix is also submitted separately as #2655; reviewers can take it from either place.Commits
use_tma_Ocrash fix (FlashAttentionForwardSm120: use_tma_O incorrectly True on SM_121 → AttributeError on tma_atom_O=None #2649) — SM120 forward crashed intma_partition(TMA O-store atom is null on SM120). Restrict TMA O-store to sm_90..sm_119.dQ_single_wgwas referenced in the bwd compile key but never defined in the SM120 branch (UnboundLocalError).__call__usedcompute_softmax_scale_log2(), which returnssoftmax_scale=Nonewhenscore_mod is None; but the kernel arg is typedFloat32and the backward rescalesdKbysoftmax_scale, so launch crashed (None to Float). Computesoftmax_scale_log2inline (foldLOG2_E) and keepsoftmax_scalea realFloat32, matching the SM100 backward.utils.atomic_add_fp32robust to thenvvm.atomicrmwbinding signature across nvidia-cutlass-dsl CUDA-toolkit builds (cu12.xres=/op=/ptr=/a=vs cu13 positional).philox.py: Philox4x32 counter-based RNG +apply_dropout(), keyed by global(batch, head, q_idx, kv_idx)so forward and the backward recompute draw the identical keep-mask regardless of fragment layout (backwardSis transposed underSdP_swapAB; handled).1/(1-p)) onPafteronline_softmax. One RNG seed per call, reused in backward viactx.P; zeroing droppedPzeroes thedV/dScontributions, matching the forward.Testing (RTX PRO 6000 Blackwell, sm_120, torch 2.12.0+cu130, cutlass-dsl 4.5.2)
dropout_p=0matches PyTorch SDPA (max err <= 9e-3 bf16) with finite grads — no-op default, no regression.|dq/dk/dv - ref| <= 1.3e-2(hdim 64/96/128, causal & non-causal).1/S/(1-p)(measured 0.02234 vs 0.02232) — statistically correct, unbiased.API
flash_attn_func(..., dropout_p=0.0). Seed is managed internally (drawn in forward, reused in backward).