Conversation
ea82bcf to
8938aeb
Compare
8938aeb to
69c2e8f
Compare
…padding
Support 1..128 Q heads in the public sparse MLA (DSA) forward and
backward paths. TMA pads heads inside the kernel without padded operand
copies in global memory. This enables training with 64-head models and
tensor-parallel shards.
Each forward tile still covers one token and one top-k gather list:
2-CTA UMMA, 128 packed Q heads, 64 rows per CTA, with K/V split across
the pair. Backward uses a 64- or 128-head tile. The dQ/dQv GEMM keeps
its 128-row tile with a dynamic head extent.
Unchanged: attention math, the 128-head TMA path, MMA tile sizes, and
pipeline synchronization.
Mechanism
---------
padded_qheads_tma_source builds a heads-first view (h, d, s, b) with a
dynamic head extent, so CuTe tiles it without requiring divisibility.
TMA zero-fills out-of-bounds loads and drops out-of-bounds stores.
regroup_padded_qheads restores the packed coordinate layout
((tile, s), d, 1, b).
LSE, row_max, and learnable-sink accesses guard the packed head index.
Without guards, the hierarchical packed layout wraps padded heads into
the next token.
TMA loads zero padded Q rows in forward and dO/P rows in backward.
These rows produce dS = 0 and contribute nothing to dK/dV. Only
scaleP/dPsum need tile-width global buffers with finite padding.
Two latent bugs
---------------
1. Backward dV staging at tile_m=64. The split-0 epilogue could
overwrite dOt/Qvt operands still used by the split-1 MMA, producing
NaNs in dV[:, 256:]. If staging equals one operand stage, each split
stages in place over its consumed stage. Otherwise, both splits take
turns using one staging area allocated after the operands.
2. Non-power-of-two preprocess tiles. A 48-row tile (2 tokens x
24 heads) produced incorrect dpsum. Padded counts use trivial packing
(qhead_per_kvhead=1, nheads_kv=H) with the existing 128-row tile. Not
all head counts divide that tile; resizing it to a non-power-of-two
is not a valid workaround.
Validation (GB200)
------------------
- Sentinel-focused run: 44 passed.
test_flash_attn_mla_sparse_bwd_sentinel covers
H in {128, 64, 96, 24, 1} x shared_kv x causal x 2 lengths.
- test_flash_attn_mla_absorbed: 96 BF16 cases passed for 1024-1024 and
255-1025, including the 128-head path.
- Probe against global-memory padding to 128 heads,
H in {1, 8, 16, 24, 48, 64, 72, 96, 100, 120}:
out/LSE/P/row_max/dQ/dQv bit-identical; dK/dV within 1e-5.
Performance at DeepSeek-V4 shapes
--------------------------------
Measurement contract: GB200, B=1, D_qk = D_v = 512, shared KV, sliding
window 128 + compressed causal top-k, via attention-gym
selected_attention. Flash uses top-k 512; pro uses top-k 1024. Times are
medians over CUDA-graph replays, in us per call (lower is better). Peak
memory is MB (lower is better).
Calculated rates (higher is better):
- fwd FLOPs = 2*T*H*K*(D_qk + D_v).
- K = valid gathered slots (top-k + window); bwd FLOPs = 2.5x fwd.
- TB/s counts gathered KV rows plus Q read and O write, not measured
DRAM traffic.
Configurations: I64 = flash H=64 in-kernel padding,
G64 = flash H=64 gmem padding to 128, F128 = flash H=128 top-k 512,
P128 = pro H=128 top-k 1024. TFf/TFb are fwd/bwd TFLOPS;
TBf is fwd TB/s. Fwd, bwd, and sum are us; peak is MB.
T cfg fwd TFf TBf bwd TFb sum peak
----- ---- ----- --- ---- ----- --- ----- -----
4096 I64 859 400 3.75 1874 458 2733 2114
4096 G64 1061 324 3.04 2835 303 3896 4557
4096 F128 918 749 4.10 2465 697 3382 4046
4096 P128 1349 917 4.38 3798 814 5147 5109
8192 I64 1666 412 3.87 3702 464 5368 4095
8192 F128 1775 774 4.24 4876 705 6651 7959
8192 P128 2650 933 4.46 7539 820 10189 10087
16384 I64 3293 417 3.91 7347 468 10640 8061
16384 F128 3506 784 4.29 9723 707 13229 15789
16384 P128 5262 940 4.49 15029 823 20292 20045
At T=4096, H=64 in-kernel padding reduces fwd+bwd time by about 30% and
peak memory by about half versus global-memory padding. H=64 forward
takes about 93% of the H=128 time: it executes the same MMA work, half
of it on padded rows, so useful TFLOPS halve. NCU on the H=64 forward
(T=4096): DRAM throughput 8%, L2 throughput 14%, tensor pipe 38% active,
issue slots 20% busy, 80% of cycles with no eligible warp, 74% of stall
cycles long-scoreboard on the per-thread cp.async gather, 1 CTA/SM
(215 KB smem). Neither memory nor the tensor pipe is saturated; the
pipeline is dependency-limited. Whether a native 64-row forward tile
would recover the padded MMA time is not established here.
Separate T=2048 comparison: GB200, top-k 512, D=512, fwd+bwd via
attention-gym selected_attention; time in us per call (lower is better).
Configuration fwd+bwd us
-------------------------- ----------
H=64 in-kernel padding 1416
Native H=128 1740
H=64 gmem padding to 128 1979
69c2e8f to
383cbcb
Compare
Codex Review SummaryThis comment shows the latest Codex review activity on this pull request.
ℹ️ About Codex in GitHubYour team has set up Codex to review pull requests in this repo. Reviews are triggered when you
Codex reacts with 👀 while any review is running, comments if it has suggestions, and reacts with 👍 once all reviews finish with no findings. |
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: 383cbcba56
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
| # Finite tile-width padding: see pack_gqa.padded_qheads_tma_source. | ||
| heads_shape = (batch_size, seqlen_q) if cu_seqlens_q is None else (total_q,) | ||
| alloc = torch.zeros if pad_qheads else torch.empty | ||
| dpsum = alloc(*heads_shape, qhead_per_kvhead, dtype=torch.float32, device=device) |
There was a problem hiding this comment.
Slice padded dpsum before reducing the sink gradient
When sparse MLA uses fewer than 64/128 heads together with learnable_sink and backward, this allocation gives dpsum the padded tile width, but the unchanged _bwd_dsink_reduce reshapes it using the real learnable_sink.shape[0] and iterates over all resulting rows while indexing the unpadded lse. This makes the reduction read past lse (and can produce an incorrect sink gradient or an illegal GPU access); pass only dpsum[..., :nheads] to the reducer or otherwise give the reducer the valid-head extent.
Useful? React with 👍 / 👎.
| self.num_stages_dV = 2 # == hdimv splits, for Umma <-> Async | ||
| self.num_epi_stages_dV = 8 # == 2 splits x 4 slots/split | ||
| # Per hdimv split: 2 warpgroup halves x 2 subtile parities. | ||
| self.num_epi_stages_dV = 4 |

