Skip to content

[bug fix] Issue 3700 sparse mla sm121 hang - #4732

Merged
kahyunnam merged 3 commits into
flashinfer-ai:mainfrom
kahyunnam:fix/issue-3700-sparse-mla-sm121-hang
Sep 1, 2026
Merged

kahyunnam merged 3 commits into
flashinfer-ai:mainfrom
kahyunnam:fix/issue-3700-sparse-mla-sm121-hang

Conversation

@kahyunnam

@kahyunnam kahyunnam commented Aug 25, 2026

Copy link
Copy Markdown
Member

📌 Description

Sparse-MLA SM120 prefill uses CTA barrier 1 as an asymmetric handshake (count 384): 256 math threads barrier.cta.arrive (non-blocking), 128 IO threads barrier.cta.sync. Math can arrive for tile ti, run ti+1, and arrive again before IO syncs ti. The PTX ISA warns against exactly that (arrive then another barrier.cta on the same id before reset). The surplus arrival completes the phase without IO; a later phase starves. Captured on GB10: IO warps parked on BAR.SYNC 1, 384 after math had exited.

Fix: alternate the handshake between ids 1 and 5 by tile parity (bar_arrive_alt / bar_sync_alt) at all five call sites. A run-ahead tile lands on the other id. Two ids suffice because math cannot start tile ti until IO has returned from the sync for ti-2, so at most one arrival is outstanding per id.

Ungated: this is an invalid arrival pattern, not an sm_121 quirk. The hang is only demonstrated on GB10; the kernel is SM12x-only.

🔍 Related Issues

Closes #3700

🚀 Pull Request Checklist

✅ Pre-commit Checks

  • I have installed pre-commit by running pip install pre-commit (or used your preferred method).
  • I have installed the hooks with pre-commit install.
  • I have run the hooks manually with pre-commit run --all-files and fixed any reported issues.

🧪 Tests

  • Tests have been added or updated as needed.
  • All tests are passing (unittest, etc.).

GB10 / sm_121: pytest tests/attention/test_sparse_mla_sm120.py309 passed. Single-process soak on heads=128, topk=1024, tokens=256: baseline 5/5 hung (usually tens of iterations, sometimes a few hundred); fix 5/5 × 30k, 0 hangs. Two concurrent pytest (-k "prefill_glm_nsa_arbitrary_fp32 or prefill_dsv4"): baseline wedged in round 3; fix 12/12 clean. All four instantiated prefill variants soaked 30k each (with assert_close); 30k bitwise-identical iterations vs a golden result (fix does not turn the hang into silent corruption). No new unit test — the race is a soak, not a unit.

Header-only change: ninja may serve a stale .so (no .d depfiles). Delete csrc_sparse_mla_sm120_prefill.cuda.o and sparse_mla_sm120.so and confirm the .so mtime advances before A/B. Wedged processes ignore SIGTERM; use timeout -s KILL.

io_bulk_gather_tile relied on the caller's __threadfence_block to separate
io_gather_scales' plain shared stores from the cp.async.bulk issues. That is a
memory fence, not an execution barrier, so the four IO warps were free to skip
past each other. On sm_121 (GB10) this intermittently desynchronizes the
barrier-1 arrive/sync handshake: the math warps exit the kernel while the IO
warps stay blocked on barrier.cta.sync 1, 384, wedging the CTA forever.

Reproduced on GB10 (5/5 trials inside 50 iterations); 150k soak iterations and
309/309 tests pass with the barrier. Left unconditional rather than arch-gated
since nothing in the failure is sm_121-specific.

AI-assisted (Cursor).

