diff --git a/python/cudnn/deepseek_sparse_attention/sparse_attention_backward/dsa_bwd_sm100.py b/python/cudnn/deepseek_sparse_attention/sparse_attention_backward/dsa_bwd_sm100.py index 122d2a2de..5d78f31f0 100644 --- a/python/cudnn/deepseek_sparse_attention/sparse_attention_backward/dsa_bwd_sm100.py +++ b/python/cudnn/deepseek_sparse_attention/sparse_attention_backward/dsa_bwd_sm100.py @@ -115,6 +115,14 @@ def __init__( barrier_id=9, num_threads=(self.num_reduce_warps + 1) * self.threads_per_warp, ) + # Orders the reduce warps dKV2/dKV3 T2R reads before the MMA warp + # overwrites those TMEM columns with the next iteration dKV0/dKV1 + # (dKV2 aliases dKV0, dKV3 aliases dKV1). barrier_id 9 is taken by + # tmem_dealloc_barrier above, so this uses the next free slot. + self.t2r_dKV23_done_barrier = pipeline.NamedBarrier( + barrier_id=10, + num_threads=(self.num_reduce_warps + 1) * self.threads_per_warp, + ) self.tmem_S_offset = 0 self.tmem_dP_offset = 0 @@ -1506,6 +1514,18 @@ def mma( mma_reduce_dKV_pipeline.producer_acquire(mma_reduce_dKV_producer_state) # dKV0 + # Wait for the reduce warps to finish the T2R of dKV2/dKV3 from + # the previous iteration before overwriting their TMEM columns: + # dKV0 aliases dKV2 and dKV1 aliases dKV3, and the producer_acquire + # above only orders against the consumer_release from two + # generations back (the dKV4 generation), not against the dKV2/dKV3 + # reads of the previous generation. + if cutlass.const_expr(not self.same_hdim_kv): + if not is_first_mma: + self.t2r_dKV23_done_barrier.arrive_and_wait() + if cutlass.const_expr(self.same_hdim_kv): + if not is_first_mma: + self.t2r_dKV4_done_barrier.arrive_and_wait() dOP_tiled_mma.set(tcgen05.Field.ACCUMULATE, False) for k_block in cutlass.range(0, cute.size(tdKVrP, mode=[2]), unroll=2): cute.gemm( @@ -1518,9 +1538,6 @@ def mma( dOP_tiled_mma.set(tcgen05.Field.ACCUMULATE, True) # dKV1 - if cutlass.const_expr(self.same_hdim_kv): - if not is_first_mma: - self.t2r_dKV4_done_barrier.arrive_and_wait() dOP_tiled_mma.set(tcgen05.Field.ACCUMULATE, False) for k_block in cutlass.range(0, cute.size(tdKVrP, mode=[2]), unroll=2): cute.gemm( @@ -1649,10 +1666,12 @@ def mma( load_mma_K_consumer_state.advance() # Gemm dKV = dO @ P part2 - # Wait for reduce warps to finish T2R of dKV4 from TMEM, - # since dKV2 shares the same TMEM offset as dKV4/dKV0. + # Wait for reduce warps to finish T2R of the TMEM columns that + # dKV2/dKV3 overwrite: dKV4 (not same_hdim) or dKV0/dKV1 (same_hdim). if cutlass.const_expr(not self.same_hdim_kv): self.t2r_dKV4_done_barrier.arrive_and_wait() + else: + self.t2r_dKV01_done_barrier.arrive_and_wait() mma_reduce_dKV_pipeline.producer_acquire(mma_reduce_dKV_producer_state) # dKV2 dOP_tiled_mma.set(tcgen05.Field.ACCUMULATE, False) @@ -1667,8 +1686,6 @@ def mma( dOP_tiled_mma.set(tcgen05.Field.ACCUMULATE, True) # dKV3 - if cutlass.const_expr(self.same_hdim_kv): - self.t2r_dKV01_done_barrier.arrive_and_wait() dOP_tiled_mma.set(tcgen05.Field.ACCUMULATE, False) for k_block in cutlass.range(0, cute.size(tdKVrP, mode=[2]), unroll=2): cute.gemm( @@ -1716,6 +1733,10 @@ def mma( if cutlass.const_expr(self.same_hdim_kv): self.t2r_dKV4_done_barrier.arrive_and_wait() + else: + # Balance the reduce warps' final t2r_dKV23_done arrive (the + # in-loop wait is skipped on the first iteration). + self.t2r_dKV23_done_barrier.arrive_and_wait() mma_compute_dQ_pipeline.producer_commit(mma_compute_dQ_producer_state) mma_compute_dQ_producer_state.advance() @@ -2183,21 +2204,15 @@ def reduce_dKV( mma_reduce_dKV_pipeline.consumer_wait(mma_reduce_dKV_consumer_state) - if cutlass.const_expr(not self.same_hdim_kv): - # Split T2R and atomic_add: load dKV0/dKV1 from TMEM to registers, - # then signal MMA warp that TMEM is free for dKV4 overlap. - rdKV0 = self.t2r_dKV(tdKVtdKV0) - rdKV1 = self.t2r_dKV(tdKVtdKV1) - cute.arch.fence_view_async_tmem_load() - self.t2r_dKV01_done_barrier.arrive_and_wait() - self.reduce_dKV_from_reg(mdKV_acc, rdKV0, rTopkIdx, 0) - self.reduce_dKV_from_reg(mdKV_acc, rdKV1, rTopkIdx, 1) - else: - rdKV1 = self.t2r_dKV(tdKVtdKV1) - cute.arch.fence_view_async_tmem_load() - self.t2r_dKV01_done_barrier.arrive_and_wait() - self.store_dKV(mdKV_acc, tdKVtdKV0, rTopkIdx, 0) - self.reduce_dKV_from_reg(mdKV_acc, rdKV1, rTopkIdx, 1) + # Split T2R and atomic_add: load dKV0/dKV1 from TMEM to registers, + # then signal MMA warp that TMEM is free to overwrite + # (dKV4 if not same_hdim_kv, dKV2/dKV3 if same_hdim_kv). + rdKV0 = self.t2r_dKV(tdKVtdKV0) + rdKV1 = self.t2r_dKV(tdKVtdKV1) + cute.arch.fence_view_async_tmem_load() + self.t2r_dKV01_done_barrier.arrive_and_wait() + self.reduce_dKV_from_reg(mdKV_acc, rdKV0, rTopkIdx, 0) + self.reduce_dKV_from_reg(mdKV_acc, rdKV1, rTopkIdx, 1) mma_reduce_dKV_pipeline.consumer_release(mma_reduce_dKV_consumer_state) mma_reduce_dKV_consumer_state.advance() @@ -2219,14 +2234,27 @@ def reduce_dKV( mma_reduce_dKV_pipeline.consumer_wait(mma_reduce_dKV_consumer_state) if cutlass.const_expr(self.same_hdim_kv): + # Split T2R and atomic_add: load dKV2/dKV3 from TMEM to registers, + # then signal MMA warp that TMEM is free for next tile's dKV0/dKV1. + rdKV2 = self.t2r_dKV(tdKVtdKV2) rdKV3 = self.t2r_dKV(tdKVtdKV3) cute.arch.fence_view_async_tmem_load() self.t2r_dKV4_done_barrier.arrive_and_wait() - self.store_dKV(mdKV_acc, tdKVtdKV2, rTopkIdx, 2) + self.reduce_dKV_from_reg(mdKV_acc, rdKV2, rTopkIdx, 2) self.reduce_dKV_from_reg(mdKV_acc, rdKV3, rTopkIdx, 3) else: - self.store_dKV(mdKV_acc, tdKVtdKV2, rTopkIdx, 2) - self.store_dKV(mdKV_acc, tdKVtdKV3, rTopkIdx, 3) + # T2R dKV2/dKV3 into registers, then signal the MMA warp that + # the TMEM columns are free for the next iteration's dKV0/dKV1 + # before doing the (slow) global atomic reduction. Reading + # them from TMEM inside store_dKV after this point would race + # the next iteration's dKV0/dKV1 overwrites, which nothing + # else orders against these reads. + rdKV2 = self.t2r_dKV(tdKVtdKV2) + rdKV3 = self.t2r_dKV(tdKVtdKV3) + cute.arch.fence_view_async_tmem_load() + self.t2r_dKV23_done_barrier.arrive_and_wait() + self.reduce_dKV_from_reg(mdKV_acc, rdKV2, rTopkIdx, 2) + self.reduce_dKV_from_reg(mdKV_acc, rdKV3, rTopkIdx, 3) mma_reduce_dKV_pipeline.consumer_release(mma_reduce_dKV_consumer_state) mma_reduce_dKV_consumer_state.advance()