diff --git a/python/sglang/srt/layers/moe/utils.py b/python/sglang/srt/layers/moe/utils.py index 414140fda..4249fad5a 100644 --- a/python/sglang/srt/layers/moe/utils.py +++ b/python/sglang/srt/layers/moe/utils.py @@ -249,6 +249,18 @@ def get_tbo_token_distribution_threshold() -> float: return TBO_TOKEN_DISTRIBUTION_THRESHOLD +def filter_moe_weight_param_global_expert(name, x, num_local_experts): + """ + Filter out for MoE expert parameters that requires global expert. + """ + return ( + not getattr(x, "_sglang_require_global_experts", False) + and not name.endswith("_blockscale_swizzled") + and x.data.ndim > 0 + and x.data.shape[0] == num_local_experts + ) + + def should_use_flashinfer_cutlass_moe_fp4_allgather(): """ Perform FP4 quantize before all-gather for flashinfer cutlass moe to reduce communication cost for high-throughput serving. diff --git a/python/sglang/srt/models/bailing_moe.py b/python/sglang/srt/models/bailing_moe.py index c353a0cd9..3a3a20883 100644 --- a/python/sglang/srt/models/bailing_moe.py +++ b/python/sglang/srt/models/bailing_moe.py @@ -63,6 +63,7 @@ from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.token_dispatcher import DeepEPDispatcher from sglang.srt.layers.moe.topk import TopK +from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.rotary_embedding import get_rope @@ -324,6 +325,9 @@ class BailingMoESparseMoeBlock(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] def _forward_shared_experts(self, hidden_states: torch.Tensor): diff --git a/python/sglang/srt/models/deepseek_v2.py b/python/sglang/srt/models/deepseek_v2.py index 748912af6..75746bfe7 100644 --- a/python/sglang/srt/models/deepseek_v2.py +++ b/python/sglang/srt/models/deepseek_v2.py @@ -103,7 +103,10 @@ from sglang.srt.layers.moe.token_dispatcher.base import ( DispatchOutput, ) from sglang.srt.layers.moe.topk import TopK, TopKOutputFormat -from sglang.srt.layers.moe.utils import RoutingMethodType +from sglang.srt.layers.moe.utils import ( + RoutingMethodType, + filter_moe_weight_param_global_expert, +) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8 import Fp8Config from sglang.srt.layers.quantization.fp8_kernel import ( @@ -587,6 +590,9 @@ class DeepseekV2MoE(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] def forward( diff --git a/python/sglang/srt/models/glm4_moe.py b/python/sglang/srt/models/glm4_moe.py index be71d0d28..5c09bc681 100644 --- a/python/sglang/srt/models/glm4_moe.py +++ b/python/sglang/srt/models/glm4_moe.py @@ -63,6 +63,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.layer import FusedMoE from sglang.srt.layers.moe.topk import TopK +from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.layers.radix_attention import RadixAttention @@ -438,6 +439,9 @@ class Glm4MoeSparseMoeBlock(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] def forward( diff --git a/python/sglang/srt/models/gpt_oss.py b/python/sglang/srt/models/gpt_oss.py index 0d79f898b..608b23cd1 100644 --- a/python/sglang/srt/models/gpt_oss.py +++ b/python/sglang/srt/models/gpt_oss.py @@ -54,6 +54,7 @@ from sglang.srt.layers.moe import get_moe_a2a_backend from sglang.srt.layers.moe.ep_moe.layer import get_moe_impl_class from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.topk import TopK +from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8_utils import dequant_mxfp4 from sglang.srt.layers.radix_attention import RadixAttention @@ -169,6 +170,9 @@ class GptOssSparseMoeBlock(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] def forward_normal( diff --git a/python/sglang/srt/models/longcat_flash.py b/python/sglang/srt/models/longcat_flash.py index 275ee9019..7ef5b4ae1 100644 --- a/python/sglang/srt/models/longcat_flash.py +++ b/python/sglang/srt/models/longcat_flash.py @@ -63,6 +63,7 @@ from sglang.srt.layers.moe.ep_moe.kernels import zero_experts_compute_triton from sglang.srt.layers.moe.ep_moe.layer import DeepEPMoE, get_moe_impl_class from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE from sglang.srt.layers.moe.topk import StandardTopKOutput, TopK +from sglang.srt.layers.moe.utils import filter_moe_weight_param_global_expert from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.quantization.fp8_kernel import is_fp8_fnuz from sglang.srt.layers.quantization.fp8_utils import ( @@ -295,6 +296,9 @@ class LongcatFlashMoE(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index 3ad9f6736..1e4b53a7d 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -58,7 +58,10 @@ from sglang.srt.layers.moe import get_moe_a2a_backend 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.utils import RoutingMethodType +from sglang.srt.layers.moe.utils import ( + RoutingMethodType, + filter_moe_weight_param_global_expert, +) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.rotary_embedding import get_rope @@ -223,6 +226,9 @@ class Qwen2MoeSparseMoeBlock(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] def _forward_shared_experts(self, hidden_states: torch.Tensor): diff --git a/python/sglang/srt/models/qwen3_moe.py b/python/sglang/srt/models/qwen3_moe.py index 3648b2cf9..d0469a62e 100644 --- a/python/sglang/srt/models/qwen3_moe.py +++ b/python/sglang/srt/models/qwen3_moe.py @@ -51,7 +51,10 @@ 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.layer import FusedMoE from sglang.srt.layers.moe.topk import TopK -from sglang.srt.layers.moe.utils import RoutingMethodType +from sglang.srt.layers.moe.utils import ( + RoutingMethodType, + filter_moe_weight_param_global_expert, +) from sglang.srt.layers.quantization.base_config import QuantizationConfig from sglang.srt.layers.radix_attention import RadixAttention from sglang.srt.layers.rotary_embedding import MRotaryEmbedding, get_rope @@ -281,6 +284,9 @@ class Qwen3MoeSparseMoeBlock(nn.Module): x.data for name, x in self.experts.named_parameters() if name not in ["correction_bias"] + and filter_moe_weight_param_global_expert( + name, x, self.experts.num_local_experts + ) ] def forward_normal(