[Model] Support standalone text-only Qwen3.5 checkpoints (#32401)
Co-authored-by: liyucheng09 <liyucheng09@gmail.com>
This commit is contained in:
co-authored by
liyucheng09
parent
ec4a7fa2b7
commit
fc8b328f5c
@@ -37,7 +37,12 @@ from sglang.srt.configs.nano_nemotron_vl import (
|
||||
)
|
||||
from sglang.srt.configs.nemotron_h import NemotronHConfig, NemotronHPuzzleConfig
|
||||
from sglang.srt.configs.olmo3 import Olmo3Config
|
||||
from sglang.srt.configs.qwen3_5 import Qwen3_5Config, Qwen3_5MoeConfig
|
||||
from sglang.srt.configs.qwen3_5 import (
|
||||
Qwen3_5Config,
|
||||
Qwen3_5MoeConfig,
|
||||
Qwen3_5MoeTextConfig,
|
||||
Qwen3_5TextConfig,
|
||||
)
|
||||
from sglang.srt.configs.qwen3_asr import Qwen3ASRConfig
|
||||
from sglang.srt.configs.qwen3_next import Qwen3NextConfig
|
||||
from sglang.srt.configs.step3_vl import (
|
||||
@@ -71,6 +76,8 @@ __all__ = [
|
||||
"Qwen3NextConfig",
|
||||
"Qwen3_5Config",
|
||||
"Qwen3_5MoeConfig",
|
||||
"Qwen3_5TextConfig",
|
||||
"Qwen3_5MoeTextConfig",
|
||||
"InternS2PreviewConfig",
|
||||
"DotsVLMConfig",
|
||||
"DotsOCRConfig",
|
||||
|
||||
@@ -647,6 +647,8 @@ class ModelConfig:
|
||||
if is_draft_model and self.hf_config.architectures[0] in [
|
||||
"Qwen3_5ForConditionalGeneration",
|
||||
"Qwen3_5MoeForConditionalGeneration",
|
||||
"Qwen3_5ForCausalLM",
|
||||
"Qwen3_5MoeForCausalLM",
|
||||
"InternS2PreviewForConditionalGeneration",
|
||||
]:
|
||||
self.hf_config.architectures[0] = "Qwen3_5ForCausalLMMTP"
|
||||
|
||||
@@ -0,0 +1,208 @@
|
||||
# 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.
|
||||
# ==============================================================================
|
||||
|
||||
import logging
|
||||
from typing import Iterable, Optional, Set, Tuple, Union
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
|
||||
from sglang.srt.distributed import get_pp_group
|
||||
from sglang.srt.eplb.expert_location import ModelConfigForExpertLocation
|
||||
from sglang.srt.layers.logits_processor import LogitsProcessor
|
||||
from sglang.srt.layers.quantization.base_config import QuantizationConfig
|
||||
from sglang.srt.layers.utils import PPMissingLayer
|
||||
from sglang.srt.layers.vocab_parallel_embedding import ParallelLMHead
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch, PPProxyTensors
|
||||
from sglang.srt.model_loader.weight_utils import default_weight_loader
|
||||
from sglang.srt.models import qwen3_5
|
||||
from sglang.srt.models.qwen2_moe import Qwen2MoeSparseMoeBlock
|
||||
from sglang.srt.runtime_context import get_server_args
|
||||
from sglang.srt.utils import LazyValue, add_prefix
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_MODEL_PREFIX = "model."
|
||||
|
||||
|
||||
class Qwen3_5ForCausalLM(nn.Module):
|
||||
body_cls = qwen3_5.Qwen3_5ForCausalLM
|
||||
|
||||
packed_modules_mapping = qwen3_5.Qwen3_5ForCausalLM.packed_modules_mapping
|
||||
supported_lora_modules = qwen3_5.Qwen3_5ForCausalLM.supported_lora_modules
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.config = config
|
||||
self.quant_config = quant_config
|
||||
self.pp_group = get_pp_group()
|
||||
|
||||
if quant_config is not None and hasattr(quant_config, "packed_modules_mapping"):
|
||||
quant_config.packed_modules_mapping = self.packed_modules_mapping
|
||||
|
||||
self.model = self.body_cls(
|
||||
config=config,
|
||||
quant_config=quant_config,
|
||||
prefix=add_prefix("model", prefix),
|
||||
)
|
||||
|
||||
if self.pp_group.is_last_rank:
|
||||
if self.pp_group.world_size == 1 and 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,
|
||||
org_num_embeddings=config.vocab_size,
|
||||
prefix=add_prefix("lm_head", prefix),
|
||||
use_attn_tp_group=get_server_args().enable_dp_lm_head,
|
||||
)
|
||||
else:
|
||||
self.lm_head = PPMissingLayer()
|
||||
|
||||
self.logits_processor = LogitsProcessor(config)
|
||||
|
||||
# Text-only checkpoints retain mrope_section, but identical position rows
|
||||
# make its rotary embedding equivalent to 1-D RoPE.
|
||||
rope_config = getattr(config, "rope_parameters", None) or getattr(
|
||||
config, "rope_scaling", None
|
||||
)
|
||||
self.is_mrope_enabled = bool(rope_config) and "mrope_section" in rope_config
|
||||
|
||||
self.capture_aux_hidden_states = False
|
||||
|
||||
@property
|
||||
def start_layer(self) -> int:
|
||||
return self.model.start_layer
|
||||
|
||||
@property
|
||||
def end_layer(self) -> int:
|
||||
return self.model.end_layer
|
||||
|
||||
def get_input_embeddings(self) -> nn.Embedding:
|
||||
return self.model.embed_tokens
|
||||
|
||||
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()
|
||||
|
||||
@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,
|
||||
**kwargs,
|
||||
) -> Union[torch.Tensor, PPProxyTensors]:
|
||||
if self.is_mrope_enabled:
|
||||
positions = forward_batch.mrope_positions
|
||||
|
||||
hidden_states = self.model(
|
||||
input_ids,
|
||||
positions,
|
||||
forward_batch,
|
||||
input_embeds,
|
||||
pp_proxy_tensors=pp_proxy_tensors,
|
||||
)
|
||||
|
||||
if not self.pp_group.is_last_rank:
|
||||
return hidden_states
|
||||
|
||||
aux_hidden_states = None
|
||||
if isinstance(hidden_states, tuple):
|
||||
hidden_states, aux_hidden_states = hidden_states
|
||||
|
||||
return self.logits_processor(
|
||||
input_ids, hidden_states, self.lm_head, forward_batch, aux_hidden_states
|
||||
)
|
||||
|
||||
def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> Set[str]:
|
||||
params_dict = dict(self.named_parameters())
|
||||
loaded_params: Set[str] = set()
|
||||
|
||||
body_weights = []
|
||||
for name, loaded_weight in weights:
|
||||
if name.startswith(_MODEL_PREFIX):
|
||||
body_weights.append((name[len(_MODEL_PREFIX) :], loaded_weight))
|
||||
elif name == "lm_head.weight":
|
||||
if self.config.tie_word_embeddings:
|
||||
continue
|
||||
if "lm_head.weight" not in params_dict:
|
||||
continue
|
||||
param = params_dict["lm_head.weight"]
|
||||
weight_loader = getattr(param, "weight_loader", default_weight_loader)
|
||||
weight_loader(param, loaded_weight)
|
||||
loaded_params.add("lm_head.weight")
|
||||
|
||||
body_loaded = self.model.load_weights(body_weights)
|
||||
loaded_params.update(f"{_MODEL_PREFIX}{n}" for n in body_loaded)
|
||||
|
||||
if self.config.tie_word_embeddings and self.pp_group.is_last_rank:
|
||||
loaded_params.add("lm_head.weight")
|
||||
|
||||
return loaded_params
|
||||
|
||||
|
||||
class Qwen3_5MoeForCausalLM(Qwen3_5ForCausalLM):
|
||||
body_cls = qwen3_5.Qwen3_5MoeForCausalLM
|
||||
|
||||
packed_modules_mapping = qwen3_5.Qwen3_5MoeForCausalLM.packed_modules_mapping
|
||||
supported_lora_modules = qwen3_5.Qwen3_5MoeForCausalLM.supported_lora_modules
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config,
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
prefix: str = "",
|
||||
) -> None:
|
||||
super().__init__(config=config, quant_config=quant_config, prefix=prefix)
|
||||
|
||||
self._routed_experts_weights_of_layer = LazyValue(
|
||||
lambda: {
|
||||
layer_id: layer.mlp.get_moe_weights()
|
||||
for layer_id, layer in enumerate(self.model.layers)
|
||||
if isinstance(layer.mlp, Qwen2MoeSparseMoeBlock)
|
||||
}
|
||||
)
|
||||
|
||||
@property
|
||||
def routed_experts_weights_of_layer(self):
|
||||
return self._routed_experts_weights_of_layer.value
|
||||
|
||||
@classmethod
|
||||
def get_model_config_for_expert_location(cls, config):
|
||||
return ModelConfigForExpertLocation(
|
||||
num_layers=config.num_hidden_layers,
|
||||
num_logical_experts=config.num_experts,
|
||||
num_groups=None,
|
||||
)
|
||||
|
||||
|
||||
EntryClass = [Qwen3_5MoeForCausalLM, Qwen3_5ForCausalLM]
|
||||
@@ -56,6 +56,8 @@ from sglang.srt.configs import (
|
||||
Olmo3Config,
|
||||
Qwen3_5Config,
|
||||
Qwen3_5MoeConfig,
|
||||
Qwen3_5MoeTextConfig,
|
||||
Qwen3_5TextConfig,
|
||||
Qwen3NextConfig,
|
||||
Step3p5Config,
|
||||
Step3p7Config,
|
||||
@@ -109,6 +111,8 @@ _CONFIG_REGISTRY: Dict[str, Type[PretrainedConfig]] = {
|
||||
DeepseekVLV2Config,
|
||||
Qwen3_5Config,
|
||||
Qwen3_5MoeConfig,
|
||||
Qwen3_5TextConfig,
|
||||
Qwen3_5MoeTextConfig,
|
||||
InternS2PreviewConfig,
|
||||
JetNemotronConfig,
|
||||
JetVLMConfig,
|
||||
|
||||
Reference in New Issue
Block a user