Closes flashinfer-ai#3700
The hang is only demonstrated on sm_121 (GB10) and its mechanism is
unresolved, so confine the extra barrier to that arch instead of changing
codegen for discrete sm_120 parts that have never been observed to fail.
Barrier 1 is arrived at non-blocking by the math warps and waited on by the IO
warps. The math side can therefore finish tile ti, arrive, and then — once tile
ti+1's mbarrier completes — run all of iteration ti+1 and arrive again before
the IO warps reach their sync for phase ti. Two arrivals in one phase is
invalid: the phase completes with the wrong participant set and a later phase
starves, deadlocking the CTA (flashinfer-ai#3700).

Double buffering bounds the math lead to exactly one tile, so alternating the
handshake between two barrier ids by tile parity makes the double arrival
impossible. Only the barrier id changes, so the instruction count is unchanged.

Supersedes the extra-IO-barrier workaround, which worked only by delaying
mbarrier completion and thus narrowing the window.
@coderabbitai

coderabbitai Bot commented Aug 25, 2026

Copy link
Copy Markdown
Contributor

Review Change Stack

No actionable comments were generated in the recent review. 🎉

ℹ️ Recent review info
⚙️ Run configuration

Configuration used: defaults

Review profile: CHILL

Plan: Pro Plus

Run ID: 296f51fe-8f6b-4de8-b2a0-3e420c237c99

📥 Commits

Reviewing files that changed from the base of the PR and between 857bc3a and c3cacb1.

📒 Files selected for processing (2)
  • include/flashinfer/attention/sparse_mla_sm120/arch/barrier.cuh
  • include/flashinfer/attention/sparse_mla_sm120/prefill_kernel.cuh

Included review availability: Your plan provides up to 8 included reviews per hour; 7 remain after this review.


📝 Walkthrough

Walkthrough

Changes

Sparse MLA barrier synchronization

Layer / File(s) Summary
Parity-selected barrier helpers
include/flashinfer/attention/sparse_mla_sm120/arch/barrier.cuh
Adds bar_arrive_alt and bar_sync_alt. Each helper selects an even or odd barrier ID from parity.
Prefill kernel integration
include/flashinfer/attention/sparse_mla_sm120/prefill_kernel.cuh
Updates single-group and multi-group IO synchronization and math-loop barrier arrival to use tile-parity-selected barriers.

Estimated code review effort: 3 (Moderate) | ~20 minutes

Merge Risk: ⚪ Minimal · up to c3cac

This localized synchronization fix changes barrier selection for sparse MLA prefill and includes reported validation; no actionable merge-blocking risk remains beyond normal checks and review.

Suggested reviewers: bkryu, lucifer1004

🚥 Pre-merge checks | ✅ 5
✅ Passed checks (5 passed)
Check name Status Explanation
Docstring Coverage ✅ Passed No functions found in the changed files to evaluate docstring coverage. Skipping docstring coverage check. Docstring coverage is scoped to functions touched by this diff. Analyzed 0 functions across 0…
Linked Issues check ✅ Passed The description explicitly states that the pull request closes issue #3700, which matches the reported sparse-MLA timeout and hang.
Out of Scope Changes check ✅ Passed The changes are limited to alternating CTA barrier helpers and their five prefill-kernel call sites. They directly support the stated hang fix.
Title check ✅ Passed The title clearly identifies the bug fix, affected Sparse MLA SM121 behavior, and related issue number. It accurately summarizes the main change.
Description check ✅ Passed The description follows the repository template and explains the barrier-arrival bug, the alternating-barrier fix, related issue, completed checks, and extensive validation results. The optional Revie…
Full details: Docstring Coverage

Explanation

No functions found in the changed files to evaluate docstring coverage. Skipping docstring coverage check. Docstring coverage is scoped to functions touched by this diff. Analyzed 0 functions across 0 files. (2 skipped: 2 unsupported.)

Full details: Description check

Explanation

The description follows the repository template and explains the barrier-arrival bug, the alternating-barrier fix, related issue, completed checks, and extensive validation results. The optional Reviewer Notes section is not required.

✨ Finishing Touches 💡 1
🛠️ Fix failing CI checks 💡
  • Create stacked PR
  • Commit on current branch
🧪 Generate unit tests (beta)
  • Create PR with unit tests

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.

❤️ Share

Comment @coderabbitai help to get the list of available commands.

@kahyunnam

Copy link
Copy Markdown
Member Author

@flashinfer-bot run

@kahyunnam
kahyunnam marked this pull request as ready for review August 26, 2026 21:12
@kahyunnam

Copy link
Copy Markdown
Member Author

/bot run tests/attention

@flashinfer-bot

Copy link
Copy Markdown
Collaborator

GitLab MR !1327 has been created, and the CI pipeline #64745524 is currently running. I'll report back once the pipeline job completes.

@kahyunnam kahyunnam changed the title Fix/issue 3700 sparse mla sm121 hang [bug fix] Issue 3700 sparse mla sm121 hang Aug 26, 2026
@kahyunnam kahyunnam self-assigned this Aug 26, 2026
@kahyunnam

Copy link
Copy Markdown
Member Author

@flashinfer-bot run

@bkryu bkryu left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Thanks @kahyunnam. Approving, but please confirm that the internal CI looks clean!

@kahyunnam

Copy link
Copy Markdown
Member Author

/bot run tests/attention

@flashinfer-bot

Copy link
Copy Markdown
Collaborator

GitLab MR !1327 has been created, and the CI pipeline #64752922 is currently running. I'll report back once the pipeline job completes.

@flashinfer-bot

Copy link
Copy Markdown
Collaborator

[FAILED] Pipeline #64752922 — 15/18 executed test jobs passed

Compared with nightly #64639043 (different CI configuration).

Unit Tests

GPU CUDA 12.9 CUDA 13.0 Notes
B300 ✅ Pass ✅ Pass
GB200 ✅ Pass ✅ Pass
GB300 ✅ Pass ✅ Pass
H100 ✅ Pass ✅ Pass
RTX Pro 6000 Blackwell ✅ Pass ❌ New New: tests.attention.test_trtllm_gen_attention_prefill (14880 failures; CUDA 13.0)
New: tests.attention.test_sliding_window (11792 failures; CUDA 13.0)
New: tests.attention.test_batch_prefill_kernels (8797 failures; CUDA 13.0)
… and 35 more
Spark 🟡 Old 🟡 Old Old: tests.attention.test_fmha_v2_prefill (52 failures; CUDA 12.9, CUDA 13.0)

✅ Pass · 🟡 Old failure · ❌ New failure · ⏱ Test timeout · ⚠️ Infrastructure · ❔ Unknown or unclassified · — Not run

Multi-GPU and Multi-Node Tests — 6/6 passed

GPU CUDA 12.9 CUDA 13.0 Notes
B300 (multi-GPU) ✅ Pass ✅ Pass
GB200 (multi-node) ✅ Pass ✅ Pass
GB300 (multi-node) ✅ Pass ✅ Pass
Failure details

New relative to nightly (attribution uncertain)

  • tests.attention.test_trtllm_gen_attention_prefill — 14880 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program start. Setting the availab…
  • tests.attention.test_sliding_window — 11792 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • failed on setup with "RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program star…
  • tests.attention.test_batch_prefill_kernels — 8797 failures on RTX Pro 6000 Blackwell / CUDA 13.0
  • tests.attention.test_hopper — 6780 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program start. Setting the availab…
  • tests.attention.test_hopper_fp8_attention — 3704 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program start. Setting the availab…
  • tests.attention.test_block_sparse — 3564 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • failed on setup with "RuntimeError: FlashInfer requires GPUs with sm75 or higher"
  • tests.attention.test_batch_decode_kernels — 3262 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • failed on setup with "RuntimeError: FlashInfer requires GPUs with sm75 or higher"
  • tests.attention.test_blackwell_fmha — 3128 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program start. Setting the availab…
  • tests.attention.test_fmha_v2_prefill — 2488 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program start. Setting the availab…
  • tests.attention.test_tensor_cores_decode — 2448 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • failed on setup with "RuntimeError: FlashInfer requires GPUs with sm75 or higher"
  • tests.attention.test_batch_invariant_fa2 — 2016 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • failed on setup with "RuntimeError: FlashInfer requires GPUs with sm75 or higher"
  • tests.attention.test_fp8_prefill — 1605 failures on RTX Pro 6000 Blackwell / CUDA 13.0
    • RuntimeError: CUDA unknown error - this may be due to an incorrectly set up environment, e.g. changing env variable CUDA_VISIBLE_DEVICES after program start. Setting the availab…
  • … and 26 more failing test groups

Pre-existing failures

  • tests.attention.test_fmha_v2_prefill — 52 failures on Spark / CUDA 12.9, Spark / CUDA 13.0
    • ValueError: fmha_v2_prefill_sm120 is only supported on SM120 GPUs; SM121 has not been validated.

@aeichler-ac

Copy link
Copy Markdown
Contributor

Independent check on GB10 (sm_121), PR head c3cacb13 vs main 39b484f1.

Hang shape from the PR description (heads=128, topk=1024, tokens=256, DSv4 prefill):

  • main: hung on seed 0 inside 1500 iters (process wedged on GPU until SIGKILL)
  • this PR: 10 seeds × 1500 iters clean; 5 seeds × 10k clean

pytest tests/attention/test_sparse_mla_sm120.py on the PR: 309 passed.

Looks good from here.

@kahyunnam
kahyunnam merged commit d209081 into flashinfer-ai:main Sep 1, 2026
42 of 58 checks passed
lucifer1004 added a commit to lucifer1004/flashinfer that referenced this pull request Sep 2, 2026
Picks up the flashinfer-ai#4732 SM121 prefill-hang fix: the asymmetric barrier-1
handshake now alternates ids 1/5 by tile parity (bar_arrive_alt /
bar_sync_alt). The old prefill_kernel.cuh was split here into
prefill_{common,mg,swapab}; the five call sites live on in
prefill_mg_kernel.cuh (SG kernel uses the Cfg:: spelling). The swapAB
kernel uses a symmetric full-CTA barrier plus mbarriers, not the
asymmetric pattern, so it does not need the alternation.

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
PetersonGuo pushed a commit to PetersonGuo/flashinfer that referenced this pull request Sep 2, 2026
<!-- .github/pull_request_template.md -->

## 📌 Description

Sparse-MLA SM120 prefill uses CTA barrier 1 as an asymmetric handshake
(count 384): 256 math threads `barrier.cta.arrive` (non-blocking), 128
IO threads `barrier.cta.sync`. Math can arrive for tile `ti`, run
`ti+1`, and arrive **again** before IO syncs `ti`. The PTX ISA warns
against exactly that (`arrive` then another `barrier.cta` on the same id
before reset). The surplus arrival completes the phase without IO; a
later phase starves. Captured on GB10: IO warps parked on `BAR.SYNC 1,
384` after math had exited.

Fix: alternate the handshake between ids **1 and 5** by tile parity
(`bar_arrive_alt` / `bar_sync_alt`) at all five call sites. A run-ahead
tile lands on the other id. Two ids suffice because math cannot start
tile `ti` until IO has returned from the sync for `ti-2`, so at most one
arrival is outstanding per id.

Ungated: this is an invalid arrival pattern, not an sm_121 quirk. The
hang is only demonstrated on GB10; the kernel is SM12x-only.

## 🔍 Related Issues

Closes flashinfer-ai#3700

## 🚀 Pull Request Checklist

### ✅ Pre-commit Checks

- [x] I have installed `pre-commit` by running `pip install pre-commit`
(or used your preferred method).
- [x] I have installed the hooks with `pre-commit install`.
- [x] I have run the hooks manually with `pre-commit run --all-files`
and fixed any reported issues.

## 🧪 Tests

- [x] Tests have been added or updated as needed.
- [x] All tests are passing (`unittest`, etc.).

GB10 / sm_121: `pytest tests/attention/test_sparse_mla_sm120.py` → **309
passed**. Single-process soak on `heads=128, topk=1024, tokens=256`:
baseline **5/5 hung** (usually tens of iterations, sometimes a few
hundred); fix **5/5 × 30k, 0 hangs**. Two concurrent pytest (`-k
"prefill_glm_nsa_arbitrary_fp32 or prefill_dsv4"`): baseline wedged in
round 3; fix **12/12 clean**. All four instantiated prefill variants
soaked 30k each (with `assert_close`); 30k bitwise-identical iterations
vs a golden result (fix does not turn the hang into silent corruption).
No new unit test — the race is a soak, not a unit.

Header-only change: ninja may serve a stale `.so` (no `.d` depfiles).
Delete `csrc_sparse_mla_sm120_prefill.cuda.o` and `sparse_mla_sm120.so`
and confirm the `.so` mtime advances before A/B. Wedged processes ignore
`SIGTERM`; use `timeout -s KILL`.
bkryu added a commit that referenced this pull request Sep 3, 2026
## Summary

Consolidated SM120 sparse-MLA rework, decode + prefill. All numbers on
RTX PRO 6000 (SM120).

- **Faster decode, small T**: T=1 14.9 → **11.3µs** (−24%, graph replay,
dual-cache 18K context); bitwise-identical outputs.
- **Calibrated dispatch**: an analytical `chunks_per_block` model +
measured decode/prefill crossover replace the per-shape autotune sweep
and the hard `T ≤ 64` cutoff (up to −61% on rerouted configs; tables
below).
- **Continuous envelopes**: decode serves any `num_heads ∈ [1,128]` and
any `topk ≥ min_topk`; prefill serves any `T ≥ 1` and any `topk % 64 ==
0` width. Runtime-topk prefill drops instantiations **75 → 55**.
- **swapAB prefill** carried from #4751 behind a per-call `prefill_impl`
override; independently re-benched at **1.12–2.37×** over MG.
- **Two new model types**: `GLM53_NOPE` (carried from #4791) and
`DOTS3_SWA` (sliding-window MLA, d_qk=1088, d_v=1024, 1160 B/token
footer-scale, padded-row KV support) — the latter also fixing five
latent bugs along the way (rope writeback overrun at D_V==D_NOPE,
flat-vs-paged addressing keyed on the wrong trait, a Python chunk-width
hardcode, an undersized amax scratch at 4 math warps, a vestigial SG
register array).
- **Public runner**: `flashinfer.mla.SparseMLASm120Wrapper` — one
persistent instance, memoized dispatch, CUDA-graph-safe (decode scratch
is routing-aware and instance-owned).
- Merged current main, incl. the #4732 SM121 prefill-hang fix.
- Also: row-strided `indices` (unblocks vllm-project/vllm#53574's
persistent-buffer narrowing) and row-strided `out_lse`; T=0 decode
returns empty instead of aborting; decode bindings now validate
`out_lse`/index dtypes/dim0.

Carries (authorship preserved): #4461 zero-token decode (rewritten;
XingSong), #4551 dispatch diagnostics +
`supported_sparse_mla_sm120_configs()` (Sam Mausberg), #4751 swapAB
(Lemon7-UP), #4791 GLM53_NOPE (lucamotz; extended with H=64/TP1 decode,
swapAB@2176, calibration coverage). Supersedes #4683: its per-shape
sweep profiles L2-resident synthetic indices, which distorts cpb when
production caches are DRAM-resident (observed on 5070 Ti) — this PR
removes the sweep instead (thanks Sam for the original analysis).

