Skip to content

[NVIDIA][CuTe,Fwd/Bwd,sm120] Attention dropout (#2436), forward + backward - #2657

Open
johnnynunez wants to merge 8 commits into
Dao-AILab:mainfrom
johnnynunez:feat/cute-dropout-2436-full
Open

johnnynunez wants to merge 8 commits into
Dao-AILab:mainfrom
johnnynunez:feat/cute-dropout-2436-full

Conversation

@johnnynunez

Copy link
Copy Markdown
Contributor

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 main checkout without depending on any other open PR). The use_tma_O fix is also submitted separately as #2655; reviewers can take it from either place.

Commits

  1. SM120 use_tma_O crash fix (FlashAttentionForwardSm120: use_tma_O incorrectly True on SM_121 → AttributeError on tma_atom_O=None #2649) — SM120 forward crashed in tma_partition (TMA O-store atom is null on SM120). Restrict TMA O-store to sm_90..sm_119.
  2. SM120 backward enablement — two bugs prevented the SM120 backward from running:
    • dQ_single_wg was referenced in the bwd compile key but never defined in the SM120 branch (UnboundLocalError).
    • __call__ used compute_softmax_scale_log2(), which returns softmax_scale=None when score_mod is None; but the kernel arg is typed Float32 and the backward rescales dK by softmax_scale, so launch crashed (None to Float). 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 across nvidia-cutlass-dsl CUDA-toolkit builds (cu12.x res=/op=/ptr=/a= vs cu13 positional).
  3. Dropout (DropOut Cute DSL #2436), forward + backward:
    • 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 (backward S is transposed under SdP_swapAB; handled).
    • Forward: inverted dropout (zero dropped, scale kept by 1/(1-p)) on P after online_softmax. One RNG seed per call, reused in backward via ctx.
    • Backward: regenerate the same mask on recomputed P; zeroing dropped P zeroes the dV/dS contributions, matching the forward.

Testing (RTX PRO 6000 Blackwell, sm_120, torch 2.12.0+cu130, cutlass-dsl 4.5.2)

  • dropout_p=0 matches PyTorch SDPA (max err <= 9e-3 bf16) with finite grads — no-op default, no regression.
  • Backward now runs and matches a PyTorch fp32 reference: |dq/dk/dv - ref| <= 1.3e-2 (hdim 64/96/128, causal & non-causal).
  • Dropout changes the output and keeps grads finite; forward+backward are deterministic for a fixed seed.
  • Forward keep-fraction 0.702 vs 0.70 target, with exact inverted scaling 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).

@johnnynunez johnnynunez changed the title [CuTe,Fwd/Bwd,sm120] Attention dropout (#2436), forward + backward [NVIDIA][CuTe,Fwd/Bwd,sm120] Attention dropout (#2436), forward + backward Jun 15, 2026
@johnnynunez

Copy link
Copy Markdown
Contributor Author

Note: the two SM120 backward-enablement fixes carried here (dQ_single_wg, softmax_scale None→Float) are now also submitted as the standalone PR #2658. This PR can be rebased on top of #2658 once it lands so the diff here is dropout-only.

@thad0ctor

Copy link
Copy Markdown
Contributor

The dropout feature itself is not in #2634. But #2657 is stacked on the same small SM120 fixes from
2655/#2658, so its non-dropout parts overlap/conflict with #2634 if that is adopted by the repo lords

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
johnnynunez force-pushed the feat/cute-dropout-2436-full branch from 1bb944f to 5a54626 Compare June 28, 2026 18:08
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.
@johnnynunez

johnnynunez commented Jul 7, 2026 •

Copy link
Copy Markdown
Contributor Author

viz @drisspg @v0i0 @tridao
To me it is an important feature to finish adoption for others repositories that uses FA C++ with dropout

@reubenconducts

Copy link
Copy Markdown
Contributor

Hi @johnnynunez, can you please add tests? Benchmarks would also be helpful for the PR.

@blake-snc

Copy link
Copy Markdown
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.

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants