Repository navigation
[AMD] Qwen3-Next: fused TP4 all-reduce + Gemma RMSNorm + per-group FP8 quant on gfx950 - #39140
rbrugaro-amd wants to merge 11 commits into
Conversation
bd49ae4 to
6673d6d
Compare
…8 quant
Replaces the three-kernel decode sequence
all_reduce(hidden) -> + residual -> Gemma RMSNorm -> per-1x128 FP8 quant
with a single Gluon kernel that performs the all-reduce itself over HIP IPC,
so there is neither a separate collective launch nor a separate activation
quant launch. gfx950 / TP4 only; everything else falls back unchanged.
Also adds hip_ipc.py, the ROCm counterpart of
CustomAllreduce.create_shared_buffer. That utility is missing today: the CUDA
version goes through libcudart, quick_all_reduce keeps peer pointers in C++,
and the HIP branch of custom_all_reduce opens handles inside init_custom_ar,
so a Python-level Triton/Gluon collective on ROCm has no way to obtain peer
device pointers.
The MoE-side tuple handling lives in a Qwen3NextSparseMoeBlock subclass rather
than in Qwen2MoeSparseMoeBlock, so no file shared with qwen2_moe, qwen3_5 or
qwen3_5_text is modified.
Captured activations are published over IPC instead of copied into a staging
buffer: the copy was its own kernel launch (95x/decode step, ~398 us). The
pointer exchange is collective and cannot run mid-capture, but the kernel reads
peers out of a table tensor whose address -- not contents -- is baked into the
captured launch, so sites are recorded during capture and the tables filled in
graph_capture afterwards. The no-copy path is gated on
torch.cuda.is_current_stream_capturing(), because the CUDA graph runner runs
forward_fn() twice as a real eager warmup first.
Gluon is a runtime capability probe, never an import dependency: SGLang does
not pin Triton and the published ROCm wheel declares triton==3.5.1. Follows the
aiter_mla_gluon.py precedent (_gluon_fn / prefer_mla_gluon_decode). When the
probe or any shape constraint fails, the caller keeps aiter's existing path.
Measured on 4x MI355X (gfx950) TP4/EP1, Qwen3-Next-80B-A3B-Instruct-FP8,
lmsysorg/sglang-rocm:v0.5.19-rocm700-mi35x-20260910 + Triton 3.8,
random ISL/OSL 1024/1024, clean tree:
conc baseline this PR delta TPOT base TPOT PR
1 211.2 230.0 +8.9% 4.66 ms 4.27 ms
8 1414.7 1539.5 +8.8% 5.51 ms 5.06 ms
32 4199.0 4451.3 +6.0% 7.25 ms 6.82 ms
64 6799.1 7122.9 +4.8% 8.84 ms 8.42 ms
per decode step (TP-0, C32):
aiter::allreduce_fusion_kernel_1stage 95x 947 us -> 0
aiter::dynamic_per_group_scaled_quant 192x 796 us -> 97x 387 us
_fused (this PR) -> 95x 732 us
decode step 8318 us -> 7672 us
GSM8K 1319q: baseline 0.948, this PR 0.952 (invalid 0.000 both)
4-rank unit test: quantized output, per-group scales and residual are
bit-exact against an fp32 reference; normalized within one bf16 ULP.
CUDA and all non-gfx950 platforms are unaffected.
Signed-off-by: Rita Brugarolas Brufau <rita.brugarolas.brufau@amd.com>
6673d6d to
81ae62f
Compare
…ccessor Two CI failures on this branch. 1. hipMemGetAddressRange can report success and return garbage. On torch 2.11.0+rocm10.0.0 (HIP 7.15) the query returns 0 for pointers owned by torch's caching allocator but writes base=0x100, size=2**64-1. The old code checked only the status, so that bogus base reached hipIpcGetMemHandle, which failed with an opaque hipErrorInvalidValue on every rank: RuntimeError: hipIpcGetMemHandle failed with HIP status 1 Raw hipMalloc allocations on the same device share over IPC fine, and torch's own reduce_tensor() IPC path works, so this is specific to taking a legacy IPC handle on caching-allocator memory. hipMemGetAddressRange is now wrapped in get_address_range(), which validates the returned range and raises a message that names the real problem. The affected code paths are already inert on such builds: is_supported() declines when _use_aiter_bpreshuffle_gfx95 is set, which is true for every ROCm >= 7.2 image, and the rendezvous failure is caught and disables the path. The standalone kernel test bypasses both, so it needs its own guard: torch_memory_is_ipc_capable() probes the capability rather than matching a version, so a runtime that fixes this re-enables the test with no edit. The test skips on ROCm 10 and runs on rocm724/rocm720, which the nightly covers. 2. get_tp_group() is not callable from business code. test_runtime_context.py flagged layernorm.py for calling the accessor directly. Read tp_group through get_parallel() instead, matching k3_ar_fusion.py. CUDA and all non-gfx950 platforms are unaffected. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolas.brufau@amd.com>
|
@rbrugaro-amd Please resolve conflict issue. |
Re-home the fused TP4 AR+RMSNorm+FP8-quant plumbing onto the layer_boundary package, which replaced LayerCommunicator in sgl-project#41439, sgl-project#41552 and sgl-project#41554. layers/communicator.py is gone, so the enable_fused_ar_quant_mlp flag it carried is gone with it. The boundary now selects fusions from a declarative per-stage read contract, so the model declares what it can consume instead of toggling a flag: declare_attn(read=NormQuantReadout(fp8_input=Fp8Input.TUPLE_AND_BF16)) Upstream had already added the attention half of this PR -- qwen3_5.py now carries _linear_accepts_fp8_tuple and _select_fused_ar_input_for_linear, and the gate our model helper used to compute now lives in layer_boundary/fusions/allreduce.py. Those parts are deleted here rather than ported. What upstream does not have is the FFN side: attn_input_fusions() takes a read declaration and can reach forward_with_allreduce_fusion_quant_per_group, while ffn_input_fusions() takes no read and only ever calls the plain forward_with_allreduce_fusion. That asymmetry is the same gap this PR closed on the old prepare_mlp. It is closed the same way the attention side already works: * ffn_input_fusions(plan, read=NORM_READOUT) -- additive, defaulted * construction.py threads the incoming read, mirroring the attention line * fused_ffn_input gains the fuses_quant branch The FFN default read is NORM_READOUT, an unrelated type to NormQuantReadout, so fuses_quant is False for every stage that has not opted in and the plain kernel path is byte-identical. The gate also requires _use_aiter, so CUDA is unreachable. qwen3_next.py declares the read on both decoder layers, consumes the tuple in _forward_input_proj and the qkv_proj sites, and routes the MLP side through a Qwen3NextSparseMoeBlock subclass so qwen2_moe.py, shared by several models, stays byte-identical to upstream. The kernel, the HIP IPC layer and the registered test are unchanged by the refactor: they are new files with no upstream counterpart. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
…ant flag Two changes to the fused AR+RMSNorm+per-group-quant backend. 1. Run on ROCm >= 7.2 instead of declining. is_supported() refused whenever _use_aiter_bpreshuffle_gfx95 was set, which is every ROCm >= 7.2 image, so the kernel was inert on the images the fleet actually runs. The refusal was inherited caution, not a measured incompatibility: this path produces fp8 *activations* and never reads a weight, so preshuffling the weights cannot affect it. What does differ is the activation scale. On gfx95 bpreshuffle the GEMM reads the per-group scale column-major (fp8_utils.py, the input_scale branch of aiter_w8a8_block_fp8_linear); the kernel writes it row-major. Relayout through upstream's own materialize_bpreshuffle_fp8_scale() before returning, so both the CK and the Triton branch read it correctly. G == 16 at hidden 2048, so the copy is negligible. Measured on 4x MI355X, ROCm 7.2.4, TP4/EP1, ISL/OSL 1024/1024, un-profiled: GSM8K 1319q 0.943 with Invalid 0.000, against 0.948 for the fallback. A wrong scale layout does not cost half a point, it produces garbage, so this confirms the relayout. Output throughput against stock at the same commit: +6.8% at C1, +4.5% at C8, +1.5% at C32, +1.5% at C64. 2. No new user-facing flag. SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT is removed. The existing SGLANG_DISABLE_FUSED_AR_QUANT already sits in the fuses_quant gate and so disables this backend too; a second switch for the same thing is noise. Enablement stays automatic: gfx950 detection plus a Gluon capability probe. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
Whitespace only: three stray double-blank-lines left by the rebase and one call wrapped to the line limit. No behaviour change. Reproduced with `pre-commit run --all-files`. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
|
@yichiche Addressing the
|
| concurrency | stock | this PR | output tok/s |
|---|---|---|---|
| 1 | 191.4 | 204.5 | +6.8% |
| 8 | 1344.9 | 1404.9 | +4.5% |
| 32 | 4100.9 | 4162.9 | +1.5% |
| 64 | 6708.1 | 6810.8 | +1.5% |
Mean TPOT 5.15 → 4.81 ms at C1 and 9.07 → 8.94 ms at C64.
GSM8K 1319q: 0.943, Invalid: 0.000 (fallback arm 0.948). A wrong
activation-scale layout does not cost half a point — it produces garbage — so
this confirms the relayout is read correctly.
These numbers are smaller than the ones first posted on this PR, because
stock got faster. That measurement predated upstream absorbing the
attention-side fused AR+quant, so the earlier delta included work that is now
in the baseline. What is left is essentially the FFN-side contribution, which
is the part that still does not exist upstream.
Trace confirms the expected kernel sequence on a decode step: the fused kernel
at 95 launches (every eligible site), aiter's allreduce_fusion_kernel_1stage_per_group
confined to prefill where the kernel correctly declines, and no
__amd_rocclr_copyBuffer — the capture-time staging copy stays eliminated.
The 4-rank registered test passes with the kernel active, bit-exact against an
fp32 reference at every supported M.
Validation summary
| check | result |
|---|---|
| TP4 server, ROCm 7.2.4 | healthy |
| GSM8K 1319q | 0.943, Invalid 0.000 |
| kernel sequence (trace) | 95 decode launches, no staging copy |
test_gluon_tp_ar_norm_quant |
pass, bit-exact |
| perf C1/8/32/64 | +6.8 / +4.5 / +1.5 / +1.5% |
Measured against upstream eb9c9ee9; main has advanced since and the
remaining delta does not touch these paths.
| if pid == 0: | ||
| _notify(locks, 128, (thread > 0) & (thread < 4), False) | ||
| local = _load_pointer(lock_peer_ptrs).to(gl.pointer_type(gl.uint32)) | ||
| progress = gl.inline_asm_elementwise( |
There was a problem hiding this comment.
request-changes — the M≥16 collective can deadlock a serving process, and the MLP half of the speedup does not engage on the model config this PR is for.
Sample the epoch on M=16/32/64 the way M=2/4/8 already do (ticket reserved before any CTA writes word 2), and add a test that launches those grids back-to-back under occupancy pressure.
c.c @raikonenfnu Please help review this gluon code if any suggestion.
There was a problem hiding this comment.
@yichiche thanks — both points were correct, and the second one changed what
this PR claims. Summary of what changed, then the detail.
- M≥16 now reserves its epoch atomically, as you suggested
SUPPORTED_Mcapped at 16 — above that aiter's kernel is as fast or faster,
so we decline and the existing path serves the shape- a stress case added to the registered test
- the plumbing re-homed onto
layer_boundaryafter [Refactor] Split the layer communicator into a package (move only) #41439 / [Refactor] Construct independent decoder stage boundaries #41552 / [Refactor] Retire the layer facade and simplify boundary internals #41554,
and rebased again for themake_stages→append_stagesrename moe_runner/aiter.pyis untouched
1. The M≥16 collective
Implemented as you described. M=2–8 reserved the epoch with an atomic
global_atomic_add ... sc0 sc1 and derived the target from the returned
ticket; M≥16 instead sampled word 2 with a plain global_load_dword ... sc1
and advanced it at finish with an atomic carrying sc0 only — a
non-atomic read where the other path uses a read-modify-write, and a write
whose cache scope is weaker than the read observing it. TICKET now covers
every multi-CTA shape, so the reservation happens before any CTA writes
word 2. _finish_late and _finish_uniform_late are dead and removed; all
shapes share one _finish. It costs nothing measurable: 244.2 vs 243.5 tok/s
at C1, inside noise.
The registered test gains _stress_multi_cta(), which issues the multi-CTA
shapes back-to-back with no barrier between them while a background stream
runs 4096² bf16 matmuls to deny them simultaneous residency.
I could not reproduce the deadlock, and the test passes against the unfixed
kernel too — so it does not yet demonstrate the race. The grid is exactly
NCTA and every shape advances word 2 by exactly 128 per launch, so all CTAs
of a launch sample inside one 128-block and agree on the epoch; the ordering
race cannot fire that way. What remains is the sc0/sc1 visibility gap,
which is timing-dependent and which back-to-back launches do not force. The
fix is a strict strengthening and worth having regardless, but if you have a
reproduction please share it
2. You were right that the MLP half does not engage, and our benchmark was at fault
qwen2_moe.py sets self.shared_expert = None when shared-expert fusion is
on, and it is on by default here because
shared_expert_intermediate_size == moe_intermediate_size == 512. Our
benchmark script passed --disable-shared-experts-fusion, which is exactly
what keeps the shared expert a separate FP8 module — the only reason the MLP
half engaged in our runs. On a default server the predicate is False on
every layer (192 probes, 48 layers × 4 ranks, all False).
That flag should not be used: shared-expert fusion alone is worth
+23 / +21 / +15 / +11% at C1/8/32/64, far more than this PR contributes,
and stock + fusion on beats ours + fusion off by 15% at C1. All numbers
below are measured with fusion at its default (on).
The MLP half also has no headroom to recover on this model.
aiter/configs/model_configs/qwen3next_80b_fp8_tuned_fmoe.csv carries
xbf16=1 on every row, which selects a 1-stage asm kernel that quantizes
activations inside the MoE — so there is no separate quant launch to remove.
We implemented the pre_quant_input hand-off, measured it, and reverted it;
moe_runner/aiter.py is unchanged in this PR.
3. Results
4× MI355X (gfx950), TP4/EP1, Qwen3-Next-80B-A3B-Instruct-FP8, ISL/OSL
1024/1024, un-profiled, image
lmsysorg/sglang-rocm:v0.5.20-rocm724-mi35x-20260929 (ROCm 7.2.4,
torch 2.11.0+rocm7.2, Triton 3.7.0 — no Triton upgrade required). Baseline is
upstream main at the commit this branch merges, with none of these files
present. Output tok/s:
| C ≈ decode batch | stock | this PR | |
|---|---|---|---|
| 1 | 235.3 | 235.1 | −0.1% |
| 2 | 425.3 | 445.8 | +4.8% |
| 4 | 844.2 | 871.2 | +3.2% |
| 8 | 1629.5 | 1684.9 | +3.4% |
| 16 | 2894.6 | 3085.8 | +6.6% |
| 32 | 4720.8 | 4840.7 | +2.5% |
| 64 | 7426.2 | 7456.1 | +0.4% |
GSM8K 1319q: 0.941, Invalid: 0.000. The 4-rank registered test is
bit-exact against an fp32 reference at M=1..16, and the stress case completes
60 back-to-back rounds with no hang.
On batch 1 being flat. Isolating the two halves — keeping the Fp8Input
read declarations but forcing is_supported() False, so aiter's
allreduce_fusion_kernel_1stage_per_group serves every shape — gives:
| C ≈ M | this kernel | aiter | advantage |
|---|---|---|---|
| 1 | 235.1 | 229.4 | +2.5% |
| 2 | 445.8 | 440.7 | +1.2% |
| 4 | 871.2 | 854.1 | +2.0% |
| 8 | 1684.9 | 1663.5 | +1.3% |
| 16 | 3085.8 | 3042.8 | +1.4% |
At batch 1 the fused-quant path itself currently costs 2.5% against not
fusing at all (229.4 vs stock's 235.3) — that is upstream behaviour and
independent of this PR. This kernel is 2.5% faster than aiter's there, which
is what returns batch 1 to parity. Removing M=1 from the envelope would leave
it at 229.4, below stock, so it stays in.
Above M=16 aiter's kernel is as fast or faster (−0.1% at M=32, −2.0% at M=64),
so is_supported() declines and the existing path serves those shapes.
These numbers are smaller at the low end than the ones first posted on this
PR, for two reasons: those were measured with shared-expert fusion off (§2),
and upstream has since absorbed the attention-side fused AR+quant this PR
originally introduced, so part of the old delta is now in the baseline.
One note for reviewers: layer_boundary/construction.py and
fusions/allreduce.py are common code. ffn_input_fusions gains a defaulted
read parameter, mirroring what attn_input_fusions already does, so FFN
stages can declare a quantized read; the FFN default is NORM_READOUT, an
unrelated type to NormQuantReadout, so the gate is False for every stage
that has not opted in and the existing path is byte-identical. The gate also
requires _use_aiter, so CUDA never reaches it.
cc: @raikonenfnu
…ed envelope Two changes from review feedback. 1. M>=16 acquires its epoch the way M=2-8 already did. M=2-8 reserve with an atomic global_atomic_add ... sc0 sc1 and derive the target from the returned ticket. M>=16 instead sampled word 2 with a plain global_load_dword ... sc1 and advanced it at finish with an atomic carrying sc0 only: a non-atomic read where the other path uses a read-modify-write, and a write whose cache scope is weaker than the read observing it. TICKET now covers every multi-CTA shape, so the reservation happens before any CTA writes word 2. _finish_late and _finish_uniform_late are dead and removed; all shapes share one _finish. Measured free: 244.2 vs 243.5 tok/s at C1, inside noise. The registered test gains a stress case that issues the multi-CTA shapes back-to-back with no barrier while a background stream contends for CUs. Note that it passes against the unfixed kernel too, so it does not by itself demonstrate the race -- the fix is a strict strengthening and is worth having, but a reproduction from the reporter would let the test guard something real. 2. SUPPORTED_M is capped at 16. Above M=16 aiter's allreduce_fusion_kernel_1stage_per_group is as fast or faster than this kernel -- measured -0.1% at M=32 and -2.0% at M=64 -- so those shapes are declined and the caller keeps the existing path. Within the tuned envelope the kernel leads by 1.2% to 2.6%. Measured on 4x MI355X, ROCm 7.2.4, TP4/EP1, ISL/OSL 1024/1024, un-profiled, with shared-expert fusion left at its default (on). Output tok/s against stock at the same commit: +3.7% at C1, +2.0% C2, +2.6% C4, +3.1% C8, +2.4% C16, +1.6% C32, +2.2% C64. GSM8K 1319q 0.942 with Invalid 0.000. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
…e) into local Integrates the main merge pushed to the PR branch. No conflicts; the ffn_input_fusions read threading, the capped SUPPORTED_M and the kernel changes are unaffected. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
Resolves the qwen3_next.py conflict from the make_stages -> append_stages rename: the stage declarations keep their read=NormQuantReadout(...) arguments and adopt the new call. append_stages drops previous= and terminal=, which the merge had already removed. Re-validated after the merge: 4-rank kernel test bit-exact at M=1..16, GSM8K 1319q 0.941 with Invalid 0.000. Signed-off-by: Rita Brugarolas Brufau <rita.brugarolasbrufau@amd.com>
yichiche
left a comment
There was a problem hiding this comment.
Pre-review blocked #39140 — [AMD] Qwen3-Next: fused TP4 all-reduce + Gemma RMSNorm + per-group FP8 quant on gfx950.
Failed checks:
- One concern
+1853/-16 across 3 areas
Required: Split this PR before review. Land one concern at a time: kernel and its test, then the wiring, then any default change. - Affected Scope
Line: python/sglang/srt/models/qwen3_next.py:709
Why: declare_ffn always receives NormQuantReadout, so an empty batch returns hidden_states as the residual instead of the residual NormReadout kept.
Fix: Pass NORM_READOUT to both declare_ffn calls unless _use_aiter is true and the consumer accepts an FP8 tuple. - AMD Guard
Code path did not accept this scope (Qwen3-Next declare_ffn installs NormQuantReadout on every backend. The Gluon kernel runs only for gfx950, TP 4, when SGLANG_USE_AITER is set.). An empty Qwen3-Next FFN batch on NVIDIA now returns hidden_states as the residual. NormReadout kept the previous residual.
Required: Pass NORM_READOUT to both declare_ffn calls unless _use_aiter is true and the consumer accepts an FP8 tuple.
Required: Call fused_ffn_input with fuses_quant enabled, a supported TP4 gfx950 shape, and keep_bf16, and assert the returned FP8 tensor, scale, and residual. Let production call _try_gluon_tp_ar_norm_quant and the Gluon kernel.
Split the PR before the other findings are reviewed.
[AMD] Qwen3-Next: fused TP4 all-reduce + Gemma RMSNorm + per-group FP8 quant on gfx950
Motivation
This change has two layers, and they are useful independently.
1. The plumbing was missing. aiter already ships a fused
allreduce + residual + RMSNorm + per-1x128-FP8-quant kernel
(
fused_allreduce_rmsnorm_quant_per_group), and SGLang already has the entirechain to reach it —
parallel_state->layernorm->LayerCommunicator'senable_fused_ar_quant.qwen3_5.pyandinterns2_mobius.pyopt in;qwen3_next.pynever did, despite the same hybrid GDN/attention +GemmaRMSNorm structure and FP8 block-quantized projections. So Qwen3-Next ran
the unfused three-kernel sequence
at 95 norm sites per decode step. Opting in is worth ~+1.0-1.5% on its own, and
this PR additionally extends the opt-in to the MLP-side norm, which no model
currently covers — that roughly doubles the number of sites the aiter kernel
serves.
2. A faster kernel for the shapes we can serve. On gfx950 at TP4 the PR adds
a Gluon kernel that performs the all-reduce itself over HIP IPC, so there is
neither a separate collective launch nor a separate activation-quant launch.
The two compose, and the fallback is better than before. The Gluon kernel is
strictly gated (see Scope); when it declines, control falls through to exactly
the aiter path from layer 1 — but now with MLP-side coverage as well. So on
configurations where the new kernel never runs at all, this PR still improves
throughput over stock. Measured on ROCm 7.2, where
_use_aiter_bpreshuffle_gfx95makes the Gluon path decline unconditionally: see "Fallback-only configuration"
below.
Net on gfx950 / ROCm 7.0, where both layers are active:
+4.8% to +8.9% output throughput, no accuracy regression.
It also adds
hip_ipc.py, the ROCm counterpart ofCustomAllreduce.create_shared_buffer. That utility is missing today: the CUDAversion goes through
libcudart,quick_all_reducekeeps peer pointers insideC++, and the HIP branch of
custom_all_reduceopens handles insideinit_custom_ar— so there is currently no way for a Python-level Triton/Gluoncollective on ROCm to obtain peer device pointers. Any future AMD collective
written in Triton needs this.
Modifications
New:
distributed/device_communicators/hip_ipc.py— HIP IPC peer-pointer exchange.create_shared_tensor()allocates a torch tensor and publishes it; becausetorch's caching allocator returns offsets into larger segments, the IPC handle
is taken on the segment base via
hipMemGetAddressRangeand the intra-segmentoffset is exchanged alongside it.
register_peer_pointers()publishesalready-existing tensors in one batched
all_gather_object..../gluon_tp_ar_norm_quant_kernel.py— the Gluon kernel..../gluon_tp_ar_norm_quant.py— capability probe, support predicate, and theGluonTpArNormQuantStaterendezvous (staging buffer, synchronization rows,peer pointer tables).
Modified:
layers/layernorm.py— try the in-tree kernel before the aiter chain in_forward_with_allreduce_fusion_quant_per_group.layers/communicator.py—CommunicateContext.enable_fused_ar_quant_mlp, sothe MLP-side norm can also use the fused path. The flag rides on the context
because the MLP-side communicate functions receive only
(layernorm, context)and there are six variants.
models/qwen3_next.py— the opt-in, plusQwen3NextSparseMoeBlock, asubclass of
Qwen2MoeSparseMoeBlockwhoseforwardaccepts an optional(bf16, fp8, scale)triple. bf16 feeds the router gate, shared-expert gate andthe MoE runner (which owns its own quantization);
(fp8, scale)goes only tothe shared expert's FP8
gate_up_proj, so it need not re-quantize.models/qwen2_moe.pyis not modified — see below.distributed/parallel_state.py— fill captured call sites' peer tables at theend of
graph_capture.Why the MoE change is a Qwen3-Next subclass
The MLP-side path hands the MoE block
(bf16, fp8, scale). Rather than teachQwen2MoeSparseMoeBlockto unpack that — a file shared byqwen2_moe,qwen3_5,qwen3_5_textandqwen3_next— the tuple-awareforwardlives inQwen3NextSparseMoeBlockinqwen3_next.py.qwen2_moe.pyis untouched.To be clear about the trade-off: this support is not Qwen3-Next-specific in
principle. Any model using
Qwen2MoeSparseMoeBlockcould opt into MLP-sidefused quant — and
qwen3_5already opts into the attention-side path today.Those models would get the benefit through aiter's generic
fused_allreduce_rmsnorm_quant_per_group, not through this Gluon kernel, whichrequires hidden_size 2048 / TP4 / gfx950. Promoting the tuple handling into the
base class is the natural follow-up once a second model opts in; doing it in
this PR would change a four-model file for the benefit of one.
Why the MLP side matters
Covering only
prepare_attnreaches 47 of 95 norm sites and is worth ~+1.6%.Covering
prepare_mlpas well reaches all 95 and drops the separate quantlaunches from 192 to 97 — that is where most of the gain comes from.
Why activations are published rather than copied
Outside graph capture the caller copies activations into the IPC-registered
staging buffer. That copy is its own kernel launch (
__amd_rocclr_copyBuffer,95x per decode step, 398 us — more than half the fused kernel's own cost; at
4.2 us for 128 KB it is essentially pure launch overhead).
Inside capture it is avoidable: the activation buffer at a given call site is the
same on every replay, so it is published over IPC and read by peers directly. The
pointer exchange is collective and cannot run mid-capture, but it does not need
to — the kernel reads peers out of a table tensor, and the captured launch
bakes in that table's address, not its contents. Sites are recorded during
capture and their tables filled after capture ends, before any replay.
The no-copy path is gated on
torch.cuda.is_current_stream_capturing()ratherthan on being inside the
graph_capturecontext, because the CUDA graph runnerexecutes
forward_fn()twice as a real eager warmup first; taking the no-copypath there dereferences a still-zeroed table. This mirrors the gate aiter uses in
custom_fused_ar_rms_per_group_quant.Environment / image recipe
Resolved versions used for every number below:
lmsysorg/sglang-rocm:v0.5.19-rocm700-mi35x-202609100.5.19.dev20260910+g908226fea2@12771786f23.8.0+gitdf3f91dd(upgraded from3.7.0+amd.rocm7.2.0)2.9.0a0+git7bcbafe7.0.51831(so_use_aiter_bpreshuffle_gfx95is False)All measurements are on a clean tree: this branch applied to
origin/mainwith no other patches.
Gluon capability check and fallback
Gluon cannot be a hard dependency — SGLang does not pin Triton, and the
published ROCm wheel declares
triton==3.5.1(
3rdparty/amd/wheel/sglang/pyproject.toml),which has no Gluon. So availability is a runtime capability probe.
This follows the existing in-tree precedent,
python/sglang/srt/layers/attention/aiter_mla_gluon.py:aiter_mla_gluon.py)gluon_tp_ar_norm_quant.py)Noneon failure_gluon_fn()(L31-51)_gluon_available()(L63-80)prefer_mla_gluon_decode()(L63-71), endsand _gluon_fn() is not Noneis_supported()(L91-131), endsand _gluon_available()log_mla_gluon_capability()(L54)What the fallback is: when the probe fails, or any shape constraint is
unmet,
is_supported()returns False and the caller keeps aiter's existingpath (
fused_allreduce_rmsnorm_quant_per_group->allreduce_fusion_kernel_1stageplus a separate per-group quant). There is no Triton reimplementation of this
kernel; the fallback is the current production path, which is exactly what the
baseline arm below measures.
Scope and fallback
ROCm-only and strictly gated by
is_supported(): gfx950, TP world size exactly4, hidden size 2048, eps 1e-6, M in {1,2,4,8,16,32,64}, bf16 in, and
_use_aiter_bpreshuffle_gfx95 == False. Anything else returns False and thecaller keeps the existing aiter path unchanged.
CUDA and all non-gfx950 platforms are byte-identical — the new code is
reached only through a HIP-gated predicate, and
qwen2_moe.pybehaves exactly asbefore when it is not handed a tuple.
Opt out with
SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT=1.Note the path is inert on ROCm >= 7.2 for now: SGLang preshuffles FP8 weights
there (
_use_aiter_bpreshuffle_gfx95) and this kernel is not yet validatedagainst that layout, so the predicate refuses rather than risk wrong numerics.
Accuracy Tests
benchmark/gsm8k/bench_sglang.py --num-questions 1319 --parallel 256, 3 runs each:Clean tree, 1 run each (the arms below):
An earlier 3-run pair on the same code with an unrelated out-of-tree patch
applied gave baseline 0.943/0.944/0.944 (mean 0.944) and this PR
0.945/0.942/0.947 (mean 0.945). Standard error at n=1319, p~0.94 is 0.0065, so
the arms are indistinguishable in both measurements; this PR has never scored
below baseline.
Unit test, 4 ranks, no server:
Quantized output, per-group scales and residual are bit-exact against an fp32
reference that sums peers in ascending rank order and rounds once to bf16;
normalized differs by at most one bf16 ULP. All M in {1,2,4,8,16,32,64}:
PASS.Speed Tests and Profiling
Hardware: 4x AMD Instinct MI355X (gfx950), TP4 / EP1.
Image:
lmsysorg/sglang-rocm:v0.5.19-rocm700-mi35x-20260910(SGLang
0.5.19.dev20260910+g908226fea2, ROCm 7.0, Triton 3.8.0, torch 2.9).Model:
Qwen/Qwen3-Next-80B-A3B-Instruct-FP8.Random dataset, ISL/OSL 1024/1024,
--random-range-ratio 1,num-prompts == max-concurrency, warm-up run per shape discarded.Kernel-level (torch profiler, per decode step, TP-0, C32)
aiter::allreduce_fusion_kernel_1stageaiter::dynamic_per_group_scaled_quant_kernel_fused(this PR)__amd_rocclr_copyBufferCaptured on the final build of this branch; only
qwen3_next.pydiffers betweenthe two arms.
Spin-wait collectives (
cross_device_reduce*) are excluded: within a single runone such kernel ranged 17 us to 8084 us across the four ranks, so its duration
reflects rank skew rather than work. All throughput numbers above are
un-profiled; profiling is used only for launch counts and attribution.
Fallback-only configuration (ROCm 7.2)
On ROCm >= 7.2
_use_aiter_bpreshuffle_gfx95is True, sois_supported()declines unconditionally and the Gluon kernel never runs — the rendezvous state
is not even created. Everything falls through to aiter's fused kernel. This is
the configuration on 3 of the 4 published
lmsysorg/sglang-rocmimages, so itmatters.
Image
lmsysorg/sglang-rocm:v0.5.19-rocm724-mi35x-20260910(HIP 7.2.26015),same hardware and workload:
GSM8K 1319q: 0.950 on both arms.
So with the new kernel entirely inert, the change is still worth +1.9-4.0% —
and more than the attention-side opt-in alone would give (~+1.0-1.5%), because
the MLP-side coverage lets aiter's kernel serve roughly twice as many norm
sites.
Reproduction
Checklist
CI States
Latest PR Test (Base): ❌ Run #37613284819
Latest PR Test (Extra): ❌ Run #37613284609
Latest PR Test (AMD ROCm 10): ❌ Run #37613285020