## Performance vs main (adc49a8)

Same GPU, fixed-seed identical inputs, CUDA-graph replay GPU-only, both
sides out-of-box (no tactic cache / no calibrated constants). Only
surfaces present on both sides listed.

| shape | main | PR | speedup |
|---|---|---|---|
| dsv4-dual-h64 (topk 128+512), T=1 | 14.80µs | 11.40µs | 1.30x |
| dsv4-dual-h64 (topk 128+512), T=8 | 19.80µs | 16.14µs | 1.23x |
| dsv4-dual-h64 (topk 128+512), T=16 | 36.66µs | 29.46µs | 1.24x |
| dsv4-dual-h64 (topk 128+512), T=64 | 97.19µs | 89.79µs | 1.08x |
| dsv4-h128 (topk 1024), T=1 | 14.70µs | 11.44µs | 1.29x |
| dsv4-h128 (topk 1024), T=64 | 231.60µs | 219.79µs | 1.05x |
| dsv3_2-h64 (topk 2048), T=1 | 14.08µs | 10.68µs | 1.32x |
| dsv3_2-h64 (topk 2048), T=64 | 216.78µs | 220.23µs | 0.98x |
| dsv3_2-h128 (topk 2048), T=1 | 16.57µs | 14.08µs | 1.18x |
| dsv3_2-h128 (topk 2048), T=64 | 324.19µs | 323.53µs | 1.00x |
| dsv4-prefill-h128 (topk 1024), T=128 | 293.82µs | 266.49µs | 1.10x |
| dsv4-prefill-h128 (topk 1024), T=2048 | 4486.89µs | 3928.68µs | 1.14x
|
| dsv4-prefill-dual-h64 (topk 128+512), T=128 | 124.66µs | 124.64µs |
1.00x |
| dsv4-prefill-dual-h64 (topk 128+512), T=2048 | 1559.56µs | 1559.35µs |
1.00x |

