Skip to content

Commit 6f75a92

Browse files
committed
[Example] Add adaptive thread selection to sparse MLA backward
Default `bwd(..., threads=...)` to None and derive the launch width from the head-block size instead of hard-coding 256. This kernel is memory-pipe bound, so its throughput is set by how many cp.async copies stay in flight, not by raw DRAM bandwidth. When the head count is not split across blocks (block_H>=64), 256 threads (8 warps) issue far more concurrent cp.async transfers than 128 (4 warps), greatly raising in-flight bytes -- the main perf lever for this shape. Smaller head blocks lack enough GEMM warp-tiling work to fill 256 threads, so they fall back to 128 and still build. For the deepseek_v32 shape (H=64, block_H=64) this resolves to 256, matching the previous behavior. Signed-off-by: Butterfingrz <13524387014@163.com>
1 parent 3b37333 commit 6f75a92

1 file changed

Lines changed: 8 additions & 1 deletion

File tree

‎examples/deepseek_v32/sparse_mla_bwd.py‎

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -91,7 +91,7 @@ def bwd(
9191
is_causal=True,
9292
block_size=32,
9393
num_stages=0,
94-
threads=256,
94+
threads=None,
9595
indices_dtype=T.int32,
9696
dtype=T.bfloat16,
9797
accum_dtype=T.float32,
@@ -122,6 +122,13 @@ def bwd(
122122
H = H_kv
123123
padded_H = max(tilelang.math.next_power_of_2(H_kv), 16)
124124
block_H = min(64, padded_H)
125+
# adaptive: this kernel is memory-pipe bound. When the head count is not split across
126+
# blocks (block_H>=64), 256 threads (8 warps) issue many more cp.async copies at once
127+
# than 128 (4 warps), greatly raising in-flight bytes -- the main perf lever here.
128+
# Smaller head blocks lack the GEMM warp-tiling work to fill 256 threads, so fall back
129+
# to 128 and still build.
130+
if threads is None:
131+
threads = 256 if block_H >= 64 else 128
125132
assert padded_H % block_H == 0
126133
NH = padded_H // block_H
127134
BS = block_size

0 commit comments

Comments
 (0)