From b417fd266d18977252ddee6fa60c74339a0bab89 Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Sun, 19 Apr 2026 05:59:33 +0000 Subject: [PATCH] [Cute,Sm120] fix Spark forward and backward regressions Restore the intended SM80-style control flow for SM120 forward and initialize the SM120 backward config so FA4 compiles and runs end to end on DGX Spark. Keep the shared SM80/SM120 backward launcher on concrete softmax-scale values to avoid DSL type errors, and add a regression test for the SM120 control-flow selection. Made-with: Cursor --- flash_attn/cute/flash_bwd.py | 10 ++++++++-- flash_attn/cute/flash_fwd_sm120.py | 7 +++++++ flash_attn/cute/interface.py | 1 + tests/cute/test_sm120.py | 20 ++++++++++++++++++++ 4 files changed, 36 insertions(+), 2 deletions(-) create mode 100644 tests/cute/test_sm120.py diff --git a/flash_attn/cute/flash_bwd.py b/flash_attn/cute/flash_bwd.py index eeb7615b1d3..2f35b9da231 100644 --- a/flash_attn/cute/flash_bwd.py +++ b/flash_attn/cute/flash_bwd.py @@ -433,7 +433,13 @@ def __call__( tile_sched_params = TileScheduler.to_underlying_arguments(tile_sched_args) grid_dim = TileScheduler.get_grid_shape(tile_sched_params) - softmax_scale_log2, softmax_scale = utils.compute_softmax_scale_log2(softmax_scale, self.score_mod) + softmax_scale_for_scoremod = softmax_scale + if cutlass.const_expr(self.score_mod is None): + softmax_scale_log2 = softmax_scale * math.log2(math.e) + else: + softmax_scale_log2, softmax_scale_for_scoremod = utils.compute_softmax_scale_log2( + softmax_scale, self.score_mod + ) self.kernel( mQ, mK, @@ -448,7 +454,7 @@ def __call__( mCuSeqlensK, mSeqUsedQ, mSeqUsedK, - softmax_scale, + softmax_scale_for_scoremod, softmax_scale_log2, self.sQ_layout, self.sK_layout, diff --git a/flash_attn/cute/flash_fwd_sm120.py b/flash_attn/cute/flash_fwd_sm120.py index 08d219acfa8..8de984c8ae3 100644 --- a/flash_attn/cute/flash_fwd_sm120.py +++ b/flash_attn/cute/flash_fwd_sm120.py @@ -7,6 +7,7 @@ import cutlass import cutlass.utils as utils_basic +from cutlass.base_dsl.arch import Arch from flash_attn.cute.flash_fwd import FlashAttentionForwardSm80 @@ -16,6 +17,12 @@ class FlashAttentionForwardSm120(FlashAttentionForwardSm80): # The compilation target is determined by the GPU at compile time, not this field. arch = 80 + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + # The base class records the physical GPU arch, but SM120 intentionally + # reuses the SM80 control-flow path and must keep the non-TMA epilogue. + self.arch = Arch.sm_80 + @staticmethod def can_implement( dtype, diff --git a/flash_attn/cute/interface.py b/flash_attn/cute/interface.py index 5b9e382d217..566f2c27cbb 100644 --- a/flash_attn/cute/interface.py +++ b/flash_attn/cute/interface.py @@ -1196,6 +1196,7 @@ def _flash_attn_bwd( causal, window_size_left, window_size_right ) + dQ_single_wg = False if arch // 10 == 12: # SM120: uses SM80 MMA with 99 KB SMEM, 128 threads (4 warps). m_block_size = 64 diff --git a/tests/cute/test_sm120.py b/tests/cute/test_sm120.py new file mode 100644 index 00000000000..c912108a9b7 --- /dev/null +++ b/tests/cute/test_sm120.py @@ -0,0 +1,20 @@ +import cutlass +from cutlass.base_dsl.arch import Arch + +from flash_attn.cute.flash_fwd_sm120 import FlashAttentionForwardSm120 + + +def test_sm120_forward_uses_sm80_control_flow(): + kernel = FlashAttentionForwardSm120( + cutlass.BFloat16, + head_dim=64, + head_dim_v=64, + qhead_per_kvhead=1, + pack_gqa=False, + tile_m=128, + tile_n=64, + num_stages=1, + num_threads=128, + ) + + assert kernel.arch == Arch.sm_80