Skip to content

[AMD] Qwen3-Next: fused TP4 all-reduce + Gemma RMSNorm + per-group FP8 quant on gfx950 - #39140

Open
rbrugaro-amd wants to merge 11 commits into
sgl-project:mainfrom
rbrugaro-amd:feat/qwen3-next-tp4-fused-collective
Open

rbrugaro-amd wants to merge 11 commits into
sgl-project:mainfrom
rbrugaro-amd:feat/qwen3-next-tp4-fused-collective

Conversation

@rbrugaro-amd

@rbrugaro-amd rbrugaro-amd commented Sep 11, 2026 •

Copy link
Copy Markdown
Contributor

[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 entire
chain to reach it — parallel_state -> layernorm -> LayerCommunicator's
enable_fused_ar_quant. qwen3_5.py and interns2_mobius.py opt in;
qwen3_next.py never did, despite the same hybrid GDN/attention +
GemmaRMSNorm structure and FP8 block-quantized projections. So Qwen3-Next ran
the unfused three-kernel sequence

all_reduce(hidden) -> + residual -> Gemma RMSNorm -> per-1x128 FP8 quant

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_gfx95
makes 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 of
CustomAllreduce.create_shared_buffer. That utility is missing today: the CUDA
version goes through libcudart, quick_all_reduce keeps peer pointers inside
C++, and the HIP branch of custom_all_reduce opens handles inside
init_custom_ar — so there is currently no way for a Python-level Triton/Gluon
collective 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; because
    torch's caching allocator returns offsets into larger segments, the IPC handle
    is taken on the segment base via hipMemGetAddressRange and the intra-segment
    offset is exchanged alongside it. register_peer_pointers() publishes
    already-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 the
    GluonTpArNormQuantState rendezvous (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, so
    the 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, plus Qwen3NextSparseMoeBlock, a
    subclass of Qwen2MoeSparseMoeBlock whose forward accepts an optional
    (bf16, fp8, scale) triple. bf16 feeds the router gate, shared-expert gate and
    the MoE runner (which owns its own quantization); (fp8, scale) goes only to
    the shared expert's FP8 gate_up_proj, so it need not re-quantize.
    models/qwen2_moe.py is not modified — see below.
  • distributed/parallel_state.py — fill captured call sites' peer tables at the
    end 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 teach
Qwen2MoeSparseMoeBlock to unpack that — a file shared by qwen2_moe,
qwen3_5, qwen3_5_text and qwen3_next — the tuple-aware forward lives in
Qwen3NextSparseMoeBlock in qwen3_next.py. qwen2_moe.py is untouched.

To be clear about the trade-off: this support is not Qwen3-Next-specific in
principle. Any model using Qwen2MoeSparseMoeBlock could opt into MLP-side
fused quant — and qwen3_5 already 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, which
requires 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_attn reaches 47 of 95 norm sites and is worth ~+1.6%.
Covering prepare_mlp as well reaches all 95 and drops the separate quant
launches 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() rather
than on being inside the graph_capture context, because the CUDA graph runner
executes forward_fn() twice as a real eager warmup first; taking the no-copy
path there dereferences a still-zeroed table. This mirrors the gate aiter uses in
custom_fused_ar_rms_per_group_quant.

Environment / image recipe

FROM lmsysorg/sglang-rocm:v0.5.19-rocm700-mi35x-20260910

# The base ships Triton 3.7.0+amd.rocm7.2.0. The Gluon kernel's launcher passes
# `llvm_fn_attrs=`, which only exists in Triton 3.8; on 3.7 the first launch
# dies with:
#   KeyError: 'Keyword argument llvm_fn_attrs was specified but unrecognised'
RUN python -m pip install --no-deps --force-reinstall --pre \
      --index-url https://download.pytorch.org/whl/nightly/rocm7.0 \
      "triton==3.8.0"

Resolved versions used for every number below:

base image lmsysorg/sglang-rocm:v0.5.19-rocm700-mi35x-20260910
sglang 0.5.19.dev20260910+g908226fea2 @ 12771786f2
triton 3.8.0+gitdf3f91dd (upgraded from 3.7.0+amd.rocm7.2.0)
torch 2.9.0a0+git7bcbafe
HIP 7.0.51831 (so _use_aiter_bpreshuffle_gfx95 is False)
GPUs 4x AMD Instinct MI355X, gfx950

All measurements are on a clean tree: this branch applied to origin/main
with 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:

role precedent (aiter_mla_gluon.py) this PR (gluon_tp_ar_norm_quant.py)
cached try/except import, None on failure _gluon_fn() (L31-51) _gluon_available() (L63-80)
predicate = shape constraints and capability prefer_mla_gluon_decode() (L63-71), ends and _gluon_fn() is not None is_supported() (L91-131), ends and _gluon_available()
startup capability log log_mla_gluon_capability() (L54) (not yet — happy to add)

What the fallback is: when the probe fails, or any shape constraint is
unmet, is_supported() returns False and the caller keeps aiter's existing
path
(fused_allreduce_rmsnorm_quant_per_group -> allreduce_fusion_kernel_1stage
plus 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 exactly
4, 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 the
caller 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.py behaves exactly as
before 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 validated
against 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):

baseline this PR
GSM8K 1319q 0.948 0.952
Invalid 0.000 0.000

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:

torchrun --nproc_per_node=4 test/srt/test_gluon_tp_ar_norm_quant.py --bench

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.

conc baseline tok/s this PR tok/s delta baseline TPOT this PR TPOT TPOT delta
1 211.2 230.0 +8.9% 4.66 ms 4.27 ms -8.4%
8 1414.7 1539.5 +8.8% 5.51 ms 5.06 ms -8.2%
32 4199.0 4451.3 +6.0% 7.25 ms 6.82 ms -6.0%
64 6799.1 7122.9 +4.8% 8.84 ms 8.42 ms -4.7%

Kernel-level (torch profiler, per decode step, TP-0, C32)

kernel baseline this PR
aiter::allreduce_fusion_kernel_1stage 95x, 947.1 us 0
aiter::dynamic_per_group_scaled_quant_kernel 192x, 796.2 us 97x, 387.4 us
_fused (this PR) 0 95x, 732.3 us
__amd_rocclr_copyBuffer 0 0
decode step 8318 us 7672 us (-7.8%)

Captured on the final build of this branch; only qwen3_next.py differs between
the two arms.

Spin-wait collectives (cross_device_reduce*) are excluded: within a single run
one 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_gfx95 is True, so is_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-rocm images, so it
matters.

Image lmsysorg/sglang-rocm:v0.5.19-rocm724-mi35x-20260910 (HIP 7.2.26015),
same hardware and workload:

conc baseline tok/s this PR tok/s delta baseline TPOT this PR TPOT
1 192.3 195.9 +1.9% 5.13 ms 5.03 ms
8 1359.8 1395.4 +2.6% 5.74 ms 5.60 ms
32 4152.7 4234.9 +2.0% 7.37 ms 7.22 ms
64 6916.7 7192.1 +4.0% 8.74 ms 8.38 ms

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

docker run -d --name sgl --device=/dev/kfd --device=/dev/dri \
  --group-add video --group-add render --ipc=host --shm-size 64g --network host \
  -v $HF_CACHE:/hf:ro \
  lmsysorg/sglang-rocm:v0.5.19-rocm700-mi35x-20260910 sleep infinity

# server (both arms)
SGLANG_USE_AITER=1 python3 -m sglang.launch_server \
  --model-path $MODEL --host 0.0.0.0 --port 8888 --trust-remote-code \
  --tensor-parallel-size 4 --kv-cache-dtype fp8_e4m3 --context-length 8192 \
  --chunked-prefill-size 16384 --ep-size 1 --dtype bfloat16 --quantization fp8 \
  --attention-backend aiter --moe-runner-backend aiter \
  --enable-aiter-allreduce-fusion --disable-radix-cache \
  --disable-shared-experts-fusion --mem-fraction-static 0.80 \
  --cuda-graph-max-bs-decode 64

# accuracy
python3 benchmark/gsm8k/bench_sglang.py --num-questions 1319 --parallel 256 --port 8888

# throughput, per concurrency C in 1 8 32 64
python3 -m sglang.bench_serving --backend sglang-oai --host 127.0.0.1 --port 8888 \
  --model $MODEL --tokenizer $MODEL --dataset-name random \
  --random-input-len 1024 --random-output-len 1024 --random-range-ratio 1 \
  --num-prompts $C --max-concurrency $C

# baseline arm: SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT=1

Checklist

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

CI States

Latest PR Test (Base): ❌ Run #37613284819
Latest PR Test (Extra): ❌ Run #37613284609
Latest PR Test (AMD ROCm 10): ❌ Run #37613285020

@github-actions github-actions Bot added the quant LLM Quantization label Sep 11, 2026
@rbrugaro-amd
rbrugaro-amd force-pushed the feat/qwen3-next-tp4-fused-collective branch 3 times, most recently from bd49ae4 to 6673d6d Compare September 14, 2026 20:12
@rbrugaro-amd
rbrugaro-amd marked this pull request as ready for review September 14, 2026 23:17
@kkHuang-amd kkHuang-amd added amd run-ci CI: run the baseline test suite on this PR labels Sep 18, 2026
…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>
@rbrugaro-amd
rbrugaro-amd force-pushed the feat/qwen3-next-tp4-fused-collective branch from 6673d6d to 81ae62f Compare September 18, 2026 21:09
rbrugaro-amd and others added 2 commits September 22, 2026 14:53
…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>
@yichiche

Copy link
Copy Markdown
Collaborator

@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>
@rbrugaro-amd

Copy link
Copy Markdown
Contributor Author

@yichiche
Rebased onto current main and re-measured. Conflict resolved — the branch is
mergeable again.

Addressing the layer_boundary refactor

layers/communicator.py was removed by #41439, #41552 and #41554, and
LayerCommunicator / CommunicateContext with it, so the
enable_fused_ar_quant_mlp flag this PR previously added no longer had a home.
Rather than port the flag, the plumbing now follows the declarative per-stage
read contract the refactor introduced: the model declares what it can consume
via declare_attn(read=NormQuantReadout(fp8_input=...)) and the boundary picks
the fusion. This PR no longer adds any constructor flag or any user-facing
environment variable
— SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT is gone too,
since the existing SGLANG_DISABLE_FUSED_AR_QUANT already gates this backend.

A good part of the original PR is now redundant and has been deleted rather
than rebased: qwen3_5.py already carries _linear_accepts_fp8_tuple and
_select_fused_ar_input_for_linear, and the enablement gate that used to live
in the model is now centralized in layer_boundary/fusions/allreduce.py.

What remains is the FFN side. attn_input_fusions(plan, read) takes a read
declaration and can reach forward_with_allreduce_fusion_quant_per_group;
ffn_input_fusions(plan) takes no read and only ever calls the plain
forward_with_allreduce_fusion. That is the same gap this PR closed on the old
prepare_mlp, and it is closed here the way the attention side already works —
read becomes a defaulted parameter, construction.py threads it exactly as
line 183 does for attention, and fused_ffn_input gains the fuses_quant
branch. The FFN default is NORM_READOUT, an unrelated type to
NormQuantReadout, so fuses_quant 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.

Note for reviewers: construction.py and fusions/allreduce.py are common
code. The change is additive and behaviour-preserving for every other model,
but I'd welcome a reviewer from the community on those two files.

The kernel, the HIP IPC layer and the registered test were untouched by the
refactor — they are new files with no upstream counterpart.

Now runs on ROCm >= 7.2

is_supported() used to decline whenever _use_aiter_bpreshuffle_gfx95 was
set, which made the kernel inert on every current image. That was inherited
caution rather than a measured incompatibility: this path produces fp8
activations and never reads a weight, so preshuffled weights cannot affect
it. What does differ is the activation scale layout, and the fix is upstream's
own materialize_bpreshuffle_fp8_scale(). The kernel is now active on
ROCm 7.2.4.

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 needed).
Baseline is a clean worktree at the same upstream commit, with none of these
files present.

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(

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.

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

@yichiche thanks — both points were correct, and the second one changed what
this PR claims. Summary of what changed, then the detail.

In this push 60f28ab 9db4834


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

@rbrugaro-amd

Pre-review blocked #39140 — [AMD] Qwen3-Next: fused TP4 all-reduce + Gemma RMSNorm + per-group FP8 quant on gfx950.

Failed checks:

  1. 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.
  2. 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.
  3. 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.

This branch has not been deployed

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

Labels

amd quant LLM Quantization 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