Skip to content

[AMD] Add opt-in MiniMax-M3 TP4 indexer context partitioning - #41488

Merged
hnyls2002 merged 7 commits into
sgl-project:mainfrom
ThomasNing:thomas/minimax-m3-indexer-cp
Oct 2, 2026
Merged

hnyls2002 merged 7 commits into
sgl-project:mainfrom
ThomasNing:thomas/minimax-m3-indexer-cp

Conversation

@ThomasNing

@ThomasNing ThomasNing commented Sep 27, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

Under TP4, MiniMax-M3's four index-query heads are split across ranks while the index-K cache is replicated. Each rank scans the full index cache for its head. This PR adds an opt-in decode path that partitions the block reads across those four ranks and scores all heads while each K tile is resident.

The measured benefit is in the indexer chain, including communication: at batch 32 / 128K context, latency falls from 109.84 to 41.63 µs (62.1% lower latency, 2.64× speedup). Short contexts and small batches regress, so the feature defaults off. Full-model serving gains have not been measured.

Modifications

  • Add SGLANG_MINIMAX_M3_INDEXER_CP=1, default 0, for gfx950 TP4 ordinary decode with attention CP/DP size 1, four index heads, four main KV heads, dimension 128, 128-token blocks, max scoring, and top-k 16. Unsupported run configurations log a reason and retain the existing path.
  • Gather existing post-RoPE queries through the attention TP communicator, then let rank r read blocks r, r+4, r+8, ... for all four heads. The cache remains replicated; this reduces read traffic, not cache capacity.
  • Use Triton kernels for shard scoring, local top-16 selection per head, and final merging. Keeping 16 candidates on every shard preserves global top-16 even when all winners belong to one shard. Initial/recent block priorities and native ROCm score/ID ordering are preserved.
  • Exchange packed score/ID keys through SGLang's communicator, using AITER custom all-gather when supported and the existing fallback otherwise. Keys use FP32 bitcast views for copying, without numeric conversion. Warm both collective paths before graph capture; merge directly from the received layout.
  • Keep the implementation in two new modules with small backend/dispatch changes. Existing top-k reuse is checked first. Prefill, index-value layers, model projections, checkpoint loading, and main sparse attention retain their existing paths.
  • Add a manual four-rank correctness/timing harness, reproduction instructions, and all raw timing samples under test/manual/minimax_m3/indexer_cp/.

The initial scope excludes speculation, TBO, HiSparse, FP8 queries, and dense sparse decode. There is no automatic batch/context crossover policy. It does not require a main-attention backend change.

Accuracy Tests

Completed on four MI355X / gfx950 GPUs:

  • 8/8 cases passed on all ranks: BF16 and FP8 E4M3 index caches × random scores, ties, winners concentrated on one shard, and mixed lengths.
  • Exact selected-ID comparisons against native SGLang and an independent FP32 PyTorch reference, including empty rows and partial blocks.
  • Changed sequence lengths after graph capture, including short, partial-block, and empty rows; replay matched both references.
  • All 17 timing shapes additionally passed native-versus-CP selected-ID parity in eager execution and graph replay.
  • Repository pre-commit hooks passed for all changed files. The scorer, runtime helper, and benchmark script are byte-identical to the files used for the successful GPU run.

Validation still outstanding: runtime feature-gate/backend dispatch through a loaded server, full-model accuracy, serving TPS/TTFT/TPOT, and measurements with the fallback collective. The harness directly invokes the indexer helper. The shape gate also accepts FP16 queries and other FP8 cache variants, which were not covered by this run.

Speed Tests and Profiling

Four MI355X GPUs, BF16 queries, FP8 E4M3 index cache, AITER custom gather available. Both gathers are included; projections and main attention are excluded. Eight indexer calls per HIP graph, 100 replays per round, seven rounds alternating baseline/CP order. Each sample uses the slowest rank; reported values are the median of seven samples.

Inputs/cache are repeatedly reused and warm, without an L2 flush or rotating-layer protocol. These are microbenchmark results and cannot be interpreted as model throughput gains.

Context Batch Native TP (µs) CP (µs) CP latency change
8K 1 9.40 15.40 +63.8%
8K 8 7.72 18.32 +137.3%
8K 32 12.12 19.58 +61.5%
32K 16 20.80 20.04 -3.7%
32K 32 30.41 24.38 -19.8%
128K 1 12.71 15.93 +25.4%
128K 8 32.85 24.02 -26.9%
128K 16 51.71 32.08 -38.0%
128K 32 109.84 41.63 -62.1%
One 128K row, remaining 1K 16 40.76 27.82 -31.7%
One 128K row, remaining 1K 32 68.89 34.62 -49.7%

Negative latency change is an improvement. All tested 8K shapes and batches 1–2 at 128K regress. The small 32K/batch-16 gain needs independent repetition. The README contains the complete 17-shape table, runtime versions, source hashes, and the reproduction command; results-gfx950.json contains every round.

Checklist

  • Format changed files with repository pre-commit hooks.
  • Add reproducible four-rank correctness and speed checks (manual GPU harness).
  • Document enablement, supported configurations, results, and limitations.
  • Provide indexer-level accuracy and speed measurements, including regressions.
  • Validate full-model accuracy and serving performance with CP off/on.
  • Integrate the multi-GPU checks into registered CI.

CI States

Latest PR Test (Base): ✅ Run #36951043442
Latest PR Test (Extra): 🚫 Run #36962666288
Latest PR Test (AMD ROCm 10): ❌ Run #36951042978

