Optimize Qwen3 Next FP8 MoE on H200 (#26204)
Co-authored-by: BBuf <xiaoyu.zhang@radixark.net>
This commit is contained in:
+162
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
+182
@@ -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
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -128,13 +128,25 @@ _is_hip = is_hip()
|
|||||||
_use_aiter = get_bool_env_var("SGLANG_USE_AITER") and _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(
|
def can_fuse_shared_expert(
|
||||||
config: PretrainedConfig,
|
config: PretrainedConfig,
|
||||||
quant_config: Optional[QuantizationConfig],
|
quant_config: Optional[QuantizationConfig],
|
||||||
) -> bool:
|
) -> 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 (
|
if (
|
||||||
get_global_server_args().disable_shared_experts_fusion is True
|
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
|
# checkpoint), fusing it into the quantized MoE weight tensor requires online
|
||||||
# quantization which is not supported. Disable fusion in this case.
|
# quantization which is not supported. Disable fusion in this case.
|
||||||
if quant_config is not None:
|
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(
|
if any(
|
||||||
"shared_expert" in layer
|
"shared_expert" in layer
|
||||||
and "shared_expert_gate" not in layer
|
and "shared_expert_gate" not in layer
|
||||||
@@ -222,6 +236,7 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
prefix: str = "",
|
prefix: str = "",
|
||||||
is_nextn: bool = False,
|
is_nextn: bool = False,
|
||||||
support_shared_expert_fusion: bool = False,
|
support_shared_expert_fusion: bool = False,
|
||||||
|
enable_cuda_shared_expert_fusion: bool = False,
|
||||||
):
|
):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
self.tp_size = get_tensor_model_parallel_world_size()
|
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}."
|
f"the number of experts {config.num_experts}."
|
||||||
)
|
)
|
||||||
self.num_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
|
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
|
self.enable_shared_expert_fusion = False # default to False
|
||||||
if _use_aiter:
|
if support_shared_expert_fusion and (
|
||||||
# enable shared expert fusion when use aiter
|
_use_aiter or (_is_cuda and enable_cuda_shared_expert_fusion)
|
||||||
|
):
|
||||||
self.enable_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)
|
and can_fuse_shared_expert(config, quant_config)
|
||||||
)
|
)
|
||||||
if self.enable_shared_expert_fusion:
|
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."""
|
"""Return sigmoid(shared_expert_gate) for fused shared expert weights."""
|
||||||
if not self.enable_shared_expert_fusion or self.shared_expert_gate is None:
|
if not self.enable_shared_expert_fusion or self.shared_expert_gate is None:
|
||||||
return None
|
return None
|
||||||
@@ -473,7 +482,11 @@ class Qwen2MoeSparseMoeBlock(nn.Module):
|
|||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
current_stream = torch.cuda.current_stream()
|
current_stream = torch.cuda.current_stream()
|
||||||
self.alt_stream.wait_stream(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 ====
|
# ===== TO BE REFACTORED ====
|
||||||
# Shared-add overlap (SGLANG_OPT_LORA_SHARED_ADD_OVERLAP): hand the add to the LoRA
|
# Shared-add overlap (SGLANG_OPT_LORA_SHARED_ADD_OVERLAP): hand the add to the LoRA
|
||||||
|
|||||||
@@ -6,10 +6,12 @@ import torch
|
|||||||
import triton
|
import triton
|
||||||
from torch import nn
|
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.configs.qwen3_next import Qwen3NextConfig
|
||||||
from sglang.srt.distributed import get_pp_group
|
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_distribution import get_global_expert_distribution_recorder
|
||||||
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
|
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.fla.layernorm_gated import RMSNorm as RMSNormGated
|
||||||
from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader
|
from sglang.srt.layers.attention.mamba.mamba import mamba_v2_sharded_weight_loader
|
||||||
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
from sglang.srt.layers.communicator import LayerCommunicator, LayerScatterModes
|
||||||
@@ -35,6 +37,11 @@ from sglang.srt.layers.vocab_parallel_embedding import (
|
|||||||
ParallelLMHead,
|
ParallelLMHead,
|
||||||
VocabParallelEmbedding,
|
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.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.model_executor.runner import get_is_capture_mode
|
from sglang.srt.model_executor.runner import get_is_capture_mode
|
||||||
from sglang.srt.model_loader.weight_utils import (
|
from sglang.srt.model_loader.weight_utils import (
|
||||||
@@ -57,14 +64,6 @@ from sglang.srt.utils import (
|
|||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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_cuda = is_cuda()
|
||||||
_is_hip = is_hip()
|
_is_hip = is_hip()
|
||||||
_is_npu = is_npu()
|
_is_npu = is_npu()
|
||||||
@@ -546,6 +545,8 @@ class Qwen3HybridLinearDecoderLayer(nn.Module):
|
|||||||
alt_stream=alt_stream,
|
alt_stream=alt_stream,
|
||||||
prefix=add_prefix("mlp", prefix.replace(".linear_attn", "")),
|
prefix=add_prefix("mlp", prefix.replace(".linear_attn", "")),
|
||||||
is_nextn=is_nextn,
|
is_nextn=is_nextn,
|
||||||
|
support_shared_expert_fusion=True,
|
||||||
|
enable_cuda_shared_expert_fusion=True,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.mlp = Qwen2MoeMLP(
|
self.mlp = Qwen2MoeMLP(
|
||||||
@@ -713,6 +714,8 @@ class Qwen3HybridAttentionDecoderLayer(nn.Module):
|
|||||||
alt_stream=alt_stream,
|
alt_stream=alt_stream,
|
||||||
prefix=add_prefix("mlp", prefix.replace(".self_attn", "")),
|
prefix=add_prefix("mlp", prefix.replace(".self_attn", "")),
|
||||||
is_nextn=is_nextn,
|
is_nextn=is_nextn,
|
||||||
|
support_shared_expert_fusion=True,
|
||||||
|
enable_cuda_shared_expert_fusion=True,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
self.mlp = Qwen2MoeMLP(
|
self.mlp = Qwen2MoeMLP(
|
||||||
@@ -1032,6 +1035,15 @@ class Qwen3NextForCausalLM(nn.Module):
|
|||||||
# For EAGLE3 support
|
# For EAGLE3 support
|
||||||
self.capture_aux_hidden_states = False
|
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(
|
self._routed_experts_weights_of_layer = LazyValue(
|
||||||
lambda: {
|
lambda: {
|
||||||
layer_id: layer.mlp.get_moe_weights()
|
layer_id: layer.mlp.get_moe_weights()
|
||||||
@@ -1044,6 +1056,14 @@ class Qwen3NextForCausalLM(nn.Module):
|
|||||||
def routed_experts_weights_of_layer(self):
|
def routed_experts_weights_of_layer(self):
|
||||||
return self._routed_experts_weights_of_layer.value
|
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()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
self,
|
self,
|
||||||
@@ -1117,7 +1137,11 @@ class Qwen3NextForCausalLM(nn.Module):
|
|||||||
ckpt_gate_proj_name="gate_proj",
|
ckpt_gate_proj_name="gate_proj",
|
||||||
ckpt_down_proj_name="down_proj",
|
ckpt_down_proj_name="down_proj",
|
||||||
ckpt_up_proj_name="up_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())
|
params_dict = dict(self.named_parameters())
|
||||||
@@ -1147,6 +1171,12 @@ class Qwen3NextForCausalLM(nn.Module):
|
|||||||
if ".self_attn." in name:
|
if ".self_attn." in name:
|
||||||
name = name.replace(".self_attn", "")
|
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:
|
# Remap modelopt FP8 KV cache scale names:
|
||||||
# checkpoint: k_proj.k_scale / v_proj.v_scale
|
# checkpoint: k_proj.k_scale / v_proj.v_scale
|
||||||
# model: attn.k_scale / attn.v_scale
|
# model: attn.k_scale / attn.v_scale
|
||||||
|
|||||||
@@ -59,7 +59,6 @@ class Qwen3NextForCausalLMMTP(Qwen3NextForCausalLM):
|
|||||||
self.quant_config = quant_config
|
self.quant_config = quant_config
|
||||||
# if not set, model load will be broken in Qwen3NextForCausalLM load_weights()
|
# if not set, model load will be broken in Qwen3NextForCausalLM load_weights()
|
||||||
self.pp_group = get_pp_group()
|
self.pp_group = get_pp_group()
|
||||||
# self.determine_num_fused_shared_experts("Qwen3NextForCausalLMMTP")
|
|
||||||
|
|
||||||
# currently based on the provided ckpt, we:
|
# 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
|
# (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,
|
use_attn_tp_group=get_global_server_args().enable_dp_lm_head,
|
||||||
)
|
)
|
||||||
self.logits_processor = LogitsProcessor(config)
|
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()
|
@torch.no_grad()
|
||||||
def forward(
|
def forward(
|
||||||
|
|||||||
Reference in New Issue
Block a user