Skip to content

MiniMax-M3: run the sparse prefill main attention through AITER Gluon paged attention - #36546

Merged
hnyls2002 merged 60 commits into
sgl-project:mainfrom
zcnrex:m3-gluon-sparse-prefill
Sep 23, 2026
Merged

hnyls2002 merged 60 commits into
sgl-project:mainfrom
zcnrex:m3-gluon-sparse-prefill

Conversation

@zcnrex

@zcnrex zcnrex commented Aug 26, 2026 •

Copy link
Copy Markdown
Collaborator

MiniMax-M3: Gluon paged-attention sparse prefill with Triton fallback

MiniMax-M3's sparse prefill runs three steps: a lightning-index attention that
produces per-query block scores, a top-k reduction over those scores, and a
sparse main attention that attends only to the selected KV blocks. Step 3 is
currently a Triton kernel (flash_prefill_with_gqa_share_sparse) that walks the
paged KV pool one 128-token block at a time via req_to_token. This PR adds a
second implementation of step 3 that calls AITER's Gluon paged attention
(pa_decode_gluon(..., ps=True)) instead, selected by
SGLANG_OPT_USE_GLUON_PREFILL. The main KV pool layout is unchanged — it stays
NHD [max_slots, 1, head_dim] — so each request's context is first gathered
into persistent SHUFFLE 5D scratch pages (GLUON_PAGE_SIZE = 64, two pages per
128-token sparse block, x = 16 // dtype.itemsize), and the per-query sparse
block table is expanded into a page table plus a per-query effective context
length. Scratch grows in 1024-page steps and is capped at 512 MiB per buffer;
past the cap the entry point raises and the caller falls back. The indexer and
the top-k reduction (steps 1 and 2) are untouched by the dispatch. Two prefill
kernels this path shares also change: a score-only index-attention kernel for
disable_index_value layers (it skips the index-value output and, when a sparse
block spans few enough pages, reads one base slot per page instead of one
req_to_token entry per token), and CDNA KV sub-tiling (SUB_K) in the Triton
sparse kernel so each QK/PV MFMA is right-sized on gfx942/gfx950.

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 |         11.21 |                     203486 |                       55.48 |
|            2 |       80000 |         11.8  |                     206810 |                      108.85 |
|            4 |       80000 |         13.24 |                     208565 |                      205.06 |
|            8 |       80000 |         15.99 |                     210130 |                      370.82 |
|           16 |       80000 |         21.66 |                     208843 |                      617.92 |
|           32 |       80000 |         32.65 |                     209254 |                      940.28 |

Output throughput +0.16%..+0.57%; input throughput +8.08%..+8.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)  |  121.9s  |  1152 tok/s  |  141K tokens

* score           =  97.46%
  stop_rate       =  99.02%
  truncated_rate  =  0.98%
  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 #35806556038
Latest PR Test (Extra): ❌ Run #35806555776
Latest PR Test (AMD ROCm 10): ❌ Run #35806556095

@zcnrex

zcnrex commented Aug 28, 2026

Copy link
Copy Markdown
Collaborator Author

/tag-and-rerun-ci

@zcnrex zcnrex added run-ci CI: run the baseline test suite on this PR amd labels Aug 28, 2026
zcnrex and others added 2 commits September 1, 2026 01:27
Route the MiniMax-M3 sparse prefill main attention (step 3) through AITER's
Gluon paged attention instead of the Triton `flash_prefill_with_gqa_share_sparse`
kernel, behind `SGLANG_OPT_USE_GLUON_PREFILL`. The main KV pool stays NHD
`[max_slots, 1, head_dim]`; each request's context is gathered into persistent
SHUFFLE 5D scratch pages (page size 64, two pages per 128-token sparse block)
and the per-query sparse block table is expanded to a page table, so
`pa_decode_gluon(..., ps=True)` can serve the whole extend batch.

A cheap static gate (`can_use_gluon_prefill`) plus a runtime shape check keep
every unsupported case on the existing Triton path; the fallback logs once.

Also in the prefill kernels this path shares:
- a score-only index-attention kernel for `disable_index_value` layers, with
  per-page K addressing (one base-slot load per page instead of one
  `req_to_token` lookup per token) when a sparse block spans few enough pages;
- CDNA KV sub-tiling in the Triton sparse kernel (`SUB_K`, 0 elsewhere) so each
  QK/PV MFMA is right-sized on gfx942/gfx950.

Co-Authored-By: Kevin Mi <mikevin920@yahoo.com>
Co-Authored-By: Alex Sun <alex.s@amd.com>
Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
- matrix_instr_nonkdim/kpack are AMD-only Triton launch kwargs that
  NVIDIA's Triton rejects; only offer those autotune configs on HIP.
- Route the arch detection for KV sub-tiling through is_gfx95_supported /
  is_gfx942_supported instead of parsing gcnArchName.
- The score-only indexer kernel and the Gluon dispatch were only validated
  on gfx950; gate them (_is_hip, and _use_aiter + gfx95 respectively) so
  other platforms keep the shared kernels. The Gluon gate also stops
  non-gfx95 ROCm from allocating scratch pages before the aiter call fails.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01GbK4exGWvsmPdMf3kmBYW4

@HaiShaw HaiShaw 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.