Human Note
Support [1, 128] Q heads in the public sparse MLA (DSA) forward and backward paths. I saw that there was a comment for supporting n_qheads = 64 not sure if you planned to do something fancier here. I did the easy thing (just pad out). This does infact come up both for dsv4 flash and when doing tp sharding. I found this when trying to integrate into torchtitan. I consider this a naive unblock but it does help :)
I thought about the multi tokens e.g. true pack-gqa but didnt want to visit the union + masking (although ill prototype this) -> same reason why I didnt support this for blocksparse impl
Agent Notes
TMA pads heads inside the kernel without padded operand copies in global memory. This enables training with 64-head models and tensor-parallel shards.
Each forward tile still covers one token and one top-k gather list: 2-CTA UMMA, 128 packed Q heads, 64 rows per CTA, with K/V split across the pair. Backward uses a 64- or 128-head tile. The dQ/dQv GEMM keeps its 128-row tile with a dynamic head extent.
Unchanged: attention math, the 128-head TMA path, MMA tile sizes, and pipeline synchronization.
Mechanism
padded_qheads_tma_sourcebuilds a heads-first view(h, d, s, b)with a dynamic head extent, so CuTe tiles it without requiring divisibility. TMA zero-fills out-of-bounds loads and drops out-of-bounds stores.regroup_padded_qheadsrestores the packed coordinate layout((tile, s), d, 1, b).dS = 0and contribute nothing to dK/dV. OnlyscaleP/dPsumneed tile-width global buffers with finite padding.Two latent bugs
tile_m=64. The split-0 epilogue could overwrite dOt/Qvt operands still used by the split-1 MMA, producing NaNs indV[:, 256:]. If staging equals one operand stage, each split now stages in place over its consumed stage. Otherwise, both splits take turns using one staging area allocated after the operands.dpsum. Padded head counts now use trivial packing (qhead_per_kvhead=1, nheads_kv=H) with the existing 128-row tile. Not all head counts divide that tile; resizing it to a non-power-of-two is not a valid workaround.Validation (GB200)
test_flash_attn_mla_sparse_bwd_sentinelcovers H ∈ {128, 64, 96, 24, 1} × shared_kv × causal × 2 lengths.test_flash_attn_mla_absorbed: 96 BF16 cases passed for1024-1024and255-1025, including the 128-head path.Performance at DeepSeek-V4 shapes
Measurement contract: GB200, B=1, D_qk = D_v = 512, shared KV, sliding window 128 + compressed causal top-k, via attention-gym
selected_attention. Flash uses top-k 512; pro uses top-k 1024. Times are medians over CUDA-graph replays, in µs per call (lower is better). Peak memory is MB (lower is better).Calculated rates (higher is better): fwd FLOPs = 2THK(D_qk + D_v), K = valid gathered slots (top-k + window); bwd FLOPs = 2.5× fwd. TB/s counts gathered KV rows plus Q read and O write.
At T=4096, H=64 in-kernel padding reduces fwd+bwd time by about 30% and peak memory by about half versus global-memory padding. H=64 forward takes about 93% of the H=128 time. NCU on the H=64 forward (T=4096): DRAM throughput 8%, L2 hit rate 74%, 74% of stall cycles are long-scoreboard waits on the per-thread cp.async gather. The forward is gather-latency-bound, so the padded MMA rows cost little; the gather itself is the lever for a follow-up.
Separate T=2048 comparison: GB200, top-k 512, D=512, fwd+bwd via attention-gym
selected_attention; time in µs per call (lower is better).