@github-actions github-actions Bot added documentation Improvements or additions to documentation jit-kernel labels Sep 27, 2026
kevin-mii pushed a commit to kevin-mii/sglang that referenced this pull request Oct 1, 2026
…onto M3-opt-0929

Applied the PR diff (head 44632b7) with a 3-way merge; environ.py kept both additions.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
kevin-mii and others added 5 commits October 2, 2026 00:42
MiniMaxM3VLConfig keeps the text config's fields on a sub-config, so
hf_config.num_key_value_heads raises AttributeError and make_indexer_cp dies
before the gate can report a reason. ModelConfig.get_total_num_kv_heads()
returns 4 for amd/MiniMax-M3-MXFP4, which is what the gate checks for.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
The gate rejected every speculative config, so CP stayed off for the whole
EAGLE3 deployment -- which is how MiniMax-M3 is served.

It does not have to. main's forward_extend funnels HIP target-verify into
forward_decode with per-row seq_lens (prefix + 1..ndt) and the request slot
repeated, because a linear EAGLE chain exposes one more KV token to each
successive query. Verify rows therefore reach the indexer as ordinary decode
queries, each carrying its own slot and causal length, which is exactly what
the context-partitioned scorer already handles -- it was never one row per
request.

Gate on that property and allowlist it, so an algorithm whose verify rows are
not independent chain rows (DSPARK's ragged lengths, NGRAM's tree-in-mask,
EAGLE with top-k > 1) disables CP rather than reading as a chain and silently
mis-scoring.

Measured on MI350X with verify rows flowing through CP (TP4, EAGLE3 real
acceptance, AgentX agentic traces, 900 s/point, 1M context, index top-k freq 1),
total tok/s/GPU: c=24 25,441 -> 28,182 (+10.8%), c=32 32,393 -> 33,952 (+4.8%),
c=40 31,571 -> 33,229 (+5.3%); ITL p50 -12 to -28%. GSM8K-500 0.862-0.872.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
…ming

results-gfx950.json is a 629-line snapshot of one machine on one day -- it pins
a torch build and a harness hash, nothing reads it, and the README already
carries the same numbers as a table. No other committed results file exists
under test/ (all five JSON files there are test inputs), so it goes; the
harness regenerates the samples.

benchmark_cp.py -> bench_cp.py: the repo has 78 bench_*.py under test/ and no
benchmark_*.py.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
The only thing verifying that CP selects the same blocks as the native TP
selector lived in test/manual, which CI never runs -- so a bit-exactness claim
(64-bit packed keys, score-descending with ID-ascending ties, the forced
init/local ladder) had nothing guarding it.

score_local_blocks, select_local_candidates and merge_candidates all take rank
as a plain argument, so all four shards run in one process on one GPU; the
all-gather a four-rank run adds is the runtime's collective, not this kernel.
That makes it a 1-GPU AMD test alongside test_minimax_rocm_verify.py rather
than a four-GPU one. Covers bf16 and fp8 at batch 1 and 4; injecting a
one-rank shift into the block stride fails it with 15 IDs differing.

bench_cp.py stays as the timing harness and keeps its four-rank parity check,
which also exercises the collectives. Its docstring carries the run command and
the "indexer microbenchmark, not model throughput" caveat, so the README went:
the algorithm and the opt-in flag are already documented in the indexer_cp
module docstrings and the environ entry, and the dated results table belongs in
the PR description, not the tree. test/manual has no other README.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
It matches no pattern in the tree. Of 78 bench_*.py, 73 are registered under
test/registered/kernels/benchmark/<group>/ with the shared run_benchmark
harness and a reduced ci_range; the 5 under test/manual are short
single-process microbenchmarks. This was a 288-line torchrun harness with
argparse, JSON output and source hashing, in a directory CI never runs.

Its correctness half is now test/registered/amd/test_minimax_indexer_cp.py,
which gets the same guarantee in CI. Its timing half cannot honestly move to
the registered 1-GPU pattern: a single-process run omits the two all-gathers,
and CP must be judged with communication included, so such a benchmark would
overstate the win. The end-to-end serving A/B is the measurement that counts
and it belongs in the PR description.

This leaves the gathers in MiniMaxIndexerCP.__call__ and graph capture of the
full chain without a dedicated harness; they are exercised by any four-rank
serving run, which is how the numbers in this PR were produced.

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

Copy link
Copy Markdown
Collaborator

/tag-and-rerun-ci

@github-actions github-actions Bot added the run-ci CI: run the baseline test suite on this PR label Oct 2, 2026
@kevin-mii kevin-mii added the run-ci-extra CI: also run the extra suite (requires run-ci) label Oct 2, 2026
@hnyls2002
hnyls2002 merged commit b69a5f2 into sgl-project:main Oct 2, 2026
258 of 308 checks passed
@bingxche

bingxche commented Oct 8, 2026 •

Copy link
Copy Markdown
Collaborator

Hi @ThomasNing , thanks for your contribution!

The test you added test_minimax_indexer_cp did not pass even in this PR checks. https://github.com/sgl-project/sglang/actions/runs/36951042978/job/110670268955#step:6:26861

And it's still failing as of today https://github.com/sgl-project/sglang/actions/runs/37709475489/job/113091796233#step:6:7682. Could you please take a look?

cc @michaelzhang-ai @HaiShaw @hnyls2002

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

documentation Improvements or additions to documentation jit-kernel run-ci CI: run the baseline test suite on this PR run-ci-extra CI: also run the extra suite (requires run-ci)

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants