Skip to content

fix(sm120): explain sparse-MLA decode dispatch misses; add config query API - #4551

Closed
SamMausberg wants to merge 1 commit into
flashinfer-ai:mainfrom
SamMausberg:fix/sparse-mla-sm120-decode-dispatch-diagnostics
Closed

SamMausberg wants to merge 1 commit into
flashinfer-ai:mainfrom
SamMausberg:fix/sparse-mla-sm120-decode-dispatch-diagnostics

Conversation

@SamMausberg

@SamMausberg SamMausberg commented Aug 16, 2026

Copy link
Copy Markdown
Contributor

Description

Follow-up to #4380 for #4541.

On SM120, sparse-MLA decode kernels are instantiated for a fixed set of
(num_heads, topk) pairs with page_block_size=64. Since #4380, a
decode-form call (num_tokens <= 64) that misses every decode instantiation
raises in Python instead of reaching the prefill orchestrator's C++
num_tokens > 64 assert, but the error is a flat parameter dump that points
at private module internals (_DECODE_DSV4_DISPATCH /
_DECODE_DSV3_2_DISPATCH), so the caller still has to open the source to
learn what to change. #4541 asks for the two remaining pieces, added here:

  1. Dispatch-miss diagnosis. The ValueError now names the actual
    mismatch and enumerates what is instantiated, e.g.:

    SM120 sparse-MLA has no decode kernel for this shape: num_tokens=5,
    num_heads=64, topk=384, d_qk=512, page_block_size=64, model_type=dsv4,
    extra_topk=0. Mismatch: topk=384 is not instantiated for num_heads=64;
    available topk: [128, 192, 256, 512, 1024]. The prefill orchestrator only
    serves num_tokens > 64, so a decode-form call must match a decode
    instantiation exactly. Query supported shapes at init time with
    flashinfer.mla.supported_sparse_mla_sm120_configs().
    

    The diagnosis distinguishes uninstantiated topk (for an instantiated
    head count), uninstantiated num_heads, both, page_block_size != 64,
    and d_qk/family mismatch. The pre-existing message prefix and shape
    summary (no decode kernel ... num_tokens=..., num_heads=..., topk=...)
    are kept so existing callers matching on it keep working.

  2. flashinfer.mla.supported_sparse_mla_sm120_configs() — a public
    query API returning a {"dsv4" | "dsv3_2" | "glm_nsa": SparseMLASm120DecodeConfig} mapping (frozen dataclass with d_qk,
    page_block_size, max_num_tokens, head_topk_pairs, plus
    supports_decode() / supported_topk() / supported_num_heads()
    helpers). This lets callers such as vLLM validate a serving configuration
    at init time instead of poking at private dispatch tables
    ([Bugfix] Make DSV4 sparse MLA work end-to-end for plain decode, MTP, and DSpark vllm-project/vllm#51538 currently has to do the latter). Exported
    lazily from flashinfer.mla following the existing prims-ts lazy-export
    pattern, and added to the docs/api/attention.rst autosummary.

No kernel or dispatch-behavior changes: every shape that dispatched to a
kernel before still dispatches identically; only the no-kernel error path
and the new query API changed.

Related Issues

Fixes #4541. Follow-up to #4380. Motivating downstream check:
vllm-project/vllm#51538.

Pull Request Checklist

Thank you for contributing to FlashInfer! Before we review your pull request, please make sure the following items are complete.

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

Tested locally on an RTX 5070 Ti (SM120, CUDA 13.2, torch 2.8, JIT build from source):

  • pytest tests/attention/test_sparse_mla_sm120_dispatch.py12 passed (new file; pure Python, no GPU required)
  • pytest tests/attention/test_sparse_mla_sm120.py310 passed (full SM12x suite: all decode/prefill correctness tests plus the extended and new dispatch-failure tests)
  • pre-commit run --all-files — all hooks pass (ruff check/format, mypy, whitespace/EOL)

Reviewer Notes

  • tests/attention/test_sparse_mla_sm120_dispatch.py is deliberately not
    gated on SM12x: it covers the config query and the message builder as pure
    Python so the diagnosis logic is exercised on any CI node. The actual raise
    inside the dispatcher stays covered by the (extended) SM12x-gated tests in
    test_sparse_mla_sm120.py.
  • supported_sparse_mla_sm120_configs() exposes the instantiation sets by
    reference as frozensets (immutable), so the public view can never drift
    from the dispatch tables.
  • The "glm_nsa" entry aliases the "dsv3_2" config object since GLM-NSA
    decode shares the DSv3.2 instantiations; if the sets ever diverge, only the
    constructor needs to split.
  • The new test file carries the same NVIDIA BSD-3 header as its sibling
    test_sparse_mla_sm120.py for consistency within the SM120 sparse-MLA
    family; happy to switch it to the Apache-2.0 header (or drop it) if you
    prefer for externally-authored tests.

Developed in combination with Claude Fable 5.

Summary by CodeRabbit

  • New Features

    • Added public configuration details for supported SM120 sparse MLA decode variants, including dimensions, page sizes, token limits, head counts, and top-k values.
    • Added validation helpers for checking supported decode configurations.
    • Improved access to these APIs through the MLA package.
  • Bug Fixes

    • Decode errors now identify unsupported dimensions, page sizes, head counts, and top-k values.
  • Documentation

    • Expanded API documentation for the configuration query and configuration type.

@coderabbitai

coderabbitai Bot commented Aug 16, 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: 6a3f9c25-0f46-4b3c-ae44-3fb3cf239223

📥 Commits

Reviewing files that changed from the base of the PR and between 38a1bde and 0466129.

📒 Files selected for processing (2)
  • flashinfer/mla/__init__.py
  • tests/attention/test_sparse_mla_sm120_dispatch.py
🚧 Files skipped from review as they are similar to previous changes (2)
  • flashinfer/mla/init.py
  • tests/attention/test_sparse_mla_sm120_dispatch.py

Included review availability: Your plan includes up to 8 reviews per rolling hour; 5 remain after this review.


📝 Walkthrough

Walkthrough

The PR adds public SM120 sparse MLA decode configuration APIs, lazy exports, API documentation, structured unsupported-shape diagnostics, and tests for configuration validation and dispatch failures.

Changes

SM120 sparse MLA configuration APIs

Layer / File(s) Summary
Configuration query and public exports
flashinfer/mla/_sparse_mla_sm120.py, flashinfer/mla/__init__.py, docs/api/attention.rst
Adds immutable decode configuration objects, supported configuration queries, lazy public exports, and API documentation.
Decode dispatch diagnostics
flashinfer/mla/_sparse_mla_sm120.py
Decode dispatch errors now identify unsupported heads, top-k, page size, or token limits and list applicable configurations.
Configuration and dispatch validation
tests/attention/test_sparse_mla_sm120_dispatch.py, tests/attention/test_sparse_mla_sm120.py
Tests cover configuration metadata, helper methods, lazy exports, invalid shapes, diagnostic messages, and unsupported page sizes before prefill.

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

Merge Risk: ⚪ Minimal · up to 04661

This change improves sparse-MLA dispatch error messages and adds a supported-configuration query API without changing successful kernel dispatch behavior. No actionable merge-blocking risk remains beyond normal checks and review.

Possibly related issues

Possibly related PRs

Suggested labels: run-ci

Suggested reviewers: sricketts, dhiraj113, aleozlx

🚥 Pre-merge checks | ✅ 4 | ❌ 1

❌ Failed checks (1 warning)

Check name Status Explanation Resolution
Docstring Coverage ⚠️ Warning Docstring coverage is 75.00% which is insufficient. The required threshold is 80.00%. Write docstrings for the functions missing them to satisfy the coverage threshold.
✅ Passed checks (4 passed)
Check name Status Explanation
Title check ✅ Passed The title clearly summarizes the two primary changes: improved SM120 sparse-MLA dispatch diagnostics and a configuration query API.
Description check ✅ Passed The description covers the changes, related issues, testing, checklist, and reviewer notes with sufficient detail.
Linked Issues check ✅ Passed The changes satisfy issue #4541 by adding actionable dispatch errors and a public API for validating supported SM120 decode shapes.
Out of Scope Changes check ✅ Passed The documentation, lazy exports, implementation, and tests directly support the linked issue and stated pull request objectives.
✨ 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.

@SamMausberg
SamMausberg force-pushed the fix/sparse-mla-sm120-decode-dispatch-diagnostics branch 2 times, most recently from 38a1bde to 0466129 Compare August 16, 2026 17:50
…ry API

Since flashinfer-ai#4380, a decode-form call (num_tokens <= 64) that matches no
decode instantiation raises in Python instead of tripping the prefill
orchestrator's C++ "num_tokens > 64" assert, but the error is a flat
parameter dump pointing at private dispatch tables, so the reader
still has to open the source to work out which parameter to change
(flashinfer-ai#4541).

Make the dispatch-miss error name the mismatch: which of topk,
num_heads, page_block_size, or d_qk missed the instantiated set, and
which values are available. The old message prefix and shape summary
are kept so callers matching on them are unaffected.

Add flashinfer.mla.supported_sparse_mla_sm120_configs() so serving
frameworks can validate (num_heads, topk, page_block_size) at init
time instead of on the first decode request; vLLM currently reads the
private tables for this (vllm-project/vllm#51538).

Dispatch behavior is unchanged: shapes that previously reached a
kernel still do; only the no-kernel error message and the new query
API changed.

Developed in combination with Claude Fable 5.

Closes flashinfer-ai#4541
@SamMausberg
SamMausberg force-pushed the fix/sparse-mla-sm120-decode-dispatch-diagnostics branch from 0466129 to 2b48965 Compare August 16, 2026 17:57
@SamMausberg

Copy link
Copy Markdown
Contributor Author

Hello, @bkryu, can you please take a look at this follow-up to #4380 when you get a chance?

lucifer1004 added a commit to lucifer1004/flashinfer that referenced this pull request Aug 28, 2026
…ry API

On a decode-form dispatch miss, name the actual mismatch (topk not
instantiated for this head count, heads not instantiated for this topk,
both, page_block_size != 64, d_qk family mismatch) and list the available
values, instead of a flat parameter dump. Add
supported_sparse_mla_sm120_configs() so callers can query the instantiated
(num_heads, topk) grid without reading private tables; the query API shares
the dispatch frozensets by reference so the two cannot drift.

Carried from flashinfer-ai#4551.

Co-authored-by: Sam Mausberg <samuelmausberg@gmail.com>
Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
ormandj pushed a commit to ormandj/flashinfer that referenced this pull request Aug 29, 2026
…ry API

On a decode-form dispatch miss, name the actual mismatch (topk not
instantiated for this head count, heads not instantiated for this topk,
both, page_block_size != 64, d_qk family mismatch) and list the available
values, instead of a flat parameter dump. Add
supported_sparse_mla_sm120_configs() so callers can query the instantiated
(num_heads, topk) grid without reading private tables; the query API shares
the dispatch frozensets by reference so the two cannot drift.

Carried from flashinfer-ai#4551.

Co-authored-by: Sam Mausberg <samuelmausberg@gmail.com>
Signed-off-by: Zihua Wu <13583761+lucifer1004@users.noreply.github.com>
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>
@bkryu

bkryu commented Sep 3, 2026

Copy link
Copy Markdown
Collaborator

Thanks @SamMausberg for this PR.

#4802 has been merged with preserved authorship so I will be closing. We look forward for more contributions from the community.

@bkryu bkryu closed this Sep 3, 2026
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.

[sm120][dsv4] Sparse MLA decode-form calls with non-instantiated topk silently fall through to the prefill orchestrator

3 participants