@zcnrex Thanks for your contributions!

Comment thread python/sglang/srt/environ.py Outdated
SGLANG_MINIMAX_M3_FUSED_MOE_COMBINE = EnvBool(False)
# Run the sparse prefill main attention through AITER's Gluon paged attention
# instead of the Triton kernel. Unsupported cases fall back to Triton.
SGLANG_OPT_USE_GLUON_PREFILL = EnvBool(True)

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.

Do we want to rename this SGLANG_MINIMAX_OPT_USE_GLUON_PREFILL?

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.

should be renamed. could you please take a look again. Thanks!

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.

Renamed in 8f2cc72 to SGLANG_OPT_USE_MINIMAX_GLUON_PREFILL, matching the other SGLANG_OPT_USE_MINIMAX_* toggles in environ.py. The variable is new in this PR, so there is no alias for the old name.

zcnrex and others added 2 commits September 3, 2026 10:52
@zcnrex
zcnrex enabled auto-merge (squash) September 3, 2026 17:53
kevin-mii and others added 6 commits September 4, 2026 23:49
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>
`ctx` reads as a context object; the value is the effective KV length the
kernel walks, which is what the comment above it already says.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
yctseng0211 and others added 4 commits September 21, 2026 10:10
Resolve conflicts with sgl-project#31446 (HiSparse for MiniMax-M3):
- topk_sparse.py: apply the HAS_LOC_MAPPING slot remap in both the SUB_K
  sub-tiled loop and the dense loop, so HiSparse stays correct on gfx950.
- minimax_sparse.py: keep both the loc_mapping and page_size/seq_lens_cpu
  params; skip the Gluon path when loc_mapping is set, since its scratch
  gather reads pool slots without the HiSparse remap (same rule main applies
  to MSA).

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
… conventions

- Fall back to Triton only on GluonPrefillUnavailableError (scratch cap or
  scratch OOM), mirroring MSAUnavailableError and AsmKernelUnavailable; a
  Gluon kernel failure now raises instead of silently running Triton. The
  topk/batch shape checks are internal invariants, so they become asserts.
- Read SGLANG_USE_AITER through envs and file SGLANG_OPT_USE_MINIMAX_GLUON_PREFILL
  with the ROCm sparse-attention toggles.
- Trim multi-line comments that narrated history, named the model's layer
  count, or repeated a docstring.
- Build page_start with itertools.accumulate, fix the fallback warning that
  said the default-on toggle "is set", and keep page_size after the scale
  kwargs in flash_prefill_with_topk_index.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…e gfx950 measured a win

The cutoff of 8 pages had no measurement behind it and turned the per-page path
on at page_size 16, where it is slower. Score-only index kernel on MI350X,
total_q 8192, KV 80000, per-page vs per-token (1 / 16 index heads):
page_size 128: 1.34x / 1.14x faster; 64: 1.10x / 1.06x; 32: tie; 16: 0.59x / 0.62x;
8 and below: 0.44x or worse. Scores are bitwise identical either way.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
@kevin-mii

Copy link
Copy Markdown
Collaborator

CI on 9f02510766: all base-a / base-b green in https://github.com/sgl-project/sglang/actions/runs/35793567703.

Validated on 4x MI350X, TP4: Gluon vs Triton parity passes (141 cases, bf16 rounding only). gsm8k 512 is 97.46% with Gluon vs 97.27% without, with no fallback on an 88k prompt.

About to push: _MAX_PER_PAGE_SLOT_UNROLL 8 → 2. The 8 was unmeasured, and on MI350X per-page slots are 1.6x slower at page_size 16, tie at 32, and win only at 64 (1.10x) and 128 (1.34x). Scores are bitwise identical.

No CI rerun needed: the constant only feeds the gfx950-gated score-only index kernel, which no CUDA or MI300 job reaches, and it only changes behavior at page_size 16/32 (default ROCm page_size is 1). The run above stays valid for everything else.

@kevin-mii

Copy link
Copy Markdown
Collaborator

/rerun-test test/registered/e2e/models/test_kimi_k3_b300.py::TestKimiK3B300MegaMoE

@kevin-mii

Copy link
Copy Markdown
Collaborator

base-c-test-8-gpu-b300 (job): TestKimiK3B300MegaMoE server died during init on an NCCL timeout (scheduler exit -6); TestKimiK3B300Balanced passed gsm8k 0.98 in the same job. Kimi-K3 CUDA path, not touched by this PR — rerunning that class.

@github-actions

github-actions Bot commented Sep 23, 2026 •

Copy link
Copy Markdown
Contributor

Results for /rerun-test test/registered/e2e/models/test_kimi_k3_b300.py::TestKimiK3B300MegaMoE:

🚀 8-gpu-b300 (1 test): ❌ View workflow run

cd test/ && python3 registered/e2e/models/test_kimi_k3_b300.py TestKimiK3B300MegaMoE

@hnyls2002
hnyls2002 merged commit 4fa2c9c into sgl-project:main Sep 23, 2026
165 of 183 checks passed
kevin-mii added a commit to zcnrex/sglang that referenced this pull request Sep 23, 2026
Resolve the environ.py conflict with sgl-project#36546 by keeping both MiniMax-M3 flags.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
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.

5 participants