[Example][DeepSeek-V3.2] Adaptive threads for sparse MLA backward - #2592
Merged
LeiWang1999 merged 1 commit intoJul 22, 2026
Merged
LeiWang1999 merged 1 commit into
LeiWang1999 merged 1 commit into
Conversation
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>
Contributor
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configurationConfiguration used: Path: .coderabbit.yaml Review profile: CHILL Plan: Pro Run ID: 📒 Files selected for processing (1)
📝 WalkthroughWalkthroughThe ChangesAdaptive thread configuration
Estimated code review effort: 2 (Simple) | ~5 minutes 🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
Contributor
Author
|
Hi , could you please take a look at this PR when you have a chance, Thanks! @LeiWang1999 @chengyupku @Rachmanino |
LeiWang1999
approved these changes
Jul 22, 2026
3 tasks done
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
Make the thread count of the sparse MLA backward kernel adaptive instead of hard-coded.
bwd(..., threads=...)now defaults toNoneand derives the launch width from the head-block size:256threads (8 warps) whenblock_H >= 64, otherwise128. This kernel is memory-pipe bound, and raising in-flightcp.asyncbytes with 8 warps is the main performance lever. For the defaultdeepseek_v32shape (H=64,block_H=64) the gate resolves to256, so behavior is unchanged.Motivation
This operator is the sole hotspot of the GLM-5 DSA (DeepSeek Sparse Attention) training backward path (
SparseMLA.autograd.Function), computing sparse MLA gradients for the latent geometry (DQKV=576,DV=512,topk=2048).ncuprofiling on H200 (sm_90) shows it is severely memory-pipe bound: L1/TEX ≈ 81–83% SOL, DRAM < 1%, achieved occupancy ≈ 12.2% (structurally pinned by registers + shared memory). Occupancy cannot be raised and DRAM is not the bottleneck, so the only legal lever is increasing parallelism to hide KV-gather latency.Changes
bwd(..., threads=...)default changed from256toNone.block_His computed, select the launch width from the head-block size:8 lines changed (+8/−1); no numerical logic is touched.
Mechanism
The bottleneck is waiting for KV-gather to land. Moving a block from 4 warps (128 threads) to 8 warps (256 threads) doubles the concurrent in-flight
cp.async/ LDGSTS transfers, i.e. it greatly raises in-flight bytes. By Little's law, effective memory-pipe throughput ≈ in-flight bytes / gather latency: the gather latency is fixed, so maximizing in-flight bytes overlaps that fixed latency more fully, and the long latency is hidden behind more in-flight warps.Constraint: an 8-warp GEMM warp-tiling needs
block_H >= 64to fill the M dimension; smaller head blocks cannot fill it and fall back to 128 threads so tiny shapes still build.Applicability (Tensor Parallelism)
num_attention_heads=64and no attention-TP split, the kernel sees the fullH=64,block_H=64lands exactly in the 256-thread tier, and in-flight bytes are maximized.Hshrinks,block_H < 64forces the 128-thread fallback, so this lever no longer applies — but the gate keeps such shapes compiling and correct.Benchmark
Kernel-only
bwd, CUPTI backend,H=64, HKV=1, DQKV=576, DV=512, topk=2048, bf16(this shape resolves to 256 threads). Comparing the adaptive256against the128fallback, the relative speedup is a stable ~3.6× on sm_90 Hopper, consistent with the reference path (~3.6×) in the H200ncuanalysis:Environment
Compatibility
deepseek_v32shape (H=64,block_H=64) the gate resolves to256, identical to the previous hard-coded value — no behavior change and no regression. The ~3.6× in the table is relative to the128fallback, and explains why256is chosen whenblock_H >= 64.128fallback for small-head shapes (block_H < 64) so tiny shapes still compile.Tests
test_sparse_mla_bwd(assert_tensors_similar,dq/dkvvs the autograd reference,eps=1e-4) passes.Summary
sparse_mla_bwd.bwdto select thread count adaptively when unspecified:256threads forblock_H >= 64128threads otherwisedeepseek_v32configuration and numerical behavior.test_sparse_mla_bwdpasses; reported H200 kernel-only benchmarking shows approximately 3.6× speedup with 256 versus 128 threads.