From 31b7fed473fd7a2024780b1461ef33fe649bf1f5 Mon Sep 17 00:00:00 2001 From: Vedaanta Agarwalla Date: Thu, 13 Aug 2026 01:55:29 -0700 Subject: [PATCH] sdpa fp8 sm107: fused LDTM row-max + row-sum-in-MMA (baked) Two Rubin levers, baked into the SM107 sibling (the Blackwell kernel is untouched; the only shared-file change is a guarded is_exclusive parameter on tile_dsl.tmem_alloc that <=512-column callers never trace): 1. Fused LDTM row-max: the unmasked softmax path loads S_acc and reduces the row max in ONE tcgen05.ld.red.f32.max per 64-column chunk (tmem_load_max_reduction_x64, the mechanism the MXFP8 kernel uses on cc10.3) instead of a manual load + software reduction; masked iters keep the software path (the fused max reduces before a mask could apply). 2. Row-sum in MMA: an N=16 ones-BMM (P x ones) accumulates the softmax denominator into 16 Sigma TMEM columns right after each sub-tile's O, where the correction warp's per-iter alpha-rescale keeps it online-softmax-consistent for free (rescale span widens 128 -> 144); the softmax warps stop reading P back through the register file to sum it, and the correction epilogue reads Sigma (one ld x1) instead of the stats sum word. B rides a constant 8-row x 128 B K-major all-ones tile written once at init (any N-split of all-ones is all-ones, so the cga2 halves are symmetric for free). Sigma now sums the fp8-QUANTIZED P, exactly matching O's numerator; empty-kv tiles guard the stale Sigma read via max == -inf (LSE = -inf convention kept). TMEM grows to the full 576-col sm107 EXCLUSIVE grant (dealloc must return the grant - a smaller dealloc leaks columns as cudaErrorTensorMemoryLeak). Same-session A/B on w2u1g (cc10.7, B8 H16 D128 e4m3->bf16 dense, vs the sibling at #579): 16k nomask 4.017 -> 3.471 ms (-13.6%), 16k causal 2.387 -> 2.070 (-13.3%), 65k nomask 60.780 -> 51.842 (-14.7%, 5430 TFLOPS eff), 65k causal 31.777 -> 27.013 (-15.0%). Board suite 34/34; SM100 box 34/34 (Blackwell untouched). Co-Authored-By: Claude Fable 5 --- python/cudnn/frost/tile_dsl/tmem.py | 29 ++-- .../fwd/kernels/prefill_d128_fp8_sm107.py | 157 ++++++++++++++---- 2 files changed, 144 insertions(+), 42 deletions(-) diff --git a/python/cudnn/frost/tile_dsl/tmem.py b/python/cudnn/frost/tile_dsl/tmem.py index 75272d4e7..9d00d9d3b 100644 --- a/python/cudnn/frost/tile_dsl/tmem.py +++ b/python/cudnn/frost/tile_dsl/tmem.py @@ -8,16 +8,25 @@ @cute.jit -def tmem_alloc(tmem_ptr_i32, num_cols: int, cta_group_kind): - # No is_exclusive: it only applies to >512-column allocations on sm_107/sm_109 - # (cute.arch.tmem.is_tmem_allocation_exclusive), so it is a no-op on the - # <=512-column sm100/sm103 allocations here, and it is fenced out of the public - # cutlass-dsl tcgen05_alloc wrapper. - nvvm.tcgen05_alloc( - tmem_ptr_i32, - cutlass.Int32(num_cols), - group=cta_group_kind, - ) +def tmem_alloc(tmem_ptr_i32, num_cols: int, cta_group_kind, is_exclusive: cutlass.Constexpr[bool] = False): + # is_exclusive is the sm_107/sm_109 >512-column mechanism (GR100's 576-col + # TMEM; cute.arch.tmem.is_tmem_allocation_exclusive) — pass True only + # there. It is fenced out of the PUBLIC cutlass-dsl tcgen05_alloc + # wrapper, so the kwarg is only emitted on the True branch + # (internal-wheel-only callers); <=512-column allocations never trace it. + if cutlass.const_expr(is_exclusive): + nvvm.tcgen05_alloc( + tmem_ptr_i32, + cutlass.Int32(num_cols), + is_exclusive=True, + group=cta_group_kind, + ) + else: + nvvm.tcgen05_alloc( + tmem_ptr_i32, + cutlass.Int32(num_cols), + group=cta_group_kind, + ) nvvm.tcgen05_relinquish_alloc_permit(group=cta_group_kind) nvvm.bar_warp_sync(cute.arch.FULL_MASK) diff --git a/python/cudnn/sdpa/fwd/kernels/prefill_d128_fp8_sm107.py b/python/cudnn/sdpa/fwd/kernels/prefill_d128_fp8_sm107.py index 1a8852c59..f12298e15 100644 --- a/python/cudnn/sdpa/fwd/kernels/prefill_d128_fp8_sm107.py +++ b/python/cudnn/sdpa/fwd/kernels/prefill_d128_fp8_sm107.py @@ -124,16 +124,15 @@ SCHED_LPT_L2, ) from cudnn.frost.tile_dsl.pointwise import ( - # SM100: no LDTM.STAT — the MASK_NONE fast path uses manual tcgen05_ld + - # row_max_reduction (see _softmax_kv_body); tmem_load_max_reduction_tile - # is not imported. - row_reduction_pair, + # SM107 fuses the MASK_NONE S load + row-max into one LDTM.STAT + # (tmem_load_max_reduction_x64); masked iters keep the software reduction. row_max_reduction, + tmem_load_max_reduction_x64, vec_scale_pair, fp32_to_fp8_pack, ) from cudnn.frost.tile_dsl.regtile import RegTile, vec_concat -from cudnn.frost.tile_dsl.mma import mma_ss, mma_ts_step +from cudnn.frost.tile_dsl.mma import mma_ss, mma_ts, mma_ts_step from cudnn.frost.tile_dsl.tma import tma_load_tile, tma_store_tile, tma_store_commit, tma_store_wait from cudnn.frost.tile_dsl.handles import MmaDesc, SmemTile, GmemTileTma, tma_slice_runtime_desc from cudnn.frost.tile_dsl.tmem import tmem_alloc, tmem_dealloc @@ -253,7 +252,33 @@ class KernelTmemLayout: STATS_STRIDE: int = 128 -LAYOUT = KernelTmemLayout() +# Row-sum-in-MMA: O widens to 144 per sub-tile (128 O + 16 Sigma); the sm107 +# EXCLUSIVE TMEM allocation grants all 576 columns regardless of the request +# and dealloc must return the grant, so TOTAL_COLS carries it. +LAYOUT = KernelTmemLayout(TOTAL_COLS=576, O1_OFF=400) + +# Sigma columns ride right after each sub-tile's O so the correction warp's +# alpha-rescale covers them in the same contiguous span. +_SUM0_OFF = LAYOUT.O0_OFF + CFG.TILE_O +_SUM1_OFF = LAYOUT.O1_OFF + CFG.TILE_O + +# fp8 1.0 bit pattern x4 for the ones tile (e4m3 0x38 / e5m2 0x3C). +_ONES_I32 = 0x38383838 if CFG.DTYPE_QKV == 0 else 0x3C3C3C3C + +# Ones-tile geometry. Rows = the per-CTA N-half of the Sigma MMA's N=16 B +# operand (parametrized on CTA_MMA so a future cga1 config supplies all 16 +# rows from real ones instead of walking the descriptor past the tile into +# neighboring SMEM — functionally masked today only because the epilogue +# reads Sigma column 0 alone; do not rely on that). Row bytes = one +# Q-swizzle atom, because the tile borrows the K-stage descriptor constants +# (SMEM_LAYOUT_QKO / STRIDE_BYTE_OFFSET_QK), which assume that width; the +# sum-MMA's K extent (TILE_N * BPE) must fit inside it — over-provisioning +# is safe by construction (all-ones beyond K is never read), under- +# provisioning never is. +_ONES_ROWS = 16 // CFG.CTA_MMA +_ONES_ROW_BYTES = CFG.Q_SWZ_BYTES +if CFG.TILE_N * CFG.BPE > _ONES_ROW_BYTES: + raise ValueError(f"prefill_d128_fp8_sm107: sum-MMA K ({CFG.TILE_N * CFG.BPE} B) exceeds the ones-tile row ({_ONES_ROW_BYTES} B swizzle atom)") # === Kernel === @@ -294,6 +319,10 @@ def _kernel( sK_raw = cutlass.Array(STORAGE_DTYPE, CFG.STAGES_KV * kBufferElems, alignment=1024, space=cutlass.AddressSpace.smem) sV_raw = cutlass.Array(STORAGE_DTYPE, CFG.STAGES_KV * vBufferElems, alignment=1024, space=cutlass.AddressSpace.smem) sO_raw = cutlass.Array(OUT_STORAGE_DTYPE, CFG.TILES_Q * oBufferElems, alignment=1024, space=cutlass.AddressSpace.smem) + # Constant all-ones B tile for the Sigma MMA — the per-CTA N-half of + # N=16 (_ONES_ROWS) x 128 B, K-major like a K stage; written once at + # init, never touched by TMA. + sOnes_raw = cutlass.Array(STORAGE_DTYPE, _ONES_ROWS * _ONES_ROW_BYTES, alignment=1024, space=cutlass.AddressSpace.smem) sQ = SmemTile( base=sQ_raw, @@ -388,6 +417,17 @@ def _kernel( bars.mb_tmem_dealloc.init() bars.mb_empty_mainloop.init() + # All 32 lanes of warp 0 fill the ones tile (v4 stores, _ONES_ROWS // 4 + # iterations at 512 B per pass); the barrier below orders it before any + # MMA reads it. + if warp_idx == 0: + _ones_vec = cutlass.Vector.from_elements( + (cutlass.Int32(_ONES_I32), cutlass.Int32(_ONES_I32), cutlass.Int32(_ONES_I32), cutlass.Int32(_ONES_I32)), cutlass.Int32 + ) + for _i in cutlass.range_constexpr((_ONES_ROWS * _ONES_ROW_BYTES) // 512): + _ones_ptr = Pointer(sOnes_raw.subview(tidx * cutlass.Int32(16) + cutlass.Int32(_i * 512)).data_ptr(), dtype=cutlass.Int32) + _ones_ptr.store(_ones_vec, alignment=16) + nvvm.fence_mbarrier_init() nvvm.barrier_cta_sync() @@ -478,6 +518,7 @@ def _kernel( sQ=sQ, sK=sK, sV=sV, + sOnes_raw=sOnes_raw, tmem_ptr_i32=tmem_ptr_i32, bars=bars, sched=sched, @@ -497,6 +538,7 @@ def _kernel( sQ=sQ, sK=sK, sV=sV, + sOnes_raw=sOnes_raw, tmem_ptr_i32=tmem_ptr_i32, bars=bars, sched=sched, @@ -880,7 +922,7 @@ def _mma_warp_quiet(tmem_ptr_i32, bars): reads through the cluster crossbar during collective MMA. """ # All-lanes warp-collective ops — NO elect_sync gating. - tmem_alloc(tmem_ptr_i32, LAYOUT.TOTAL_COLS, CTA_GROUP_KIND) + tmem_alloc(tmem_ptr_i32, LAYOUT.TOTAL_COLS, CTA_GROUP_KIND, is_exclusive=True) # Must match lead MMA's +1 arrive on softmax (id=1) and correction (id=2) bars. nvvm.barrier_cta_arrive(1, 32 * (CFG.SOFTMAX_WARPGROUPS * CFG.SOFTMAX_WG_WARPS + 1)) nvvm.barrier_cta_arrive(2, 32 * (CFG.CORRECTION_WARPS + 1)) @@ -896,6 +938,7 @@ def _mma_warp_group( sQ, sK, sV, + sOnes_raw, tmem_ptr_i32, bars, sched, @@ -910,7 +953,7 @@ def _mma_warp_group( Spill-free in OTHER_REGS (40). Non-leader cga2 CTA uses _mma_warp_quiet. """ - tmem_alloc(tmem_ptr_i32, LAYOUT.TOTAL_COLS, CTA_GROUP_KIND) + tmem_alloc(tmem_ptr_i32, LAYOUT.TOTAL_COLS, CTA_GROUP_KIND, is_exclusive=True) nvvm.barrier_cta_arrive(1, 32 * (CFG.SOFTMAX_WARPGROUPS * CFG.SOFTMAX_WG_WARPS + 1)) nvvm.barrier_cta_arrive(2, 32 * (CFG.CORRECTION_WARPS + 1)) @@ -967,6 +1010,40 @@ def _mma_warp_group( desc_Q0 = sQ[0].desc() desc_Q1 = sQ[1].desc() + # Row-sum in MMA: an N=16 ones-BMM (P x ones) accumulates the softmax + # denominator into the Sigma columns beside O. B rides the constant + # K-major ones tile (any N-split of an all-ones matrix is all-ones, so + # the cga2 halves are symmetric for free). + sOnes = SmemTile( + base=sOnes_raw, + elems_per_stage=_ONES_ROWS * _ONES_ROW_BYTES, + stages=1, + leading_byte_offset=LEADING_BYTE_OFFSET_QK, + stride_byte_offset=STRIDE_BYTE_OFFSET_QK, + layout=SMEM_LAYOUT_QKO, + ) + desc_ones = sOnes[0].desc() + idesc_sum = prims.Tcgen05InstrDesc.build( + c_dtype=cutlass.Float32, + a_dtype=STORAGE_DTYPE, + b_dtype=STORAGE_DTYPE, + n_dim=16, + m_dim=CFG.TILE_M * CFG.CTA_MMA, + k_dim=1, + ) + sum_desc = MmaDesc( + M=CFG.TILE_M * CFG.CTA_MMA, + N=16, + K=CFG.TILE_N, + bpe_a=CFG.BPE, + bpe_b=CFG.BPE, + tile_k_hw=CFG.TILE_K_HW_BMM2, + btranspose=False, + cta_group=CFG.CTA_MMA, + idesc=idesc_sum, + kind=MMA_KIND, + ) + if cutlass.const_expr(CFG.MASK_FLAGS == 0): kv_left = cutlass.Int32(0) kv_right = seqlen_kv // cutlass.Int32(CFG.TILE_N) @@ -1067,6 +1144,7 @@ def _mma_warp_group( NUM_KPHASES_PV_PER_CHUNK + local_k, cutlass.Boolean(True), ) + mma_ts(sum_desc, (tmem_raw.subview(LAYOUT.P0_OFF)), desc_ones, (tmem_raw.subview(_SUM0_OFF)), accumulate=is_not_first_bmm2) elect_p = nvvm.elect_sync() bars.mb_bmm2_done[0].arrive(mcast_mask=mcast_mask, cta_group=CFG.CTA_MMA, pred=elect_p) @@ -1092,6 +1170,7 @@ def _mma_warp_group( NUM_KPHASES_PV_PER_CHUNK + local_k, cutlass.Boolean(True), ) + mma_ts(sum_desc, (tmem_raw.subview(LAYOUT.P1_OFF)), desc_ones, (tmem_raw.subview(_SUM1_OFF)), accumulate=is_not_first_bmm2) elect_p = nvvm.elect_sync() bars.mb_bmm2_done[1].arrive(mcast_mask=mcast_mask, cta_group=CFG.CTA_MMA, pred=elect_p) bars.mb_v_empty[old_state.idx].arrive(mcast_mask=mcast_mask, cta_group=CFG.CTA_MMA, pred=elect_p) @@ -1128,6 +1207,7 @@ def _mma_warp_group( NUM_KPHASES_PV_PER_CHUNK + local_k, cutlass.Boolean(True), ) + mma_ts(sum_desc, (tmem_raw.subview(LAYOUT.P0_OFF)), desc_ones, (tmem_raw.subview(_SUM0_OFF)), accumulate=is_not_first_bmm2_epi) elect_p = nvvm.elect_sync() bars.mb_bmm2_done[0].arrive(mcast_mask=mcast_mask, cta_group=CFG.CTA_MMA, pred=elect_p) @@ -1147,6 +1227,7 @@ def _mma_warp_group( NUM_KPHASES_PV_PER_CHUNK + local_k, cutlass.Boolean(True), ) + mma_ts(sum_desc, (tmem_raw.subview(LAYOUT.P1_OFF)), desc_ones, (tmem_raw.subview(_SUM1_OFF)), accumulate=is_not_first_bmm2_epi) elect_p = nvvm.elect_sync() bars.mb_bmm2_done[1].arrive(mcast_mask=mcast_mask, cta_group=CFG.CTA_MMA, pred=elect_p) bars.mb_v_empty[kv_state.idx].arrive(mcast_mask=mcast_mask, cta_group=CFG.CTA_MMA, pred=elect_p) @@ -1266,17 +1347,13 @@ def _softmax_kv_body( for m in chunks_max[1:]: current_max_unscaled = cute.math.max(current_max_unscaled, m) else: - # SM100: manual row-max (no LDTM.STAT / tmem_load_max_reduction_tile) — - # the masked path's pattern sans mask. S_acc is FP32 regardless of dtype. - raw_chunks = [ - nvvm.tcgen05_ld( - "32x32b", - nvvm.make_tmem_ptr(s_addr_base + cutlass.Int32(c * CHUNK), cutlass.Float32), - num=CHUNK, - ) - for c in range(N_CHUNKS) - ] - chunks_max = [row_max_reduction(raw_chunks[c]) for c in range(N_CHUNKS)] + # SM107 has LDTM.STAT: one tcgen05.ld.red.f32.max per 64-col chunk does + # the S_acc load AND the row-max in one op (data regs + max at index + # CHUNK). Unmasked iters only — the fused max reduces before a mask + # could be applied (masked iters take the software path above). + res_chunks = [tmem_load_max_reduction_x64(s_addr_base + cutlass.Int32(c * CHUNK)) for c in range(N_CHUNKS)] + raw_chunks = [cutlass.Vector.from_elements(tuple(r[:CHUNK]), cutlass.Int32).bitcast(cutlass.Float32) for r in res_chunks] + chunks_max = [cutlass.Vector.from_elements((r[CHUNK],), cutlass.Int32).bitcast(cutlass.Float32)[0] for r in res_chunks] reg_S_vec = vec_concat(raw_chunks) current_max_unscaled = chunks_max[0] for m in chunks_max[1:]: @@ -1322,8 +1399,8 @@ def _softmax_kv_body( chunk_S_0 = reg_S[0:CHUNK].vec chunk_P_0 = cute.math.exp2(chunk_S_0, fastmath=True) - # Hoist chunk-0 sum before cast to overlap with cast's FFMA chain. - hoisted_sum = row_reduction_pair(chunk_P_0) + # No RF row-sum: the denominator accumulates in the Sigma TMEM columns + # via the ones-MMA (P is read straight from TMEM by the MMA pipe). chunk_P_0_fp8 = chunk_P_0.to(STORAGE_DTYPE) nvvm.tcgen05_st( "32x32b", @@ -1333,11 +1410,9 @@ def _softmax_kv_body( nvvm.tcgen05_wait(kind=nvvm.Tcgen05Wait.STORE) bars.mb_bmm2_ready[sub_tile_id * N_CHUNKS + 0].arrive(leader_cta_id=leader_cta_id, cta_group=CFG.CTA_MMA) - deferred_P_1 = None if cutlass.const_expr(N_CHUNKS == 2): chunk_S_1 = reg_S[CHUNK : 2 * CHUNK].vec - deferred_P_1 = cute.math.exp2(chunk_S_1, fastmath=True) - chunk_P_1_fp8 = deferred_P_1.to(STORAGE_DTYPE) + chunk_P_1_fp8 = cute.math.exp2(chunk_S_1, fastmath=True).to(STORAGE_DTYPE) nvvm.tcgen05_st( "32x32b", nvvm.make_tmem_ptr( @@ -1352,12 +1427,8 @@ def _softmax_kv_body( if sub_tile_id == 0: nvvm.barrier_cta_sync(barrier_id=8, thread_count=256) - new_p_sum_pair = hoisted_sum - if cutlass.const_expr(N_CHUNKS == 2): - new_p_sum_pair = new_p_sum_pair + row_reduction_pair(deferred_P_1) - alpha_pair = cutlass.Vector.from_elements((alpha, alpha), cutlass.Float32) - total_sum = total_sum * alpha_pair + new_p_sum_pair - + # total_sum stays zeros — the final stats store's sum word is dead + # (correction reads the Sigma column instead). return total_max, total_sum @@ -1595,6 +1666,10 @@ def _correction_warp_group( # O_CHUNK=16 keeps live range short — O_CHUNK=32 spilled correction-warp regs. O_CHUNK = 16 N_CHUNKS_O = CFG.TILE_O // O_CHUNK + # The alpha-rescale must also cover the 16 Sigma columns right after O + # (the running denominator rescales exactly like O); the epilogue store + # loops keep the real TILE_O width. + N_CHUNKS_O_RESCALE = (CFG.TILE_O + 16) // O_CHUNK O_CHUNK_EPI = 64 N_CHUNKS_O_EPI = CFG.TILE_O // O_CHUNK_EPI # D_BLOCK_SIZE must use O_SWZ_B not V_SWZ_B (under cga2 V may drop swizzle). @@ -1670,7 +1745,7 @@ def _correction_warp_group( # vec_scale_pair emits mul_packed_f32x2 → SASS FMUL2. O_HALF = O_CHUNK // 2 if ~all_alpha_one: - for chunk_idx in cutlass.range_constexpr(N_CHUNKS_O): + for chunk_idx in cutlass.range_constexpr(N_CHUNKS_O_RESCALE): o_addr_a = tmem_base_iter + cutlass.Int32(tmem_O_off + chunk_idx * O_CHUNK) o_addr_b = tmem_base_iter + cutlass.Int32(tmem_O_off + chunk_idx * O_CHUNK + O_HALF) o_a = nvvm.tcgen05_ld( @@ -1712,7 +1787,25 @@ def _correction_warp_group( ) nvvm.tcgen05_wait(kind=nvvm.Tcgen05Wait.LOAD) total_max_scaled = stats_vec[0] # log2-units - total_sum = stats_vec[1] + # The denominator lives in the Sigma column (alpha-consistent via + # the shared rescale); the stats word 1 is dead. bmm2_done above + # ordered the last ones-MMA. Empty-kv tiles never ran a BMM2, so + # Sigma is stale there — max == -inf marks them; fall back to 0 + # (the LSE = -inf convention). + _sum_vec = nvvm.tcgen05_ld( + "32x32b", + nvvm.make_tmem_ptr(tmem_base_epi + cutlass.Int32(tmem_O_off + CFG.TILE_O), cutlass.Float32), + num=1, + ) + nvvm.tcgen05_wait(kind=nvvm.Tcgen05Wait.LOAD) + _NEG_INF_EPI = cutlass.Float32(-3.4028235e38) + total_sum = cutlass.Float32( + arith.select( + (total_max_scaled == _NEG_INF_EPI).ir_value(), + cutlass.Float32(0.0).ir_value(), + _sum_vec[0].ir_value(), + ) + ) bars.mb_stat_empty[qs].arrive() # Release MMA's NEXT-tile prologue BMM1 into this S_acc slot: the