Skip to content

[CuTe, SM100] Support 1..128 Q heads in sparse MLA via in-kernel TMA padding - #2882

Closed
drisspg wants to merge 1 commit into
mainfrom
drisspg/sparse-mla-pad-qheads
Closed

drisspg wants to merge 1 commit into
mainfrom
drisspg/sparse-mla-pad-qheads

Conversation

@drisspg

@drisspg drisspg commented Sep 12, 2026 •

Copy link
Copy Markdown
Collaborator

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_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 now 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 × 24 heads) produced incorrect 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)

  • Sentinel-focused run: 44 passed. test_flash_attn_mla_sparse_bwd_sentinel covers H ∈ {128, 64, 96, 24, 1} × shared_kv × causal × 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 ∈ {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 µ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.

T config fwd us fwd TFLOPS fwd TB/s bwd us bwd TFLOPS fwd+bwd us peak MB
4096 flash H=64 in-kernel pad 859 400 3.75 1874 458 2733 2114
4096 flash H=64 gmem pad->128 1061 324 3.04 2835 303 3896 4557
4096 H=128, flash topk 512 918 749 4.10 2465 697 3382 4046
4096 pro H=128 topk 1024 1349 917 4.38 3798 814 5147 5109
8192 flash H=64 in-kernel pad 1666 412 3.87 3702 464 5368 4095
8192 H=128, flash topk 512 1775 774 4.24 4876 705 6651 7959
8192 pro H=128 topk 1024 2650 933 4.46 7539 820 10189 10087
16384 flash H=64 in-kernel pad 3293 417 3.91 7347 468 10640 8061
16384 H=128, flash topk 512 3506 784 4.29 9723 707 13229 15789
16384 pro H=128 topk 1024 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. 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).

Configuration fwd+bwd us
H=64 in-kernel padding 1416
Native H=128 1740
H=64 gmem padding to 128 1979

…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
@drisspg
drisspg force-pushed the drisspg/sparse-mla-pad-qheads branch from 69c2e8f to 383cbcb Compare September 12, 2026 17:30
@drisspg
drisspg marked this pull request as ready for review September 12, 2026 17:44
@drisspg
drisspg requested a review from jayhshah September 12, 2026 17:44
@chatgpt-codex-connector

chatgpt-codex-connector Bot commented Sep 12, 2026 •

Copy link
Copy Markdown

Codex Review Summary

This comment shows the latest Codex review activity on this pull request.

Review Status Commit Review trigger
📝 Code Review ✅ Completed 2026-09-12T17:47:29.992564Z 383cbcb Draft marked ready
ℹ️ 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" or "@codex security review".

Codex reacts with 👀 while any review is running, comments if it has suggestions, and reacts with 👍 once all reviews finish with no findings.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 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)

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge 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

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Image

asked fables to make a viz for. why this change

@drisspg drisspg closed this Sep 12, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant