Skip to content

[AMD] Small-M MXFP4 fused-MoE kernel for gfx950 (Qwen) - #40204

Merged
HaiShaw merged 5 commits into
sgl-project:mainfrom
zijiecode:pr/smallm-moe-gfx950
Sep 24, 2026
Merged

HaiShaw merged 5 commits into
sgl-project:mainfrom
zijiecode:pr/smallm-moe-gfx950

Conversation

@zijiecode

@zijiecode zijiecode commented Sep 18, 2026 •

Copy link
Copy Markdown
Collaborator

Motivation

For decode batches of 1 to ~40 tokens per rank, the MXFP4 MoE is the biggest cost per step on MI355X, and AITER
fused_moe only uses a small part of the HBM bandwidth there. This PR adds a HIP kernel pair for these shapes; above 40 tokens AITER is still used.

Modifications

  • New python/sglang/kernels/ops/moe/smallm_moe_gfx950/ (__init__.py, smallm_moe.hip): built with hipcc on
    first use and launched through ctypes, no build-time or wheel change.
  • Phase 1: gate/up GEMM per (expert slot, 16-column slice), FP4 expanded with v_cvt_scalef32_pk_bf16_fp4, MFMA
    16x16x32 bf16, silu(g) * u written as bf16. Phase 2: down GEMM per (token, 64-row slice) with v_dot2_f32_bf16,
    top-k weights applied in fp32.
  • Uses AITER's existing weight and scale layouts, so loading and preprocessing are untouched. Activations stay bf16.
  • moe_runner/aiter.py: one guarded call before fused_moe, taken for FP4x2 weights, bf16 hidden 4096, per-rank
    intermediate 256 (the TP4 shape), 10 or 11 slots, silu, tok <= 40, and none of expert_mask / doweight_stage1 /
    bias / a1_scale / no_combine / num_local_tokens. Everything else is unchanged. The TP2 shape (intermediate 512)
    crosses over to AITER much earlier (~20 tokens) and is left to a follow-up.
  • On by default on gfx950; SGLANG_ROCM_SMALLM_MOE=0 turns it off.
  • Registered AMD unit test test/registered/amd/test_smallm_moe_gfx950.py: dequantized torch reference (rel L2 < 5e-3
    at 1 / 3 / 16 / 40 tokens), fallback above 40 and for the TP2 shape, the off switch, graph replay == eager. Passes on
    v0.5.19-rocm720-mi35x-20260915 and v0.5.18-rocm10-mi35x-20260902.

Accuracy Tests

amd/Qwen3.5-397B-A17B-MXFP4-AttnFP8-V2 (revision e17e5f0), MI355X TP4, InferenceX recipe (FP8 KV cache, EAGLE MTP
3 steps / 4 draft tokens, real draft), lm-eval local-chat-completions with chat template, 5-shot GSM8K, 16384-token
output budget, temperature 0.

Configuration Strict-match Flexible-extract
Baseline (AITER fused_moe) 97.50% (1286/1319) 97.57% (1287/1319)
This PR 97.73% (1289/1319) 97.57% (1287/1319)

GPQA diamond (198 questions, repeat 8), same server, thinking on, temperature 0.6 / top_p 0.95 / top_k 20, 32768-token
output budget:

Configuration Score
Baseline (AITER fused_moe) 0.869 (1377/1584)
This PR 0.875 (1386/1584)
Reference (Qwen/Qwen3.5-397B-A17B model card) 0.884

Benchmarking and Profiling

MI355X, TP4 / EP1, InferenceX MI355X recipe for this model (image lmsysorg/sglang-rocm:v0.5.19-rocm720-mi35x-20260915,
--kv-cache-dtype fp8_e4m3 --mem-fraction-static 0.80, EAGLE MTP 3 / 1 / 4, --speculative-draft-model-quantization quark_mxfp4 in both configurations), tree = main e4cbb28e + #39901 + #39902 (+ this PR), full 3600 s aiperf runs on
the same GPUs, one configuration after the other.

