[kernel] sm_100a backward: drop tcgen05.wait::ld after the TMEM loads
The destination registers of tcgen05.ld are scoreboarded, so the consumers wait on their own; the four waits (dV/dK merge, S^T, dP^T and dQ loads) are removed, the fence stays. Same change as the reference kernel in books/. Verified: CPU-reference checks incl. ragged variable_block_sizes, zero-count kv blocks and STRESS_N=5 bitwise reruns; tests/test_block_sparse_bwd_sm100a.py 15 passed, tests/test_block_sparse_sm100a_dispatch.py 10 passed, tests/test_block_sparse_sm100a.py 37 passed (extension rebuilt from this tree). Interleaved A/B 4k..524k: -1.2..+0.9%, neutral. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
This commit is contained in:
@@ -552,7 +552,6 @@ __global__ void __cluster_dims__(1, 1, 1) __launch_bounds__(N_WARPS * 32, 1) vsa
|
||||
stage_tile + (size_t)col_half * (KV_TILE * SUB_COLS_BF16) + kv_row * SUB_COLS_BF16;
|
||||
uint32_t acc_regs[HALF_COLS];
|
||||
tcgen05_ld_32x32b_x64(tmem_acc + tmem_f32_offset, acc_regs);
|
||||
tcgen05_wait_ld();
|
||||
tcgen05_fence_before_thread_sync();
|
||||
const float2* acc2 = reinterpret_cast<const float2*>(acc_regs);
|
||||
const float2 scale2 = f32x2_splat(sm_scale);
|
||||
@@ -638,7 +637,6 @@ __global__ void __cluster_dims__(1, 1, 1) __launch_bounds__(N_WARPS * 32, 1) vsa
|
||||
|
||||
uint32_t st_regs[HALF_COLS];
|
||||
tcgen05_ld_32x32b_x64(tmem_st + tmem_f32_offset, st_regs);
|
||||
tcgen05_wait_ld();
|
||||
tcgen05_fence_before_thread_sync();
|
||||
|
||||
float2* pt_fp32 = reinterpret_cast<float2*>(st_regs);
|
||||
@@ -668,7 +666,6 @@ __global__ void __cluster_dims__(1, 1, 1) __launch_bounds__(N_WARPS * 32, 1) vsa
|
||||
|
||||
uint32_t dpt_regs[HALF_COLS];
|
||||
tcgen05_ld_32x32b_x64(tmem_dpt + tmem_f32_offset, dpt_regs);
|
||||
tcgen05_wait_ld();
|
||||
tcgen05_fence_before_thread_sync();
|
||||
|
||||
const float2* dpt2 = reinterpret_cast<const float2*>(dpt_regs);
|
||||
@@ -787,7 +784,6 @@ __global__ void __cluster_dims__(1, 1, 1) __launch_bounds__(N_WARPS * 32, 1) vsa
|
||||
tcgen05_ld_32x32b_x64(tmem_dq + tmem_lane_base + (uint32_t)(c * 64),
|
||||
reinterpret_cast<uint32_t (&)[64]>(dq_regs[c * 64]));
|
||||
}
|
||||
tcgen05_wait_ld();
|
||||
tcgen05_fence_before_thread_sync();
|
||||
if (elect_one_sync()) {
|
||||
mbarrier_arrive(smem_ptr_u32(empty_bar_dq));
|
||||
|
||||
Reference in New Issue
Block a user