Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 8 additions & 2 deletions flash_attn/cute/flash_bwd.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -448,7 +454,7 @@ def __call__(
mCuSeqlensK,
mSeqUsedQ,
mSeqUsedK,
softmax_scale,
softmax_scale_for_scoremod,
softmax_scale_log2,
self.sQ_layout,
self.sK_layout,
Expand Down
7 changes: 7 additions & 0 deletions flash_attn/cute/flash_fwd_sm120.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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,
Expand Down
1 change: 1 addition & 0 deletions flash_attn/cute/interface.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
20 changes: 20 additions & 0 deletions tests/cute/test_sm120.py
Original file line number Diff line number Diff line change
@@ -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