[AMD] Fuse shared_expert_gate GEMV into the MoE append kernel (HIP/aiter) (#28666)

Co-authored-by: sogalin_codegen <39478626+sogalin@users.noreply.github.com>
This commit is contained in:
jacky.cheng
2026-08-13 21:32:27 -07:00
committed by GitHub
co-authored by sogalin_codegen
parent 0a6bbbe128
commit 240a12b302
4 changed files with 670 additions and 38 deletions
@@ -1342,6 +1342,8 @@ def _fused_append_shared_experts_with_weights_kernel(
shared_weights_ptr,
out_ids_ptr,
out_weights_ptr,
hidden_ptr,
wgate_ptr,
N_BASE,
scale,
K: tl.constexpr,
@@ -1349,6 +1351,9 @@ def _fused_append_shared_experts_with_weights_kernel(
BLOCK_K: tl.constexpr,
BLOCK_S: tl.constexpr,
APPLY_SIGMOID: tl.constexpr,
FUSE_GATE: tl.constexpr,
HIDDEN: tl.constexpr,
BLOCK_H: tl.constexpr,
):
pid = tl.program_id(0)
@@ -1366,12 +1371,20 @@ def _fused_append_shared_experts_with_weights_kernel(
offs_s = tl.arange(0, BLOCK_S)
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
if FUSE_GATE:
offs_h = tl.arange(0, BLOCK_H)
mask_h = offs_h < HIDDEN
h = tl.load(hidden_ptr + pid * HIDDEN + offs_h, mask=mask_h, other=0.0).to(
tl.float32
)
w = tl.load(wgate_ptr + offs_h, mask=mask_h, other=0.0).to(tl.float32)
logit = tl.sum(h * w)
shared_val = tl.sigmoid(logit) * scale
shared_ws = tl.zeros((BLOCK_S,), dtype=tl.float32) + shared_val
else:
shared_ws = tl.load(shared_weights_ptr + pid * S + offs_s, mask=mask_s)
if APPLY_SIGMOID:
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)
@@ -1384,31 +1397,65 @@ def fused_append_shared_experts_with_weights(
num_fused_shared_experts,
N=None,
apply_sigmoid=False,
fuse_gate=False,
hidden_states=None,
gate_weight=None,
scale=1.0,
):
"""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.
Two optional in-kernel fusions are supported (both default off → legacy
behavior is preserved byte-for-byte):
- ``apply_sigmoid=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.
- ``fuse_gate=True``: the shared_expert_gate GEMV
(``hidden_states @ gate_weight.T``) + sigmoid + ``scale`` are computed
*inside* the kernel, eliminating the standalone gate GEMM launch.
``shared_weights`` is ignored; ``hidden_states`` ([M, HIDDEN]) and
``gate_weight`` ([1, HIDDEN] or [HIDDEN]) must be provided. This subsumes
``apply_sigmoid`` (the sigmoid is intrinsic), so the two are mutually
exclusive.
"""
assert not (
fuse_gate and apply_sigmoid
), "fuse_gate already applies sigmoid in-kernel; do not also set apply_sigmoid"
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
# 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:
shared_weights_2d = shared_weights_2d.expand(m, s)
shared_weights_2d = shared_weights_2d.contiguous()
if fuse_gate:
assert (
hidden_states is not None and gate_weight is not None
), "fuse_gate=True requires hidden_states and gate_weight"
hidden_arg = hidden_states.contiguous()
wgate_arg = gate_weight.reshape(-1).contiguous()
hidden_dim = hidden_arg.shape[1]
block_h = triton.next_power_of_2(hidden_dim)
shared_arg = topk_weights
num_warps = 8
else:
# When fusing sigmoid in-kernel (apply_sigmoid), 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:
shared_weights_2d = shared_weights_2d.expand(m, s)
shared_arg = shared_weights_2d.contiguous()
# hidden_ptr / wgate_ptr are unused; pass placeholders.
hidden_arg = topk_weights
wgate_arg = topk_weights
hidden_dim = 1
block_h = 1
num_warps = 1
out_ids = torch.empty((m, k + s), dtype=topk_ids.dtype, device=topk_ids.device)
out_weights = torch.empty(
@@ -1421,9 +1468,11 @@ def fused_append_shared_experts_with_weights(
_fused_append_shared_experts_with_weights_kernel[(m,)](
topk_ids,
topk_weights,
shared_weights_2d,
shared_arg,
out_ids,
out_weights,
hidden_arg,
wgate_arg,
N_BASE=N,
scale=scale,
K=k,
@@ -1431,6 +1480,9 @@ def fused_append_shared_experts_with_weights(
BLOCK_K=block_k,
BLOCK_S=block_s,
APPLY_SIGMOID=apply_sigmoid,
num_warps=1,
FUSE_GATE=fuse_gate,
HIDDEN=hidden_dim,
BLOCK_H=block_h,
num_warps=num_warps,
)
return out_ids, out_weights
+46 -16
View File
@@ -394,34 +394,64 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
return F.sigmoid(shared_logits) * scale, 1.0
return shared_logits, scale
def _shared_expert_scale(self) -> float:
"""1/ep_size pre-scale for the fused shared-expert routing weight.
Mirrors the scaling applied in _get_shared_expert_weights; see that
method for the allreduce-EP rationale.
"""
moe_ep_size = get_parallel().moe_ep_size
if moe_ep_size > 1 and not is_deepep_class_backend():
return 1.0 / float(moe_ep_size)
return 1.0
def _append_shared_to_topk_output(
self,
topk_output: StandardTopKOutput,
hidden_states: torch.Tensor,
) -> StandardTopKOutput:
"""Append shared expert ids and weights to topk output before fused MoE."""
if not self.enable_shared_expert_fusion:
if not self.enable_shared_expert_fusion or self.shared_expert_gate is None:
return topk_output
shared = self._get_shared_expert_weights(hidden_states)
if shared is None:
return topk_output
shared_weights, shared_scale = shared
from sglang.kernels.ops.moe.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,
)
if _use_aiter:
# HIP/aiter: fuse the shared_expert_gate GEMV + sigmoid + scale into
# the append kernel, eliminating the standalone gate GEMM launch.
# This subsumes the sigmoid-only fusion: there is no separate gate
# GEMM and no _get_shared_expert_weights call on this path.
fused_topk_ids, fused_topk_weights = (
fused_append_shared_experts_with_weights(
topk_output.topk_ids,
topk_output.topk_weights,
None,
self.num_fused_shared_experts,
N=self.num_experts,
fuse_gate=True,
hidden_states=hidden_states,
gate_weight=self.shared_expert_gate.weight,
scale=self._shared_expert_scale(),
)
)
else:
# CUDA: _get_shared_expert_weights returns pre-activated weights
# (sigmoid + scale already folded in) → legacy append, no fusion.
shared = self._get_shared_expert_weights(hidden_states)
if shared is None:
return topk_output
shared_weights, _ = shared
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,
)
)
return StandardTopKOutput(
topk_weights=fused_topk_weights,
topk_ids=fused_topk_ids,