Skip to content

MiniMax-M3: share the sparse index top-k across layers and reuse the decode top-k buffer - #36527

Merged
Fridge003 merged 7 commits into
sgl-project:mainfrom
zcnrex:m3-sparse-decode-topk
Sep 8, 2026
Merged

Fridge003 merged 7 commits into
sgl-project:mainfrom
zcnrex:m3-sparse-decode-topk

Conversation

@zcnrex

@zcnrex zcnrex commented Aug 26, 2026 •

Copy link
Copy Markdown
Collaborator

MiniMax-M3: share the sparse index top-k across layers and reuse the decode top-k buffer

MiniMax-M3's sparse attention runs a lightning-indexer pass in every sparse layer
to choose the top-k KV blocks that the main attention then reads. On the
value-disabled sparse layers that indexer has no output of its own — idx_o is
None there, only the block selection is consumed — so its result can be
computed once and shared by a run of consecutive layers. This PR adds
SGLANG_MINIMAX_M3_INDEX_TOPK_FREQ=N (default 2), which groups every N such
layers: the group's first layer computes and stores the reduced top-k, and the
remaining layers reuse it, skipping both the flash-index attention and the
group-reduce. In decode the selection is written straight into a persistent
per-batch-size int32 buffer (topk_out=) rather than a freshly allocated
tensor per layer, which removes an allocation from every sparse decode layer and
keeps the destination address fixed so CUDA-graph capture never allocates. The
buffer is only useful because of the sharing, so the two changes ship together.

Sharing the selection changes which KV blocks the skip layers attend to — see
## Correctness below. It is therefore gated: index_topk_freq is pinned to 1
on every non-ROCm platform and whenever two-batch overlap is enabled, and at
freq 1 no buffer is allocated and every new argument is None/False, so the
default CUDA path is byte-for-byte main's.

Two smaller changes in the same files ride along:

  • The layer-invariant prefill seqblock metadata (cu_seqblocks_q,
    max_seqblock_q, all_seqblock_q) is computed once per ForwardBatch in the
    backend and passed down, instead of being recomputed inside each sparse layer.
    Same values. The owning ForwardBatch is part of the cache key because a
    single metadata init can be followed by more than one ForwardBatch reaching
    the layers (two-batch overlap splits into two children with different
    extend_seq_lens), so a hit requires the same object, not just a live cache.
  • forward_decode gets the _is_sparse_kv_cached_by_fusion guard that the
    extend path already has on main, so the decode sparse-KV store is not repeated
    when the fused qk-norm-rope kernel already wrote it. Idempotent write either
    way; this only removes the launch.

The decode score kernel's autotune space now drops configs whose token tile is
smaller than a sparse block (BLOCK_SIZE_N < block_size), and on ROCm the sweep
is replaced by one fixed config plus a BLOCK_SIZE_N heuristic, because runtime
autotune can fire under CUDA-graph capture and mistune there.

Performance

Measured on 8x MI350X (gfx950), MiniMax-M3-MXFP8, TP8, 80k input / 600 output.
Two bench reps per boot; the warm (second) rep is reported and the two agreed
within 0.1%. Noise floor is +/-0.4%. Tables are the harness output verbatim,
trimmed to the first five columns. Baseline is upstream/main.
Baseline (upstream/main)

Input lens: [80000]. Output lens: [600]. Cache hit rate: 90.0%.
|   batch size |   input len |   latency (s) |   input throughput (tok/s) |   output throughput (tok/s) |
|--------------|-------------|---------------|----------------------------|-----------------------------|
|            1 |       80000 |         11.26 |                     187984 |                       55.35 |
|            2 |       80000 |         11.88 |                     191222 |                      108.68 |
|            4 |       80000 |         13.39 |                     192965 |                      204.62 |
|            8 |       80000 |         16.27 |                     193687 |                      370.13 |
|           16 |       80000 |         22.2  |                     193188 |                      616.28 |
|           32 |       80000 |         33.78 |                     193245 |                      934.92 |

With this PR

Input lens: [80000]. Output lens: [600]. Cache hit rate: 90.0%.
|   batch size |   input len |   latency (s) |   input throughput (tok/s) |   output throughput (tok/s) |
|--------------|-------------|---------------|----------------------------|-----------------------------|
|            1 |       80000 |         10.44 |                     209588 |                       59.67 |
|            2 |       80000 |         11.14 |                     210417 |                      115.61 |
|            4 |       80000 |         12.35 |                     213259 |                      221.18 |
|            8 |       80000 |         14.41 |                     215085 |                      419.8  |
|           16 |       80000 |         19.37 |                     214509 |                      716.39 |
|           32 |       80000 |         28.94 |                     215010 |                     1127.4  |

Output throughput +6.38%..+20.59%; input throughput +10.04%..+11.49%.

Accuracy

gsm8k, 512 examples, max_tokens 2048, temperature 0, seed 0, against the
same server boot as the perf run.
Baseline

== gsm8k ==
512 examples (single-shot)  |  109.0s  |  1338 tok/s  |  146K tokens

* score           =  97.07%
  stop_rate       =  98.44%
  truncated_rate  =  1.56%  [warn: hitting max_tokens]
  error_rate      =  0.00%

With this PR

== gsm8k ==
512 examples (single-shot)  |  97.7s  |  1463 tok/s  |  143K tokens

