From 06e0df5899aa1ca4d5526250a33e6ef22b418c4a Mon Sep 17 00:00:00 2001 From: Xiaoyu Zhang <1182563586@qq.com> Date: Thu, 11 Jun 2026 22:18:43 +0800 Subject: [PATCH] Optimize Qwen3 Next FP8 MoE on H200 (#26204) Co-authored-by: BBuf --- ...dtype=fp8_w8a8,block_shape=[128, 128].json | 162 ++++++++++++++++ ...=fp8_w8a8,block_shape=[128, 128]_down.json | 182 ++++++++++++++++++ python/sglang/srt/models/qwen2_moe.py | 49 +++-- python/sglang/srt/models/qwen3_next.py | 48 ++++- python/sglang/srt/models/qwen3_next_mtp.py | 14 +- 5 files changed, 427 insertions(+), 28 deletions(-) create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128].json create mode 100644 python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128]_down.json diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128].json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128].json new file mode 100644 index 000000000..e4e9f48a9 --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128].json @@ -0,0 +1,162 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4 + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "80": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "256": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "512": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "1024": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "1536": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "2048": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "3072": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "4096": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + }, + "8192": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3 + } +} diff --git a/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128]_down.json b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128]_down.json new file mode 100644 index 000000000..cd97b36fd --- /dev/null +++ b/python/sglang/srt/layers/moe/moe_runner/triton_utils/configs/triton_3_6_0/E=513,N=512,device_name=NVIDIA_H200,dtype=fp8_w8a8,block_shape=[128, 128]_down.json @@ -0,0 +1,182 @@ +{ + "1": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "USE_TMA": true + }, + "2": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "USE_TMA": true + }, + "4": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "USE_TMA": true + }, + "8": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "USE_TMA": true + }, + "16": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 64, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 4, + "USE_TMA": true + }, + "24": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "32": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "48": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "64": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "80": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "96": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "128": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "256": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "512": { + "BLOCK_SIZE_M": 16, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "1024": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "1536": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "2048": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "3072": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "4096": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + }, + "8192": { + "BLOCK_SIZE_M": 64, + "BLOCK_SIZE_N": 128, + "BLOCK_SIZE_K": 128, + "GROUP_SIZE_M": 32, + "num_warps": 4, + "num_stages": 3, + "USE_TMA": true + } +} diff --git a/python/sglang/srt/models/qwen2_moe.py b/python/sglang/srt/models/qwen2_moe.py index a2147e750..6fa7bb549 100644 --- a/python/sglang/srt/models/qwen2_moe.py +++ b/python/sglang/srt/models/qwen2_moe.py @@ -128,13 +128,25 @@ _is_hip = is_hip() _use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _is_hip +def get_num_shared_experts(config: PretrainedConfig) -> int: + n_shared_experts = getattr(config, "n_shared_experts", None) + if n_shared_experts is not None: + return n_shared_experts + if ( + hasattr(config, "shared_expert_intermediate_size") + and config.shared_expert_intermediate_size > 0 + ): + return 1 + return 0 + + def can_fuse_shared_expert( config: PretrainedConfig, quant_config: Optional[QuantizationConfig], ) -> bool: - """Whether the shared expert may be fused as an extra MoE expert (Qwen3.5 + Aiter). + """Whether the shared expert may be fused as an extra MoE expert. - Caller must still gate on support_shared_expert_fusion and _use_aiter. + Caller must still gate on the model/backend support flag. """ if ( get_global_server_args().disable_shared_experts_fusion is True @@ -148,7 +160,9 @@ def can_fuse_shared_expert( # checkpoint), fusing it into the quantized MoE weight tensor requires online # quantization which is not supported. Disable fusion in this case. if quant_config is not None: - exclude_layers = getattr(quant_config, "exclude_layers", []) + exclude_layers = getattr(quant_config, "exclude_layers", None) + if exclude_layers is None: + exclude_layers = getattr(quant_config, "ignored_layers", []) if any( "shared_expert" in layer and "shared_expert_gate" not in layer @@ -222,6 +236,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module): prefix: str = "", is_nextn: bool = False, support_shared_expert_fusion: bool = False, + enable_cuda_shared_expert_fusion: bool = False, ): super().__init__() self.tp_size = get_tensor_model_parallel_world_size() @@ -233,23 +248,15 @@ class Qwen2MoeSparseMoeBlock(nn.Module): f"the number of experts {config.num_experts}." ) self.num_experts = config.num_experts - self.num_shared_experts = 0 + self.num_shared_experts = get_num_shared_experts(config) 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 + if support_shared_expert_fusion and ( + _use_aiter or (_is_cuda and enable_cuda_shared_expert_fusion) + ): self.enable_shared_expert_fusion = ( - support_shared_expert_fusion + self.num_shared_experts > 0 and can_fuse_shared_expert(config, quant_config) ) if self.enable_shared_expert_fusion: @@ -345,7 +352,9 @@ class Qwen2MoeSparseMoeBlock(nn.Module): ) ] - def _get_shared_expert_weights(self, hidden_states: torch.Tensor) -> torch.Tensor: + def _get_shared_expert_weights( + self, hidden_states: torch.Tensor + ) -> Optional[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 @@ -473,7 +482,11 @@ class Qwen2MoeSparseMoeBlock(nn.Module): ) -> torch.Tensor: current_stream = torch.cuda.current_stream() self.alt_stream.wait_stream(current_stream) - shared_output = self._forward_shared_experts(hidden_states.clone()) + shared_output = ( + self._forward_shared_experts(hidden_states.clone()) + if self.shared_expert is not None + else None + ) # ===== TO BE REFACTORED ==== # Shared-add overlap (SGLANG_OPT_LORA_SHARED_ADD_OVERLAP): hand the add to the LoRA diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index 3d1231e4a..8139bc8ac 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -6,10 +6,12 @@ import torch import triton from torch import nn +from sglang.jit_kernel.triton.gdn_fused_proj import fused_qkvzba_split_reshape_cat from sglang.srt.configs.qwen3_next import Qwen3NextConfig from sglang.srt.distributed import get_pp_group from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation +from sglang.srt.layers.attention.fla.fused_norm_gate import FusedRMSNormGated from sglang.srt.layers.attention.fla.layernorm_gated import RMSNorm as RMSNormGated from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes @@ -35,6 +37,11 @@ from sglang.srt.layers.vocab_parallel_embedding import ( ParallelLMHead, VocabParallelEmbedding, ) +from sglang.srt.model_executor.cuda_graph_config import ( + Backend, + Phase, + check_cuda_graph_backend, +) from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.model_executor.runner import get_is_capture_mode from sglang.srt.model_loader.weight_utils import ( @@ -57,14 +64,6 @@ from sglang.srt.utils import ( logger = logging.getLogger(__name__) -from sglang.jit_kernel.triton.gdn_fused_proj import fused_qkvzba_split_reshape_cat -from sglang.srt.layers.attention.fla.fused_norm_gate import FusedRMSNormGated -from sglang.srt.model_executor.cuda_graph_config import ( - Backend, - Phase, - check_cuda_graph_backend, -) - _is_cuda = is_cuda() _is_hip = is_hip() _is_npu = is_npu() @@ -546,6 +545,8 @@ class Qwen3HybridLinearDecoderLayer(nn.Module): alt_stream=alt_stream, prefix=add_prefix("mlp", prefix.replace(".linear_attn", "")), is_nextn=is_nextn, + support_shared_expert_fusion=True, + enable_cuda_shared_expert_fusion=True, ) else: self.mlp = Qwen2MoeMLP( @@ -713,6 +714,8 @@ class Qwen3HybridAttentionDecoderLayer(nn.Module): alt_stream=alt_stream, prefix=add_prefix("mlp", prefix.replace(".self_attn", "")), is_nextn=is_nextn, + support_shared_expert_fusion=True, + enable_cuda_shared_expert_fusion=True, ) else: self.mlp = Qwen2MoeMLP( @@ -1032,6 +1035,15 @@ class Qwen3NextForCausalLM(nn.Module): # For EAGLE3 support self.capture_aux_hidden_states = False + self.num_fused_shared_experts = self._get_num_fused_shared_experts() + if self.num_fused_shared_experts > 1: + raise ValueError( + "Qwen3-Next shared expert fusion currently supports exactly one " + "shared expert because checkpoint weight remapping maps it into " + "a single fused MoE expert slot." + ) + self.enable_shared_expert_fusion = self.num_fused_shared_experts > 0 + self._routed_experts_weights_of_layer = LazyValue( lambda: { layer_id: layer.mlp.get_moe_weights() @@ -1044,6 +1056,14 @@ class Qwen3NextForCausalLM(nn.Module): def routed_experts_weights_of_layer(self): return self._routed_experts_weights_of_layer.value + def _get_num_fused_shared_experts(self) -> int: + if not hasattr(self.model, "layers"): + return 0 + for layer in self.model.layers: + if isinstance(layer.mlp, Qwen2MoeSparseMoeBlock): + return layer.mlp.num_fused_shared_experts + return 0 + @torch.no_grad() def forward( self, @@ -1117,7 +1137,11 @@ class Qwen3NextForCausalLM(nn.Module): 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=( + self.config.num_experts + if not self.enable_shared_expert_fusion + else self.config.num_experts + self.num_fused_shared_experts + ), ) params_dict = dict(self.named_parameters()) @@ -1147,6 +1171,12 @@ class Qwen3NextForCausalLM(nn.Module): if ".self_attn." in name: name = name.replace(".self_attn", "") + if self.enable_shared_expert_fusion and "mlp.shared_expert." in name: + name = name.replace( + "mlp.shared_expert.", + f"mlp.experts.{self.config.num_experts}.", + ) + # Remap modelopt FP8 KV cache scale names: # checkpoint: k_proj.k_scale / v_proj.v_scale # model: attn.k_scale / attn.v_scale diff --git a/python/sglang/srt/models/qwen3_next_mtp.py b/python/sglang/srt/models/qwen3_next_mtp.py index 3bade2141..c9e8064f5 100644 --- a/python/sglang/srt/models/qwen3_next_mtp.py +++ b/python/sglang/srt/models/qwen3_next_mtp.py @@ -59,7 +59,6 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): self.quant_config = quant_config # if not set, model load will be broken in Qwen3NextForCausalLM load_weights() self.pp_group = get_pp_group() - # self.determine_num_fused_shared_experts("Qwen3NextForCausalLMMTP") # currently based on the provided ckpt, we: # (1) do not use_dedicated_mtp_embeddings provided in ckpt since not provided and directly use the target model embeddings @@ -86,6 +85,19 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM): use_attn_tp_group=get_global_server_args().enable_dp_lm_head, ) self.logits_processor = LogitsProcessor(config) + # Mirror Qwen3NextForCausalLM.__init__'s shared-expert fusion setup so + # the inherited load_weights() can find the attribute on the MTP path. + # We compute it from the actual MTP MoE layer (1 layer with is_nextn=True), + # not hardcode it — when the layer's MoE pre-fuses the shared expert, + # load_weights must remap mlp.shared_expert.* into the fused slot. + self.num_fused_shared_experts = self._get_num_fused_shared_experts() + if self.num_fused_shared_experts > 1: + raise ValueError( + "Qwen3-Next MTP shared expert fusion currently supports exactly one " + "shared expert because checkpoint weight remapping maps it into " + "a single fused MoE expert slot." + ) + self.enable_shared_expert_fusion = self.num_fused_shared_experts > 0 @torch.no_grad() def forward(