[AMD] Enable share expert fusion with router experts for Qwen3.5 BF16 & FP8 (#20736)

Co-authored-by: Chen, Todd <zhenchen@amd.com>
Co-authored-by: jacky.cheng <yichiche@amd.com>
This commit is contained in:
Chen, Zhentao
2026-04-14 18:52:36 -07:00
committed by GitHub
co-authored by Chen, Todd jacky.cheng
parent 46c8a597ef
commit ea05ea5abe
2 changed files with 218 additions and 8 deletions
+108 -5
View File
@@ -62,7 +62,7 @@ from sglang.srt.layers.moe import (
)
from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class
from sglang.srt.layers.moe.fused_moe_triton import FusedMoE
from sglang.srt.layers.moe.topk import TopK
from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK, TopKOutputChecker
from sglang.srt.layers.moe.utils import (
RoutingMethodType,
filter_moe_weight_param_global_expert,
@@ -88,8 +88,10 @@ from sglang.srt.server_args import get_global_server_args
from sglang.srt.utils import (
add_prefix,
cpu_has_amx_support,
get_bool_env_var,
is_cpu,
is_cuda,
is_hip,
make_layers,
use_intel_amx_backend,
)
@@ -100,6 +102,25 @@ logger = logging.getLogger(__name__)
_is_cuda = is_cuda()
_is_cpu = is_cpu()
_is_cpu_amx_available = cpu_has_amx_support()
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
def can_fuse_shared_expert(
config: PretrainedConfig,
) -> bool:
"""Whether the shared expert may be fused as an extra MoE expert (Qwen3.5 + Aiter).
Caller must still gate on ``support_shared_expert_fusion`` and ``_use_aiter``.
"""
if (
get_global_server_args().disable_shared_experts_fusion is True
or getattr(config, "shared_expert_intermediate_size", 0) <= 0
or config.shared_expert_intermediate_size != config.moe_intermediate_size
or get_moe_a2a_backend().is_deepep()
):
return False
return True
class Qwen2MoeMLP(nn.Module):
@@ -163,6 +184,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
alt_stream: Optional[torch.cuda.Stream] = None,
prefix: str = "",
is_nextn: bool = False,
support_shared_expert_fusion: bool = False,
):
super().__init__()
self.tp_size = get_tensor_model_parallel_world_size()
@@ -173,6 +195,27 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
f"Tensor parallel size {self.tp_size} is greater than "
f"the number of experts {config.num_experts}."
)
self.num_experts = config.num_experts
self.num_shared_experts = 0
self.num_fused_shared_experts = 0
if hasattr(config, "n_shared_experts"):
# config defines the number of shared experts
self.num_shared_experts = config.n_shared_experts
elif (
hasattr(config, "shared_expert_intermediate_size")
and config.shared_expert_intermediate_size > 0
):
# n_shared_experts is not defined, but shared_expert_intermediate_size is defined, so we use 1 as the number of shared experts
self.num_shared_experts = 1
self.enable_shared_expert_fusion = False # default to False
if _use_aiter:
# enable shared expert fusion when use aiter
self.enable_shared_expert_fusion = (
support_shared_expert_fusion and can_fuse_shared_expert(config)
)
if self.enable_shared_expert_fusion:
self.num_fused_shared_experts = self.num_shared_experts
self.topk = TopK(
top_k=config.num_experts_per_tok,
@@ -182,14 +225,24 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
self.experts = get_moe_impl_class(quant_config)(
layer_id=self.layer_id,
top_k=config.num_experts_per_tok,
num_experts=config.num_experts
+ get_global_server_args().ep_num_redundant_experts,
top_k=(
config.num_experts_per_tok
if not self.enable_shared_expert_fusion
else config.num_experts_per_tok + self.num_fused_shared_experts
),
num_experts=(
config.num_experts + get_global_server_args().ep_num_redundant_experts
if not self.enable_shared_expert_fusion
else config.num_experts
+ get_global_server_args().ep_num_redundant_experts
+ self.num_fused_shared_experts
),
hidden_size=config.hidden_size,
intermediate_size=config.moe_intermediate_size,
quant_config=quant_config,
prefix=add_prefix("experts", prefix),
routing_method_type=RoutingMethodType.RenormalizeNaive,
num_fused_shared_experts=self.num_fused_shared_experts,
)
self.gate = ReplicatedLinear(
@@ -199,7 +252,13 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
quant_config=None,
prefix=add_prefix("gate", prefix),
)
if config.shared_expert_intermediate_size > 0:
# When enable_shared_expert_fusion, the shared expert runs inside the MoE kernel
# (via _append_shared_to_topk_output); a separate shared_expert MLP would
# double-count. If fusion is off (num_fused_shared_experts == 0), keep shared_expert.
if (
config.shared_expert_intermediate_size > 0
and not self.enable_shared_expert_fusion
):
self.shared_expert = Qwen2MoeMLP(
hidden_size=config.hidden_size,
intermediate_size=config.shared_expert_intermediate_size,
@@ -245,6 +304,46 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
)
]
def _get_shared_expert_weights(self, hidden_states: torch.Tensor) -> torch.Tensor:
"""Return sigmoid(shared_expert_gate) for fused shared expert weights."""
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
return F.sigmoid(shared_logits)
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:
return topk_output
shared_weights = self._get_shared_expert_weights(hidden_states)
if shared_weights is None:
return topk_output
M = topk_output.topk_ids.shape[0]
shared_expert_id = self.num_experts
shared_ids = torch.full(
(M, self.num_fused_shared_experts),
shared_expert_id,
dtype=topk_output.topk_ids.dtype,
device=topk_output.topk_ids.device,
)
shared_weights = shared_weights.expand(M, self.num_fused_shared_experts).to(
topk_output.topk_weights.dtype
)
fused_topk_ids = torch.cat([topk_output.topk_ids, shared_ids], dim=-1)
fused_topk_weights = torch.cat(
[topk_output.topk_weights, shared_weights], dim=-1
)
return StandardTopKOutput(
topk_weights=fused_topk_weights,
topk_ids=fused_topk_ids,
router_logits=topk_output.router_logits,
)
def _forward_shared_experts(self, hidden_states: torch.Tensor):
shared_output = None
if self.shared_expert is not None:
@@ -300,6 +399,10 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
# router_logits: (num_tokens, n_experts)
router_logits, _ = self.gate(hidden_states)
topk_output = self.topk(hidden_states, router_logits)
if self.enable_shared_expert_fusion and TopKOutputChecker.format_is_standard(
topk_output
):
topk_output = self._append_shared_to_topk_output(topk_output, hidden_states)
return self.experts(hidden_states, topk_output)
def forward_normal_dual_stream(
+110 -3
View File
@@ -86,9 +86,11 @@ from sglang.srt.utils import (
LazyValue,
add_prefix,
cpu_has_amx_support,
get_bool_env_var,
is_cpu,
is_cuda,
is_gfx95_supported,
is_hip,
is_npu,
make_layers,
set_weight_attrs,
@@ -100,6 +102,8 @@ _is_cuda = is_cuda()
_is_npu = is_npu()
_is_cpu = is_cpu()
_is_gfx95 = is_gfx95_supported()
_is_hip = is_hip()
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip
_is_amx_available = cpu_has_amx_support()
@@ -528,6 +532,7 @@ class Qwen3_5LinearDecoderLayer(nn.Module):
alt_stream=alt_stream,
prefix=add_prefix("mlp", prefix.replace(".linear_attn", "")),
is_nextn=is_nextn,
support_shared_expert_fusion=True,
)
is_layer_sparse = True
is_previous_layer_sparse = True
@@ -739,6 +744,7 @@ class Qwen3_5AttentionDecoderLayer(nn.Module):
alt_stream=alt_stream,
prefix=add_prefix("mlp", prefix.replace(".self_attn", "")),
is_nextn=is_nextn,
support_shared_expert_fusion=True,
)
is_layer_sparse = True
is_previous_layer_sparse = True
@@ -1494,6 +1500,20 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
self.is_mrope_enabled = "mrope_section" in rope_config
self.deepstack_visual_indexes = self.visual.deepstack_visual_indexes
self.num_fused_shared_experts = 0
if _use_aiter:
self.num_fused_shared_experts = self._get_num_fused_shared_experts()
self.enable_shared_expert_fusion = self.num_fused_shared_experts > 0
def _get_num_fused_shared_experts(self):
if not (
hasattr(self.model, "layers")
and len(self.model.layers) > 0
and hasattr(self.model.layers[0].mlp, "num_fused_shared_experts")
):
return 0
return self.model.layers[0].mlp.num_fused_shared_experts
def get_embed_and_head(self):
embed = self.model.embed_tokens.weight if self.pp_group.is_first_rank else None
@@ -1525,13 +1545,19 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
("in_proj_ba.", "in_proj_a.", 1),
]
num_experts = self.config.num_experts
# Params for weights, fp8 weight scales, fp8 activation scales
# (param_name, weight_name, expert_id, shard_id)
expert_params_mapping = FusedMoE.make_expert_params_mapping(
ckpt_gate_proj_name="gate_proj",
ckpt_down_proj_name="down_proj",
ckpt_up_proj_name="up_proj",
num_experts=self.config.num_experts,
num_experts=(
num_experts
if not self.enable_shared_expert_fusion
else num_experts + self.num_fused_shared_experts
),
)
# Skip loading extra parameters for GPTQ/modelopt models.
@@ -1552,7 +1578,40 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
("experts.w2_weight", "experts.down_proj", 0, "w2"),
]
num_experts = self.config.num_experts
if self.enable_shared_expert_fusion:
"""
When shared experts are fused, we need to map the shared experts to routed experts.
mlp.share_expert.gate_up_proj.weight --> experts.512.gate_up_proj.weight -> experts.w13_weight, expert_id = 512
mlp.share_expert.down_proj.weight --> experts.512.down_proj.weight -> experts.w2_weight, expert_id = 512
"""
fused_expert_params_mapping += [
(
"experts.w13_",
f"experts.{num_experts}.gate_up_proj.",
num_experts,
"w1",
),
(
"experts.w2_",
f"experts.{num_experts}.down_proj.",
num_experts,
"w2",
),
## shared experts may contain gate_proj and up_proj instead of gate_up_proj
(
"experts.w13_",
f"experts.{num_experts}.gate_proj.",
num_experts,
"w1",
),
(
"experts.w13_",
f"experts.{num_experts}.up_proj.",
num_experts,
"w3",
),
]
def load_fused_expert_weights(
name: str,
@@ -1609,6 +1668,14 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
):
continue
if self.enable_shared_expert_fusion:
if "mlp.shared_expert." in name:
# Firstly map mlp.shared_expert.xx_proj to mlp.experts.512.xx_proj
name = name.replace(
"mlp.shared_expert.",
f"mlp.experts.{num_experts}.",
)
for param_name, weight_name, shard_id in stacked_params_mapping:
if name.endswith("experts.gate_up_proj") or name.endswith(
"experts.down_proj"
@@ -1657,7 +1724,10 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
is_expert_weight = True
name_mapped = name.replace(weight_name, param_name)
if is_fused_expert:
# is_fused_expert is True, the checkpoint contains gate_up_proj and down_proj for each expert
if "experts.gate_up_proj" in name:
# experts.gate_up_proj contains all 512 routed experts, excluding shared experts
# split into w1 and w3
loaded_weight = loaded_weight.chunk(2, dim=-2)
load_fused_expert_weights(
name_mapped,
@@ -1673,7 +1743,8 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
"w3",
num_experts,
)
else:
elif "experts.down_proj" in name:
# experts.down_proj contains all 512 routed experts, excluding shared experts
load_fused_expert_weights(
name_mapped,
params_dict,
@@ -1681,6 +1752,42 @@ class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration):
shard_id,
num_experts,
)
elif self.enable_shared_expert_fusion:
# shared experts should be loaded to experts.w13_weight and experts.w2_weight
param = params_dict[name_mapped]
weight_loader = getattr(
param, "weight_loader", default_weight_loader
)
param = params_dict[name_mapped]
if f"{num_experts}.gate_up_proj" in name:
# split into w1 and w3
loaded_weight = loaded_weight.chunk(2, dim=-2)
# load to experts.w13_weight, shard_id = w1, expert_id = 512
weight_loader(
param,
loaded_weight[0],
name_mapped,
"w1",
expert_id,
)
# load to experts.w13_weight, shard_id = w3, expert_id = 512
weight_loader(
param,
loaded_weight[1],
name_mapped,
"w3",
expert_id,
)
else:
# load down_proj to experts.w2_weight, shard_id = w2, expert_id = 512
# Or load gate_proj and up_proj to experts.w13_weight, shard_id = w1/w3, expert_id = 512
weight_loader(
param,
loaded_weight,
name_mapped,
shard_id,
expert_id,
)
else:
# Skip loading extra parameters for GPTQ models.
if (