Concurrency 1:

Metric Baseline This PR Change
TPOT p50 (ms/token) 3.66 3.22 −12.0%
TPOT p90 (ms/token) 4.01 3.71 −7.5%
TPOT mean (ms/token) 3.72 3.34 −10.2%
ITL p95 (ms) 4.12 3.85 −6.6%
TTFT p50 / p90 (ms) 406 / 897 398 / 888 −1.9% / −1.0%
Requests completed in 3600 s 559 594 +6.3%
E2E normalized interactivity p50 / p90 (tok/s/user) 186 / 107 202 / 119 +8.4% / +11%

Concurrency 8:

Metric Baseline This PR Change
TPOT p50 (ms/token) 4.37 4.11 −5.9%
TPOT p90 (ms/token) 5.09 4.82 −5.3%
TPOT mean (ms/token) 4.76 4.30 −9.7%
ITL p95 (ms) 5.62 5.10 −9.3%
TTFT p50 / p90 (ms) 771 / 1266 672 / 1155 −12.9% / −8.8%
Requests / generated tokens in 3600 s 1455 / 1.352M 1476 / 1.421M +1.4% / +5.1%
E2E normalized interactivity p50 / p90 (tok/s/user) 150.6 / 81.3 165.9 / 89.6 +10% / +10%

Kernel microbenchmark (MI355X, HIP-graph replay with the MALL flushed between calls, Qwen3.5-397B TP4 shape, E=513, 11 slots):

Tokens AITER fused_moe (µs) This kernel (µs) Change
1 50–60 29–33 −42%
4 57 32 −44%
16 88 66 −25%
32 117 104 −11%
48 132 133 parity

Checklist

  • Format your code according to the Format code with pre-commit guide.
  • Add unit tests.
  • Update documentation.
  • Provide accuracy and speed benchmark results.
  • Follow the SGLang code style guidance.

CI States

Latest PR Test (Base): ⏳ Run #35960012293
Latest PR Test (Extra): ❌ Run #35960011935
Latest PR Test (AMD ROCm 10): ⏳ Run #35960012296

Two HIP kernels for the MXFP4 MoE at 1-40 tokens per rank on MI355X (hidden 4096,
per-rank intermediate 256/512, 10 or 11 expert slots), built with hipcc on first use
and launched through ctypes from the aiter MoE runner. Same weight/scale layouts as
aiter fused_moe, bf16 activations, on by default on gfx950 (SGLANG_ROCM_SMALLM_MOE=0
turns it off), falls back to aiter on any build/load/launch failure and above 40 tokens.
Adds a registered AMD unit test against a dequantized torch reference.
ROCm 10 images ship a second libamdhip64 in _rocm_sdk_devel; the unversioned
CDLL("libamdhip64.so") picked it and every launch on a torch stream failed with
709 (context is destroyed). Initialize CUDA first and bind to torch's copy via
find_loaded_library(). Unit test passes on rocm720 and rocm10 MI35x images.
The intermediate-512 (TP2) shape crosses over to aiter's flydsl path much
earlier (~20 tokens vs ~44), so it stays on aiter for now and gets its own
cap in a follow-up. The kernel entry points for 512 remain in the source.
@HaiShaw
HaiShaw merged commit 32290dd into sgl-project:main Sep 24, 2026
93 of 133 checks passed
kevin-mii added a commit to zcnrex/sglang that referenced this pull request Sep 24, 2026
Conflict in moe_runner/aiter.py: main (sgl-project#40204) added the gfx950 small-M
MXFP4 fused-MoE early return just before the fused_moe call; this branch
picks the fused_moe activation (popping one passed via fused_moe_kwargs)
at the same spot. Kept main's early return first and this branch's
activation selection after it, since only the fused_moe call reads it.
The small-M path requires hidden 4096 and topk 10/11, so MiniMax-M3
(hidden 6144, topk 5) never takes it.

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.

4 participants