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..36f95d1b189c --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant.py @@ -0,0 +1,346 @@ +# 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}. +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. 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``. +""" + +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 is_hip + +logger = logging.getLogger(__name__) + +_is_hip = is_hip() + +TP_SIZE = 4 +HIDDEN_SIZE = 2048 +EPS = 1.0e-6 +# 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. +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 + 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 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 + 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, + ) + ) + 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 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..356d516563ee --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/gluon_tp_ar_norm_quant_kernel.py @@ -0,0 +1,478 @@ +"""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 _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 + # 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: + 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, + ) + _finish(local, target, last) + + +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..14ca03379e5b --- /dev/null +++ b/python/sglang/srt/distributed/device_communicators/hip_ipc.py @@ -0,0 +1,365 @@ +# 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 functools import lru_cache +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 +# Sentinel written by hipMemGetAddressRange on builds where the query is broken. +_SIZE_T_MAX = (1 << 64) - 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_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( + 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) + + data_ptr = tensor.data_ptr() + 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) + 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 + + +@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. + + 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: + data_ptr = t.data_ptr() + base_value, _ = lib.get_address_range(data_ptr) + local.append( + ( + lib.get_ipc_handle(ctypes.c_void_p(base_value)), + 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 7f403b85dc9c..c1cbe003302f 100644 --- a/python/sglang/srt/distributed/parallel_state.py +++ b/python/sglang/srt/distributed/parallel_state.py @@ -683,7 +683,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/layer_boundary/construction.py b/python/sglang/srt/layers/layer_boundary/construction.py index f426ec14a60e..4c9d8533c8e2 100644 --- a/python/sglang/srt/layers/layer_boundary/construction.py +++ b/python/sglang/srt/layers/layer_boundary/construction.py @@ -213,7 +213,11 @@ def __init__( if kind is StageKind.ATTENTION else () ) - fused = ffn_input_fusions(self) if kind is StageKind.FFN else () + fused = ( + ffn_input_fusions(self, next(iter(self.edges.values())).incoming.need.read) + if kind is StageKind.FFN + else () + ) self.paths = {} for variant, edges in self.edges.items(): if is_branch: diff --git a/python/sglang/srt/layers/layer_boundary/fusions/allreduce.py b/python/sglang/srt/layers/layer_boundary/fusions/allreduce.py index f3a2341d2660..dcd08c837035 100644 --- a/python/sglang/srt/layers/layer_boundary/fusions/allreduce.py +++ b/python/sglang/srt/layers/layer_boundary/fusions/allreduce.py @@ -24,6 +24,7 @@ from sglang.srt.layers.layer_boundary.output import UnreducedOutput from sglang.srt.layers.layer_boundary.residual.add_norm import ( NORM_QUANT_READOUT, + NORM_READOUT, Fp8Input, NormQuantReadout, aiter_ar_fusion_applies, @@ -111,7 +112,7 @@ def fused_attn_input( ) -def ffn_input_fusions(plan) -> Tuple["FfnInputFusion", ...]: +def ffn_input_fusions(plan, read=NORM_READOUT) -> Tuple["FfnInputFusion", ...]: """The fused kernels that can take the attention -> FFN steps, in the order they are tried. They add the residual plainly; the boundary tries them only when the update it writes in is a plain add. A backend's come @@ -119,11 +120,27 @@ def ffn_input_fusions(plan) -> Tuple["FfnInputFusion", ...]: given = plan.fusions.ffn_input_fusions(plan) if plan.fusions else () if not hasattr(plan.norm, "forward_with_allreduce_fusion"): return given + # The same gate the attention side applies. An FFN stage that has not + # declared an Fp8Input keeps NORM_READOUT, an unrelated type, so this is + # False and the plain kernel below is reached exactly as before. + fuses_quant = ( + isinstance(read, NormQuantReadout) + and read.fp8_input is not None + and _use_aiter + and not get_bool_env_var("SGLANG_DISABLE_FUSED_AR_QUANT", default="false") + and get_exec().comm.enable_aiter_allreduce_fusion + and hasattr(plan.norm, "forward_with_allreduce_fusion_quant_per_group") + ) return ( *given, FfnInputFusion( completes=SumGroup.ATTN_TP, - run=partial(fused_ffn_input, plan), + run=partial( + fused_ffn_input, + plan, + fuses_quant=fuses_quant, + keep_bf16=fuses_quant and read.fp8_input is Fp8Input.TUPLE_AND_BF16, + ), preserves_residual=( flashinfer_preserves_residual if isinstance(plan.norm, (RMSNorm, GemmaRMSNorm)) @@ -152,14 +169,29 @@ def fused_ffn_input( hidden_states: torch.Tensor, residual: torch.Tensor, forward_batch: ForwardBatch, + *, + fuses_quant: bool = False, + keep_bf16: bool = False, ) -> Optional[Tuple[torch.Tensor, torch.Tensor]]: """The attention-TP all-reduce, residual add and norm in one aiter or - flashinfer kernel; None when neither takes the batch.""" + flashinfer kernel; None when neither takes the batch. The optional FP8 + result follows the consumer read declaration, as on the attention side.""" if not ( aiter_ar_fusion_applies(hidden_states, forward_batch) or flashinfer_ar_fusion_applies(hidden_states.shape[0]) ): return None + if fuses_quant: + # Falls back to AR+RMSNorm + separate quant internally when the + # fully-fused kernel cannot service the shape. + quant_result = plan.norm.forward_with_allreduce_fusion_quant_per_group( + hidden_states, + residual, + use_attn_tp_group=True, + keep_bf16=keep_bf16, + ) + if quant_result is not None: + return quant_result return plan.norm.forward_with_allreduce_fusion( hidden_states, residual, use_attn_tp_group=True ) diff --git a/python/sglang/srt/layers/layernorm.py b/python/sglang/srt/layers/layernorm.py index 6068e8f2095f..9e6c7bdd4318 100644 --- a/python/sglang/srt/layers/layernorm.py +++ b/python/sglang/srt/layers/layernorm.py @@ -267,6 +267,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.runtime_context import get_parallel + + tp_group = get_parallel().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, @@ -325,6 +389,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 cf7444960a99..77873d9d173c 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 @@ -28,6 +29,10 @@ declare_ffn, ) from sglang.srt.layers.layer_boundary.residual import batch as residual_batch +from sglang.srt.layers.layer_boundary.residual.add_norm import ( + Fp8Input, + NormQuantReadout, +) from sglang.srt.layers.layernorm import GemmaRMSNorm from sglang.srt.layers.linear import ( ColumnParallelLinear, @@ -91,6 +96,138 @@ 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 ( + 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 member of a fused-AR output tuple that ``linear`` can consume.""" + 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" + ) + + +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)``. 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( + self, + hidden_states, + forward_batch: Optional[ForwardBatch] = 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 + shared = getattr(self, "shared_expert", None) + 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)), + ) + else: + return super().forward(hs_bf16, forward_batch, defer_finalize) + + try: + return super().forward(hs_bf16, forward_batch, defer_finalize) + finally: + 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 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) + 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): def __init__( self, @@ -399,7 +536,20 @@ def _forward_input_proj(self, hidden_states: torch.Tensor): else: DUAL_STREAM_TOKEN_THRESHOLD = 1024 - seq_len, _ = hidden_states.shape + # With Fp8Input.TUPLE_AND_BF16 declared, the boundary hands down a + # (bf16, fp8, scale) triple: an FP8 in_proj_qkvz takes (fp8, scale) + # directly, while the small in_proj_ba gating projection is not + # 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_ba = _select_fused_ar_input_for_linear(hidden_states, self.in_proj_ba) + else: + hs_shape = hs_qkvz = hs_ba = hidden_states + + seq_len, _ = hs_shape.shape if ( self.alt_stream is not None and get_is_capture_mode() @@ -407,13 +557,13 @@ def _forward_input_proj(self, hidden_states: torch.Tensor): ): current_stream = torch.cuda.current_stream() self.alt_stream.wait_stream(current_stream) - projected_states_qkvz, _ = self.in_proj_qkvz(hidden_states) + projected_states_qkvz, _ = self.in_proj_qkvz(hs_qkvz) with torch.cuda.stream(self.alt_stream): - projected_states_ba, _ = self.in_proj_ba(hidden_states) + projected_states_ba, _ = self.in_proj_ba(hs_ba) current_stream.wait_stream(self.alt_stream) else: - projected_states_qkvz, _ = self.in_proj_qkvz(hidden_states) - projected_states_ba, _ = self.in_proj_ba(hidden_states) + projected_states_qkvz, _ = self.in_proj_qkvz(hs_qkvz) + projected_states_ba, _ = self.in_proj_ba(hs_ba) return projected_states_qkvz, projected_states_ba def forward( @@ -520,7 +670,7 @@ def __init__( self.layer_id = layer_id if self.is_layer_sparse: - self.mlp = Qwen2MoeSparseMoeBlock( + self.mlp = Qwen3NextSparseMoeBlock( layer_id=layer_id, config=config, quant_config=quant_config, @@ -544,12 +694,30 @@ def __init__( self.post_attention_layernorm = GemmaRMSNorm( config.hidden_size, eps=config.rms_norm_eps ) + # Declare the fused AR+RMSNorm+per-group-quant read only when the + # consuming projection can take the (fp8, scale) pair; otherwise the + # boundary keeps the plain AR+RMSNorm path. + accepts_fp8_input = _linear_accepts_fp8_tuple(self.linear_attn.in_proj_qkvz) self.attn_boundary, self.ffn_boundary = append_stages( - (declare_attn(), self.input_layernorm), + ( + declare_attn( + read=NormQuantReadout( + fp8_input=Fp8Input.TUPLE_AND_BF16 if accepts_fp8_input else None + ) + ), + self.input_layernorm, + ), ( declare_ffn( sparse=self.is_layer_sparse, next_layer_sparse=is_next_layer_sparse, + # Only the Qwen3Next MoE subclass unpacks the tuple; a dense + # Qwen2MoeMLP layer keeps the plain AR+RMSNorm read. + read=NormQuantReadout( + fp8_input=Fp8Input.TUPLE_AND_BF16 + if self.is_layer_sparse and _moe_accepts_fp8_tuple(self.mlp) + else None + ), ), self.post_attention_layernorm, ), @@ -677,7 +845,7 @@ def __init__( is_next_layer_sparse = True if self.is_layer_sparse: - self.mlp = Qwen2MoeSparseMoeBlock( + self.mlp = Qwen3NextSparseMoeBlock( layer_id=layer_id, config=config, quant_config=quant_config, @@ -705,12 +873,30 @@ def __init__( self.q_norm = GemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps) self.k_norm = GemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps) + # Declare the fused AR+RMSNorm+per-group-quant read only when the + # consuming projection can take the (fp8, scale) pair; otherwise the + # boundary keeps the plain AR+RMSNorm path. + accepts_fp8_input = _linear_accepts_fp8_tuple(self.qkv_proj) self.attn_boundary, self.ffn_boundary = append_stages( - (declare_attn(), self.input_layernorm), + ( + declare_attn( + read=NormQuantReadout( + fp8_input=Fp8Input.TUPLE if accepts_fp8_input else None + ) + ), + self.input_layernorm, + ), ( declare_ffn( sparse=self.is_layer_sparse, next_layer_sparse=is_next_layer_sparse, + # Only the Qwen3Next MoE subclass unpacks the tuple; a dense + # Qwen2MoeMLP layer keeps the plain AR+RMSNorm read. + read=NormQuantReadout( + fp8_input=Fp8Input.TUPLE_AND_BF16 + if self.is_layer_sparse and _moe_accepts_fp8_tuple(self.mlp) + else None + ), ), self.post_attention_layernorm, ), @@ -741,7 +927,9 @@ def _apply_qk_norm( return q, k def forward_prepare_native(self, positions, hidden_states): - qkv, _ = self.qkv_proj(hidden_states) + qkv, _ = self.qkv_proj( + _select_fused_ar_input_for_linear(hidden_states, self.qkv_proj) + ) if self.attn_output_gate: q_gate, k, v = qkv.split( [self.q_size * 2, self.kv_size, self.kv_size], dim=-1 @@ -760,7 +948,9 @@ def forward_prepare_native(self, positions, hidden_states): return q, k, v, gate def forward_prepare_npu(self, positions, hidden_states, forward_batch): - qkv, _ = self.qkv_proj(hidden_states) + qkv, _ = self.qkv_proj( + _select_fused_ar_input_for_linear(hidden_states, self.qkv_proj) + ) # Calculate first full attention layer ID based on config if self.attn.layer_id == (self.config.full_attention_interval - 1): self.rotary_emb.get_cos_sin_with_position(positions) 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..53601a01c8ce --- /dev/null +++ b/test/registered/amd/test_gluon_tp_ar_norm_quant.py @@ -0,0 +1,298 @@ +"""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 + + +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) + 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 _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"]) + 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 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):") + 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()