From 96368a5f77bccd6f3560366dcefe809b48731595 Mon Sep 17 00:00:00 2001 From: "jacky.cheng" Date: Wed, 8 Jul 2026 17:34:20 +0800 Subject: [PATCH] [AMD] Fuse shared-expert sigmoid + bf16->fp32 cast into the MoE append kernel (3 kernels -> 1) (#28658) --- .../triton_utils/fused_moe_triton_kernels.py | 31 +++++++++++++++++-- python/sglang/srt/models/qwen2_moe.py | 30 +++++++++++++----- 2 files changed, 50 insertions(+), 11 deletions(-) diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_kernels.py b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_kernels.py index e727e0f6d..82f88e465 100644 --- a/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_kernels.py +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/fused_moe_triton_kernels.py @@ -1330,10 +1330,12 @@ def _fused_append_shared_experts_with_weights_kernel( out_ids_ptr, out_weights_ptr, N_BASE, + scale, K: tl.constexpr, S: tl.constexpr, BLOCK_K: tl.constexpr, BLOCK_S: tl.constexpr, + APPLY_SIGMOID: tl.constexpr, ): pid = tl.program_id(0) @@ -1352,22 +1354,43 @@ def _fused_append_shared_experts_with_weights_kernel( mask_s = offs_s < S shared_ids = tl.cast(N_BASE + offs_s, ids.dtype) shared_ws = tl.load(shared_weights_ptr + pid * S + offs_s, mask=mask_s) + if APPLY_SIGMOID: + # Fuse sigmoid(shared_gate) + dtype upcast (+ optional 1/ep_size scale) + # in-register so the raw bf16 logits stream straight into the fp32 + # output, eliminating the standalone sigmoid and bf16->fp32 copy kernels. + shared_ws = tl.sigmoid(shared_ws.to(tl.float32)) * scale tl.store(out_ids_ptr + out_row_ptr + K + offs_s, shared_ids, mask=mask_s) tl.store(out_weights_ptr + out_row_ptr + K + offs_s, shared_ws, mask=mask_s) def fused_append_shared_experts_with_weights( - topk_ids, topk_weights, shared_weights, num_fused_shared_experts, N=None + topk_ids, + topk_weights, + shared_weights, + num_fused_shared_experts, + N=None, + apply_sigmoid=False, + scale=1.0, ): - """Like fused_append_shared_experts but accepts per-token shared weights tensor.""" + """Like fused_append_shared_experts but accepts per-token shared weights tensor. + + When ``apply_sigmoid`` is True, ``shared_weights`` are treated as raw gate + logits: the kernel applies ``sigmoid`` (in fp32) and the optional ``scale`` + in-register, so the caller can skip the separate ``sigmoid`` activation and + the bf16->fp32 cast. When False the legacy behavior is preserved exactly. + """ assert N is not None, "N (shared expert base id) must be provided" m, k = topk_ids.shape s = int(num_fused_shared_experts) if s <= 0: return topk_ids, topk_weights - shared_weights_2d = shared_weights.to(topk_weights.dtype) + # When fusing sigmoid in-kernel, keep the raw logits dtype (the kernel emits + # fp32 directly); otherwise match the output weight dtype as before. + shared_weights_2d = ( + shared_weights if apply_sigmoid else shared_weights.to(topk_weights.dtype) + ) if shared_weights_2d.ndim == 1: shared_weights_2d = shared_weights_2d.unsqueeze(-1) if shared_weights_2d.shape[1] < s: @@ -1389,10 +1412,12 @@ def fused_append_shared_experts_with_weights( out_ids, out_weights, N_BASE=N, + scale=scale, K=k, S=s, BLOCK_K=block_k, BLOCK_S=block_s, + APPLY_SIGMOID=apply_sigmoid, num_warps=1, ) return out_ids, out_weights diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index fc18c3408..54d2d3c36 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -357,14 +357,18 @@ class Qwen2MoeSparseMoeBlock(nn.Module): def _get_shared_expert_weights( self, hidden_states: torch.Tensor - ) -> Optional[torch.Tensor]: - """Return sigmoid(shared_expert_gate) for fused shared expert weights.""" + ) -> Optional[Tuple[torch.Tensor, float]]: + """Return the shared_expert_gate weights and the 1/ep_size scale. + + On the AMD AITER path the sigmoid activation and the scale are applied + (in fp32) inside the fused append kernel, so this returns the raw gate + logits to avoid a standalone activation kernel + cast. On the CUDA path + the legacy eager ``sigmoid(logits) * scale`` is returned unchanged. + """ if not self.enable_shared_expert_fusion or self.shared_expert_gate is None: return None shared_out = self.shared_expert_gate(hidden_states) shared_logits = shared_out[0] if isinstance(shared_out, tuple) else shared_out - w = F.sigmoid(shared_logits) - # This block runs only on the AMD AITER shared_expert_fusion path # Allreduce-EP path: the fused shared expert occupies a single global # slot loaded onto every EP rank (see FusedMoE.__init__: num_shared_slots # == num_fused_shared_experts when not is_deepep_class_backend()). Every @@ -372,10 +376,15 @@ class Qwen2MoeSparseMoeBlock(nn.Module): # post-experts all_reduce sums it ep_size times. Pre-scale the per-token # routing weight by 1/ep_size to cancel this, mirroring DeepSeek-V2's # fused_shared_experts_scaling_factor pattern. + scale = 1.0 moe_ep_size = get_parallel().moe_ep_size if moe_ep_size > 1 and not is_deepep_class_backend(): - w = w / float(moe_ep_size) - return w + scale = 1.0 / float(moe_ep_size) + # Only AITER fuses sigmoid + cast in-kernel; on CUDA keep the legacy + # eager activation so the NVIDIA path behavior is unchanged. + if not _use_aiter: + return F.sigmoid(shared_logits) * scale, 1.0 + return shared_logits, scale def _append_shared_to_topk_output( self, @@ -385,20 +394,25 @@ class Qwen2MoeSparseMoeBlock(nn.Module): """Append shared expert ids and weights to topk output before fused MoE.""" if not self.enable_shared_expert_fusion: return topk_output - shared_weights = self._get_shared_expert_weights(hidden_states) - if shared_weights is None: + shared = self._get_shared_expert_weights(hidden_states) + if shared is None: return topk_output + shared_weights, shared_scale = shared from sglang.srt.layers.moe.moe_runner.triton_utils.fused_moe_triton_kernels import ( fused_append_shared_experts_with_weights, ) + # AITER returns raw logits + scale for in-kernel sigmoid fusion; CUDA + # returns pre-activated weights (scale already folded in) → no fusion. fused_topk_ids, fused_topk_weights = fused_append_shared_experts_with_weights( topk_output.topk_ids, topk_output.topk_weights, shared_weights, self.num_fused_shared_experts, N=self.num_experts, + apply_sigmoid=_use_aiter, + scale=shared_scale, ) return StandardTopKOutput( topk_weights=fused_topk_weights,