From 81ae62f21b9a3b6e48a483c62542139cf34d78cf Mon Sep 17 00:00:00 2001 From: Rita Brugarolas Brufau Date: Fri, 11 Sep 2026 22:33:19 +0000 Subject: [PATCH 1/5] [AMD] Qwen3-Next: fused TP4 all-reduce + Gemma RMSNorm + per-group FP8 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 --- .../gluon_tp_ar_norm_quant.py | 347 +++++++++++ .../gluon_tp_ar_norm_quant_kernel.py | 541 ++++++++++++++++++ .../device_communicators/hip_ipc.py | 310 ++++++++++ .../sglang/srt/distributed/parallel_state.py | 46 +- python/sglang/srt/layers/communicator.py | 38 +- python/sglang/srt/layers/layernorm.py | 80 +++ python/sglang/srt/models/qwen3_next.py | 160 +++++- .../amd/test_gluon_tp_ar_norm_quant.py | 228 ++++++++ 8 files changed, 1742 insertions(+), 8 deletions(-) create mode 100644 python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py create mode 100644 python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py create mode 100644 python/sglang/srt/distributed/device_communicators/hip_ipc.py create mode 100644 test/registered/amd/test_gluon_tp_ar_norm_quant.py diff --git a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py new file mode 100644 index 000000000000..83787a993f47 --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py @@ -0,0 +1,347 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Fused TP all-reduce + residual add + Gemma RMSNorm + per-1x128 FP8 quant. + +Single-kernel replacement for the three-kernel decode sequence + + all_reduce(hidden) -> + residual -> Gemma RMSNorm -> per-1x128 FP8 quant + +on gfx950 at TP4. The kernel reads peers' staging buffers directly over HIP IPC +and performs its own arrival/completion handshake, so there is no separate +collective launch. + +Scope, enforced by :func:`is_supported`: + * gfx950 (CDNA4) only -- the kernel uses cdna4 ``buffer_load``/``buffer_store`` + and wave64 GCN assembly. + * TP world size exactly 4 -- the reduction is hand-unrolled over 4 peers and + the epoch protocol counts in units of 3 remote peers. + * hidden size exactly 2048, eps 1e-6, and M in {1,2,4,8,16,32,64}. + * ``_use_aiter_bpreshuffle_gfx95`` must be False. On ROCm >= 7.2 SGLang + physically preshuffles FP8 weights into the gfx95 bpreshuffle layout, which + this kernel's consumers do not expect. Returning False here makes the caller + fall back rather than produce wrong numbers. + +Anything outside that envelope returns False and the caller keeps the stock +path. Opt out entirely with ``SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT=1``. + +The kernel body lives in ``gluon_tp_ar_norm_quant_kernel.py``. It is tuned per +token count M; see ``SUPPORTED_M``. +""" + +from __future__ import annotations + +import logging +from typing import List, Optional, Tuple + +import torch +import torch.distributed as dist +from torch.distributed import ProcessGroup + +from sglang.srt.distributed.device_communicators.hip_ipc import ( + close_shared_tensor, + create_shared_tensor, + register_peer_pointers, +) +from sglang.srt.utils import get_bool_env_var, is_hip + +logger = logging.getLogger(__name__) + +_is_hip = is_hip() + +TP_SIZE = 4 +HIDDEN_SIZE = 2048 +EPS = 1.0e-6 +SUPPORTED_M = (1, 2, 4, 8, 16, 32, 64) +GROUP_SIZE = 128 +# Words 0/1 count arrivals / completed reads; word 2 is local CTA progress; +# word 3 is unused. They must start at zero and persist across calls. +SYNC_WORDS = 4 +# One lock row per CUDA-graph call site, reserved before capture so the pointer +# table baked into the captured launch stays valid. +CAPTURE_SITE_CAPACITY = 2048 + + +def _gluon_available() -> bool: + """Probe for Gluon + the CDNA4 intrinsics the kernel needs. + + Mirrors the probe idiom in ``aiter_mla_gluon.py``: SGLang does not pin a + Triton version that guarantees Gluon (the published ROCm wheel declares + triton 3.5.1), so this must be a runtime capability check, never an import + dependency. + """ + try: + import triton.experimental.gluon # noqa: F401 + from triton.experimental.gluon.language.amd.cdna4 import ( # noqa: F401 + buffer_load, + buffer_store, + ) + except Exception as exc: # pragma: no cover - depends on installed triton + logger.debug("Gluon TP AR+norm+quant unavailable: %s", exc) + return False + return True + + +def _arch_is_gfx950(device: torch.device) -> bool: + try: + name = torch.cuda.get_device_properties(device).gcnArchName + except Exception: # pragma: no cover + return False + return name.split(":", 1)[0] == "gfx950" + + +def is_available(world_size: int, device: torch.device) -> bool: + """Shape-independent half of :func:`is_supported`. + + Lets the CUDA-graph capture path decide whether to build the rendezvous + state *before* capture starts. The state must exist by then: it is created + lazily on first use, and a state created mid-capture would not have its + capture bookkeeping armed, silently falling back to the staging copy. + """ + if not _is_hip or world_size != TP_SIZE: + return False + if get_bool_env_var("SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT", default="false"): + return False + from sglang.srt.layers.quantization.fp8_utils import ( + _use_aiter_bpreshuffle_gfx95, + ) + + if _use_aiter_bpreshuffle_gfx95: + return False + return _arch_is_gfx950(device) and _gluon_available() + + +def is_supported( + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + weight: torch.Tensor, + eps: float, + world_size: int, + group_size: int = GROUP_SIZE, +) -> bool: + """Support predicate. False means "caller should use the stock path".""" + if not _is_hip or residual is None: + return False + if get_bool_env_var("SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT", default="false"): + return False + if world_size != TP_SIZE: + return False + if hidden_states.dim() != 2 or hidden_states.shape[-1] != HIDDEN_SIZE: + return False + if hidden_states.shape[0] not in SUPPORTED_M: + return False + if group_size != GROUP_SIZE: + return False + if hidden_states.dtype is not torch.bfloat16: + return False + if residual.shape != hidden_states.shape or residual.dtype is not torch.bfloat16: + return False + if weight.numel() != HIDDEN_SIZE: + return False + # The kernel bakes eps into the launch; only the Qwen3-Next value is tuned. + if abs(eps - EPS) > 1e-12: + return False + # On ROCm >= 7.2 SGLang preshuffles FP8 weights; this path is not validated + # against that layout and would silently produce wrong results. + from sglang.srt.layers.quantization.fp8_utils import ( + _use_aiter_bpreshuffle_gfx95, + ) + + if _use_aiter_bpreshuffle_gfx95: + return False + if not _arch_is_gfx950(hidden_states.device): + return False + return _gluon_available() + + +class GluonTpArNormQuantState: + """Rendezvous + peer pointer tables for the fused collective. + + Allocates, once per TP group: + * a staging buffer that peers read this rank's pre-reduction activations from + * one 4-word synchronization row for eager calls + * ``CAPTURE_SITE_CAPACITY`` further rows, one per CUDA-graph call site + + and publishes the peer pointer tables the kernel indexes. + + Two table orderings, both load-bearing: + * ``input_ptrs`` is in **global rank order** -- entry i is rank i's staging + buffer, and the kernel sums entries 0..3 in that order to preserve the + production reduction tree. + * ``lock_ptrs`` is rotated **local-rank-first** -- entry 0 is always this + rank's own lock row, entries 1..3 the three remote peers. The kernel + relies on this: it notifies lanes 1..3 (the peers that are not self) and + polls entry 0 for its own arrival counter. + """ + + def __init__(self, group: ProcessGroup, device: torch.device, max_rows: int): + self.group = group + self.device = device + self.rank = dist.get_rank(group=group) + self.world_size = dist.get_world_size(group=group) + if self.world_size != TP_SIZE: + raise ValueError(f"expected TP{TP_SIZE}, got {self.world_size}") + self.max_rows = max_rows + self._closed = False + lock_bytes = SYNC_WORDS * torch.int32.itemsize + + self.staging, staging_peers, b0 = create_shared_tensor( + (max_rows, HIDDEN_SIZE), torch.bfloat16, device, group=group + ) + self._lock, lock_peers, b1 = create_shared_tensor( + (SYNC_WORDS,), torch.int32, device, group=group + ) + self._capture_locks, capture_peers, b2 = create_shared_tensor( + (CAPTURE_SITE_CAPACITY, SYNC_WORDS), torch.int32, device, group=group + ) + self._opened_bases = b0 + b1 + b2 + + self.input_ptr_table = torch.tensor( + staging_peers, dtype=torch.uint64, device=device + ) + self.lock_ptr_table = torch.tensor( + self._rotate(lock_peers), dtype=torch.uint64, device=device + ) + # Row `site` of every peer's capture-lock allocation, rotated the same way. + self.capture_lock_tables = torch.tensor( + [ + self._rotate([base + site * lock_bytes for base in capture_peers]) + for site in range(CAPTURE_SITE_CAPACITY) + ], + dtype=torch.uint64, + device=device, + ) + # Per-call-site peer pointer rows for captured activations. Contents are + # written after capture; the captured launch only bakes in the address. + self.capture_input_tables = torch.zeros( + (CAPTURE_SITE_CAPACITY, TP_SIZE), dtype=torch.uint64, device=device + ) + self._next_site = 0 + self._capturing = False + self._pending = [] + self._captured_inputs = [] + torch.cuda.synchronize(device) + + def _rotate(self, ptrs: List[int]) -> Tuple[int, ...]: + """Reorder global-rank-ordered pointers to local-rank-first.""" + return tuple( + ptrs[(self.rank + offset) % self.world_size] + for offset in range(self.world_size) + ) + + def reserve_site(self) -> int: + """Claim a lock row for one CUDA-graph call site.""" + if self._next_site >= CAPTURE_SITE_CAPACITY: + raise RuntimeError("Gluon TP AR+norm+quant graph site capacity exceeded") + site = self._next_site + self._next_site += 1 + return site + + # ---- CUDA-graph capture ------------------------------------------------- + # + # Outside capture the caller must copy its activations into `staging`, which + # peers already have mapped. That copy is a separate kernel launch per call + # (~4.2 us, ~95x per decode step) and dominates the fused kernel's own cost. + # + # Inside capture we avoid it: the activation tensor at a given call site is + # the *same buffer* on every replay, so we can publish that buffer over IPC + # and have peers read it directly. The exchange is a collective and cannot + # run mid-capture, but it does not need to: the kernel reads its peer + # pointers out of a table *tensor*, and the captured launch bakes in that + # table's address, not its contents. So we record pointers during capture + # and fill the table after capture ends, before any replay. + + def begin_capture(self) -> None: + if self._capturing: + raise RuntimeError("nested Gluon TP AR+norm+quant capture") + self._capturing = True + self._pending = [] + + def abort_capture(self) -> None: + self._capturing = False + self._pending = [] + + def record_site(self, hidden_states: torch.Tensor) -> int: + """Claim a site for this captured call and remember its input buffer.""" + site = self.reserve_site() + self._pending.append((site, hidden_states)) + return site + + def finish_capture(self) -> None: + """Exchange the captured buffers' pointers and fill their table rows.""" + pending, self._pending = self._pending, [] + self._capturing = False + if not pending: + return + rows, opened = register_peer_pointers([t for _, t in pending], group=self.group) + self._opened_bases.extend(opened) + for (site, tensor), row in zip(pending, rows): + if row[self.rank] != tensor.data_ptr(): + raise RuntimeError( + "local IPC pointer does not alias the captured input" + ) + self.capture_input_tables[site].copy_( + torch.tensor(row, dtype=torch.uint64, device=self.device) + ) + # Peers must see the filled tables before the first replay. + torch.cuda.synchronize(self.device) + self._captured_inputs.extend(t for _, t in pending) + + @property + def capturing(self) -> bool: + return self._capturing + + def close(self) -> None: + if self._closed: + return + self._closed = True + close_shared_tensor(self._opened_bases) + + +def fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( + state: GluonTpArNormQuantState, + hidden_states: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + site: Optional[int] = None, +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Run the fused collective. + + Returns ``(fp8_out, scale_out, residual_out, bf16_out)``. ``bf16_out`` is the + pre-quantization normed activation, which GDN-style layers need for their + bf16 gating projection (``keep_bf16=True`` in the caller); it is produced by + the same kernel, so obtaining it costs nothing extra. + """ + from sglang.srt.distributed.device_communicators.gluon_tp_ar_norm_quant_kernel import ( + tp4_allreduce_add_gemma_rmsnorm_group_fp8_quant_gluon, + ) + + rows = hidden_states.shape[0] + # `state.capturing` only says we are inside the graph_capture *context*. + # That context also brackets real eager warmup executions (the CUDA graph + # runner calls forward_fn() twice before torch.cuda.graph), and those must + # use the staging copy: the per-site peer table is not filled until capture + # finishes, so taking the no-copy path there dereferences a zeroed table. + # Gate on the actual stream-capture state, the same way aiter's + # custom_fused_ar_rms_* entry points do. + if state.capturing and torch.cuda.is_current_stream_capturing(): + # Publish this call site's own activation buffer instead of copying into + # staging; peers will read it directly on every replay. + site = state.record_site(hidden_states) + local_input = hidden_states + input_table = state.capture_input_tables[site] + lock_table = state.capture_lock_tables[site] + else: + local_input = state.staging[:rows] + local_input.copy_(hidden_states) + input_table = state.input_ptr_table + lock_table = state.lock_ptr_table + normalized, residual_out, quantized, scales, _reduced = ( + tp4_allreduce_add_gemma_rmsnorm_group_fp8_quant_gluon( + local_input, + input_table, + lock_table, + residual, + weight, + eps=EPS, + ) + ) + return quantized, scales, residual_out, normalized diff --git a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py new file mode 100644 index 000000000000..dc469597ad10 --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py @@ -0,0 +1,541 @@ +"""TP4 all-reduce, Gemma RMSNorm, and BF16-to-group-FP8 conversion. + +M2 overlaps scalar pointer-table reads and keeps its acquisition address scalar. +M16/M32 use scalar completion polling after the existing final-CTA election. +All shapes preserve the production arithmetic and caller-owned epoch protocol. +""" + +import torch +import triton.experimental.gluon as gluon +import triton.experimental.gluon.language as gl +from triton.experimental.gluon.language.amd.cdna4 import buffer_load, buffer_store + + +@gluon.jit +def _notify(locks, increment: gl.constexpr, active, publish: gl.constexpr): + """One wave pushes integer notifications to the three peer-owned counters.""" + if publish: + gl.inline_asm_elementwise( + """ + v_cmp_ne_u32 vcc, 0, $3 + s_and_saveexec_b64 $0, vcc + s_cbranch_execz 1f + s_waitcnt vmcnt(0) lgkmcnt(0) + global_atomic_add $1, $2, off sc1 + 1: + s_or_saveexec_b64 $0, $0 + """, + "=&s,v,v,v,~{vcc},~{scc},~{memory}", + [locks, increment, active.to(gl.int32)], + dtype=gl.uint64, + is_pure=False, + pack=1, + ) + else: + gl.inline_asm_elementwise( + """ + v_cmp_ne_u32 vcc, 0, $3 + s_and_saveexec_b64 $0, vcc + global_atomic_add $1, $2, off sc1 + s_or_saveexec_b64 $0, $0 + """, + "=&s,v,v,v,~{vcc},~{scc},~{memory}", + [locks, increment, active.to(gl.int32)], + dtype=gl.uint64, + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _scalar_acquire(counter, target, thread): + """Only wave zero polls; the caller subsequently joins the whole CTA.""" + address = counter.to(gl.uint64) + gl.inline_asm_elementwise( + """ + v_cmp_eq_u32 vcc, 0, $5 + s_cbranch_vccz 2f + v_readfirstlane_b32 s0, $2 + v_readfirstlane_b32 s1, $3 + v_readfirstlane_b32 $1, $4 + 1: + s_load_dword $0, s[0:1], 0x0 glc + s_waitcnt lgkmcnt(0) + s_sub_u32 $0, $0, $1 + s_cmp_ge_i32 $0, 0 + s_cbranch_scc0 1b + 2: + """, + "=&s,=&s,v,v,v,v,~{s0},~{s1},~{scc},~{vcc},~{memory}", + [address.to(gl.uint32), (address >> 32).to(gl.uint32), target, thread], + dtype=(gl.uint32, gl.uint32), + is_pure=False, + pack=1, + ) + + +# M2 already has a scalar local pointer; retain it through arrival acquisition. +@gluon.jit +def _scalar_acquire_uniform(counter, target, thread): + gl.inline_asm_elementwise( + """ + v_cmp_eq_u32 vcc, 0, $4 + s_cbranch_vccz 2f + v_readfirstlane_b32 $1, $3 + 1: + s_load_dword $0, $2, 0x0 glc + s_waitcnt lgkmcnt(0) + s_sub_u32 $0, $0, $1 + s_cmp_ge_i32 $0, 0 + s_cbranch_scc0 1b + 2: + """, + "=&s,=&s,s,v,v,~{scc},~{vcc},~{memory}", + [counter.to(gl.uint64), target, thread], + dtype=(gl.uint32, gl.uint32), + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _vector_acquire(counter, target, thread, materialize: gl.constexpr): + condition = (thread == 0).to(gl.uint32) if materialize else thread + compare: gl.constexpr = ( + "v_cmp_ne_u32 vcc, 0, $4" if materialize else "v_cmp_eq_u32 vcc, 0, $4" + ) + gl.inline_asm_elementwise( + compare + + """ + s_and_saveexec_b64 $0, vcc + s_cbranch_execz 2f + 1: + global_load_dword $1, $2, off sc1 + s_waitcnt vmcnt(0) + v_sub_u32 $1, $1, $3 + v_cmp_ge_i32 vcc, $1, 0 + s_cbranch_vccz 1b + 2: + s_or_saveexec_b64 $0, $0 + """, + "=&s,=&v,v,v,v,~{vcc},~{scc},~{memory}", + [counter, target, condition], + dtype=(gl.uint64, gl.uint32), + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _reserve_epoch(local, thread, step: gl.constexpr): + """Reserve a CTA ticket. Only lane zero consumes its returned epoch.""" + _, ticket = gl.inline_asm_elementwise( + """ + v_mov_b32 $1, 0 + v_cmp_eq_u32 vcc, 0, $4 + s_and_saveexec_b64 $0, vcc + s_cbranch_execz 1f + global_atomic_add $1, $2, $3, off offset:8 sc0 sc1 + s_waitcnt vmcnt(0) + 1: + s_or_saveexec_b64 $0, $0 + """, + "=&s,=&v,v,v,v,~{vcc},~{scc},~{memory}", + [local, step, thread], + dtype=(gl.uint64, gl.uint32), + is_pure=False, + pack=1, + ) + return ticket + + +@gluon.jit +def _advance_single_row(local, next_progress, thread): + """A one-CTA invocation needs no ticket-election atomic.""" + gl.inline_asm_elementwise( + """ + v_cmp_eq_u32 vcc, 0, $3 + s_and_saveexec_b64 $0, vcc + global_store_dword $1, $2, off offset:8 + s_or_saveexec_b64 $0, $0 + """, + "=&s,v,v,v,~{vcc},~{scc},~{memory}", + [local, next_progress, thread], + dtype=gl.uint64, + is_pure=False, + pack=1, + ) + return (thread == 0).to(gl.uint32) + + +@gluon.jit +def _finish(local, target, last): + gl.inline_asm_elementwise( + """ + v_cmp_ne_u32 vcc, 0, $4 + s_and_saveexec_b64 $0, vcc + s_cbranch_execz 2f + 1: + global_load_dword $1, $2, off offset:4 sc1 + s_waitcnt vmcnt(0) + v_sub_u32 $1, $1, $3 + v_cmp_ge_i32 vcc, $1, 0 + s_cbranch_vccz 1b + 2: + s_or_saveexec_b64 $0, $0 + """, + "=&s,=&v,v,v,v,~{vcc},~{scc},~{memory}", + [local, target, last], + dtype=(gl.uint64, gl.uint32), + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _finish_late(local, target, thread, step: gl.constexpr): + gl.inline_asm_elementwise( + """ + v_cmp_eq_u32 vcc, 0, $6 + s_and_saveexec_b64 $0, vcc + s_cbranch_execz 3f + global_atomic_add $1, $3, $4, off offset:8 sc0 + s_waitcnt vmcnt(0) + v_add_u32 $1, $4, $1 + v_mul_lo_u32 $1, 3, $1 + v_cmp_eq_u32 vcc, $1, $5 + s_cbranch_vccz 3f + 2: + global_load_dword $2, $3, off offset:4 sc1 + s_waitcnt vmcnt(0) + v_sub_u32 $2, $2, $5 + v_cmp_ge_i32 vcc, $2, 0 + s_cbranch_vccz 2b + 3: + s_or_saveexec_b64 $0, $0 + """, + "=&s,=&v,=&v,v,v,v,v,~{vcc},~{scc},~{memory}", + [local, step, target, thread], + dtype=(gl.uint64, gl.uint32, gl.uint32), + is_pure=False, + pack=1, + ) + + +# Only the final CTA polls. Earlier CTAs remain free to retire on a single CU. +@gluon.jit +def _finish_uniform_late(local, target, thread, step: gl.constexpr): + gl.inline_asm_elementwise( + """ + v_cmp_eq_u32 vcc, 0, $6 + s_and_saveexec_b64 $0, vcc + s_cbranch_execz 3f + global_atomic_add $1, $3, $4, off offset:8 sc0 + s_waitcnt vmcnt(0) + v_add_u32 $1, $4, $1 + v_mul_lo_u32 $1, 3, $1 + v_cmp_eq_u32 vcc, $1, $5 + s_cbranch_vccz 3f + v_readfirstlane_b32 s0, $5 + 2: + s_load_dword $2, $7, 0x4 glc + s_waitcnt lgkmcnt(0) + s_sub_u32 $2, $2, s0 + s_cmp_ge_i32 $2, 0 + s_cbranch_scc0 2b + 3: + s_or_saveexec_b64 $0, $0 + """, + "=&s,=&v,=&s,v,v,v,v,s,~{s0},~{vcc},~{scc},~{memory}", + [local, step, target, thread, local.to(gl.uint64)], + dtype=(gl.uint64, gl.uint32, gl.uint32), + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _load_pointer(table): + return gl.inline_asm_elementwise( + "s_load_dwordx2 $0, $1, 0x0\n s_waitcnt lgkmcnt(0)", + "=s,s,~{memory}", + [table.to(gl.uint64)], + dtype=gl.uint64, + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _load_inputs(table, owner): + return gl.inline_asm_elementwise( + """ + s_load_dwordx2 $0, $4, $5 + s_load_dwordx2 $1, $4, $6 + s_load_dwordx2 $2, $4, $7 + s_load_dwordx2 $3, $4, $8 + s_waitcnt lgkmcnt(0) + """, + "=&s,=&s,=&s,=&s,s,s,s,s,s,~{memory}", + [ + table.to(gl.uint64), + owner * 8, + ((owner + 1) % 4) * 8, + ((owner + 2) % 4) * 8, + ((owner + 3) % 4) * 8, + ], + dtype=(gl.uint64, gl.uint64, gl.uint64, gl.uint64), + is_pure=False, + pack=1, + ) + + +# Issue the five independent M2 pointer reads before a single scalar-memory wait. +@gluon.jit +def _load_tables(lock_table, input_table): + return gl.inline_asm_elementwise( + """ + s_load_dwordx2 $0, $5, 0x0 + s_load_dwordx2 $1, $6, 0x0 + s_load_dwordx2 $2, $6, 0x8 + s_load_dwordx2 $3, $6, 0x10 + s_load_dwordx2 $4, $6, 0x18 + s_waitcnt lgkmcnt(0) + """, + "=&s,=&s,=&s,=&s,=&s,s,s,~{memory}", + [lock_table.to(gl.uint64), input_table.to(gl.uint64)], + dtype=(gl.uint64, gl.uint64, gl.uint64, gl.uint64, gl.uint64), + is_pure=False, + pack=1, + ) + + +@gluon.jit +def _store_local(pointer, offsets, values, use_buffer: gl.constexpr): + if use_buffer: + buffer_store(values.to(pointer.dtype.element_ty), pointer, offsets) + else: + gl.store(pointer + offsets, values) + + +@gluon.jit +def _fused( + input_peer_ptrs, + lock_peer_ptrs, + residual_ptr, + weight_ptr, + normalized_ptr, + residual_out_ptr, + quantized_ptr, + scales_ptr, + reduced_ptr, + M: gl.constexpr, + ROUND: gl.constexpr, + FP8_MAX: gl.constexpr, + EPS: gl.constexpr, +): + PAIR: gl.constexpr = M == 64 + NCTA: gl.constexpr = M // 2 if PAIR else M + TICKET: gl.constexpr = M >= 2 and M <= 8 + NW: gl.constexpr = 16 if PAIR else 8 + pid = gl.program_id(0) + if TICKET: + thread = gl.arange(0, NW * 64, layout=gl.BlockedLayout([1], [64], [NW], [0])) + peer = thread % 4 + locks = gl.load(lock_peer_ptrs + peer).to(gl.pointer_type(gl.uint32)) + if pid == 0: + _notify(locks, 128, (thread > 0) & (thread < 4), M <= 4) + if M == 2: + local_address, a0, a1, a2, a3 = _load_tables( + lock_peer_ptrs, input_peer_ptrs + ) + local = local_address.to(gl.pointer_type(gl.uint32)) + else: + local = _load_pointer(lock_peer_ptrs).to(gl.pointer_type(gl.uint32)) + # Returned tickets stay valid even when another CTA advances word 2. + # Only lane zero uses target/last; arrival acquisition joins the CTA. + ticket = _reserve_epoch(local, thread, 128 // NCTA) + target = ((ticket & 0xFFFFFF80) + 128) * 3 + last = ((thread == 0) & ((ticket & 127) == 128 - 128 // NCTA)).to(gl.uint32) + elif M == 1: + local = _load_pointer(lock_peer_ptrs).to(gl.pointer_type(gl.uint32)) + # Local progress advances by 128 per call, regardless of grid size. + # The epoch cannot change until every CTA has sampled it. + progress = gl.inline_asm_elementwise( + "global_load_dword $0, $1, off offset:8 sc1\n s_waitcnt vmcnt(0)", + "=v,v,~{memory}", + [local], + dtype=gl.uint32, + is_pure=False, + pack=1, + ) + next_progress = (progress & 0xFFFFFF80) + 128 + target = next_progress * 3 + thread = gl.arange(0, NW * 64, layout=gl.BlockedLayout([1], [64], [NW], [0])) + peer = thread % 4 + locks = gl.load(lock_peer_ptrs + peer).to(gl.pointer_type(gl.uint32)) + if pid == 0: + _notify(locks, 128, (thread > 0) & (thread < 4), True) + else: + thread = gl.arange(0, NW * 64, layout=gl.BlockedLayout([1], [64], [NW], [0])) + peer = thread % 4 + locks = gl.load(lock_peer_ptrs + peer).to(gl.pointer_type(gl.uint32)) + 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( + "global_load_dword $0, $1, off offset:8 sc1\n s_waitcnt vmcnt(0)", + "=v,v,~{memory}", + [local], + dtype=gl.uint32, + is_pure=False, + pack=1, + ) + next_progress = (progress & 0xFFFFFF80) + 128 + target = next_progress * 3 + + # Four adjacent elements per lane preserve the production RMS sum tree. + L: gl.constexpr = gl.BlockedLayout( + [1, 1, 4], [1, 2, 32], [2 if PAIR else 1, 8, 1], [2, 1, 0] + ) + rows = pid * (2 if PAIR else 1) + gl.arange( + 0, 2 if PAIR else 1, layout=gl.SliceLayout(1, gl.SliceLayout(2, L)) + ) + groups = gl.arange(0, 16, layout=gl.SliceLayout(0, gl.SliceLayout(2, L))) + within = gl.arange(0, 128, layout=gl.SliceLayout(0, gl.SliceLayout(1, L))) + cols = groups[None, :, None] * 128 + within[None, None, :] + offsets = rows[:, None, None] * 2048 + cols + if M == 2 or M == 16: + residual = buffer_load(residual_ptr, offsets).to(gl.float32) + weight = buffer_load(weight_ptr, cols).to(gl.float32) + else: + residual = gl.load(residual_ptr + offsets).to(gl.float32) + weight = gl.load(weight_ptr + cols).to(gl.float32) + owner = pid // 8 if PAIR else 0 + if M != 2: + a0, a1, a2, a3 = _load_inputs(input_peer_ptrs, owner) + p0 = a0.to(gl.pointer_type(gl.bfloat16)) + p1 = a1.to(gl.pointer_type(gl.bfloat16)) + p2 = a2.to(gl.pointer_type(gl.bfloat16)) + p3 = a3.to(gl.pointer_type(gl.bfloat16)) + if M == 2: + _scalar_acquire_uniform(local, target, thread) + elif M <= 4: + _scalar_acquire(local, target, thread) + else: + _vector_acquire(local, target, thread, M == 8) + gl.barrier() + x0 = gl.load(p0 + offsets, cache_modifier=".ca").to(gl.float32) + x1 = gl.load(p1 + offsets, cache_modifier=".ca").to(gl.float32) + x2 = gl.load(p2 + offsets, cache_modifier=".ca").to(gl.float32) + x3 = gl.load(p3 + offsets, cache_modifier=".ca").to(gl.float32) + reduced = (((x0 + x1) + x2) + x3).to(gl.bfloat16, fp_downcast_rounding=ROUND) + if M <= 2: + reduced = gl.inline_asm_elementwise( + "", + "=v,0,~{memory}", + [reduced], + dtype=gl.bfloat16, + is_pure=False, + pack=2, + ) + else: + _store_local(reduced_ptr, offsets, reduced, M <= 2 or M == 16) + reduced = gl.inline_asm_elementwise( + "", + "=v,0,~{memory}", + [reduced], + dtype=gl.bfloat16, + is_pure=False, + pack=2, + ) + # Publish only after every wave has consumed all four peer payloads. + # The final ticket holder waits; other CTAs can retire on a single-CU stream. + gl.barrier() + if M == 1: + last = _advance_single_row(local, next_progress, thread) + _notify(locks + 1, 128 // NCTA, (thread > 0) & (thread < 4), False) + if M <= 2: + _store_local(reduced_ptr, offsets, reduced, M <= 2 or M == 16) + + value = reduced.to(gl.float32) + residual + residual_out = value.to(gl.bfloat16, fp_downcast_rounding=ROUND) + variance = gl.sum(gl.sum(value * value, 2), 1) / 2048.0 + normalized = (value * gl.rsqrt(variance[:, None, None] + EPS)) * weight + normalized = normalized.to(gl.bfloat16, fp_downcast_rounding=ROUND) + _store_local(normalized_ptr, offsets, normalized, M <= 2 or M == 16) + _store_local(residual_out_ptr, offsets, residual_out, M <= 2 or M == 16) + values = normalized.to(gl.float32) + maximum = gl.maximum(gl.max(gl.abs(values), 2), 1.0e-10) + scales = maximum * (1.0 / FP8_MAX) + if M == 4: + _store_local(scales_ptr, rows[:, None] * 16 + groups[None, :], scales, False) + quantized = gl.clamp(values * (1.0 / scales[:, :, None]), -FP8_MAX, FP8_MAX) + _store_local(quantized_ptr, offsets, quantized, M <= 2 or M == 16) + if M != 4: + _store_local( + scales_ptr, + rows[:, None] * 16 + groups[None, :], + scales, + M <= 2 or M == 16, + ) + if M <= 8: + _finish(local, target, last) + elif M == 16 or M == 32: + _finish_uniform_late(local, target, thread, 128 // NCTA) + else: + _finish_late(local, target, thread, 128 // NCTA) + + +def tp4_allreduce_add_gemma_rmsnorm_group_fp8_quant_gluon( + local_input: torch.Tensor, + input_peer_ptrs: torch.Tensor, + lock_peer_ptrs: torch.Tensor, + residual: torch.Tensor, + gemma_weight: torch.Tensor, + *, + eps: float = 1.0e-6, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: + """Return normalized, residual, FP8, scales, and the all-reduce witness. + + Caller-owned synchronization words start at zero and persist across ordered + invocations. Words 0/1 count arrivals/completed reads in 384-unit epochs; + word 2 tracks local CTA progress in 128-unit epochs. M2--M8 atomically + reserve tickets to acquire the epoch and elect the final CTA together. + M1 needs only one progress store; M16--M64 retain late completion election. + Signed differences handle 32-bit rollover. Word 3 is unused. No tensors + or host state are cached. + """ + assert eps == 1.0e-6 + m, width = local_input.shape + assert width == 2048 and m in (1, 2, 4, 8, 16, 32, 64) + arch = torch.cuda.get_device_properties(local_input.device).gcnArchName.split( + ":", 1 + )[0] + assert arch == "gfx950" + rounding = "rtne" + fp8_dtype = torch.float8_e4m3fn + fp8_max = 448.0 + normalized = torch.empty_like(local_input) + residual_out = torch.empty_like(local_input) + quantized = torch.empty_like(local_input, dtype=fp8_dtype) + scales = torch.empty((m, 16), dtype=torch.float32, device=local_input.device) + reduced = torch.empty_like(local_input) + _fused[(m // 2 if m == 64 else m,)]( + input_peer_ptrs, + lock_peer_ptrs, + residual, + gemma_weight, + normalized, + residual_out, + quantized, + scales, + reduced, + m, + rounding, + fp8_max, + eps, + num_warps=16 if m == 64 else 8, + enable_fp_fusion=False, + ) + return normalized, residual_out, quantized, scales, reduced diff --git a/python/sglang/srt/distributed/device_communicators/hip_ipc.py b/python/sglang/srt/distributed/device_communicators/hip_ipc.py new file mode 100644 index 000000000000..6e45b8fa3579 --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/hip_ipc.py @@ -0,0 +1,310 @@ +# SPDX-License-Identifier: Apache-2.0 +"""HIP IPC shared buffers: the ROCm counterpart of +``CustomAllreduce.create_shared_buffer``. + +``CustomAllreduce.create_shared_buffer`` hands Python a list of peer device +pointers, which is what a Triton/Gluon collective needs in order to address +other ranks' memory directly. It is CUDA-only: it goes through +``CudaRTLibrary``, which loads ``libcudart``. + +On ROCm the existing collectives keep their peer pointers inside C++ +(``quick_all_reduce`` exchanges handles via ``ops.qr_{get,open}_handles``; the +HIP branch of ``custom_all_reduce`` opens them inside ``init_custom_ar``), so +there is no in-tree way for a Python-level kernel to obtain them. This module +fills that gap with the same API shape as the CUDA path. + +Pointers are raw ``hipMalloc`` allocations, not torch tensors, so they are +stable for the process lifetime and are never moved by the caching allocator — +a requirement for CUDA-graph capture, where the peer pointer table is baked +into the captured launch. +""" + +from __future__ import annotations + +import ctypes +import logging +from typing import List, Optional + +import torch.distributed as dist +from torch.distributed import ProcessGroup + +logger = logging.getLogger(__name__) + +# sizeof(hipIpcMemHandle_t); matches CUDA's cudaIpcMemHandle_t. +_HANDLE_BYTES = 64 +# hipIpcMemLazyEnablePeerAccess +_LAZY_ENABLE_PEER_ACCESS = 1 + + +class HipIpcMemHandle(ctypes.Structure): + # c_ubyte, not c_char: a c_char array is treated as a NUL-terminated string + # by ctypes, so both reading and assigning would truncate the 64-byte handle + # at its first zero byte. + _fields_ = [("reserved", ctypes.c_ubyte * _HANDLE_BYTES)] + + +class HipRTLibrary: + """Minimal ctypes binding for the HIP runtime calls we need.""" + + _instance: Optional[HipRTLibrary] = None + + def __new__(cls) -> HipRTLibrary: + if cls._instance is None: + cls._instance = super().__new__(cls) + cls._instance._init_lib() + return cls._instance + + def _init_lib(self) -> None: + self.lib = ctypes.CDLL("libamdhip64.so") + self.lib.hipMalloc.argtypes = [ + ctypes.POINTER(ctypes.c_void_p), + ctypes.c_size_t, + ] + self.lib.hipMalloc.restype = ctypes.c_int + self.lib.hipFree.argtypes = [ctypes.c_void_p] + self.lib.hipFree.restype = ctypes.c_int + self.lib.hipMemset.argtypes = [ + ctypes.c_void_p, + ctypes.c_int, + ctypes.c_size_t, + ] + self.lib.hipMemset.restype = ctypes.c_int + self.lib.hipIpcGetMemHandle.argtypes = [ + ctypes.POINTER(HipIpcMemHandle), + ctypes.c_void_p, + ] + self.lib.hipIpcGetMemHandle.restype = ctypes.c_int + self.lib.hipIpcOpenMemHandle.argtypes = [ + ctypes.POINTER(ctypes.c_void_p), + HipIpcMemHandle, + ctypes.c_uint, + ] + self.lib.hipIpcOpenMemHandle.restype = ctypes.c_int + self.lib.hipIpcCloseMemHandle.argtypes = [ctypes.c_void_p] + self.lib.hipIpcCloseMemHandle.restype = ctypes.c_int + self.lib.hipMemGetAddressRange.argtypes = [ + ctypes.POINTER(ctypes.c_void_p), + ctypes.POINTER(ctypes.c_size_t), + ctypes.c_void_p, + ] + self.lib.hipMemGetAddressRange.restype = ctypes.c_int + + @staticmethod + def _check(status: int, op: str) -> None: + if status != 0: + raise RuntimeError(f"{op} failed with HIP status {status}") + + def malloc(self, size_in_bytes: int) -> ctypes.c_void_p: + ptr = ctypes.c_void_p() + self._check(self.lib.hipMalloc(ctypes.byref(ptr), size_in_bytes), "hipMalloc") + return ptr + + def free(self, ptr: int) -> None: + self._check(self.lib.hipFree(ctypes.c_void_p(ptr)), "hipFree") + + def memset(self, ptr: ctypes.c_void_p, value: int, size_in_bytes: int) -> None: + self._check(self.lib.hipMemset(ptr, value, size_in_bytes), "hipMemset") + + def get_ipc_handle(self, ptr: ctypes.c_void_p) -> bytes: + handle = HipIpcMemHandle() + self._check( + self.lib.hipIpcGetMemHandle(ctypes.byref(handle), ptr), + "hipIpcGetMemHandle", + ) + return ctypes.string_at(ctypes.addressof(handle), _HANDLE_BYTES) + + def open_ipc_handle(self, raw: bytes) -> int: + if len(raw) != _HANDLE_BYTES: + raise ValueError( + f"IPC handle must be {_HANDLE_BYTES} bytes, got {len(raw)}" + ) + handle = HipIpcMemHandle() + ctypes.memmove(ctypes.addressof(handle), raw, _HANDLE_BYTES) + opened = ctypes.c_void_p() + self._check( + self.lib.hipIpcOpenMemHandle( + ctypes.byref(opened), handle, _LAZY_ENABLE_PEER_ACCESS + ), + "hipIpcOpenMemHandle", + ) + if not opened.value: + raise RuntimeError("hipIpcOpenMemHandle returned NULL") + return opened.value + + def close_ipc_handle(self, ptr: int) -> None: + self._check( + self.lib.hipIpcCloseMemHandle(ctypes.c_void_p(ptr)), + "hipIpcCloseMemHandle", + ) + + +def create_shared_buffer( + size_in_bytes: int, + group: Optional[ProcessGroup] = None, + zero_fill: bool = False, +) -> List[int]: + """Allocate ``size_in_bytes`` on every rank and return all peer pointers. + + The returned list is in **global rank order**: entry ``i`` addresses rank + ``i``'s allocation, and entry ``dist.get_rank(group)`` is this rank's own. + Mirrors ``CustomAllreduce.create_shared_buffer``. + + ``zero_fill`` is required for synchronization words, whose protocols assume + counters start at zero. + """ + lib = HipRTLibrary() + pointer = lib.malloc(size_in_bytes) + if zero_fill: + lib.memset(pointer, 0, size_in_bytes) + + handle = lib.get_ipc_handle(pointer) + world_size = dist.get_world_size(group=group) + rank = dist.get_rank(group=group) + + handles: List[Optional[bytes]] = [None] * world_size + dist.all_gather_object(handles, handle, group=group) + + pointers: List[int] = [] + for i, h in enumerate(handles): + pointers.append(pointer.value if i == rank else lib.open_ipc_handle(h)) + return pointers + + +def create_shared_tensor( + shape, + dtype, + device, + group: Optional[ProcessGroup] = None, + zero_fill: bool = True, +): + """Allocate a torch tensor and return ``(tensor, peer_pointers)``. + + Unlike :func:`create_shared_buffer`, the local allocation is an ordinary + torch tensor, so callers can ``copy_()`` into it and pass it to kernels. + Because torch's caching allocator hands out offsets inside larger segments, + the IPC handle is taken on the *segment base* (via ``hipMemGetAddressRange``) + and the intra-segment offset is exchanged alongside it; peers reopen the base + and re-apply the offset. + + ``peer_pointers`` is in global rank order; entry ``rank`` aliases + ``tensor.data_ptr()``. + """ + import torch + + lib = HipRTLibrary() + # Always zero-initialised: the synchronisation rows require counters to + # start at zero, and a zeroed staging buffer is harmless. + tensor = torch.zeros(shape, dtype=dtype, device=device) + + base = ctypes.c_void_p() + size = ctypes.c_size_t() + data_ptr = tensor.data_ptr() + lib._check( + lib.lib.hipMemGetAddressRange( + ctypes.byref(base), ctypes.byref(size), ctypes.c_void_p(data_ptr) + ), + "hipMemGetAddressRange", + ) + offset = data_ptr - base.value + handle = lib.get_ipc_handle(base) + + world_size = dist.get_world_size(group=group) + rank = dist.get_rank(group=group) + payload = [None] * world_size + dist.all_gather_object(payload, (handle, offset), group=group) + + pointers: List[int] = [] + opened_bases: List[int] = [] + for i, (h, off) in enumerate(payload): + if i == rank: + pointers.append(data_ptr) + else: + peer_base = lib.open_ipc_handle(h) + opened_bases.append(peer_base) + pointers.append(peer_base + off) + if pointers[rank] != data_ptr: + raise RuntimeError("local IPC pointer does not alias the source tensor") + # `opened_bases` must be handed to close_shared_tensor(); closing + # base+offset is invalid, and the local side is owned by torch. + return tensor, pointers, opened_bases + + +def register_peer_pointers(tensors, group: Optional[ProcessGroup] = None): + """Publish existing tensors over IPC and return their peer pointers. + + Unlike :func:`create_shared_tensor`, the tensors already exist and are owned + by someone else (here: activations captured inside a CUDA graph). One + batched ``all_gather_object`` covers the whole list, so this costs a single + collective no matter how many call sites were captured. + + Returns ``(rows, opened_bases)`` where ``rows[i]`` is the global-rank-ordered + peer pointer list for ``tensors[i]``. + + Every rank must pass the same number of tensors in the same order. + """ + lib = HipRTLibrary() + local = [] + for t in tensors: + base = ctypes.c_void_p() + size = ctypes.c_size_t() + data_ptr = t.data_ptr() + lib._check( + lib.lib.hipMemGetAddressRange( + ctypes.byref(base), ctypes.byref(size), ctypes.c_void_p(data_ptr) + ), + "hipMemGetAddressRange", + ) + local.append((lib.get_ipc_handle(base), data_ptr - base.value)) + + world_size = dist.get_world_size(group=group) + rank = dist.get_rank(group=group) + gathered = [None] * world_size + dist.all_gather_object(gathered, local, group=group) + if any(len(g) != len(local) for g in gathered): + raise RuntimeError("ranks captured different numbers of collective sites") + + rows, opened_bases = [], [] + # Reopening the same segment handle repeatedly is wasteful and can exhaust + # the mapping table, so cache per (rank, handle). + cache = {} + for i in range(len(local)): + row = [] + for r in range(world_size): + handle, offset = gathered[r][i] + if r == rank: + row.append(tensors[i].data_ptr()) + continue + key = (r, handle) + if key not in cache: + peer_base = lib.open_ipc_handle(handle) + cache[key] = peer_base + opened_bases.append(peer_base) + row.append(cache[key] + offset) + rows.append(row) + return rows, opened_bases + + +def close_shared_tensor(opened_bases: List[int]) -> None: + """Close peer mappings created by :func:`create_shared_tensor`.""" + lib = HipRTLibrary() + for base in opened_bases: + try: + lib.close_ipc_handle(base) + except RuntimeError as exc: # teardown must not mask the real error + logger.debug("hipIpcCloseMemHandle failed during teardown: %s", exc) + + +def free_shared_buffer( + pointers: List[int], group: Optional[ProcessGroup] = None +) -> None: + """Close peer mappings and free this rank's own allocation.""" + lib = HipRTLibrary() + rank = dist.get_rank(group=group) + for i, ptr in enumerate(pointers): + if ptr and i != rank: + try: + lib.close_ipc_handle(ptr) + except RuntimeError as exc: # teardown must not mask the real error + logger.debug("hipIpcCloseMemHandle failed during teardown: %s", exc) + if pointers and pointers[rank]: + lib.free(pointers[rank]) diff --git a/python/sglang/srt/distributed/parallel_state.py b/python/sglang/srt/distributed/parallel_state.py index e8b061d73756..18140a817c17 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -650,7 +650,51 @@ def graph_capture( else: maybe_pymscclpp_context = pymscclpp_comm.change_state(enable=True) with maybe_pynccl_context, maybe_pymscclpp_context: - yield graph_capture_context + # The Gluon TP AR+norm collective publishes each captured call + # site's activation buffer over IPC rather than copying into a + # staging buffer. That pointer exchange is collective and cannot + # run mid-capture, but it does not need to: the kernel reads its + # peers out of a table *tensor*, and the captured launch bakes in + # that table's address, not its contents. So record sites during + # capture and fill the tables here, before any replay. + gluon_ar = getattr(self, "_gluon_tp_ar_state", None) + if gluon_ar is None: + # Build it now rather than lazily on first use: first use may + # itself be inside this capture, and a state created then + # would not have its capture bookkeeping armed, silently + # falling back to the staging copy for every captured site. + try: + from sglang.srt.distributed.device_communicators import ( + gluon_tp_ar_norm_quant as _g, + ) + + dev = torch.device("cuda", torch.cuda.current_device()) + if _g.is_available(self.world_size, dev): + gluon_ar = _g.GluonTpArNormQuantState( + group=self.device_group, + device=dev, + max_rows=max(_g.SUPPORTED_M), + ) + self._gluon_tp_ar_state = gluon_ar + except Exception as exc: + logger.warning( + "Gluon TP AR+norm pre-capture setup failed, disabling: %s", + exc, + ) + self._gluon_tp_ar_state = False + gluon_ar = None + if not gluon_ar: # None, or False when permanently disabled + gluon_ar = None + if gluon_ar is not None: + gluon_ar.begin_capture() + try: + yield graph_capture_context + except BaseException: + if gluon_ar is not None: + gluon_ar.abort_capture() + raise + if gluon_ar is not None: + gluon_ar.finish_capture() def all_reduce(self, input_: torch.Tensor) -> torch.Tensor: """ diff --git a/python/sglang/srt/layers/communicator.py b/python/sglang/srt/layers/communicator.py index a60b70f22d9d..e377c328b305 100644 --- a/python/sglang/srt/layers/communicator.py +++ b/python/sglang/srt/layers/communicator.py @@ -529,6 +529,7 @@ def __init__( force_layernorm_before_dp_gather: bool = False, enable_fused_ar_quant: bool = False, fused_ar_quant_keep_bf16: bool = False, + enable_fused_ar_quant_mlp: bool = False, _is_sp_variant: bool = False, ): self.layer_scatter_modes = layer_scatter_modes @@ -545,6 +546,11 @@ def __init__( self._context.force_layernorm_before_dp_gather = ( force_layernorm_before_dp_gather ) + # Separate opt-in from the attention-side one: turning this on makes + # prepare_mlp hand the MLP a (bf16, fp8, scale) triple, which the + # consuming block must be prepared for. Only models whose MoE block + # unpacks that tuple may enable it. + self._context.enable_fused_ar_quant_mlp = enable_fused_ar_quant_mlp self._post_init_communicate() self._speculative_algo = SpeculativeAlgorithm.from_string( get_spec().speculative_algorithm @@ -1003,6 +1009,10 @@ class CommunicateContext: cache = None tp_rank: int force_layernorm_before_dp_gather: bool = False + # Mirror of LayerCommunicator's opt-in so the MLP-side communicate + # functions, which receive only (layernorm, context), can reach it. The + # attention side reads the flags off ``self`` directly in prepare_attn. + enable_fused_ar_quant_mlp: bool = False def is_same_group_size(self, a: ScatterMode, b: ScatterMode): return self.process_group_sizes[a] == self.process_group_sizes[b] @@ -1296,9 +1306,31 @@ def _gather_hidden_states_and_residual( apply_aiter_all_reduce_fusion(hidden_states) or apply_flashinfer_allreduce_fusion(hidden_states.shape[0]) ) and hasattr(layernorm, "forward_with_allreduce_fusion"): - hidden_states, residual = layernorm.forward_with_allreduce_fusion( - hidden_states, residual, use_attn_tp_group=True - ) + # Prefer the fused AR+RMSNorm+per-group-quant kernel, which also + # absorbs the separate activation-quant launch the MoE block + # would otherwise do. keep_bf16 is always on here: the MoE block + # has bf16 consumers (router gate, shared-expert gate) alongside + # the FP8 shared-expert projection. Returns None when the model + # has not opted in or the shape is not serviceable, in which + # case we fall back to plain AR+RMSNorm below. + quant_result = None + if context.enable_fused_ar_quant_mlp and hasattr( + layernorm, "forward_with_allreduce_fusion_quant_per_group" + ): + quant_result = ( + layernorm.forward_with_allreduce_fusion_quant_per_group( + hidden_states, + residual, + use_attn_tp_group=True, + keep_bf16=True, + ) + ) + if quant_result is not None: + hidden_states, residual = quant_result + else: + hidden_states, residual = layernorm.forward_with_allreduce_fusion( + hidden_states, residual, use_attn_tp_group=True + ) handled = True if not handled: diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 5bb45333a398..1e7d60ac1904 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -255,6 +255,70 @@ def _forward_with_allreduce_fusion( return norm_module.forward(x, residual, post_residual_addition) +# Same attribute GroupCoordinator.graph_capture looks for. +_GLUON_TP_AR_STATE_ATTR = "_gluon_tp_ar_state" + + +def _try_gluon_tp_ar_norm_quant( + x: torch.Tensor, + residual: torch.Tensor, + weight: torch.Tensor, + eps: float, + world_size: int, + group_size: int, + keep_bf16: bool, +): + """Run the in-tree Gluon TP all-reduce+norm+quant kernel, or return None. + + The rendezvous state (IPC staging buffer, synchronization rows, peer pointer + tables) is created once per TP group and cached on the group object, the same + way ``ca_comm`` is. Returns the same tuple shapes as the aiter path so the + caller is agnostic to which backend ran. + """ + try: + from sglang.srt.distributed.device_communicators import ( + gluon_tp_ar_norm_quant as _gluon_ar, + ) + except Exception: + return None + + if not _gluon_ar.is_supported(x, residual, weight, eps, world_size, group_size): + return None + + from sglang.srt.distributed import get_tp_group + + tp_group = get_tp_group() + state = getattr(tp_group, _GLUON_TP_AR_STATE_ATTR, None) + if state is None: + try: + state = _gluon_ar.GluonTpArNormQuantState( + group=tp_group.device_group, + device=x.device, + max_rows=max(_gluon_ar.SUPPORTED_M), + ) + except Exception as exc: + # Rendezvous is collective: every rank must reach the same verdict, + # so a failure here permanently disables the path on all ranks + # rather than leaving them out of step. + logger.warning( + "Gluon TP AR+norm+quant rendezvous failed, disabling: %s", exc + ) + setattr(tp_group, _GLUON_TP_AR_STATE_ATTR, False) + return None + setattr(tp_group, _GLUON_TP_AR_STATE_ATTR, state) + if state is False: + return None + + fp8_out, scale_out, residual_out, bf16_out = ( + _gluon_ar.fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( + state, x, residual, weight + ) + ) + if keep_bf16: + return (bf16_out, fp8_out, scale_out), residual_out + return (fp8_out, scale_out), residual_out + + def _forward_with_allreduce_fusion_quant_per_group( norm_module, x: torch.Tensor, @@ -313,6 +377,22 @@ def _forward_with_allreduce_fusion_quant_per_group( if world_size <= 1: return None + # Preferred backend on gfx950/TP4: a single Gluon kernel that performs the + # all-reduce itself over HIP IPC, so there is no separate collective launch. + # is_supported() is strict (arch, TP size, hidden size, M, eps, and + # non-bpreshuffle only); anything else drops through to the aiter chain. + fused = _try_gluon_tp_ar_norm_quant( + x, + residual, + weight, + norm_module.variance_epsilon, + world_size, + group_size, + keep_bf16, + ) + if fused is not None: + return fused + # ``transpose_scale=use_bpreshuffle`` asks the fused kernel to emit the # per-group scale directly in the column-major layout the gfx95 bpreshuffle # GEMM consumes (identical to ``materialize_bpreshuffle_fp8_scale``), so the diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index 43da47e00813..07a0236ba4b1 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -3,6 +3,7 @@ from typing import Any, Iterable, Optional, Set, Tuple import torch +import torch.nn.functional as F import triton from torch import nn @@ -49,11 +50,17 @@ sharded_weight_loader, ) from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock -from sglang.srt.runtime_context import get_forward, get_parallel, get_stream +from sglang.srt.runtime_context import ( + get_exec, + get_forward, + get_parallel, + get_stream, +) from sglang.srt.utils import ( LazyValue, add_prefix, cpu_has_amx_support, + get_bool_env_var, is_cpu, is_cuda, is_hip, @@ -69,6 +76,52 @@ _is_npu = is_npu() _is_cpu = is_cpu() _is_amx_available = cpu_has_amx_support() +_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip + + +def _enable_qwen3_next_fused_ar_quant() -> bool: + """Gate the fused AR+RMSNorm+per-group-FP8-quant path for Qwen3-Next. + + Mirrors ``_enable_qwen35_fused_ar_quant`` in ``qwen3_5.py``. Qwen3-Next has + the same hybrid GDN/attention + GemmaRMSNorm structure, so the same fused + epilogue applies: it replaces the ``--enable-aiter-allreduce-fusion`` + 3-kernel sequence (AR -> RMSNorm -> per-group quant) with a single fused + aiter kernel, or with a 2-kernel path when the fully-fused variant is not + eligible. ``LayerCommunicator`` falls back to plain AR+RMSNorm whenever the + helper returns ``None``, so enabling this never regresses the existing + AR+RMSNorm fusion. + + Opt-out: ``SGLANG_DISABLE_FUSED_AR_QUANT=1``. + """ + if not _use_aiter: + return False + if get_bool_env_var("SGLANG_DISABLE_FUSED_AR_QUANT", default="false"): + return False + return bool(get_exec().comm.enable_aiter_allreduce_fusion) + + +def _linear_accepts_fp8_tuple(linear: nn.Module) -> bool: + quant_method = getattr(linear, "quant_method", None) + return quant_method.__class__.__name__ == "Fp8LinearMethod" and ( + getattr(quant_method, "block_quant", False) + or getattr(quant_method, "use_mxfp8", False) + ) + + +def _select_fused_ar_input_for_linear(hidden_states, linear: nn.Module): + """Pick the right member of a fused-AR output tuple for ``linear``.""" + if not isinstance(hidden_states, tuple): + return hidden_states + if len(hidden_states) == 3: + hs_bf16, hs_fp8, hs_scale = hidden_states + if _linear_accepts_fp8_tuple(linear): + return (hs_fp8, hs_scale) + return hs_bf16 + if len(hidden_states) == 2 and _linear_accepts_fp8_tuple(linear): + return hidden_states + raise TypeError( + f"{linear.__class__.__name__} cannot consume fused AR quant tuple input" + ) if _is_npu: @@ -374,7 +427,23 @@ def fix_query_key_value_ordering( return query, key, value, z, b, a - def _forward_input_proj(self, hidden_states: torch.Tensor): + def _forward_input_proj(self, hidden_states): + # The fused AR+RMSNorm+per-group-quant path hands down a tuple. + # in_proj_qkvz is FP8 block-quantized and consumes (fp8, scale) + # directly; in_proj_ba is a small bf16 projection on the same normed + # activations, so it needs the unquantized bf16 side-output rather + # than a lossy dequantization. + if isinstance(hidden_states, tuple): + hs_bf16 = hidden_states[0] + hs_qkvz = _select_fused_ar_input_for_linear( + hidden_states, self.in_proj_qkvz + ) + projected_states_qkvz, _ = self.in_proj_qkvz(hs_qkvz) + projected_states_ba, _ = self.in_proj_ba(hs_bf16) + return projected_states_qkvz, projected_states_ba + return self._forward_input_proj_tensor(hidden_states) + + def _forward_input_proj_tensor(self, hidden_states: torch.Tensor): if ( _is_cpu or _is_npu @@ -504,6 +573,67 @@ def _apply_qwen3_next_mlp( return hidden_states, residual +class Qwen3NextSparseMoeBlock(Qwen2MoeSparseMoeBlock): + """Qwen2MoeSparseMoeBlock that can consume a fused-AR quant tuple. + + When the MLP-side fused AR+RMSNorm+per-group-quant epilogue is enabled, the + LayerCommunicator hands down ``(bf16, fp8, scale)`` instead of a tensor. + Everything here runs on the bf16 view -- the router gate and the + shared-expert gate are unquantized projections, and the MoE runner owns its + own quantization -- except the shared expert's FP8 ``gate_up_proj``, which + consumes ``(fp8, scale)`` directly and so avoids re-quantizing activations + that were already quantized upstream. + + This lives here rather than in ``Qwen2MoeSparseMoeBlock`` to keep the shared + block byte-identical for the other models that use it (qwen2_moe, qwen3_5, + qwen3_5_text). The support is not Qwen3-Next-specific in principle: any of + those models could opt into MLP-side fused quant through aiter's generic + kernel. Promoting this to the base class is the natural follow-up once a + second model opts in. + """ + + def forward( + self, + hidden_states, + forward_batch=None, + defer_finalize: bool = False, + ) -> torch.Tensor: + if not isinstance(hidden_states, tuple): + return super().forward(hidden_states, forward_batch, defer_finalize) + + hs_bf16, hs_fp8, hs_scale = hidden_states + if not _linear_accepts_fp8_tuple( + getattr(self.shared_expert, "gate_up_proj", None) + ): + # Shared expert cannot take fp8; nothing to gain, stay on bf16. + return super().forward(hs_bf16, forward_batch, defer_finalize) + + original = self._forward_shared_experts + + def shared_with_fp8(hidden, apply_gate: bool = True): + # Gates stay on bf16; only the FP8 projection sees (fp8, scale). + shared_output = self.shared_expert((hs_fp8, hs_scale)) + if self.shared_expert_gate is not None and apply_gate: + from sglang.srt.models.qwen2_moe import _is_hip as _qm_is_hip + + gate = self.shared_expert_gate(hidden) + if _qm_is_hip: + from sglang.kernels.ops.moe.triton_sigmoid_gate_mul import ( + sigmoid_gate_mul_broadcast, + ) + + shared_output = sigmoid_gate_mul_broadcast(shared_output, gate) + else: + shared_output = F.sigmoid(gate) * shared_output + return shared_output + + self._forward_shared_experts = shared_with_fp8 + try: + return super().forward(hs_bf16, forward_batch, defer_finalize) + finally: + self._forward_shared_experts = original + + class Qwen3HybridLinearDecoderLayer(nn.Module): def __init__( self, @@ -535,7 +665,7 @@ def __init__( ) if self.is_layer_sparse: - self.mlp = Qwen2MoeSparseMoeBlock( + self.mlp = Qwen3NextSparseMoeBlock( layer_id=layer_id, config=config, quant_config=quant_config, @@ -562,6 +692,16 @@ def __init__( input_layernorm=self.input_layernorm, post_attention_layernorm=self.post_attention_layernorm, allow_reduce_scatter=True, + # GDN layers need both the bf16 normed output (for the small bf16 + # in_proj_ba gating projection) and the (fp8, scale) pair, so the + # fused kernel only helps when in_proj_qkvz can consume fp8. + enable_fused_ar_quant=_enable_qwen3_next_fused_ar_quant() + and _linear_accepts_fp8_tuple(self.linear_attn.in_proj_qkvz), + fused_ar_quant_keep_bf16=_enable_qwen3_next_fused_ar_quant() + and _linear_accepts_fp8_tuple(self.linear_attn.in_proj_qkvz), + # MLP side: hand the MoE block (bf16, fp8, scale) so the shared + # expert's FP8 gate_up_proj can skip re-quantizing activations. + enable_fused_ar_quant_mlp=_enable_qwen3_next_fused_ar_quant(), ) def forward( @@ -703,7 +843,7 @@ def __init__( ) if self.is_layer_sparse: - self.mlp = Qwen2MoeSparseMoeBlock( + self.mlp = Qwen3NextSparseMoeBlock( layer_id=layer_id, config=config, quant_config=quant_config, @@ -734,6 +874,14 @@ def __init__( input_layernorm=self.input_layernorm, post_attention_layernorm=self.post_attention_layernorm, allow_reduce_scatter=True, + # Full-attention layers have a single FP8 qkv_proj consumer, so no + # bf16 side-output is needed. + enable_fused_ar_quant=_enable_qwen3_next_fused_ar_quant() + and _linear_accepts_fp8_tuple(self.qkv_proj), + fused_ar_quant_keep_bf16=False, + # MLP side: hand the MoE block (bf16, fp8, scale) so the shared + # expert's FP8 gate_up_proj can skip re-quantizing activations. + enable_fused_ar_quant_mlp=_enable_qwen3_next_fused_ar_quant(), ) self.alt_stream = alt_stream @@ -761,6 +909,10 @@ def _apply_qk_norm( return q, k def forward_prepare_native(self, positions, hidden_states): + if _use_aiter and isinstance(hidden_states, tuple): + hidden_states = _select_fused_ar_input_for_linear( + hidden_states, self.qkv_proj + ) qkv, _ = self.qkv_proj(hidden_states) if self.attn_output_gate: q_gate, k, v = qkv.split( diff --git a/test/registered/amd/test_gluon_tp_ar_norm_quant.py b/test/registered/amd/test_gluon_tp_ar_norm_quant.py new file mode 100644 index 000000000000..a7ef03c69b07 --- /dev/null +++ b/test/registered/amd/test_gluon_tp_ar_norm_quant.py @@ -0,0 +1,228 @@ +"""Correctness of the fused TP all-reduce + residual add + Gemma RMSNorm + +per-1x128 FP8 quant Gluon kernel. + +Runs as a 4-rank torchrun job and compares every supported token count M +against an fp32 reference. Outside the serving stack, so it is fast and gives a +per-call cost that does not depend on the torch profiler -- collectives absorb +rank skew under profiling, which makes trace durations untrustworthy here. + +Requires 4 gfx950 GPUs; skipped everywhere else. + + python -m unittest test.registered.amd.test_gluon_tp_ar_norm_quant + # or directly: + torchrun --nproc_per_node=4 test/registered/amd/test_gluon_tp_ar_norm_quant.py [--bench] +""" + +import os +import subprocess +import sys +import unittest + +import torch + +from sglang.test.ci.ci_register import register_amd_ci +from sglang.test.test_utils import CustomTestCase + +register_amd_ci(est_time=180, suite="stage-c-test-large-8-gpu-amd-mi35x") + +WORLD_SIZE = 4 + + +def _gfx950_gpus() -> int: + """Number of visible gfx950 devices (0 if CUDA/HIP unavailable).""" + if not torch.cuda.is_available(): + return 0 + n = 0 + for i in range(torch.cuda.device_count()): + try: + arch = torch.cuda.get_device_properties(i).gcnArchName + except Exception: + return 0 + if arch.split(":", 1)[0] == "gfx950": + n += 1 + return n + + +class TestGluonTpArNormQuant(CustomTestCase): + @unittest.skipUnless( + _gfx950_gpus() >= WORLD_SIZE, + f"needs {WORLD_SIZE} gfx950 GPUs", + ) + def test_bit_exact_against_fp32_reference(self): + """Spawn the 4-rank worker; it exits non-zero on any mismatch.""" + env = dict(os.environ) + env.setdefault( + "HIP_VISIBLE_DEVICES", ",".join(str(i) for i in range(WORLD_SIZE)) + ) + proc = subprocess.run( + [ + sys.executable, + "-m", + "torch.distributed.run", + f"--nproc_per_node={WORLD_SIZE}", + os.path.abspath(__file__), + "--worker", + ], + env=env, + capture_output=True, + text=True, + timeout=1800, + ) + self.assertEqual( + proc.returncode, + 0, + f"4-rank kernel check failed\n--- stdout ---\n{proc.stdout}\n" + f"--- stderr ---\n{proc.stderr}", + ) + self.assertIn("RESULT: PASS", proc.stdout) + + +# -------------------------------------------------------------------------- +# 4-rank worker (this file is also the torchrun entry point) +# -------------------------------------------------------------------------- + +import argparse + +import torch.distributed as dist + +from sglang.srt.distributed.device_communicators import gluon_tp_ar_norm_quant as G + +HIDDEN = G.HIDDEN_SIZE +EPS = G.EPS +FP8_MAX = 448.0 + + +def reference(x_local, residual, weight, group_size=128): + """Eager reference: all-reduce -> +residual -> Gemma RMSNorm -> per-group FP8. + + The reduction is done by gathering every rank's contribution and summing in + fp32 **in ascending rank order**, then rounding once to bf16 -- exactly what + the kernel does (``(((x0+x1)+x2)+x3)`` over the global-rank-ordered peer + table). Using ``dist.all_reduce`` here instead would introduce a different + accumulation order/precision and show up as a spurious mismatch. + + The round-to-bf16 before the residual add is deliberate and matches aiter's + fused kernel, so the fused path equals the unfused one bit for bit. + """ + world = dist.get_world_size() + gathered = [torch.empty_like(x_local) for _ in range(world)] + dist.all_gather(gathered, x_local) + acc = gathered[0].to(torch.float32) + for i in range(1, world): + acc = acc + gathered[i].to(torch.float32) + reduced = acc.to(torch.bfloat16) + + value = reduced.to(torch.float32) + residual.to(torch.float32) + residual_out = value.to(torch.bfloat16) + + var = (value * value).mean(dim=-1, keepdim=True) + normed = value * torch.rsqrt(var + EPS) * weight.to(torch.float32) + normed_bf16 = normed.to(torch.bfloat16) + + v = normed_bf16.to(torch.float32).view(-1, HIDDEN // group_size, group_size) + amax = v.abs().amax(dim=-1).clamp_min(1.0e-10) + scales = amax / FP8_MAX + q = torch.clamp(v / scales[..., None], -FP8_MAX, FP8_MAX) + # The kernel stores fp8_e4m3; round the reference the same way, otherwise + # every comparison is off by up to one fp8 ULP (32 near the 448 ceiling). + q = q.to(torch.float8_e4m3fn).to(torch.float32) + return q.view(-1, HIDDEN), scales, residual_out, normed_bf16 + + +def main(): + ap = argparse.ArgumentParser() + ap.add_argument("--bench", action="store_true") + ap.add_argument("--iters", type=int, default=200) + args = ap.parse_args() + + rank = int(os.environ["RANK"]) + torch.cuda.set_device(rank) + dist.init_process_group("nccl") + device = torch.device("cuda", rank) + torch.manual_seed(1234 + rank) + + weight = (torch.randn(HIDDEN, device=device) * 0.02 + 1.0).to(torch.bfloat16) + dist.broadcast(weight, src=0) + + state = G.GluonTpArNormQuantState( + group=dist.group.WORLD, device=device, max_rows=max(G.SUPPORTED_M) + ) + if rank == 0: + print(f"rendezvous OK; peer input table = {state.input_ptr_table.tolist()}") + + ok = True + for m in G.SUPPORTED_M: + x = (torch.randn(m, HIDDEN, device=device) * 0.5).to(torch.bfloat16) + residual = (torch.randn(m, HIDDEN, device=device) * 0.5).to(torch.bfloat16) + dist.broadcast(residual, src=0) # residual is rank-identical in serving + + ref_q, ref_s, ref_r, ref_n = reference(x, residual, weight) + dist.barrier() + q, s, r, n = G.fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( + state, x, residual, weight + ) + dist.barrier() + + dq = (q.to(torch.float32) - ref_q).abs().max().item() + ds = (s.to(torch.float32) - ref_s).abs().max().item() + dr = (r.to(torch.float32) - ref_r.to(torch.float32)).abs().max().item() + dn = (n.to(torch.float32) - ref_n.to(torch.float32)).abs().max().item() + # fp8 values are integers in [-448,448]; scales are ~1e-3. Tolerances are + # loose on q (1 ulp of fp8) and tight on the bf16 outputs. + good = ( + dq <= 1.001 + and ds <= 2e-3 * max(1.0, ref_s.abs().max().item()) + and dr <= 3e-2 + and dn <= 3e-2 + ) + ok &= good + if rank == 0: + print( + f"M={m:3d} dq={dq:8.4f} ds={ds:10.3e} dres={dr:8.4f} dnorm={dn:8.4f}" + f" {'OK' if good else 'MISMATCH'}" + ) + + if args.bench and ok: + if rank == 0: + print("\nper-call cost (no profiler):") + for m in G.SUPPORTED_M: + x = (torch.randn(m, HIDDEN, device=device) * 0.5).to(torch.bfloat16) + residual = (torch.randn(m, HIDDEN, device=device) * 0.5).to(torch.bfloat16) + for _ in range(20): + G.fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( + state, x, residual, weight + ) + dist.barrier() + torch.cuda.synchronize() + start = torch.cuda.Event(enable_timing=True) + end = torch.cuda.Event(enable_timing=True) + start.record() + for _ in range(args.iters): + G.fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( + state, x, residual, weight + ) + end.record() + torch.cuda.synchronize() + us = start.elapsed_time(end) * 1000.0 / args.iters + t = torch.tensor([us], device=device) + dist.all_reduce(t, op=dist.ReduceOp.MAX) + if rank == 0: + print(f" M={m:3d} {t.item():7.2f} us/call (max over ranks)") + + state.close() + dist.barrier() + if rank == 0: + print("\nRESULT:", "PASS" if ok else "FAIL") + dist.destroy_process_group() + sys.exit(0 if ok else 1) + + +if __name__ == "__main__": + # CI executes registered files as `python3 `, so the default entry + # point must run the tests. The 4-rank worker is behind an explicit flag and + # is what the test case re-invokes under torchrun. + if "--worker" in sys.argv: + sys.argv.remove("--worker") + main() + else: + unittest.main() From c4fd739e8c5728c8ada5b916c83dd8b4200bfe5f Mon Sep 17 00:00:00 2001 From: Rita Brugarolas Brufau Date: Wed, 23 Sep 2026 07:10:08 +0000 Subject: [PATCH 2/5] [AMD] Qwen3-Next: probe HIP IPC capability, use the runtime context accessor 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 --- .../device_communicators/hip_ipc.py | 91 +++++++++++++++---- python/sglang/srt/layers/layernorm.py | 4 +- .../amd/test_gluon_tp_ar_norm_quant.py | 21 +++++ 3 files changed, 96 insertions(+), 20 deletions(-) diff --git a/python/sglang/srt/distributed/device_communicators/hip_ipc.py b/python/sglang/srt/distributed/device_communicators/hip_ipc.py index 6e45b8fa3579..14ca03379e5b 100644 --- a/python/sglang/srt/distributed/device_communicators/hip_ipc.py +++ b/python/sglang/srt/distributed/device_communicators/hip_ipc.py @@ -23,6 +23,7 @@ import ctypes import logging +from functools import lru_cache from typing import List, Optional import torch.distributed as dist @@ -34,6 +35,8 @@ _HANDLE_BYTES = 64 # hipIpcMemLazyEnablePeerAccess _LAZY_ENABLE_PEER_ACCESS = 1 +# Sentinel written by hipMemGetAddressRange on builds where the query is broken. +_SIZE_T_MAX = (1 << 64) - 1 class HipIpcMemHandle(ctypes.Structure): @@ -105,6 +108,39 @@ def free(self, ptr: int) -> None: def memset(self, ptr: ctypes.c_void_p, value: int, size_in_bytes: int) -> None: self._check(self.lib.hipMemset(ptr, value, size_in_bytes), "hipMemset") + def get_address_range(self, ptr: int) -> tuple[int, int]: + """Segment base and size containing ``ptr``, validated. + + ``hipMemGetAddressRange`` is not merely allowed to fail here -- on some + ROCm builds it returns success while writing nonsense for pointers owned + by torch's caching allocator (observed: base ``0x100``, size ``2**64-1`` + on ``torch 2.11.0+rocm10.0.0`` / HIP 7.15). Trusting the status code + alone forwards that garbage to ``hipIpcGetMemHandle``, which then fails + with an opaque ``hipErrorInvalidValue``. Validate the payload instead. + """ + base = ctypes.c_void_p() + size = ctypes.c_size_t() + self._check( + self.lib.hipMemGetAddressRange( + ctypes.byref(base), ctypes.byref(size), ctypes.c_void_p(ptr) + ), + "hipMemGetAddressRange", + ) + base_value = base.value or 0 + size_value = size.value + if not 0 < base_value <= ptr or size_value in (0, _SIZE_T_MAX): + raise RuntimeError( + "hipMemGetAddressRange reported success but returned an " + f"implausible range for 0x{ptr:x}: base=0x{base_value:x} " + f"size={size_value}. This memory cannot be shared over HIP IPC." + ) + if ptr - base_value >= size_value: + raise RuntimeError( + f"hipMemGetAddressRange range [0x{base_value:x}, +{size_value}) " + f"does not contain 0x{ptr:x}" + ) + return base_value, size_value + def get_ipc_handle(self, ptr: ctypes.c_void_p) -> bytes: handle = HipIpcMemHandle() self._check( @@ -196,17 +232,10 @@ def create_shared_tensor( # start at zero, and a zeroed staging buffer is harmless. tensor = torch.zeros(shape, dtype=dtype, device=device) - base = ctypes.c_void_p() - size = ctypes.c_size_t() data_ptr = tensor.data_ptr() - lib._check( - lib.lib.hipMemGetAddressRange( - ctypes.byref(base), ctypes.byref(size), ctypes.c_void_p(data_ptr) - ), - "hipMemGetAddressRange", - ) - offset = data_ptr - base.value - handle = lib.get_ipc_handle(base) + base_value, _ = lib.get_address_range(data_ptr) + offset = data_ptr - base_value + handle = lib.get_ipc_handle(ctypes.c_void_p(base_value)) world_size = dist.get_world_size(group=group) rank = dist.get_rank(group=group) @@ -229,6 +258,34 @@ def create_shared_tensor( return tensor, pointers, opened_bases +@lru_cache(maxsize=1) +def torch_memory_is_ipc_capable() -> bool: + """Can a torch-allocated tensor be published over HIP IPC on this build? + + Not every ROCm build supports this. On ``torch 2.11.0+rocm10.0.0`` (HIP + 7.15) ``hipMemGetAddressRange`` returns success but yields a bogus range for + caching-allocator pointers, and ``hipIpcGetMemHandle`` rejects them -- + while raw ``hipMalloc`` allocations on the same device share fine. The + capability is therefore probed, not inferred from a version number, so a + later runtime that fixes it is picked up with no code change. + + Single-process and side-effect free; the result is cached per process. + """ + import torch + + if not torch.cuda.is_available(): + return False + try: + probe = torch.zeros(1024, dtype=torch.uint8, device="cuda") + lib = HipRTLibrary() + base_value, _ = lib.get_address_range(probe.data_ptr()) + lib.get_ipc_handle(ctypes.c_void_p(base_value)) + except Exception as exc: # noqa: BLE001 - any failure means "unsupported" + logger.debug("HIP IPC on torch memory unavailable: %s", exc) + return False + return True + + def register_peer_pointers(tensors, group: Optional[ProcessGroup] = None): """Publish existing tensors over IPC and return their peer pointers. @@ -245,16 +302,14 @@ def register_peer_pointers(tensors, group: Optional[ProcessGroup] = None): lib = HipRTLibrary() local = [] for t in tensors: - base = ctypes.c_void_p() - size = ctypes.c_size_t() data_ptr = t.data_ptr() - lib._check( - lib.lib.hipMemGetAddressRange( - ctypes.byref(base), ctypes.byref(size), ctypes.c_void_p(data_ptr) - ), - "hipMemGetAddressRange", + base_value, _ = lib.get_address_range(data_ptr) + local.append( + ( + lib.get_ipc_handle(ctypes.c_void_p(base_value)), + data_ptr - base_value, + ) ) - local.append((lib.get_ipc_handle(base), data_ptr - base.value)) world_size = dist.get_world_size(group=group) rank = dist.get_rank(group=group) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index fc2ab80df98c..43e76a49f24c 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -285,9 +285,9 @@ def _try_gluon_tp_ar_norm_quant( if not _gluon_ar.is_supported(x, residual, weight, eps, world_size, group_size): return None - from sglang.srt.distributed import get_tp_group + from sglang.srt.runtime_context import get_parallel - tp_group = get_tp_group() + tp_group = get_parallel().tp_group state = getattr(tp_group, _GLUON_TP_AR_STATE_ATTR, None) if state is None: try: diff --git a/test/registered/amd/test_gluon_tp_ar_norm_quant.py b/test/registered/amd/test_gluon_tp_ar_norm_quant.py index a7ef03c69b07..d7380f2e7b18 100644 --- a/test/registered/amd/test_gluon_tp_ar_norm_quant.py +++ b/test/registered/amd/test_gluon_tp_ar_norm_quant.py @@ -43,11 +43,32 @@ def _gfx950_gpus() -> int: return n +def _torch_memory_is_ipc_capable() -> bool: + """Probe HIP IPC on torch memory, which this kernel's rendezvous requires. + + Deliberately a capability probe and not a version check: some ROCm builds + (e.g. torch 2.11.0+rocm10.0.0 / HIP 7.15) cannot share caching-allocator + memory over IPC, and a runtime that later fixes this re-enables the test + with no edit here. + """ + if _gfx950_gpus() < 1: + return False + from sglang.srt.distributed.device_communicators.hip_ipc import ( + torch_memory_is_ipc_capable, + ) + + return torch_memory_is_ipc_capable() + + class TestGluonTpArNormQuant(CustomTestCase): @unittest.skipUnless( _gfx950_gpus() >= WORLD_SIZE, f"needs {WORLD_SIZE} gfx950 GPUs", ) + @unittest.skipUnless( + _torch_memory_is_ipc_capable(), + "HIP IPC on torch-allocated memory is unsupported on this ROCm build", + ) def test_bit_exact_against_fp32_reference(self): """Spawn the 4-rank worker; it exits non-zero on any mismatch.""" env = dict(os.environ) From 0c27af3680d81a6d8fdebabc9beed6608fd0fc04 Mon Sep 17 00:00:00 2001 From: Rita Brugarolas Brufau Date: Thu, 1 Oct 2026 15:34:44 +0000 Subject: [PATCH 3/5] [AMD] Qwen3-Next: support the gfx95 bpreshuffle layout, drop a redundant 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 --- .../gluon_tp_ar_norm_quant.py | 44 +++++++++---------- 1 file changed, 20 insertions(+), 24 deletions(-) diff --git a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py index 83787a993f47..71c69a26ce54 100644 --- a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py @@ -15,13 +15,15 @@ * TP world size exactly 4 -- the reduction is hand-unrolled over 4 peers and the epoch protocol counts in units of 3 remote peers. * hidden size exactly 2048, eps 1e-6, and M in {1,2,4,8,16,32,64}. - * ``_use_aiter_bpreshuffle_gfx95`` must be False. On ROCm >= 7.2 SGLang - physically preshuffles FP8 weights into the gfx95 bpreshuffle layout, which - this kernel's consumers do not expect. Returning False here makes the caller - fall back rather than produce wrong numbers. +ROCm >= 7.2 is supported: there SGLang preshuffles FP8 weights and the gfx95 +bpreshuffle GEMM reads the activation scale column-major, so the scale this +kernel writes row-major is relayed out through upstream's +``materialize_bpreshuffle_fp8_scale`` before it is returned. The weights +themselves are never touched by this path. Anything outside that envelope returns False and the caller keeps the stock -path. Opt out entirely with ``SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT=1``. +path. The existing ``SGLANG_DISABLE_FUSED_AR_QUANT`` opts out of the whole +fused AR+quant path, this backend included; no separate flag is introduced. The kernel body lives in ``gluon_tp_ar_norm_quant_kernel.py``. It is tuned per token count M; see ``SUPPORTED_M``. @@ -41,7 +43,7 @@ create_shared_tensor, register_peer_pointers, ) -from sglang.srt.utils import get_bool_env_var, is_hip +from sglang.srt.utils import is_hip logger = logging.getLogger(__name__) @@ -98,14 +100,6 @@ def is_available(world_size: int, device: torch.device) -> bool: """ if not _is_hip or world_size != TP_SIZE: return False - if get_bool_env_var("SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT", default="false"): - return False - from sglang.srt.layers.quantization.fp8_utils import ( - _use_aiter_bpreshuffle_gfx95, - ) - - if _use_aiter_bpreshuffle_gfx95: - return False return _arch_is_gfx950(device) and _gluon_available() @@ -120,8 +114,6 @@ def is_supported( """Support predicate. False means "caller should use the stock path".""" if not _is_hip or residual is None: return False - if get_bool_env_var("SGLANG_DISABLE_GLUON_TP_AR_NORM_QUANT", default="false"): - return False if world_size != TP_SIZE: return False if hidden_states.dim() != 2 or hidden_states.shape[-1] != HIDDEN_SIZE: @@ -139,14 +131,6 @@ def is_supported( # The kernel bakes eps into the launch; only the Qwen3-Next value is tuned. if abs(eps - EPS) > 1e-12: return False - # On ROCm >= 7.2 SGLang preshuffles FP8 weights; this path is not validated - # against that layout and would silently produce wrong results. - from sglang.srt.layers.quantization.fp8_utils import ( - _use_aiter_bpreshuffle_gfx95, - ) - - if _use_aiter_bpreshuffle_gfx95: - return False if not _arch_is_gfx950(hidden_states.device): return False return _gluon_available() @@ -344,4 +328,16 @@ def fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( eps=EPS, ) ) + from sglang.srt.layers.quantization.fp8_utils import ( + _use_aiter_bpreshuffle_gfx95, + materialize_bpreshuffle_fp8_scale, + ) + + if _use_aiter_bpreshuffle_gfx95: + # On ROCm >= 7.2 the gfx95 bpreshuffle GEMM reads the activation scale + # column-major; the kernel writes it row-major. Relayout with + # upstream's own helper so both the CK and the Triton branch of + # w8a8_block_fp8_linear interpret it correctly. G == 16 at hidden + # 2048, so this copy is negligible. + scales = materialize_bpreshuffle_fp8_scale(scales) return quantized, scales, residual_out, normalized From a34d600249fe2c7d7135340e3e037b12730a9522 Mon Sep 17 00:00:00 2001 From: Rita Brugarolas Brufau Date: Thu, 1 Oct 2026 16:05:46 +0000 Subject: [PATCH 4/5] [AMD] Qwen3-Next: apply ruff-format to qwen3_next.py 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 --- python/sglang/srt/models/qwen3_next.py | 7 +++---- 1 file changed, 3 insertions(+), 4 deletions(-) diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index 3a123b86b1ea..fcb57b4e3fc0 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -94,7 +94,6 @@ fused_qkvzba_split_reshape_cat = fused_qkvzba_split_reshape_cat_npu - def _linear_accepts_fp8_tuple(linear: nn.Module) -> bool: quant_method = getattr(linear, "quant_method", None) return quant_method.__class__.__name__ == "Fp8LinearMethod" and ( @@ -119,7 +118,6 @@ def _select_fused_ar_input_for_linear(hidden_states, linear: nn.Module): ) - class Qwen3NextSparseMoeBlock(Qwen2MoeSparseMoeBlock): """MoE block that can consume the fused AR+norm+quant tuple. @@ -165,7 +163,6 @@ def shared_with_fp8(hidden, apply_gate: bool = True): self._forward_shared_experts = original - def _moe_accepts_fp8_tuple(mlp) -> bool: """True when this MLP is the subclass that unpacks (bf16, fp8, scale).""" shared = getattr(mlp, "shared_expert", None) @@ -482,7 +479,9 @@ def _forward_input_proj(self, hidden_states: torch.Tensor): # block-quantized and takes the bf16 side-output. if isinstance(hidden_states, tuple): hs_shape = hidden_states[0] - hs_qkvz = _select_fused_ar_input_for_linear(hidden_states, self.in_proj_qkvz) + hs_qkvz = _select_fused_ar_input_for_linear( + hidden_states, self.in_proj_qkvz + ) hs_ba = _select_fused_ar_input_for_linear(hidden_states, self.in_proj_ba) else: hs_shape = hs_qkvz = hs_ba = hidden_states From 60f28ab74a5d62b80e38d2d7efeb8a7fe6091eff Mon Sep 17 00:00:00 2001 From: Rita Brugarolas Brufau Date: Mon, 5 Oct 2026 06:57:22 +0000 Subject: [PATCH 5/5] [AMD] Qwen3-Next: reserve the epoch atomically for M>=16, cap the tuned 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 --- .../gluon_tp_ar_norm_quant.py | 5 +- .../gluon_tp_ar_norm_quant_kernel.py | 75 +----------- python/sglang/srt/models/qwen3_next.py | 108 +++++++++++++----- .../amd/test_gluon_tp_ar_norm_quant.py | 49 ++++++++ 4 files changed, 141 insertions(+), 96 deletions(-) diff --git a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py index 71c69a26ce54..36f95d1b189c 100644 --- a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py @@ -52,7 +52,10 @@ TP_SIZE = 4 HIDDEN_SIZE = 2048 EPS = 1.0e-6 -SUPPORTED_M = (1, 2, 4, 8, 16, 32, 64) +# Tuned decode shapes. Above M=16 aiter's allreduce_fusion_kernel_1stage_per_group +# is faster than this kernel (measured -0.1% at M=32, -2.0% at M=64), so those +# shapes are declined and the caller keeps the existing path. +SUPPORTED_M = (1, 2, 4, 8, 16) GROUP_SIZE = 128 # Words 0/1 count arrivals / completed reads; word 2 is local CTA progress; # word 3 is unused. They must start at zero and persist across calls. diff --git a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py index dc469597ad10..356d516563ee 100644 --- a/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py @@ -192,68 +192,6 @@ def _finish(local, target, last): ) -@gluon.jit -def _finish_late(local, target, thread, step: gl.constexpr): - gl.inline_asm_elementwise( - """ - v_cmp_eq_u32 vcc, 0, $6 - s_and_saveexec_b64 $0, vcc - s_cbranch_execz 3f - global_atomic_add $1, $3, $4, off offset:8 sc0 - s_waitcnt vmcnt(0) - v_add_u32 $1, $4, $1 - v_mul_lo_u32 $1, 3, $1 - v_cmp_eq_u32 vcc, $1, $5 - s_cbranch_vccz 3f - 2: - global_load_dword $2, $3, off offset:4 sc1 - s_waitcnt vmcnt(0) - v_sub_u32 $2, $2, $5 - v_cmp_ge_i32 vcc, $2, 0 - s_cbranch_vccz 2b - 3: - s_or_saveexec_b64 $0, $0 - """, - "=&s,=&v,=&v,v,v,v,v,~{vcc},~{scc},~{memory}", - [local, step, target, thread], - dtype=(gl.uint64, gl.uint32, gl.uint32), - is_pure=False, - pack=1, - ) - - -# Only the final CTA polls. Earlier CTAs remain free to retire on a single CU. -@gluon.jit -def _finish_uniform_late(local, target, thread, step: gl.constexpr): - gl.inline_asm_elementwise( - """ - v_cmp_eq_u32 vcc, 0, $6 - s_and_saveexec_b64 $0, vcc - s_cbranch_execz 3f - global_atomic_add $1, $3, $4, off offset:8 sc0 - s_waitcnt vmcnt(0) - v_add_u32 $1, $4, $1 - v_mul_lo_u32 $1, 3, $1 - v_cmp_eq_u32 vcc, $1, $5 - s_cbranch_vccz 3f - v_readfirstlane_b32 s0, $5 - 2: - s_load_dword $2, $7, 0x4 glc - s_waitcnt lgkmcnt(0) - s_sub_u32 $2, $2, s0 - s_cmp_ge_i32 $2, 0 - s_cbranch_scc0 2b - 3: - s_or_saveexec_b64 $0, $0 - """, - "=&s,=&v,=&s,v,v,v,v,s,~{s0},~{vcc},~{scc},~{memory}", - [local, step, target, thread, local.to(gl.uint64)], - dtype=(gl.uint64, gl.uint32, gl.uint32), - is_pure=False, - pack=1, - ) - - @gluon.jit def _load_pointer(table): return gl.inline_asm_elementwise( @@ -336,7 +274,11 @@ def _fused( ): PAIR: gl.constexpr = M == 64 NCTA: gl.constexpr = M // 2 if PAIR else M - TICKET: gl.constexpr = M >= 2 and M <= 8 + # Every multi-CTA shape reserves its epoch atomically. Sampling word 2 + # with a plain load instead (as M>=16 used to) races the finish-time + # writes: a CTA can observe a counter from a neighbouring epoch and + # wait on a target that is never signalled. + TICKET: gl.constexpr = M >= 2 NW: gl.constexpr = 16 if PAIR else 8 pid = gl.program_id(0) if TICKET: @@ -479,12 +421,7 @@ def _fused( scales, M <= 2 or M == 16, ) - if M <= 8: - _finish(local, target, last) - elif M == 16 or M == 32: - _finish_uniform_late(local, target, thread, 128 // NCTA) - else: - _finish_late(local, target, thread, 128 // NCTA) + _finish(local, target, last) def tp4_allreduce_add_gemma_rmsnorm_group_fp8_quant_gluon( diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index fcb57b4e3fc0..5dbc99174e43 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -1,5 +1,6 @@ import enum import logging +from functools import partial from typing import Any, Iterable, Optional, Set, Tuple import torch @@ -121,11 +122,18 @@ def _select_fused_ar_input_for_linear(hidden_states, linear: nn.Module): class Qwen3NextSparseMoeBlock(Qwen2MoeSparseMoeBlock): """MoE block that can consume the fused AR+norm+quant tuple. - With ``Fp8Input.TUPLE_AND_BF16`` declared on the FFN stage, the boundary - hands down ``(bf16, fp8, scale)``. Only the shared expert's FP8 - ``gate_up_proj`` consumes ``(fp8, scale)``; the router gate, the - shared-expert gate and the MoE runner all stay on bf16. Subclassing keeps - ``qwen2_moe.py``, which several models share, byte-identical to upstream. + With ``Fp8Input.TUPLE_AND_BF16`` declared on the FFN stage the boundary + hands down ``(bf16, fp8, scale)``. Where the fp8 goes depends on whether + the shared expert was fused into the routed experts: + + * fusion off -- the separate shared expert's FP8 ``gate_up_proj`` takes + ``(fp8, scale)``; + * fusion on -- the routed experts take it via ``pre_quant_input``, + which covers every expert rather than just the shared one. + + The router gate, the shared-expert gate and the MoE's own bf16 inputs are + untouched. Subclassing keeps ``qwen2_moe.py``, shared by several models, + byte-identical to upstream. """ def forward( @@ -139,36 +147,84 @@ def forward( hs_bf16, hs_fp8, hs_scale = hidden_states shared = getattr(self, "shared_expert", None) - if shared is None or not _linear_accepts_fp8_tuple( - getattr(shared, "gate_up_proj", None) - ): - # Nothing downstream can take the fp8; drop it rather than requantize. - return super().forward(hs_bf16, forward_batch, defer_finalize) - - original = self._forward_shared_experts - - def shared_with_fp8(hidden, apply_gate: bool = True): - # Gates stay on bf16; only the FP8 projection sees (fp8, scale). - shared_output = self.shared_expert((hs_fp8, hs_scale)) - if self.shared_expert_gate is not None and apply_gate: - shared_output = ( - torch.sigmoid(self.shared_expert_gate(hidden)) * shared_output + restore = [] + + if shared is not None: + if not _linear_accepts_fp8_tuple(getattr(shared, "gate_up_proj", None)): + return super().forward(hs_bf16, forward_batch, defer_finalize) + original = self._forward_shared_experts + + def shared_with_fp8(hidden, apply_gate: bool = True): + # Gates stay on bf16; only the FP8 projection sees (fp8, scale). + shared_output = self.shared_expert((hs_fp8, hs_scale)) + if self.shared_expert_gate is not None and apply_gate: + shared_output = ( + torch.sigmoid(self.shared_expert_gate(hidden)) * shared_output + ) + return shared_output + + self._forward_shared_experts = shared_with_fp8 + restore.append(("_forward_shared_experts", original)) + elif _routed_experts_accept_pre_quant(self): + # The runner declines the hand-off (wrong quant layout, router + # weights pre-applied) by ignoring pre_quant_input, so this is safe + # to offer unconditionally once the backend is known to read it. + experts = self.experts + for name in ("forward", "forward_deferred_finalize"): + bound = getattr(experts, name, None) + if bound is None: + continue + restore.append((name, bound, experts)) + setattr( + experts, + name, + partial(bound, pre_quant_input=(hs_fp8, hs_scale)), ) - return shared_output + else: + return super().forward(hs_bf16, forward_batch, defer_finalize) - self._forward_shared_experts = shared_with_fp8 try: return super().forward(hs_bf16, forward_batch, defer_finalize) finally: - self._forward_shared_experts = original + for entry in restore: + if len(entry) == 2: + setattr(self, entry[0], entry[1]) + else: + setattr(entry[2], entry[0], entry[1]) + + +def _routed_experts_accept_pre_quant(mlp) -> bool: + """Whether the routed-expert runner can take caller-quantized activations. + + Always False for now. The plumbing works -- FusedMoE.pre_quant_input reaches + the runner and the aiter standard pre-permute can forward it -- but aiter's + heuristic then selects an asm kernel that this gfx950 build does not carry: + + fmoe_fp8_blockscale_g1u1 failed: get_heuristic_kernel not find kernel + gfx950_..._fmoe_bf16_blockscaleBf16_g1u1_vs_pf2_silu_16x128 + + aiter's own mori path carries an upscale fallback for exactly this class of + gap. Until a build ships that kernel, declining keeps the fp8 unproduced + rather than produced and discarded, which measured as a net regression. + """ + return False def _moe_accepts_fp8_tuple(mlp) -> bool: - """True when this MLP is the subclass that unpacks (bf16, fp8, scale).""" + """True when this MLP can consume (bf16, fp8, scale) somewhere downstream. + + Two shapes, depending on shared-expert fusion: + * fusion off -- a separate ``shared_expert`` whose FP8 ``gate_up_proj`` + takes ``(fp8, scale)`` directly; + * fusion on -- no separate shared expert, so the fp8 goes to the routed + experts via ``pre_quant_input``. + """ + if not isinstance(mlp, Qwen3NextSparseMoeBlock): + return False shared = getattr(mlp, "shared_expert", None) - return isinstance(mlp, Qwen3NextSparseMoeBlock) and _linear_accepts_fp8_tuple( - getattr(shared, "gate_up_proj", None) - ) + if shared is not None: + return _linear_accepts_fp8_tuple(getattr(shared, "gate_up_proj", None)) + return _routed_experts_accept_pre_quant(mlp) class Qwen3GatedDeltaNet(nn.Module): diff --git a/test/registered/amd/test_gluon_tp_ar_norm_quant.py b/test/registered/amd/test_gluon_tp_ar_norm_quant.py index d7380f2e7b18..53601a01c8ce 100644 --- a/test/registered/amd/test_gluon_tp_ar_norm_quant.py +++ b/test/registered/amd/test_gluon_tp_ar_norm_quant.py @@ -150,10 +150,56 @@ def reference(x_local, residual, weight, group_size=128): return q.view(-1, HIDDEN), scales, residual_out, normed_bf16 +def _stress_multi_cta(state, weight, device, rank, rounds: int): + """Hammer the multi-CTA shapes back-to-back while the CUs are contended. + + M>=16 launches NCTA = 16..32 workgroups that must all agree on one epoch. + If a CTA can observe a word-2 value from a neighbouring epoch, the elected + CTA waits on a target no peer will ever signal and the rank hangs. The race + needs the CTAs of one launch to be spread over time, so a background stream + is kept busy to deny them simultaneous residency, and the shapes are issued + back-to-back so consecutive epochs overlap in flight. + """ + pressure = torch.cuda.Stream() + hog_a = torch.randn(4096, 4096, device=device, dtype=torch.bfloat16) + hog_b = torch.randn(4096, 4096, device=device, dtype=torch.bfloat16) + shapes = [m for m in G.SUPPORTED_M if m >= 16] + ok = True + for r in range(rounds): + with torch.cuda.stream(pressure): + for _ in range(8): + hog_a = torch.mm(hog_a, hog_b) + for m in shapes: + x = (torch.randn(m, HIDDEN, device=device) * 0.5).to(torch.bfloat16) + residual = (torch.randn(m, HIDDEN, device=device) * 0.5).to(torch.bfloat16) + dist.broadcast(residual, src=0) + ref_q, ref_s, ref_r, _ = reference(x, residual, weight) + # No barrier between shapes: consecutive epochs must stay in flight. + q, s, res, _ = G.fused_tp_ar_add_gemma_rmsnorm_group_fp8_quant( + state, x, residual, weight + ) + if r == rounds - 1: + dq = (q.to(torch.float32) - ref_q).abs().max().item() + dr = ( + (res.to(torch.float32) - ref_r.to(torch.float32)).abs().max().item() + ) + if not (dq <= 1.001 and dr <= 3e-2): + ok = False + if rank == 0: + print(f" stress M={m}: MISMATCH dq={dq} dr={dr}") + torch.cuda.current_stream().wait_stream(pressure) + torch.cuda.synchronize() + dist.barrier() + if rank == 0: + print(f" stress: {rounds} rounds x {shapes} completed, no hang") + return ok + + def main(): ap = argparse.ArgumentParser() ap.add_argument("--bench", action="store_true") ap.add_argument("--iters", type=int, default=200) + ap.add_argument("--stress-rounds", type=int, default=60) args = ap.parse_args() rank = int(os.environ["RANK"]) @@ -203,6 +249,9 @@ def main(): f" {'OK' if good else 'MISMATCH'}" ) + if ok: + ok = _stress_multi_cta(state, weight, device, rank, rounds=args.stress_rounds) + if args.bench and ok: if rank == 0: print("\nper-call cost (no profiler):")