diff --git a/benchmark/kernels/fused_moe_triton/common_utils.py b/benchmark/kernels/fused_moe_triton/common_utils.py index 5f2d9aa8a..a1fe4e0a3 100644 --- a/benchmark/kernels/fused_moe_triton/common_utils.py +++ b/benchmark/kernels/fused_moe_triton/common_utils.py @@ -66,6 +66,7 @@ def get_model_config( "Qwen3MoeForCausalLM", "Qwen3NextForCausalLM", "Qwen3VLMoeForConditionalGeneration", + "Qwen3_5MoeForConditionalGeneration", ]: E = config.num_experts // ep_size topk = config.num_experts_per_tok diff --git a/python/sglang/srt/configs/__init__.py b/python/sglang/srt/configs/__init__.py index 865131927..f84fbb9ce 100644 --- a/python/sglang/srt/configs/__init__.py +++ b/python/sglang/srt/configs/__init__.py @@ -18,6 +18,7 @@ from sglang.srt.configs.longcat_flash import LongcatFlashConfig from sglang.srt.configs.nano_nemotron_vl import NemotronH_Nano_VL_V2_Config from sglang.srt.configs.nemotron_h import NemotronHConfig from sglang.srt.configs.olmo3 import Olmo3Config +from sglang.srt.configs.qwen3_5 import Qwen3_5Config, Qwen3_5MoeConfig from sglang.srt.configs.qwen3_next import Qwen3NextConfig from sglang.srt.configs.step3_vl import ( Step3TextConfig, @@ -43,6 +44,8 @@ __all__ = [ "KimiLinearConfig", "KimiK25Config", "Qwen3NextConfig", + "Qwen3_5Config", + "Qwen3_5MoeConfig", "DotsVLMConfig", "DotsOCRConfig", "FalconH1Config", diff --git a/python/sglang/srt/configs/model_config.py b/python/sglang/srt/configs/model_config.py index e76c15705..bdd091e66 100644 --- a/python/sglang/srt/configs/model_config.py +++ b/python/sglang/srt/configs/model_config.py @@ -319,6 +319,13 @@ class ModelConfig: self.hf_config.architectures[0] = "Qwen3NextForCausalLMMTP" self.hf_config.num_nextn_predict_layers = 1 + if is_draft_model and self.hf_config.architectures[0] in [ + "Qwen3_5ForConditionalGeneration", + "Qwen3_5MoeForConditionalGeneration", + ]: + self.hf_config.architectures[0] = "Qwen3_5ForCausalLMMTP" + self.hf_config.num_nextn_predict_layers = 1 + if is_draft_model and self.hf_config.architectures[0] == "ExaoneMoEForCausalLM": self.hf_config.architectures[0] = "ExaoneMoEForCausalLMMTP" self.hf_config.num_nextn_predict_layers = 1 @@ -1193,6 +1200,8 @@ multimodal_model_archs = [ "Qwen2_5_VLForConditionalGeneration", "Qwen3VLForConditionalGeneration", "Qwen3VLMoeForConditionalGeneration", + "Qwen3_5ForConditionalGeneration", + "Qwen3_5MoeForConditionalGeneration", "Qwen3OmniMoeForConditionalGeneration", "KimiVLForConditionalGeneration", "InternVLChatModel", diff --git a/python/sglang/srt/configs/qwen3_5.py b/python/sglang/srt/configs/qwen3_5.py new file mode 100644 index 000000000..cce393161 --- /dev/null +++ b/python/sglang/srt/configs/qwen3_5.py @@ -0,0 +1,113 @@ +from transformers import PretrainedConfig + +from sglang.srt.configs.qwen3_next import Qwen3NextConfig +from sglang.srt.configs.qwen3_vl import Qwen3VLVisionConfig + + +class Qwen3_5VisionConfig(Qwen3VLVisionConfig): + model_type = "qwen3_5" + base_config_key = "vision_config" + + +class Qwen3_5TextConfig(Qwen3NextConfig): + model_type = "qwen3_5_text" + base_config_key = "text_config" + + def __init__( + self, + **kwargs, + ): + super().__init__(**kwargs) + if self.rope_scaling is None: + self.rope_scaling = {} + + +class Qwen3_5Config(PretrainedConfig): + r""" + This is the configuration class to store the configuration of a [`Qwen3_5Model`]. It is used to instantiate a + Qwen3.5 model according to the specified arguments, defining the model architecture. Instantiating a configuration + with the defaults will yield a similar configuration to that of + Qwen3.5. + + Configuration objects inherit from [`PretrainedConfig`] and can be used to control the model outputs. Read the + documentation from [`PretrainedConfig`] for more information. + + + Args: + text_config (`Union[PreTrainedConfig, dict]`, *optional*, defaults to `Qwen3_5TextConfig`): + The config object or dictionary of the text backbone. + vision_config (`Union[PreTrainedConfig, dict]`, *optional*, defaults to `Qwen3_5VisionConfig`): + The config object or dictionary of the vision backbone. + image_token_id (`int`, *optional*, defaults to 151655): + The image token index to encode the image prompt. + video_token_id (`int`, *optional*, defaults to 151656): + The video token index to encode the image prompt. + vision_start_token_id (`int`, *optional*, defaults to 151652): + The start token index to encode the image prompt. + vision_end_token_id (`int`, *optional*, defaults to 151653): + The end token index to encode the image prompt. + tie_word_embeddings (`bool`, *optional*, defaults to `False`): + Whether to tie the word embeddings. + + ```python + >>> from transformers import Qwen3_5ForConditionalGeneration, Qwen3_5Config + + >>> # Initializing a Qwen3.5 style configuration + >>> configuration = Qwen3_5Config() + + >>> # Initializing a model from the Qwen3.5 style configuration + >>> model = Qwen3_5ForConditionalGeneration(configuration) + + >>> # Accessing the model configuration + >>> configuration = model.config + ```""" + + model_type = "qwen3_5" + sub_configs = { + "vision_config": Qwen3_5VisionConfig, + "text_config": Qwen3_5TextConfig, + } + keys_to_ignore_at_inference = ["past_key_values"] + + def __init__( + self, + text_config=None, + vision_config=None, + image_token_id=151655, + video_token_id=151656, + vision_start_token_id=151652, + vision_end_token_id=151653, + tie_word_embeddings=False, + **kwargs, + ): + if isinstance(vision_config, dict): + self.vision_config = self.sub_configs["vision_config"](**vision_config) + elif vision_config is None: + self.vision_config = self.sub_configs["vision_config"]() + + if isinstance(text_config, dict): + self.text_config = self.sub_configs["text_config"](**text_config) + elif text_config is None: + self.text_config = self.sub_configs["text_config"]() + + self.image_token_id = image_token_id + self.video_token_id = video_token_id + self.vision_start_token_id = vision_start_token_id + self.vision_end_token_id = vision_end_token_id + super().__init__(**kwargs, tie_word_embeddings=tie_word_embeddings) + + +class Qwen3_5MoeVisionConfig(Qwen3_5VisionConfig): + model_type = "qwen3_5_moe" + + +class Qwen3_5MoeTextConfig(Qwen3_5TextConfig): + model_type = "qwen3_5_moe_text" + + +class Qwen3_5MoeConfig(Qwen3_5Config): + model_type = "qwen3_5_moe" + sub_configs = { + "vision_config": Qwen3_5MoeVisionConfig, + "text_config": Qwen3_5MoeTextConfig, + } diff --git a/python/sglang/srt/layers/logits_processor.py b/python/sglang/srt/layers/logits_processor.py index d0d7c9344..8664bbb17 100644 --- a/python/sglang/srt/layers/logits_processor.py +++ b/python/sglang/srt/layers/logits_processor.py @@ -104,6 +104,8 @@ class LogitsProcessorOutput: ## Part 5: Customized Info customized_info: Optional[Dict[str, List[Any]]] = None + mm_input_embeds: Optional[torch.Tensor] = None + @dataclasses.dataclass class LogitsMetadata: @@ -146,6 +148,8 @@ class LogitsMetadata: # Whether this batch is prefill-only (no token generation needed) is_prefill_only: bool = False + mm_input_embeds: Optional[torch.Tensor] = None + @classmethod def from_forward_batch(cls, forward_batch: ForwardBatch): if ( @@ -196,6 +200,7 @@ class LogitsMetadata: global_num_tokens_for_logprob_cpu=forward_batch.global_num_tokens_for_logprob_cpu, global_num_tokens_for_logprob_gpu=forward_batch.global_num_tokens_for_logprob_gpu, dp_padding_mode=DpPaddingMode.SUM_LEN, + mm_input_embeds=forward_batch.mm_input_embeds, ) def compute_dp_attention_metadata(self): @@ -341,6 +346,7 @@ class LogitsProcessor(nn.Module): return LogitsProcessorOutput( next_token_logits=sampled_logits, hidden_states=hidden_states_to_store, + mm_input_embeds=logits_metadata.mm_input_embeds, ) # Start to process input logprobs @@ -386,6 +392,7 @@ class LogitsProcessor(nn.Module): input_top_logprobs_idx=logprobs_result.input_top_logprobs_idx, input_token_ids_logprobs_val=logprobs_result.input_token_ids_logprobs_val, input_token_ids_logprobs_idx=logprobs_result.input_token_ids_logprobs_idx, + mm_input_embeds=logits_metadata.mm_input_embeds, ) def _get_pruned_states( @@ -1067,6 +1074,10 @@ class LogitsProcessor(nn.Module): input_top_logprobs_idx=input_top_logprobs_idx, input_token_ids_logprobs_val=input_token_ids_logprobs_val, input_token_ids_logprobs_idx=input_token_ids_logprobs_idx, + # FIXME: These fields are not logits-related but are passed through here as a + # workaround since ForwardBatch is local to forward_batch_generation(). + # They should be moved to GenerationBatchResult to keep this class clean. + mm_input_embeds=logits_metadata.mm_input_embeds, ) diff --git a/python/sglang/srt/layers/rotary_embedding.py b/python/sglang/srt/layers/rotary_embedding.py index 6db5f0987..3980037a6 100644 --- a/python/sglang/srt/layers/rotary_embedding.py +++ b/python/sglang/srt/layers/rotary_embedding.py @@ -1825,7 +1825,9 @@ class MRotaryEmbedding(RotaryEmbedding): **kwargs, ) if ( - model_type.startswith("qwen3_vl") or model_type.startswith("qwen3_vl_moe") + model_type.startswith("qwen3_vl") + or model_type.startswith("qwen3_vl_moe") + or model_type.startswith("qwen3_5") ) and video_grid_thw is not None: video_grid_thw = torch.repeat_interleave( video_grid_thw, video_grid_thw[:, 0], dim=0 @@ -1925,6 +1927,8 @@ class MRotaryEmbedding(RotaryEmbedding): "qwen2_vl", "qwen3_vl", "qwen3_vl_moe", + "qwen3_5", + "qwen3_5_moe", ): t_index = ( torch.arange(llm_grid_t, device=position_ids.device) diff --git a/python/sglang/srt/managers/mm_utils.py b/python/sglang/srt/managers/mm_utils.py index fdad40e34..672fe2384 100644 --- a/python/sglang/srt/managers/mm_utils.py +++ b/python/sglang/srt/managers/mm_utils.py @@ -1121,6 +1121,7 @@ def general_mm_embed_routine( if isinstance(feature, torch.Tensor) and feature.is_cuda: mm_item.feature = feature.to("cpu", non_blocking=True) forward_batch.mm_inputs = None + forward_batch.mm_input_embeds = input_embeds else: input_embeds = embed_tokens(input_ids) # Copy to pre-allocated buffer if available (for CUDA graph address stability) diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index f96651f40..086680de9 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -350,6 +350,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # Speculative decoding spec_info: Optional[SpecInput] = None spec_algorithm: SpeculativeAlgorithm = None + mm_input_embeds: Optional[torch.Tensor] = None capture_hidden_mode: CaptureHiddenMode = None # For padding diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 3fe264646..ecbdc987b 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -38,6 +38,8 @@ from sglang.srt.configs import ( Lfm2Config, NemotronH_Nano_VL_V2_Config, NemotronHConfig, + Qwen3_5Config, + Qwen3_5MoeConfig, Qwen3NextConfig, ) from sglang.srt.configs.device_config import DeviceConfig @@ -1548,8 +1550,15 @@ class ModelRunner(ModelRunnerKVCacheMixin): @property def hybrid_gdn_config(self): - config = self.model_config.hf_config - if isinstance(config, Qwen3NextConfig | JetNemotronConfig | JetVLMConfig): + config = self.model_config.hf_config.get_text_config() + if isinstance( + config, + Qwen3NextConfig + | Qwen3_5Config + | Qwen3_5MoeConfig + | JetNemotronConfig + | JetVLMConfig, + ): return config return None @@ -2532,7 +2541,9 @@ class ModelRunner(ModelRunnerKVCacheMixin): def model_is_mrope(self) -> bool: """Detect if the model has "mrope" rope_scaling type. mrope requires keep "rope_deltas" between prompt and decoding phases.""" - rope_scaling = getattr(self.model_config.hf_text_config, "rope_scaling", {}) + rope_scaling = getattr( + self.model_config.hf_text_config, "rope_parameters", None + ) or getattr(self.model_config.hf_text_config, "rope_scaling", {}) if rope_scaling is None: return False is_mrope_enabled = "mrope_section" in rope_scaling diff --git a/python/sglang/srt/models/qwen3_5.py b/python/sglang/srt/models/qwen3_5.py new file mode 100644 index 000000000..c6c71247c --- /dev/null +++ b/python/sglang/srt/models/qwen3_5.py @@ -0,0 +1,1310 @@ +# Copyright 2025 Qwen Team +# Copyright 2025 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== +"""Inference-only Qwen3.5 model and Qwen3.5 MoE model compatible with HuggingFace weights.""" +import logging +from functools import lru_cache +from typing import Iterable, Optional, Set, Tuple, Union + +import torch +import torch.nn as nn +from einops import rearrange + +# Model Executor +from sglang.srt.compilation.piecewise_context_manager import get_forward_context + +# Configs +from sglang.srt.configs.qwen3_5 import ( + Qwen3_5Config, + Qwen3_5MoeConfig, + Qwen3_5TextConfig, +) + +# Distributed +from sglang.srt.distributed import get_pp_group +from sglang.srt.eplb.expert_distribution import get_global_expert_distribution_recorder + +# Layers - Attention +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 +from sglang.srt.layers.dp_attention import ( + get_attention_tp_rank, + get_attention_tp_size, + is_dp_attention_enabled, +) + +# Layers - Others +from sglang.srt.layers.layernorm import GemmaRMSNorm + +# Layers - Linear +from sglang.srt.layers.linear import ( + ColumnParallelLinear, + MergedColumnParallelLinear, + QKVParallelLinear, + RowParallelLinear, +) +from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE +from sglang.srt.layers.quantization.base_config import QuantizationConfig +from sglang.srt.layers.radix_attention import RadixAttention +from sglang.srt.layers.radix_linear_attention import RadixLinearAttention +from sglang.srt.layers.rotary_embedding import get_rope +from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding +from sglang.srt.model_executor.cuda_graph_runner import get_is_capture_mode +from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors +from sglang.srt.model_loader.weight_utils import ( + default_weight_loader, + sharded_weight_loader, +) +from sglang.srt.models.qwen2_moe import Qwen2MoeMLP, Qwen2MoeSparseMoeBlock + +# Models +from sglang.srt.models.qwen3_next import gdn_with_output +from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration + +# Utils +from sglang.srt.utils import add_prefix, is_cuda, is_npu, make_layers, set_weight_attrs +from sglang.srt.utils.hf_transformers_utils import get_processor + +logger = logging.getLogger(__name__) +_is_cuda = is_cuda() +_is_npu = is_npu() + +cached_get_processor = lru_cache(get_processor) + + +class Qwen3_5GatedDeltaNet(nn.Module): + def __init__( + self, + config: Qwen3_5TextConfig, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + alt_stream: Optional[torch.cuda.Stream] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.attn_tp_rank = get_attention_tp_rank() + self.attn_tp_size = get_attention_tp_size() + self.hidden_size = config.hidden_size + self.num_v_heads = config.linear_num_value_heads + self.num_k_heads = config.linear_num_key_heads + self.head_k_dim = config.linear_key_head_dim + self.head_v_dim = config.linear_value_head_dim + self.key_dim = self.head_k_dim * self.num_k_heads + self.value_dim = self.head_v_dim * self.num_v_heads + self.alt_stream = alt_stream + + self.conv_kernel_size = config.linear_conv_kernel_dim + self.layer_id = layer_id + self.activation = config.hidden_act + self.layer_norm_epsilon = config.rms_norm_eps + + # Conv1d layer + self.conv_dim = self.key_dim * 2 + self.value_dim + self.conv1d = ColumnParallelLinear( + input_size=self.conv_kernel_size, + output_size=self.conv_dim, + bias=False, + quant_config=None, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("conv1d", prefix), + ) + self.conv1d.weight.data = self.conv1d.weight.data.unsqueeze(1) + + # Split projection layers (following vLLM's implementation) + # Instead of fused in_proj_qkvz and in_proj_ba, use separate layers + self.in_proj_qkv = MergedColumnParallelLinear( + input_size=self.hidden_size, + output_sizes=[self.key_dim, self.key_dim, self.value_dim], + bias=False, + quant_config=quant_config, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("in_proj_qkv", prefix), + ) + self.in_proj_z = ColumnParallelLinear( + input_size=self.hidden_size, + output_size=self.value_dim, + bias=False, + quant_config=quant_config, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("in_proj_z", prefix), + ) + self.in_proj_b = ColumnParallelLinear( + input_size=self.hidden_size, + output_size=self.num_v_heads, + bias=False, + quant_config=quant_config, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("in_proj_b", prefix), + ) + self.in_proj_a = ColumnParallelLinear( + input_size=self.hidden_size, + output_size=self.num_v_heads, + bias=False, + quant_config=quant_config, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("in_proj_a", prefix), + ) + + # Conv1d weight loader setup + query_key_settings = (self.key_dim, 0, False) + value_settings = (self.value_dim, 0, False) + + delattr(self.conv1d.weight, "weight_loader") + set_weight_attrs( + self.conv1d.weight, + { + "weight_loader": mamba_v2_sharded_weight_loader( + [ + query_key_settings, + query_key_settings, + value_settings, + ], + self.attn_tp_size, + self.attn_tp_rank, + ) + }, + ) + + # State parameters + self.dt_bias = nn.Parameter( + torch.ones(self.num_v_heads // self.attn_tp_size), + ) + self.A_log = nn.Parameter( + torch.empty(self.num_v_heads // self.attn_tp_size), + ) + + set_weight_attrs(self.A_log, {"weight_loader": sharded_weight_loader(0)}) + set_weight_attrs(self.dt_bias, {"weight_loader": sharded_weight_loader(0)}) + + conv_weights = self.conv1d.weight.view( + self.conv1d.weight.size(0), self.conv1d.weight.size(2) + ) + # RadixLinearAttention layer + self.attn = RadixLinearAttention( + layer_id=layer_id, + num_q_heads=self.num_k_heads // self.attn_tp_size, + num_k_heads=self.num_k_heads // self.attn_tp_size, + num_v_heads=self.num_v_heads // self.attn_tp_size, + head_q_dim=self.head_k_dim, + head_k_dim=self.head_k_dim, + head_v_dim=self.head_v_dim, + conv_weights=conv_weights, + bias=self.conv1d.bias, + activation=self.activation, + A_log=self.A_log, + dt_bias=self.dt_bias, + ) + + # Normalization layer + self.norm = RMSNormGated( + self.head_v_dim, + eps=self.layer_norm_epsilon, + group_size=None, + norm_before_gate=True, + device=torch.get_device_module().current_device(), + dtype=config.torch_dtype, + ) + + # Output projection + self.out_proj = RowParallelLinear( + self.value_dim, + self.hidden_size, + bias=False, + input_is_parallel=True, + reduce_results=False, + quant_config=quant_config, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("out_proj", prefix), + ) + + def fix_query_key_value_ordering( + self, + mixed_qkv, + z, + b, + a, + ): + raise NotImplementedError( + "Qwen3.5 Series dont need to fix query key value ordering" + ) + + def forward( + self, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ): + output = torch.empty_like(hidden_states) + if forward_batch.forward_mode.is_extend() and get_forward_context() is not None: + gdn_with_output( + hidden_states, + output, + self.layer_id, + ) + return output + else: + return self._forward(hidden_states, forward_batch) + + def _forward( + self, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ): + """ + Forward pass with three parts: + 1. Input projection + 2. Core attention (custom op) + 3. Output projection + """ + seq_len, _ = hidden_states.shape + + mixed_qkv, _ = self.in_proj_qkv(hidden_states) + z, _ = self.in_proj_z(hidden_states) + z = z.reshape(z.size(0), -1, self.head_v_dim) + b, _ = self.in_proj_b(hidden_states) + a, _ = self.in_proj_a(hidden_states) + + b = b.contiguous() + a = a.contiguous() + + core_attn_out = self.attn.forward( + forward_batch=forward_batch, + mixed_qkv=mixed_qkv, + a=a, + b=b, + ) + + z_shape_og = z.shape + core_attn_out = core_attn_out.reshape(-1, core_attn_out.shape[-1]) + z = z.reshape(-1, z.shape[-1]) + core_attn_out = self.norm(core_attn_out, z) + core_attn_out = core_attn_out.reshape(z_shape_og) + core_attn_out = rearrange(core_attn_out, "... h d -> ... (h d)") + output, _ = self.out_proj(core_attn_out) + return output + + +class Qwen3_5LinearDecoderLayer(nn.Module): + """Qwen3.5 Decoder Layer with Linear Attention (GatedDeltaNet).""" + + def __init__( + self, + config: Qwen3_5TextConfig, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + alt_stream: Optional[torch.cuda.Stream] = None, + ) -> None: + super().__init__() + self.config = config + self.layer_id = layer_id + self.linear_attn = Qwen3_5GatedDeltaNet( + config, layer_id, quant_config, alt_stream, prefix + ) + + # NOTE: Determine the MLP type based on the model type + # Qwen3.5 use all layers for MLP / Qwen3.5-MoE use sparse MoE blocks + if config.model_type == "qwen3_5_moe_text": + self.mlp = Qwen2MoeSparseMoeBlock( + layer_id=layer_id, + config=config, + quant_config=quant_config, + alt_stream=alt_stream, + prefix=add_prefix("mlp", prefix.replace(".self_attn", "")), + ) + elif config.model_type == "qwen3_5_text": + self.mlp = Qwen2MoeMLP( + hidden_size=config.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=add_prefix("mlp", prefix.replace(".self_attn", "")), + ) + else: + raise ValueError(f"Invalid model type: {config.model_type}") + + self.layer_scatter_modes = LayerScatterModes.init_new( + layer_id=layer_id, + num_layers=config.num_hidden_layers, + is_layer_sparse=False, + is_previous_layer_sparse=False, + is_next_layer_sparse=False, + ) + + self.input_layernorm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = GemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.layer_communicator = LayerCommunicator( + layer_scatter_modes=self.layer_scatter_modes, + input_layernorm=self.input_layernorm, + post_attention_layernorm=self.post_attention_layernorm, + allow_reduce_scatter=True, + ) + + def forward( + self, + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + **kwargs, + ): + forward_batch = kwargs.get("forward_batch", None) + + hidden_states, residual = self.layer_communicator.prepare_attn( + hidden_states, residual, forward_batch + ) + + if not forward_batch.forward_mode.is_idle(): + hidden_states = self.linear_attn( + hidden_states, + forward_batch, + ) + + # Fully Connected + hidden_states, residual = self.layer_communicator.prepare_mlp( + hidden_states, residual, forward_batch + ) + + use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( + forward_batch + ) + hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) + + hidden_states, residual = self.layer_communicator.postprocess_layer( + hidden_states, residual, forward_batch + ) + + return hidden_states, residual + + +class Qwen3_5AttentionDecoderLayer(nn.Module): + """Qwen3.5 Decoder Layer with Full Attention.""" + + def __init__( + self, + config: Qwen3_5TextConfig, + layer_id: int, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + alt_stream: Optional[torch.cuda.Stream] = None, + ) -> None: + super().__init__() + self.config = config + self.hidden_size = config.hidden_size + self.attn_tp_rank = get_attention_tp_rank() + self.attn_tp_size = get_attention_tp_size() + self.total_num_heads = config.num_attention_heads + assert self.total_num_heads % self.attn_tp_size == 0 + self.num_heads = self.total_num_heads // self.attn_tp_size + self.total_num_kv_heads = config.num_key_value_heads + if self.total_num_kv_heads >= self.attn_tp_size: + assert self.total_num_kv_heads % self.attn_tp_size == 0 + else: + assert self.attn_tp_size % self.total_num_kv_heads == 0 + self.num_kv_heads = max(1, self.total_num_kv_heads // self.attn_tp_size) + self.head_dim = config.head_dim or (self.hidden_size // self.num_heads) + self.q_size = self.num_heads * self.head_dim + self.kv_size = self.num_kv_heads * self.head_dim + self.scaling = self.head_dim**-0.5 + self.max_position_embeddings = getattr(config, "max_position_embeddings", 8192) + + if hasattr(config, "rope_parameters"): + self.rope_scaling = getattr(config, "rope_parameters", None) + else: + self.rope_scaling = getattr(config, "rope_scaling", None) + + self.rope_theta = self.rope_scaling.get("rope_theta", 10000) + self.partial_rotary_factor = self.rope_scaling.get("partial_rotary_factor", 1.0) + self.layer_id = layer_id + + self.attn_output_gate = getattr(config, "attn_output_gate", True) + if self.attn_output_gate: + logger.warning_once("using attn output gate!") + + self.rotary_emb = get_rope( + head_size=self.head_dim, + rotary_dim=self.head_dim, + max_position=self.max_position_embeddings, + rope_scaling=self.rope_scaling, + base=self.rope_theta, + partial_rotary_factor=self.partial_rotary_factor, + is_neox_style=True, + dtype=torch.get_default_dtype(), + ) + + self.qkv_proj = QKVParallelLinear( + config.hidden_size, + self.head_dim, + self.total_num_heads * (1 + self.attn_output_gate), + self.total_num_kv_heads, + bias=False, + quant_config=quant_config, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("qkv_proj", prefix), + ) + + self.o_proj = RowParallelLinear( + self.total_num_heads * self.head_dim, + config.hidden_size, + bias=False, + quant_config=quant_config, + reduce_results=False, + tp_rank=self.attn_tp_rank, + tp_size=self.attn_tp_size, + prefix=add_prefix("o_proj", prefix), + ) + + self.attn = RadixAttention( + self.num_heads, + self.head_dim, + self.scaling, + num_kv_heads=self.num_kv_heads, + layer_id=layer_id, + prefix=f"{prefix}.attn", + ) + + # Dense MLP for non-MoE variant + if config.model_type == "qwen3_5_text": + self.mlp = Qwen2MoeMLP( + hidden_size=config.hidden_size, + intermediate_size=config.intermediate_size, + hidden_act=config.hidden_act, + quant_config=quant_config, + prefix=add_prefix("mlp", prefix.replace(".self_attn", "")), + ) + elif config.model_type == "qwen3_5_moe_text": + self.mlp = Qwen2MoeSparseMoeBlock( + layer_id=layer_id, + config=config, + quant_config=quant_config, + alt_stream=alt_stream, + prefix=add_prefix("mlp", prefix.replace(".self_attn", "")), + ) + else: + raise ValueError(f"Invalid model type: {config.model_type}") + + self.layer_scatter_modes = LayerScatterModes.init_new( + layer_id=layer_id, + num_layers=config.num_hidden_layers, + is_layer_sparse=False, + is_previous_layer_sparse=False, + is_next_layer_sparse=False, + ) + + self.input_layernorm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = GemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + self.q_norm = GemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps) + self.k_norm = GemmaRMSNorm(self.head_dim, eps=config.rms_norm_eps) + + self.layer_communicator = LayerCommunicator( + layer_scatter_modes=self.layer_scatter_modes, + input_layernorm=self.input_layernorm, + post_attention_layernorm=self.post_attention_layernorm, + allow_reduce_scatter=True, + ) + + self.alt_stream = alt_stream + + def _apply_qk_norm( + self, q: torch.Tensor, k: torch.Tensor + ) -> Tuple[torch.Tensor, torch.Tensor]: + """Apply Q/K normalization with optional alt_stream overlap.""" + if self.alt_stream is not None and get_is_capture_mode(): + current_stream = torch.cuda.current_stream() + self.alt_stream.wait_stream(current_stream) + q_by_head = q.reshape(-1, self.head_dim) + q_by_head = self.q_norm(q_by_head) + with torch.cuda.stream(self.alt_stream): + k_by_head = k.reshape(-1, self.head_dim) + k_by_head = self.k_norm(k_by_head) + current_stream.wait_stream(self.alt_stream) + else: + q_by_head = q.reshape(-1, self.head_dim) + q_by_head = self.q_norm(q_by_head) + k_by_head = k.reshape(-1, self.head_dim) + k_by_head = self.k_norm(k_by_head) + q = q_by_head.view(q.shape) + k = k_by_head.view(k.shape) + return q, k + + def self_attention( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + forward_batch: ForwardBatch, + ) -> torch.Tensor: + """Full attention forward pass.""" + qkv, _ = self.qkv_proj(hidden_states) + + if self.attn_output_gate: + q_gate, k, v = qkv.split( + [self.q_size * 2, self.kv_size, self.kv_size], dim=-1 + ) + orig_shape = q_gate.shape[:-1] + q_gate = q_gate.view(*orig_shape, self.num_heads, -1) + q, gate = torch.chunk(q_gate, 2, dim=-1) + q = q.reshape(*orig_shape, -1) + gate = gate.reshape(*orig_shape, -1) + else: + q, k, v = qkv.split([self.q_size, self.kv_size, self.kv_size], dim=-1) + + q, k = self._apply_qk_norm(q, k) + q, k = self.rotary_emb(positions, q, k) + attn_output = self.attn(q, k, v, forward_batch) + + if self.attn_output_gate: + gate = torch.sigmoid(gate) + attn_output = attn_output * gate + + output, _ = self.o_proj(attn_output) + return output + + def forward( + self, + positions: torch.Tensor, + hidden_states: torch.Tensor, + residual: Optional[torch.Tensor], + forward_batch: ForwardBatch, + **kwargs, + ): + hidden_states, residual = self.layer_communicator.prepare_attn( + hidden_states, residual, forward_batch + ) + + if not forward_batch.forward_mode.is_idle(): + hidden_states = self.self_attention( + positions=positions, + hidden_states=hidden_states, + forward_batch=forward_batch, + ) + + # Fully Connected + hidden_states, residual = self.layer_communicator.prepare_mlp( + hidden_states, residual, forward_batch + ) + use_reduce_scatter = self.layer_communicator.should_use_reduce_scatter( + forward_batch + ) + hidden_states = self.mlp(hidden_states, forward_batch, use_reduce_scatter) + + hidden_states, residual = self.layer_communicator.postprocess_layer( + hidden_states, residual, forward_batch + ) + + return hidden_states, residual + + +ALL_DECODER_LAYER_TYPES = { + "attention": Qwen3_5AttentionDecoderLayer, + "linear_attention": Qwen3_5LinearDecoderLayer, +} + + +class Qwen3_5ForCausalLM(nn.Module): + """Qwen3.5 Model with support for dense variant.""" + + def __init__( + self, + config: Qwen3_5TextConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__() + self.config = config + self.hidden_size = config.hidden_size + self.pp_group = get_pp_group() + + alt_stream = torch.cuda.Stream() if _is_cuda else None + + # Embedding layer + if self.pp_group.is_first_rank: + self.embed_tokens = VocabParallelEmbedding( + config.vocab_size, + config.hidden_size, + org_num_embeddings=config.vocab_size, + enable_tp=not is_dp_attention_enabled(), + ) + + # Decoder layers + def get_layer(idx: int, prefix: str): + layer_type = config.layers_block_type[idx] + layer_class = ALL_DECODER_LAYER_TYPES[layer_type] + if layer_type == "attention": + prefix = add_prefix("self_attn", prefix) + else: + prefix = add_prefix("linear_attn", prefix) + return layer_class( + config=config, + layer_id=idx, + quant_config=quant_config, + prefix=prefix, + alt_stream=alt_stream, + ) + + self.layers = make_layers( + config.num_hidden_layers, + get_layer, + prefix=f"{prefix}.layers", + ) + + # Final normalization + if self.pp_group.is_last_rank: + self.norm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + + def get_input_embeddings(self) -> nn.Embedding: + return self.embed_tokens + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor] = None, + pp_proxy_tensors: Optional[PPProxyTensors] = None, + input_deepstack_embeds: Optional[torch.Tensor] = None, + ) -> Union[torch.Tensor, PPProxyTensors]: + # Initialize hidden states + if self.pp_group.is_first_rank: + if input_embeds is None: + hidden_states = self.embed_tokens(input_ids) + else: + hidden_states = input_embeds + residual = None + else: + assert pp_proxy_tensors is not None + hidden_states = pp_proxy_tensors["hidden_states"] + residual = pp_proxy_tensors["residual"] + + # Pass through decoder layers + for layer_idx in range(len(self.layers)): + layer = self.layers[layer_idx] + with get_global_expert_distribution_recorder().with_current_layer( + layer_idx + ): + hidden_states, residual = layer( + positions=positions, + hidden_states=hidden_states, + residual=residual, + forward_batch=forward_batch, + ) + + # Process deepstack embeddings if provided + if ( + input_deepstack_embeds is not None + and input_deepstack_embeds.numel() > 0 + and layer_idx < 3 + ): + sep = self.hidden_size * layer_idx + hidden_states.add_( + input_deepstack_embeds[:, sep : sep + self.hidden_size] + ) + + # Return intermediate tensors for pipeline parallelism + if not self.pp_group.is_last_rank: + return PPProxyTensors( + { + "hidden_states": hidden_states, + "residual": residual, + } + ) + + # Apply final normalization + if hidden_states.shape[0] != 0: + if residual is None: + hidden_states = self.norm(hidden_states) + else: + hidden_states, _ = self.norm(hidden_states, residual) + + return hidden_states + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + loaded_params: Set[str] = set() + params_dict = dict(self.named_parameters(remove_duplicate=False)) + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + if "mtp" in name: + continue + if "visual" in name: + continue + if "language_model" in name: + name = name.replace(r"model.language_model.", r"model.") + if ".self_attn." in name: + name = name.replace(".self_attn", "") + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + + if "mlp.experts" in name: + continue + + name = name.replace(weight_name, param_name) + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue + # Skip layers on other devices. + # if is_pp_missing_parameter(name, self): + # continue + if name not in params_dict: + continue + param = params_dict[name] + weight_loader = getattr(param, "weight_loader") + weight_loader(param, loaded_weight, shard_id) + break + else: + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + logger.warning(f"Parameter {name} not found in params_dict") + continue + param = params_dict[name] + + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + loaded_params.add(name) + return loaded_params + + +class Qwen3_5MoeForCausalLM(Qwen3_5ForCausalLM): + def __init__( + self, + config: Qwen3_5TextConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + ) -> None: + super().__init__(config=config, quant_config=quant_config, prefix=prefix) + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + # 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, + ) + + # Skip loading extra parameters for GPTQ/modelopt models. + ignore_suffixes = ( + ".bias", + "_bias", + ".k_scale", + "_k_scale", + ".v_scale", + "_v_scale", + ".weight_scale", + "_weight_scale", + ".input_scale", + "_input_scale", + ) + + is_fused_expert = False + fused_expert_params_mapping = [ + ("experts.w13_weight", "experts.gate_up_proj", 0, "w1"), + ("experts.w2_weight", "experts.down_proj", 0, "w2"), + ] + + num_experts = self.config.num_experts + + def load_fused_expert_weights( + name: str, + params_dict: dict, + loaded_weight: torch.Tensor, + shard_id: str, + num_experts: int, + ): + param = params_dict[name] + weight_loader = param.weight_loader + # let ep moe layer to gracefully handle expert_ids that do not belong to local moe rank + for expert_id in range(num_experts): + curr_expert_weight = loaded_weight[expert_id] + weight_loader( + param, + curr_expert_weight, + name, + shard_id, + expert_id, + ) + return True + + loaded_params: Set[str] = set() + params_dict = dict(self.named_parameters(remove_duplicate=False)) + + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + if "mtp" in name: + continue + if "visual" in name: + continue + if "language_model" in name: + name = name.replace(r"model.language_model.", r"model.") + if ".self_attn." in name: + name = name.replace(".self_attn", "") + + for param_name, weight_name, shard_id in stacked_params_mapping: + if "experts.gate_up_proj" in name or "experts.down_proj" in name: + is_fused_expert = True + expert_params_mapping = fused_expert_params_mapping + + # Skip non-stacked layers and experts (experts handled below). + if weight_name not in name: + continue + + # We have mlp.experts[0].gate_proj in the checkpoint. + # Since we handle the experts below in expert_params_mapping, + # we need to skip here BEFORE we update the name, otherwise + # name will be updated to mlp.experts[0].gate_up_proj, which + # will then be updated below in expert_params_mapping + # for mlp.experts[0].gate_gate_up_proj, which breaks load. + if "mlp.experts" in name: + continue + name = name.replace(weight_name, param_name) + # Skip loading extra parameters for GPTQ/modelopt models. + if name.endswith(ignore_suffixes) and name not in params_dict: + continue + + if name not in params_dict: + continue + + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + # Track if this is an expert weight to enable early skipping + is_expert_weight = False + + for mapping in expert_params_mapping: + param_name, weight_name, expert_id, shard_id = mapping + if weight_name not in name: + continue + # Anyway, this is an expert weight and should not be + # attempted to load as other weights later + is_expert_weight = True + name_mapped = name.replace(weight_name, param_name) + if is_fused_expert: + if "experts.gate_up_proj" in name: + loaded_weight = loaded_weight.chunk(2, dim=-2) + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight[0], + "w1", + num_experts, + ) + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight[1], + "w3", + num_experts, + ) + else: + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight, + shard_id, + num_experts, + ) + else: + # Skip loading extra parameters for GPTQ/modelopt models. + if ( + name_mapped.endswith(ignore_suffixes) + and name_mapped not in params_dict + ): + continue + param = params_dict[name_mapped] + # We should ask the weight loader to return success or + # not here since otherwise we may skip experts with + # # other available replicas. + weight_loader = param.weight_loader + weight_loader( + param, + loaded_weight, + name_mapped, + shard_id=shard_id, + expert_id=expert_id, + ) + name = name_mapped + break + else: + if is_expert_weight: + # This is an expert weight but not mapped to this rank, skip all remaining processing + continue + + # Skip loading extra parameters for GPTQ/modelopt models. + if name.endswith(ignore_suffixes) and name not in params_dict: + continue + + if name in params_dict.keys(): + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + else: + logger.warning(f"Parameter {name} not found in params_dict") + loaded_params.add(name) + + return loaded_params + + +class Qwen3_5ForConditionalGeneration(Qwen3VLForConditionalGeneration): + def __init__( + self, + config: Qwen3_5Config, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + language_model_cls=Qwen3_5ForCausalLM, + ): + super().__init__(config, quant_config, prefix, language_model_cls) + + rope_config = getattr(self.config, "rope_parameters", None) or getattr( + self.config, "rope_scaling", {} + ) + self.is_mrope_enabled = "mrope_section" in rope_config + + self.deepstack_visual_indexes = self.visual.deepstack_visual_indexes + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + def set_embed_and_head(self, embed, head): + del self.model.embed_tokens.weight + del self.lm_head.weight + self.model.embed_tokens.weight = embed + self.lm_head.weight = head + torch.cuda.empty_cache() + torch.cuda.synchronize() + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + loaded_params: Set[str] = set() + params_dict = dict(self.named_parameters(remove_duplicate=False)) + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + if "mtp" in name: + continue + if "language_model" in name: + name = name.replace(r"model.language_model.", r"model.") + if ".self_attn." in name: + name = name.replace(".self_attn", "") + + for param_name, weight_name, shard_id in stacked_params_mapping: + if weight_name not in name: + continue + + if "visual" in name or "mlp.experts" in name: + continue + + name = name.replace(weight_name, param_name) + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue + # Skip layers on other devices. + # if is_pp_missing_parameter(name, self): + # continue + if name not in params_dict: + continue + param = params_dict[name] + weight_loader = getattr(param, "weight_loader") + weight_loader(param, loaded_weight, shard_id) + break + else: + if "visual" in name: + # adapt to VisionAttention + name = name.replace(r"attn.qkv.", r"attn.qkv_proj.") + name = name.replace(r"model.visual.", r"visual.") + + # print(name, loaded_weight.shape) + # Skip loading extra bias for GPTQ models. + if name.endswith(".bias") and name not in params_dict: + continue + if name not in params_dict: + logger.warning(f"Parameter {name} not found in params_dict") + continue + param = params_dict[name] + + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight) + loaded_params.add(name) + return loaded_params + + +class Qwen3_5MoeForConditionalGeneration(Qwen3VLForConditionalGeneration): + """Qwen3.5 MoE Vision-Language Model.""" + + def __init__( + self, + config: Qwen3_5MoeConfig, + quant_config: Optional[QuantizationConfig] = None, + prefix: str = "", + language_model_cls=Qwen3_5MoeForCausalLM, + ) -> None: + super().__init__(config, quant_config, prefix, language_model_cls) + rope_config = getattr(self.config, "rope_parameters", None) or getattr( + self.config, "rope_scaling", {} + ) + self.is_mrope_enabled = "mrope_section" in rope_config + + self.deepstack_visual_indexes = self.visual.deepstack_visual_indexes + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + def set_embed_and_head(self, embed, head): + del self.model.embed_tokens.weight + del self.lm_head.weight + self.model.embed_tokens.weight = embed + self.lm_head.weight = head + torch.cuda.empty_cache() + torch.cuda.synchronize() + + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + # 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, + ) + + # Skip loading extra parameters for GPTQ/modelopt models. + ignore_suffixes = ( + ".bias", + "_bias", + ".k_scale", + "_k_scale", + ".v_scale", + "_v_scale", + ".weight_scale", + "_weight_scale", + ".input_scale", + "_input_scale", + ) + + is_fused_expert = False + fused_expert_params_mapping = [ + ("experts.w13_weight", "experts.gate_up_proj", 0, "w1"), + ("experts.w2_weight", "experts.down_proj", 0, "w2"), + ] + + num_experts = self.config.num_experts + + def load_fused_expert_weights( + name: str, + params_dict: dict, + loaded_weight: torch.Tensor, + shard_id: str, + num_experts: int, + ): + param = params_dict[name] + weight_loader = param.weight_loader + # let ep moe layer to gracefully handle expert_ids that do not belong to local moe rank + for expert_id in range(num_experts): + curr_expert_weight = loaded_weight[expert_id] + weight_loader( + param, + curr_expert_weight, + name, + shard_id, + expert_id, + ) + return True + + loaded_params: Set[str] = set() + params_dict = dict(self.named_parameters(remove_duplicate=False)) + + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + if "mtp" in name: + continue + if "language_model" in name: + name = name.replace(r"model.language_model.", r"model.") + if ".self_attn." in name: + name = name.replace(".self_attn", "") + + for param_name, weight_name, shard_id in stacked_params_mapping: + if "experts.gate_up_proj" in name or "experts.down_proj" in name: + is_fused_expert = True + expert_params_mapping = fused_expert_params_mapping + + # Skip non-stacked layers and experts (experts handled below). + if weight_name not in name: + continue + if "visual" in name: + continue + + # We have mlp.experts[0].gate_proj in the checkpoint. + # Since we handle the experts below in expert_params_mapping, + # we need to skip here BEFORE we update the name, otherwise + # name will be updated to mlp.experts[0].gate_up_proj, which + # will then be updated below in expert_params_mapping + # for mlp.experts[0].gate_gate_up_proj, which breaks load. + if "mlp.experts" in name: + continue + name = name.replace(weight_name, param_name) + # Skip loading extra parameters for GPTQ/modelopt models. + if name.endswith(ignore_suffixes) and name not in params_dict: + continue + + if name not in params_dict: + continue + + param = params_dict[name] + weight_loader = param.weight_loader + weight_loader(param, loaded_weight, shard_id) + break + else: + # Track if this is an expert weight to enable early skipping + is_expert_weight = False + + for mapping in expert_params_mapping: + param_name, weight_name, expert_id, shard_id = mapping + if weight_name not in name: + continue + if "visual" in name or self.config.encoder_only: + continue + # Anyway, this is an expert weight and should not be + # attempted to load as other weights later + is_expert_weight = True + name_mapped = name.replace(weight_name, param_name) + if is_fused_expert: + if "experts.gate_up_proj" in name: + loaded_weight = loaded_weight.chunk(2, dim=-2) + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight[0], + "w1", + num_experts, + ) + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight[1], + "w3", + num_experts, + ) + else: + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight, + shard_id, + num_experts, + ) + else: + # Skip loading extra parameters for GPTQ/modelopt models. + if ( + name_mapped.endswith(ignore_suffixes) + and name_mapped not in params_dict + ): + continue + param = params_dict[name_mapped] + # We should ask the weight loader to return success or + # not here since otherwise we may skip experts with + # # other available replicas. + weight_loader = param.weight_loader + weight_loader( + param, + loaded_weight, + name_mapped, + shard_id=shard_id, + expert_id=expert_id, + ) + name = name_mapped + break + else: + if is_expert_weight: + # This is an expert weight but not mapped to this rank, skip all remaining processing + continue + + if "visual" in name: + # adapt to VisionAttention + name = name.replace(r"attn.qkv.", r"attn.qkv_proj.") + name = name.replace(r"model.visual.", r"visual.") + + # Skip loading extra parameters for GPTQ/modelopt models. + if name.endswith(ignore_suffixes) and name not in params_dict: + continue + + if name in params_dict.keys(): + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + else: + logger.warning(f"Parameter {name} not found in params_dict") + loaded_params.add(name) + + return loaded_params + + +EntryClass = [Qwen3_5MoeForConditionalGeneration, Qwen3_5ForConditionalGeneration] diff --git a/python/sglang/srt/models/qwen3_5_mtp.py b/python/sglang/srt/models/qwen3_5_mtp.py new file mode 100644 index 000000000..c8d21cb3e --- /dev/null +++ b/python/sglang/srt/models/qwen3_5_mtp.py @@ -0,0 +1,415 @@ +# Copyright 2023-2024 SGLang Team +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ============================================================================== + +"""Inference-only Qwen3_5 MTP model.""" +import logging +from typing import Iterable, Optional, Tuple + +import torch +from torch import nn +from transformers import PretrainedConfig + +from sglang.srt.distributed import get_pp_group, get_tensor_model_parallel_world_size +from sglang.srt.layers.layernorm import GemmaRMSNorm +from sglang.srt.layers.logits_processor import LogitsProcessor +from sglang.srt.layers.moe.fused_moe_triton.layer import FusedMoE +from sglang.srt.layers.vocab_parallel_embedding import ( + ParallelLMHead, + VocabParallelEmbedding, +) +from sglang.srt.model_executor.forward_batch_info import ForwardBatch +from sglang.srt.model_loader.weight_utils import default_weight_loader +from sglang.srt.models.qwen3_5 import Qwen3_5AttentionDecoderLayer +from sglang.srt.utils import add_prefix + +logger = logging.getLogger(__name__) + + +class Qwen3_5MultiTokenPredictor(nn.Module): + def __init__(self, config: PretrainedConfig, quant_config=None, prefix: str = ""): + super().__init__() + + self.config = config + + self.vocab_size = config.vocab_size + + self.mtp_start_layer_idx = config.num_hidden_layers + self.num_mtp_layers = getattr(config, "mtp_num_hidden_layers", 1) + + self.embed_tokens = VocabParallelEmbedding( + self.vocab_size, + config.hidden_size, + ) + + self.fc = nn.Linear(2 * config.hidden_size, config.hidden_size, bias=False) + + config.full_attention_interval = 1 + self.layers = torch.nn.ModuleList( + [ + Qwen3_5AttentionDecoderLayer( + config, + idx, + quant_config, + prefix=add_prefix(f"layers.{idx}", prefix), + ) + for idx in range(self.num_mtp_layers) + ] + ) + + self.norm = GemmaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.pre_fc_norm_hidden = GemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + self.pre_fc_norm_embedding = GemmaRMSNorm( + config.hidden_size, eps=config.rms_norm_eps + ) + + def embed_input_ids(self, input_ids: torch.Tensor) -> torch.Tensor: + return self.embed_tokens(input_ids) + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: torch.Tensor, + input_embeds: Optional[torch.Tensor] = None, + **kwargs, + ): + # if get_pp_group().is_first_rank: + assert input_embeds is None + input_embeds = forward_batch.mm_input_embeds + if ( + forward_batch.forward_mode.is_extend() + and forward_batch.contains_mm_inputs() + and not forward_batch.forward_mode.is_draft_extend() + ): + assert input_embeds is not None + input_embeds = torch.cat( + [input_embeds[:-1], self.embed_tokens(input_ids[-1].unsqueeze(0))] + ) + + if input_embeds is None: + input_embeds = self.embed_tokens(input_ids) + + hidden_states = forward_batch.spec_info.hidden_states + + # Some idle batch has 0 batch size. GemmaRMSNorm.forward would fail due to bs=0. + if not forward_batch.forward_mode.is_idle(): + input_embeds = self.pre_fc_norm_embedding(input_embeds) + hidden_states = self.pre_fc_norm_hidden(hidden_states) + hidden_states = torch.cat([input_embeds, hidden_states], dim=-1) + + hidden_states = self.fc(hidden_states) + residual = None + + if self.num_mtp_layers == 1: + hidden_states, residual = self.layers[0]( + positions=positions, + hidden_states=hidden_states, + residual=residual, + forward_batch=forward_batch, + ) + else: + raise ("not implementation for other mtp layers[self.num_mtp_layers > 1]") + + if not get_pp_group().is_last_rank: + # For pipeline parallel, return intermediate tensors + return hidden_states + + hidden_states, _ = self.norm(hidden_states, residual) + return hidden_states + + +class Qwen3_5ForCausalLMMTP(nn.Module): + + def __init__( + self, + config: PretrainedConfig, + quant_config=None, + prefix: str = "", + ) -> None: + super().__init__() + + self.is_multimodal = hasattr(config, "text_config") + if self.is_multimodal: + config = config.text_config + + self.config = config + self.tp_size = get_tensor_model_parallel_world_size() + self.quant_config = quant_config + self.pp_group = get_pp_group() + + self.model = Qwen3_5MultiTokenPredictor( + config, quant_config, prefix=add_prefix("mtp", prefix) + ) + + if get_pp_group().is_last_rank: + if config.tie_word_embeddings: + self.lm_head = self.model.embed_tokens + else: + self.lm_head = ParallelLMHead( + config.vocab_size, + config.hidden_size, + quant_config=quant_config, + prefix=add_prefix("lm_head", prefix), + ) + else: + # For pipeline parallel, create a placeholder layer + self.lm_head = nn.Linear(1, 1, bias=False) + + self.logits_processor = LogitsProcessor(config) + + def get_embed_and_head(self): + return self.model.embed_tokens.weight, self.lm_head.weight + + def set_embed_and_head(self, embed, head): + del self.model.embed_tokens.weight + if not self.config.tie_word_embeddings: + del self.lm_head.weight + + self.model.embed_tokens.weight = embed + self.lm_head.weight = head + torch.cuda.empty_cache() + torch.cuda.synchronize() + + @torch.no_grad() + def forward( + self, + input_ids: torch.Tensor, + positions: torch.Tensor, + forward_batch: ForwardBatch, + input_embeds: Optional[torch.Tensor] = None, + **kwargs, + ): + hidden_states = self.model( + input_ids, + positions, + forward_batch, + input_embeds, + ) + + if not get_pp_group().is_last_rank: + # For pipeline parallel, return intermediate results + return hidden_states + + return self.logits_processor( + input_ids, hidden_states, self.lm_head, forward_batch + ) + + def load_weights( + self, weights: Iterable[Tuple[str, torch.Tensor]], is_mtp: bool = False + ): + stacked_params_mapping = [ + # (param_name, shard_name, shard_id) + ("qkv_proj", "q_proj", "q"), + ("qkv_proj", "k_proj", "k"), + ("qkv_proj", "v_proj", "v"), + ("gate_up_proj", "gate_proj", 0), + ("gate_up_proj", "up_proj", 1), + ] + + # Params for MoE experts (non-fused/fused) + num_experts = getattr(self.config, "num_experts", None) + if num_experts is not None: + 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=num_experts, + ) + else: + expert_params_mapping = [] + + # Skip loading extra parameters for GPTQ/modelopt models. + ignore_suffixes = ( + ".bias", + "_bias", + ".k_scale", + "_k_scale", + ".v_scale", + "_v_scale", + ".weight_scale", + "_weight_scale", + ".input_scale", + "_input_scale", + ) + + # fused experts: experts.w13_weight / experts.w2_weight + is_fused_expert = False + fused_expert_params_mapping = [ + ("experts.w13_weight", "experts.gate_up_proj", 0, "w1"), + ("experts.w2_weight", "experts.down_proj", 0, "w2"), + ] + + def load_fused_expert_weights( + name: str, + params_dict: dict, + loaded_weight: torch.Tensor, + shard_id: str, + num_experts: int, + ): + param = params_dict[name] + weight_loader = param.weight_loader + # Let EP MoE layer handle expert_ids that do not belong to local moe rank + for expert_id in range(num_experts): + curr_expert_weight = loaded_weight[expert_id] + weight_loader( + param, + curr_expert_weight, + name, + shard_id, + expert_id, + ) + return True + + params_dict = dict(self.named_parameters()) + loaded_params: set[str] = set() + + for name, loaded_weight in weights: + if "rotary_emb.inv_freq" in name: + continue + + # Only process MTP branch weights + if "mtp" not in name: + continue + + # Some checkpoints use model.language_model.mtp.* prefix + if "language_model" in name: + name = name.replace(r"model.language_model.", r"model.") + + if name.startswith("mtp."): + # Remove the mtp. prefix for processing + name = name.replace("mtp.", "model.") + + if ".self_attn." in name: + name = name.replace(".self_attn", "") + + # 1) Process stacked parameters (q_proj/k_proj/v_proj & gate_proj/up_proj) + for param_name, weight_name, shard_id in stacked_params_mapping: + # Check if this is a fused expert weight + if "experts.gate_up_proj" in name or "experts.down_proj" in name: + is_fused_expert = True + expert_params_mapping = fused_expert_params_mapping + + # Skip non-matching weights + if weight_name not in name: + continue + + # Skip MoE experts.* here, handled separately below + if "mlp.experts" in name: + continue + + name_mapped = name.replace(weight_name, param_name) + + # Skip loading extra parameters for GPTQ/modelopt models. + if ( + name_mapped.endswith(ignore_suffixes) + and name_mapped not in params_dict + ): + continue + + if name_mapped not in params_dict: + continue + + param = params_dict[name_mapped] + weight_loader = getattr(param, "weight_loader", default_weight_loader) + weight_loader(param, loaded_weight, shard_id) + name = name_mapped + break + else: + # 2) Process MoE expert weights (including fused experts) + is_expert_weight = False + + for mapping in expert_params_mapping: + param_name, weight_name, expert_id, shard_id = mapping + if weight_name not in name: + continue + + is_expert_weight = True + name_mapped = name.replace(weight_name, param_name) + + # Fused experts: single checkpoint weight contains multiple experts + if is_fused_expert and num_experts is not None: + if "experts.gate_up_proj" in name: + # gate_up_proj fused: split into w1 / w3 + loaded_w1, loaded_w3 = loaded_weight.chunk(2, dim=-2) + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_w1, + "w1", + num_experts, + ) + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_w3, + "w3", + num_experts, + ) + else: + # down_proj fused: distribute entire weight + load_fused_expert_weights( + name_mapped, + params_dict, + loaded_weight, + shard_id, + num_experts, + ) + else: + # Non-fused expert, load by expert_id/shard + if ( + name_mapped.endswith(ignore_suffixes) + and name_mapped not in params_dict + ): + continue + if name_mapped not in params_dict: + break + param = params_dict[name_mapped] + weight_loader = param.weight_loader + weight_loader( + param, + loaded_weight, + name_mapped, + shard_id=shard_id, + expert_id=expert_id, + ) + name = name_mapped + break + else: + # Skip expert weight if not handled by current rank + if is_expert_weight: + continue + + # 3) Regular non-stacked / non-expert parameters, use default loader + if name.endswith(ignore_suffixes) and name not in params_dict: + continue + + if name in params_dict: + param = params_dict[name] + weight_loader = getattr( + param, "weight_loader", default_weight_loader + ) + weight_loader(param, loaded_weight) + else: + logger.warning_once( + f"Parameter {name} not found in params_dict, skip loading" + ) + + loaded_params.add(name) + return loaded_params + + +EntryClass = [Qwen3_5ForCausalLMMTP] diff --git a/python/sglang/srt/models/qwen3_next.py b/python/sglang/srt/models/qwen3_next.py index 880201a7d..a013dca53 100644 --- a/python/sglang/srt/models/qwen3_next.py +++ b/python/sglang/srt/models/qwen3_next.py @@ -617,7 +617,10 @@ class Qwen3HybridAttentionDecoderLayer(nn.Module): self.scaling = self.head_dim**-0.5 self.rope_theta = getattr(config, "rope_theta", 10000) self.max_position_embeddings = getattr(config, "max_position_embeddings", 8192) - self.rope_scaling = getattr(config, "rope_scaling", None) + if "rope_parameters" in config: + self.rope_scaling = getattr(config, "rope_parameters", None) + else: + self.rope_scaling = getattr(config, "rope_scaling", None) self.partial_rotary_factor = config.partial_rotary_factor self.layer_id = layer_id diff --git a/python/sglang/srt/multimodal/processors/qwen_vl.py b/python/sglang/srt/multimodal/processors/qwen_vl.py index eb648542d..9e7d0346f 100644 --- a/python/sglang/srt/multimodal/processors/qwen_vl.py +++ b/python/sglang/srt/multimodal/processors/qwen_vl.py @@ -7,6 +7,7 @@ from typing import List, Union import numpy as np import torch import torchvision +from decord import VideoReader from PIL import Image from torchvision.transforms import InterpolationMode @@ -15,6 +16,10 @@ from sglang.srt.layers.rotary_embedding import MRotaryEmbedding from sglang.srt.managers.schedule_batch import Modality, MultimodalDataItem from sglang.srt.models.qwen2_5_vl import Qwen2_5_VLForConditionalGeneration from sglang.srt.models.qwen2_vl import Qwen2VLForConditionalGeneration +from sglang.srt.models.qwen3_5 import ( + Qwen3_5ForConditionalGeneration, + Qwen3_5MoeForConditionalGeneration, +) from sglang.srt.models.qwen3_omni_moe import Qwen3OmniMoeForConditionalGeneration from sglang.srt.models.qwen3_vl import Qwen3VLForConditionalGeneration from sglang.srt.models.qwen3_vl_moe import Qwen3VLMoeForConditionalGeneration @@ -148,6 +153,9 @@ async def preprocess_video( image_factor: int = IMAGE_FACTOR, video_config: dict = {}, ) -> torch.Tensor: + # preprocessed video + if not isinstance(vr, VideoReader): + return vr entry_time = time.perf_counter() total_frames, video_fps = len(vr), vr.get_avg_fps() @@ -226,6 +234,8 @@ class QwenVLImageProcessor(SGLangBaseProcessor): Qwen2_5_VLForConditionalGeneration, Qwen3VLForConditionalGeneration, Qwen3VLMoeForConditionalGeneration, + Qwen3_5ForConditionalGeneration, + Qwen3_5MoeForConditionalGeneration, Qwen3OmniMoeForConditionalGeneration, ] @@ -326,7 +336,12 @@ class QwenVLImageProcessor(SGLangBaseProcessor): preprocess_time = time.perf_counter() # NOTE: for qwen3-vl, video_meta need to be passed in, since do_sample_frames is already done in preprocess_video - if self.hf_config.model_type in ("qwen3_vl", "qwen3_vl_moe"): + if self.hf_config.model_type in ( + "qwen3_vl", + "qwen3_vl_moe", + "qwen3_5", + "qwen3_5_moe", + ): mm_items, input_ids, ret = self.process_and_combine_mm_data( base_output, self.mm_tokens, diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index 9e9a2c626..ecee665ad 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -1558,7 +1558,11 @@ class ServerArgs: "Use flashinfer_trtllm as MoE runner backend on sm100 for " f"{model_arch}" ) - elif model_arch in ["Qwen3NextForCausalLM"]: + elif model_arch in [ + "Qwen3NextForCausalLM", + "Qwen3_5MoeForConditionalGeneration", + "Qwen3_5ForConditionalGeneration", + ]: if is_sm100_supported(): quant_method = get_quantization_config(hf_config) if self.quantization is None and quant_method is not None: @@ -1573,7 +1577,7 @@ class ServerArgs: ): self.moe_runner_backend = "flashinfer_trtllm" logger.info( - "Use flashinfer_trtllm as MoE runner backend on sm100 for Qwen3NextForCausalLM" + f"Use flashinfer_trtllm as MoE runner backend on sm100 for {model_arch}" ) self._handle_mamba_radix_cache( model_arch=model_arch, diff --git a/python/sglang/srt/speculative/eagle_worker.py b/python/sglang/srt/speculative/eagle_worker.py index ac689cbf5..333c206b2 100644 --- a/python/sglang/srt/speculative/eagle_worker.py +++ b/python/sglang/srt/speculative/eagle_worker.py @@ -291,7 +291,11 @@ class EAGLEWorker(TpModelWorker): self.draft_model_runner.tp_group ), speculative_moe_backend_context(), speculative_moe_a2a_backend_context(): self.forward_draft_extend( - batch, logits_output.hidden_states, next_token_ids, seq_lens_cpu + batch, + logits_output.hidden_states, + next_token_ids, + seq_lens_cpu, + logits_output.mm_input_embeds, ) return GenerationBatchResult( logits_output=logits_output, @@ -856,6 +860,7 @@ class EAGLEWorker(TpModelWorker): hidden_states: torch.Tensor, next_token_ids: torch.Tensor, seq_lens_cpu: Optional[torch.Tensor], + mm_input_embeds: Optional[torch.Tensor] = None, ): """Run draft model extend. This API modifies the states of the batch. @@ -880,6 +885,8 @@ class EAGLEWorker(TpModelWorker): model_worker_batch, self.draft_model_runner ) forward_batch.return_logprob = False + if mm_input_embeds is not None: + forward_batch.mm_input_embeds = mm_input_embeds logits_output = self.draft_model_runner.forward(forward_batch).logits_output if self.enable_nan_detection: detect_nan(logits_output) diff --git a/python/sglang/srt/utils/common.py b/python/sglang/srt/utils/common.py index 9281163ef..0c8a80be0 100644 --- a/python/sglang/srt/utils/common.py +++ b/python/sglang/srt/utils/common.py @@ -1000,6 +1000,8 @@ def load_video(video_file: Union[str, bytes], use_gpu: bool = True): tmp_file.write(video_bytes) tmp_file.close() vr = VideoReader(tmp_file.name, ctx=ctx) + elif isinstance(video_file, (list, tuple, torch.Tensor, np.ndarray)): + vr = video_file else: raise ValueError(f"Unsupported video input type: {type(video_file)}") diff --git a/python/sglang/srt/utils/hf_transformers_utils.py b/python/sglang/srt/utils/hf_transformers_utils.py index 2efb1f7b5..b4fd6c734 100644 --- a/python/sglang/srt/utils/hf_transformers_utils.py +++ b/python/sglang/srt/utils/hf_transformers_utils.py @@ -62,6 +62,8 @@ from sglang.srt.configs import ( NemotronH_Nano_VL_V2_Config, NemotronHConfig, Olmo3Config, + Qwen3_5Config, + Qwen3_5MoeConfig, Qwen3NextConfig, Step3p5Config, Step3VLConfig, @@ -93,6 +95,8 @@ _CONFIG_REGISTRY: List[Type[PretrainedConfig]] = [ NemotronH_Nano_VL_V2_Config, NemotronHConfig, DeepseekVLV2Config, + Qwen3_5Config, + Qwen3_5MoeConfig, JetNemotronConfig, JetVLMConfig, KimiK25Config,