Decode gains concentrate at small T (launch-bound); the two decode
commits behind them: `quantize_q_to_smem` rewritten as a vectorized
single pass (3 `bar.sync` → 1), and the decode-dsv4 IO gather reads each
candidate's index once instead of twice. T=64 decode and dual-cache
prefill are unchanged within noise.

## swapAB prefill (#4751)

Re-benched on the PRO 6000 (#4751's table was measured on a PRO 5000),
same grid, MG↔swapAB cross-checked at 5e-2 on identical inputs, `auto`
bitwise-identical to forced swapAB:

| shape | MG | swapAB | speedup |
|---|---|---|---|
| H=64, T=128 | 250.8µs | 159.7µs | 1.57× |
| H=64, T=512 | 798.7µs | 565.6µs | 1.41× |
| H=64, T=2048 | 2948.1µs | 2158.6µs | 1.37× |
| H=64, T=8192 | 11673.6µs | 8607.7µs | 1.36× |
| H=128, T=128 | 349.6µs | 267.9µs | 1.30× |
| H=128, T=512 | 1348.2µs | 840.0µs | 1.60× |
| H=128, T=2048 | 5330.9µs | 3011.6µs | 1.77× |
| H=128, T=8192 | 21156.9µs | 11847.7µs | 1.79× |

Wins everywhere; the H=64 large-T plateau (~1.4×, one CTA per token
saturates ~1280 GB/s vs ~1860 at H=128) is a flat asymptote out to
T=32768, so no dispatch range limit. KV layout and all parameters
unchanged; both scale formats, sinks, and variable `topk_length`
supported. `prefill_impl`: `"auto"` (default) / `"swapab"` / `"mg"`;
forcing swapab at an ineligible shape raises.

## Dispatch: cpb model + crossover

**cpb model** — analytical pick over gather bandwidth/latency, per-block
overhead, and the exact list-scheduling makespan of the split grid, with
an L2-footprint guard rail (at topk=1024+2176 dual the heuristic picks a
single 50-chunk block at 2.7× L2 — ncu: L2 hit 69.7% vs 86.8%, costing
33%; the guard recovers it to 1.02×). Calibrated once per device inside
`autotune()` tuning mode (6 fixed measurements over a ~2 GiB pool, timed
as queued batches over rotating fresh index sets — launch latency
overlaps execution, and the batch length keeps each set's reuse distance
past an L2 turnover; small numpy LM fit; any failure = silent fallback
to the C++ heuristic, so the new path can't be worse than status quo).
Offline pick error vs exhaustive sweep (DRAM-cold protocol): **mean
1.011× / max 1.061×**; beats the heuristic by up to **1.37×** at mid
shapes. A GPU accuracy-guard test fails loudly if a future kernel change
breaks the model's assumptions, measured with the same protocol the
calibration runs. Host cost ~8µs/call, memoized; zero per-replay under
CUDA graphs.

**Per-shape refinement** — the model's residual pick error concentrates
at mid-T wave-quantization shapes (measured up to **1.35×**, e.g.
DOTS3_SWA T=32: 78.0µs → 57.8µs). tuning-mode decode-form calls time the
model pick ±6 candidates with the calibration protocol and persist the
measured best as a per-shape override in the same tuning cache;
`_resolve_cpb` consults overrides first, then the model. Across 12
production bucket shapes (T=16..64, three families, two-pass re-timing):
**never worse than the model (12/12), closes every pocket to ≤1.03×**.
Shapes never warmed (off-graph calls, arbitrary T, dual-cache) stay on
the model. Capture-time calls only read the table/model and freeze — no
measurement ever runs under graph capture or in serving.

**Crossover** — per-config `decode_max_tokens` measured during the same
tuning pass (probe T ∈ {4..64}, both paths, DRAM-faithful fresh indices;
decode wins iff ≤ 0.95× prefill). Uncalibrated behavior is unchanged.
Measured examples:

| config | `decode_max_tokens` | Σ T∈{24,32,48,64}: old policy →
calibrated |
|---|---|---|
| DSv3.2 H=128 topk=2048 (swapAB side) | 8 | 1271.4 → 494.0 µs (−61%) |
| DSv4 H=64 topk=512 | 24 | 292.3 → 216.7 µs (−26%) |
| DSv4 H=64 topk=128 | 16 | 132.3 → 96.7 µs (−27%) |
| DSv4 H=8 topk=1024 | 64 (decode dominates) | no rerouting |

Full per-probe data for all 71 calibrated configs: kernel-bench
`crossover-v5` baseline. A public `calibrate_sparse_mla_sm120(device,
heads=, topks=, families=, force=)` makes any envelope shape tunable
outside tuning mode (idempotent skip-existing; `force=True`
re-measures).

## Runtime envelopes (head counts and topk widths)

- **Decode**: any H ∈ [1,128] — dedicated instantiations on the
production grid (0.9–2.5% faster), one runtime-H instance otherwise,
**40/40 bitwise-identical** between the two. Any `topk ≥ min_topk` (1;
513 for DOTS3_SWA so the window fits). The `_DECODE_*_DISPATCH` objects
vLLM probes are membership predicates with exactly this meaning;
`supported_sparse_mla_sm120_configs()` exposes the envelopes for
init-time validation. Off-grid example: H=80 T=16 is 1.14× faster than
the pad-to-128 workaround callers needed before.
- **Prefill**: same topk rule across SG / MG / dual / swapAB. One
deliberate residual asymmetry: **decode serves ragged widths (partial
tail chunk, tested at topk=500); prefill requires whole 64-wide index
tiles** — all production topk widths qualify, tail support needs
predicated gathers + tail masking across the IO and math paths, and is
deferred until a model needs it. This is safe at the routing layer: a
ragged decode-form call has no prefill envelope and simply stays on
decode (no crossover), and a ragged T>64 call fails loudly at the
binding. 50-config parity vs the pinned build: worst **+0.94%**. One
variant needed kernel-side help: DOTS3_SWA SG's BI=32 tiles are too
short to cover the index→rope address-chain latency once the
compile-time trip count disappeared (+24% `long_scoreboard` in NCU). The
SG loop now stages the three per-tile index reads one tile ahead in
registers, `if constexpr`-scoped to short tiles (unconditional staging
taxed BI=64 SG +2.3%). Net: **374.6µs vs the pinned build's 380.7µs** at
H=64/T=256, registers flat, `long_scoreboard` back to parity.

## Plan layer

All dispatch policy lives in one memoized Python planner
(`_sparse_mla_sm120_plan.py`): each variant declares its envelope once,
`plan()` picks by envelope + crossover + `prefill_impl`. The C++ side is
a policy-free launcher registry (the old `dispatch_v32` chain is
deleted). Single-sourcing surfaced two latent upstream bugs, fixed here:
prefill launchers never checked `page_block_size` against the compiled
64 (silent wrong-stride launch), and dual-cache decode-form
DSv3.2-family calls silently ignored the secondary cache.

## Runner and CUDA graphs

`SparseMLASm120Wrapper` holds buffers persistently: LSE pre-sized at
construction, decode split-K scratch allocated only when the call
actually routes to decode and cached for the instance's lifetime (a
per-call temporary's freed block can be recycled into a later capture
while an older graph replays into it). Capture contract: construct and
warm up every captured shape before capture (or pass `out_lse`/scratch
explicitly); replay is pure graph replay with zero Python. Both routing
variants are correct for any T, so a crossover inside a padding bucket
is at worst suboptimal, never wrong. GPU tests pin capture/replay for
crossover dispatch and for runner-internal scratch.

## Compatibility

Public Python API: unchanged except additive kwargs; `flashinfer.mla`
exports purely additive; no-constants path behaves exactly as today.
Deliberate behavior changes:

- Per-shape tactic caches (`sparse_mla_sm120_decode_dsv{4,3_2}.json`)
are ignored; the new calibration file is schema-versioned (v1),
unrecognized versions treated as absent and recalibrated.
- `autotune(True)` runs a one-time-per-device calibration (~2 GiB
transient pool) instead of profiling each new shape; honors
`skip_ops={"sparse_mla_sm120"}`; refuses to run under CUDA graph
capture; cache writes serialized with a FileLock.
- With calibration present, decode-form calls beyond the measured
crossover route to prefill (the point of the feature).
- T ≤ 64 shapes outside the old fixed grid now take the runtime decode
instantiation instead of raising.
- Prefill serves any `topk % 64 == 0` (≥ 513 for DOTS3_SWA); ragged
widths fail at the binding.
- Inline-scale (DSv3.2/GLM) KV caches must be contiguous through the
paged entry (prefill flat-addresses the cache and crossover makes
routing dynamic); contiguous padded-row caches remain decode-served and
fail loudly only if prefill-routed.
- `indices`/`out_lse` may be row-strided views (widening); the decode
binding previously corrupted a strided `out_lse` silently.
- C++ launcher entries gained row-stride parameters — internal to the
JIT module, no stable ABI consumers.

Out of scope (tracked follow-ups): H=64 swapAB bandwidth at large T; a
pinned-topk fast path à la decode-H for DOTS3_SWA SG (locked clocks show
~2% there, boost clocks show nothing — not worth the instantiation axis
on current evidence).

## Test plan

All on RTX PRO 6000: **658 passed** across
`test_sparse_mla_sm120{,_dispatch,_cpb_model}.py` and
`test_autotuner_core.py`, pre-commit clean — including the 68-config
small-T prefill matrix vs the reference (T ∈ {1..64} × SG/MG/swapAB/dual
× sink/truncation), 27 C++⟺Python envelope-consistency probes,
runtime-H/topk parity gates (bitwise where required), crossover routing
+ CUDA-graph capture/replay tests, runner scratch routing/lifetime
tests, and the review-round regression tests (row-strided `out_lse`, cpb
save/publish/FileLock, grid-completeness gating, padded-cache rejection,
skip_ops/capture guards).

This PR was prepared with AI assistance; all changes reviewed and tested
locally by the submitter.

---------

Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
Co-authored-by: XingSong <sunwenhan@xfusion.com>
Co-authored-by: Sam Mausberg <samuelmausberg@gmail.com>
Co-authored-by: Lemon7-UP <fearless192@163.com>
Co-authored-by: Luca Motz <321921718+lucamotz@users.noreply.github.com>
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
Co-authored-by: Brian K. Ryu <bryu@nvidia.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[Bug]Flaky issue on Spark: test_sparse_mla_sm120 timeout with no results

4 participants