* score           =  96.68%
  stop_rate       =  98.44%
  truncated_rate  =  1.56%  [warn: hitting max_tokens]
  error_rate      =  0.00%

Note on the launch command

--moe-runner-backend aiter is not honoured on current main for mxfp8. Arg resolution logs:

mxfp8 quantization supports only cutlass, deep_gemm, flashinfer_trtllm,
flashinfer_trtllm_routed, triton backends. Overriding 'aiter'.

and falls back to triton. Both sides of this A/B therefore ran the Triton MoE runner, so the comparison is apples-to-apples and the deltas stand — but AITER_CONFIG_FMOE is inert under this configuration and the aiter MoE path is not what was measured. The flag is kept in the command above only because it matches the invocation actually used; anyone reproducing will get triton and should see the same numbers.


CI States

Latest PR Test (Base): ⏳ Run #34259568986
Latest PR Test (Extra): ❌ Run #34259568651
Latest PR Test (AMD ROCm 7.2): ❌ Run #34259568790

@zcnrex
zcnrex force-pushed the m3-sparse-decode-topk branch from f1494dc to 5bdab43 Compare August 27, 2026 05:54
@zcnrex zcnrex added amd run-ci CI: run the baseline test suite on this PR labels Aug 28, 2026
zcnrex and others added 2 commits September 1, 2026 02:18
…decode top-k buffer

MiniMax-M3's sparse attention runs a lightning-indexer pass per sparse layer to
pick the top-k KV blocks the main attention then reads. On the value-disabled
sparse layers that indexer produces no output of its own -- only the block
selection -- so its result can be computed once and shared by a run of
consecutive layers. SGLANG_MINIMAX_M3_INDEX_TOPK_FREQ=N groups every N such
layers: the group's first layer computes and stores the reduced top-k, the rest
reuse it and skip both the flash-index attention and the group reduce. The
decode path writes that selection straight into a persistent per-batch-size
int32 buffer (`topk_out=`) instead of allocating a fresh tensor per layer, which
also keeps the address fixed so CUDA-graph capture never allocates.

Sharing the selection changes which KV blocks the skip layers attend to, so it
is gated: the freq is forced to 1 on every non-ROCm platform and whenever
two-batch overlap is enabled, and freq=1 restores main's per-layer behaviour
exactly (no buffer is allocated and every new argument is None/False).

Two smaller changes ride along in the same files:

- The layer-invariant prefill seqblock metadata (cu_seqblocks_q,
  max_seqblock_q, all_seqblock_q) is computed once per ForwardBatch in the
  backend and passed down, instead of being recomputed inside every sparse
  layer. Same values; the owning ForwardBatch is part of the cache key because
  two-batch overlap sends two children through one metadata init.
- forward_decode gets the same `_is_sparse_kv_cached_by_fusion` guard the
  extend path already has on main, so the decode KV store is not repeated when
  the fused qk-norm-rope kernel already wrote it.

The decode score kernel's autotune space drops configs whose token tile is
smaller than a sparse block, and on ROCm the sweep is replaced by a fixed
config plus a BLOCK_SIZE_N heuristic, because runtime autotune can fire under
CUDA-graph capture and mistune.

Co-Authored-By: Kevin Mi <mikevin920@yahoo.com>
Sharing changes which KV blocks the skip layers attend to and was only
accuracy-gated on gfx950; every other platform keeps freq 1, i.e. main's
per-layer selection.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
@kevin-mii
kevin-mii force-pushed the m3-sparse-decode-topk branch from f03d128 to be34340 Compare September 1, 2026 02:19
@zcnrex
zcnrex enabled auto-merge (squash) September 3, 2026 17:48
@Fridge003
Fridge003 disabled auto-merge September 8, 2026 23:38
@Fridge003
Fridge003 merged commit a25bbca into sgl-project:main Sep 8, 2026
121 of 150 checks passed
kevin-mii pushed a commit to zcnrex/sglang that referenced this pull request Sep 10, 2026
sgl-project#36527 landed the shared index top-k, which wraps Step 1 and Step 2 in a
cached_topk_idx branch and threads cu_seqblocks_q/cached_topk_idx through the
backend. Resolutions: keep both parameter sets; thread this PR's page_size into
the Step 1 call inside main's else branch; graft the Gluon Step 3 onto main's
structure so the caching and the Gluon path compose -- Gluon runs after the
top-k is resolved either way, and the MSA/Triton fallback stays behind
`o is None`.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
kevin-mii pushed a commit to zcnrex/sglang that referenced this pull request Sep 10, 2026
sgl-project#36527 added a MiniMax-M3 sparse-attention section to environ.py in the same
region. Kept both knobs and moved this PR's fp8 index-cache toggle inside that
section, since it is the same class of ROCm MiniMax-M3 knob.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
rkarhila-amd added a commit to rkarhila-amd/sglang that referenced this pull request Sep 11, 2026
Resolve minimax_sparse_backend.py against sgl-project#36527: keep seqblock/index-cache
prefill on main, route ROCm TARGET_VERIFY through decode-verify with uniform
metadata instead of extend prefill.
rkarhila-amd added a commit to rkarhila-amd/sglang that referenced this pull request Sep 11, 2026
Resolve minimax_sparse_backend.py against sgl-project#36527: keep seqblock/index-cache
prefill on main, route gfx950 TARGET_VERIFY through decode-verify with uniform
metadata instead of extend prefill.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

amd jit-kernel run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants