Skip to content

[AMD] One-launch small-M MoE router for Qwen3.5 on gfx950 - #41133

Merged
HaiShaw merged 5 commits into
sgl-project:mainfrom
chuyeh:amd/qwen35-smallm-router
Oct 8, 2026
Merged

HaiShaw merged 5 commits into
sgl-project:mainfrom
chuyeh:amd/qwen35-smallm-router

Conversation

@chuyeh

@chuyeh chuyeh commented Sep 24, 2026 •

Copy link
Copy Markdown
Contributor

Motivation

On MI355X at low concurrency, each Qwen3.5-397B MoE layer routes a verify batch of M = 4 or 8 tokens (4 MTP verify tokens × 1–2 running requests) through three launches before the small-M MoE kernel from #40204:

  • the router gate GEMM ([M, 4096] x [512, 4096])
  • _router_triton_kernel (softmax over 512 experts, top-10, renormalize; one wave per token doing 10 serial argmax passes)
  • _fused_append_shared_experts_with_weights_kernel (the shared-expert gate GEMV and slot 10)

All three are latency-bound at these sizes. Together they take 19.9 µs per layer in a verify step (eager trace, M = 4). Across 60 layers that is about 1.2 ms per verify step, and the MTP layer adds more in draft and draft extend. This PR replaces them with one launch.

Modifications

  • New python/sglang/kernels/ops/moe/smallm_router_gfx950/ (__init__.py, smallm_router.hip). It is built with hipcc on first use and launched through ctypes, reusing [AMD] Small-M MXFP4 fused-MoE kernel for gfx950 (Qwen) #40204's _hipcc, _hip_lib, _check and _Kernel. No build-time or wheel change.
  • One launch of 257 or 513 blocks, depending on M. Each block computes one or two of the 513 gate rows (512 routed experts plus the shared-expert gate, read in place from gate.weight and shared_expert_gate.weight) with fp32 accumulation. The last block to finish, elected by a self-resetting ticket counter so the kernel is graph-safe, rounds the routed logits to bf16, then runs top-10 (lowest id wins ties, as in _router_triton_kernel), softmax and renormalize. It writes slot 10 = (512, sigmoid(shared logit) * scale). The output is exactly the topk_ids [M, 11] int32 / topk_weights [M, 11] fp32 that smallm_moe_gfx950 already consumes.
  • qwen2_moe.py: one guarded call at the top of _forward_router_experts (5 lines). The guard requires gfx95, 1 <= M <= 8, bf16 hidden 4096, gate [512, 4096] and shared gate [1, 4096] in bf16, exactly one fused shared expert, EP 1, the auto/aiter MoE runner, top-10 with renormalize, no EPLB remap or simulated routing, and not inside a piecewise CUDA graph. Everything else takes today's path. The routed-expert capture and expert-distribution recorder hooks still run.
  • The kernel is used for M <= 8 only. The single-block top-k tail makes it slower than today's path from 9 tokens up (0.80x at 9, 0.45x at 40), so c4 verify batches (16 tokens) keep today's path.
  • On by default on gfx950; SGLANG_ROCM_SMALLM_ROUTER=0 turns it off. A build or load failure disables it for the process.
  • Registered AMD unit test test/registered/amd/test_smallm_router_gfx950.py. It covers ids and weights against today's moe_fused_gate + fused_append_shared_experts_with_weights at 1 / 3 / 4 / 8 tokens on exactly representable inputs (many ties), fallback at 9 and 41 tokens and for unsupported dtype and shape, the off switch, and graph replay == eager. It passes on v0.5.20-rocm720-mi35x-20260923.

Accuracy Tests

Router logits are rounded to bf16 before top-k, as today. Against an fp64-then-bf16 reference on real activations (10,022 token rows), this kernel agrees on 100% of expert choices. Today's path agrees on 98.3%, because its bf16 gate GEMM is not correctly rounded at M >= 2.

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 (InferenceX
task), thinking disabled, 16384-token output budget, temperature 0. At most 2 running requests, so every verify batch (<= 8 tokens) is inside the kernel's range.

Configuration Strict-match Flexible-extract
Baseline (main 1417345f5f) 97.73% (1289/1319) 97.73% (1289/1319)
This PR 97.57% (1287/1319) 97.57% (1287/1319)

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

Configuration Score
Baseline (main 1417345f5f) 0.871 (1379/1584)
This PR 0.881 (1396/1584)
Reference (Qwen/Qwen3.5-397B-A17B model card) 0.884

Same GSM8K task at TP2 (thinking disabled, at most 2 running requests):

Configuration Strict-match Flexible-extract
Baseline (main 1417345f5f) 97.57% (1287/1319) 97.57% (1287/1319)
This PR 97.65% (1288/1319) 97.65% (1288/1319)

Benchmarking and Profiling

Kernel microbenchmarks

MI355X, one GPU, HIP-graph replay, 512 MB memset before each call to flush L2 / MALL, routing for one MoE layer, in µs:

Tokens 3 kernels today This kernel Change
1 12.39 8.54 −31%
2 12.82 8.80 −31%
4 12.90 9.05 −30%
8 13.26 11.66 −12%

In the live server the step saves more than this table suggests: today's _router_triton_kernel runs at 11.2 µs per layer in verify and 23.6 µs in draft extend (about 7 µs in isolation), while the new kernel costs about the same everywhere.

AgentX benchmark

Setup:

  • Hardware: MI355X, TP4 / EP1.
  • Model and recipe: the InferenceX MI355X agentic recipe for this model, image rocm/sgl-dev:v0.5.20-rocm720-mi35x-20260923.
  • Server flags: --kv-cache-dtype fp8_e4m3 --mem-fraction-static 0.80, EAGLE MTP 3 / 1 / 4 with simulated acceptance length 3.39, BF16 MTP layer as in the recipe.
  • Client: full 3600 s aiperf AgentX trace replay.
  • Order: one configuration at a time on the same four GPUs, nothing else on the node.

Baseline is main 1417345f5f (includes #39901, #39902 and #40204); this PR is the same commit plus this PR.

Concurrency 1:

Metric Baseline This PR Change
TPOT median (ms/token) 2.60 2.48 −4.6%
TPOT p90 (ms/token) 3.04 2.94 −3.3%
ITL p95 (ms) 3.14 3.03 −3.5%
TTFT p50 / p90 (ms) 372 / 795 392 / 797 +5.2% / +0.2%
Requests completed in 3600 s 657 684 +4.1%
Interactivity (1 / TPOT) p50 / p90 (tok/s/user) 384.0 / 328.4 403.6 / 339.9 +5.1% / +3.5%
E2E normalized interactivity p50 / p90 (tok/s/user) 238.3 / 137.6 241.8 / 140.8 +1.5% / +2.3%

Concurrency 8, average of two full runs per configuration (the second pass in reverse order). This is a no-regression check: verify batches are 4 x running requests, so up to 32 tokens, above this kernel's M <= 8 guard; the router change is inactive once more than two requests run:

Metric Baseline This PR Change
TPOT median (ms/token) 3.40 3.29 −3.4%
TPOT p90 (ms/token) 4.18 4.19 +0.4%
ITL p95 (ms) 4.62 4.63 +0.3%
TTFT p50 / p90 (ms) 602 / 1049 591 / 1032 −1.8% / −1.6%
Requests / generated tokens in 3600 s 1502 / 1.481M 1506 / 1.490M +0.2% / +0.7%
Interactivity (1 / TPOT) p50 / p90 (tok/s/user) 294.4 / 239.5 304.6 / 238.8 +3.5% / −0.3%
E2E normalized interactivity p50 / p90 (tok/s/user) 194.4 / 101.1 202.6 / 103.5 +4.2% / +2.4%

TPOT is reported as median and p90 rather than mean: in this replay a few short-output requests wait behind another request's long chunked prefill (up to 4.3 s for a 2-token answer), and a single such request moves the mean of a 3600 s run by up to 3 ms/token. Matching every request across the six concurrency-8 runs (baseline, #41133, #41134, two runs each), no request stalled in both runs of a PR without also stalling in a baseline run.

At TP2 with the official HiCache high-concurrency setup this router runs only at 8 or fewer tokens, so it is idle once more than two requests run. Concurrency 20 TPOT median is 6.94 ms versus 6.97 ms for the baseline. At concurrency 40 the median stays within about 1 ms of the baseline runs (14.7–15.3 ms), and the same tree with SGLANG_ROCM_SMALLM_ROUTER=0 lands on the same median.

Decode step time

Concurrency 1, one request, same image and recipe, 15 requests per configuration over three alternating server blocks. Baseline is this PR with SGLANG_ROCM_SMALLM_ROUTER=0:

Context Baseline This PR Change
72k tokens 8.79 ms 8.07 ms −8.2%
200k tokens 9.06 ms 8.38 ms −7.5%

Checklist


CI States

Latest PR Test (Base): ✅ Run #36722177674
Latest PR Test (Extra): ❌ Run #36722177232
Latest PR Test (AMD ROCm 10): ❌ Run #36722177864

At 1-8 tokens the router GEMM, softmax top-k and shared-expert append were three latency-bound launches ahead of the small-M MoE kernel; one kernel now writes the same 11-slot ids and weights.

Co-authored-by: Cursor <cursoragent@cursor.com>
@chuyeh
chuyeh marked this pull request as ready for review September 28, 2026 05:54
@yichiche yichiche added the run-ci CI: run the baseline test suite on this PR label Sep 29, 2026
yichiche and others added 4 commits September 29, 2026 12:28
…router

New SGLANG_* flags belong in environ.py; the router's two DPP reductions now share one helper (bit-identical output).

Co-authored-by: Cursor <cursoragent@cursor.com>

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

Guarded by _use_aiter checks plus a process-wide disable flag, this fuses the gate GEMM, softmax top-10, and shared-expert append into one launch, used only when M<=8, with a clear perf gain at low concurrency. LGTM.

@HaiShaw
HaiShaw merged commit b2cd65b into sgl-project:main Oct 8, 2026
300 of 361 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

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.

3 participants