diff --git a/python/sglang/srt/model_loader/auto_loader.py b/python/sglang/srt/model_loader/auto_loader.py index 259c07dd4961..81e7b9f0da38 100644 --- a/python/sglang/srt/model_loader/auto_loader.py +++ b/python/sglang/srt/model_loader/auto_loader.py @@ -38,6 +38,7 @@ from torch.nn import Parameter from sglang.srt.layers.utils.common import get_layer_id +from sglang.srt.model_loader.weight_utils import default_weight_loader from sglang.srt.models.utils import AutoWeightsLoader, WeightsMapper __all__ = [ @@ -48,6 +49,7 @@ "STANDARD_GATE_UP_MAPPING", "STANDARD_STACKED_MAPPING", "LLAMA_STACKED_MAPPING", + "load_with_stacked_dispatch", "filter_pp_weights", "register_weight_remap", "get_weight_remap", @@ -148,6 +150,30 @@ def try_load( ) +def load_with_stacked_dispatch( + module: nn.Module, + weights: Iterable[tuple[str, torch.Tensor]], + mapping: StackedParamsDispatch, + *, + ignore_unexpected_suffixes: tuple[str, ...] = (".bias", ".kv_scale"), +) -> set[str]: + """Load submodule weights via stacked dispatch, then direct param loaders.""" + loaded: set[str] = set() + params_dict = dict(module.named_parameters()) + for name, tensor in weights: + target = mapping.try_load(name, tensor, params_dict) + if target is not None: + loaded.add(target) + continue + if name in params_dict: + wl = getattr(params_dict[name], "weight_loader", default_weight_loader) + wl(params_dict[name], tensor) + loaded.add(name) + elif not any(name.endswith(suffix) for suffix in ignore_unexpected_suffixes): + pass + return loaded + + # --------------------------------------------------------------------------- # Pipeline Parallel Weight Filter # --------------------------------------------------------------------------- diff --git a/python/sglang/srt/models/gemma2.py b/python/sglang/srt/models/gemma2.py index 8df231622d1d..3f04343260ca 100644 --- a/python/sglang/srt/models/gemma2.py +++ b/python/sglang/srt/models/gemma2.py @@ -92,6 +92,14 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: x, _ = self.down_proj(x) return x + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_GATE_UP_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING) + class Gemma2Attention(nn.Module): def __init__( @@ -200,6 +208,14 @@ def forward( output, _ = self.o_proj(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_QKV_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING) + class Gemma2DecoderLayer(nn.Module): def __init__( @@ -458,6 +474,13 @@ def get_attention_sliding_window_size(self): return get_attention_sliding_window_size(self.config) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) ("qkv_proj", "q_proj", "q"), @@ -498,5 +521,26 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): weight_loader(param, loaded_weight) loaded_params.add(name) + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import AutoWeightsLoader + + params_dict = dict(self.named_parameters()) + + def _prepare( + src: Iterable[Tuple[str, torch.Tensor]], + ) -> Iterable[Tuple[str, torch.Tensor]]: + for name, loaded_weight in src: + remapped = maybe_remap_kv_scale_name(name, params_dict) + if remapped is None: + continue + yield remapped, loaded_weight + + loader = AutoWeightsLoader( + self, + skip_prefixes=["lm_head."], + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + return loader.load_weights(_prepare(weights)) + EntryClass = Gemma2ForCausalLM diff --git a/python/sglang/srt/models/glm4.py b/python/sglang/srt/models/glm4.py index f1feda9e9ec3..c430e8d02d16 100644 --- a/python/sglang/srt/models/glm4.py +++ b/python/sglang/srt/models/glm4.py @@ -97,6 +97,14 @@ def forward( x, _ = self.down_proj(x) return x + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_GATE_UP_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING) + class Glm4Attention(nn.Module): def __init__( @@ -192,6 +200,14 @@ def forward( output, _ = self.o_proj(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_QKV_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING) + class Glm4DecoderLayer(nn.Module): """A single transformer layer. @@ -549,6 +565,13 @@ def end_layer(self): return self.model.end_layer def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) (".qkv_proj", ".q_proj", "q"), @@ -611,6 +634,41 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): else: logger.warning(f"Parameter {name} not found in params_dict") + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + AutoWeightsLoader, + filter_pp_weights, + ) + + if hasattr(self.model, "start_layer"): + weights = filter_pp_weights( + weights, self.model.start_layer, self.model.end_layer + ) + + skip_prefixes = [] + if self.config.tie_word_embeddings: + skip_prefixes.append("lm_head.") + + loader = AutoWeightsLoader( + self, + skip_prefixes=skip_prefixes, + skip_substrs=["projector"], + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + loaded = loader.load_weights(weights) + + if self.config.tie_word_embeddings: + params_dict = dict(self.named_parameters()) + if "lm_head.weight" in params_dict: + embed = dict(self.model.named_parameters()).get("embed_tokens.weight") + if embed is not None: + lm_head = params_dict["lm_head.weight"] + wl = getattr(lm_head, "weight_loader", default_weight_loader) + wl(lm_head, embed.data) + loaded.add("lm_head.weight") + + return loaded + def get_embed_and_head(self): return self.model.embed_tokens.weight, self.lm_head.weight diff --git a/python/sglang/srt/models/granite.py b/python/sglang/srt/models/granite.py index c5102822d675..e4b01e8fef07 100644 --- a/python/sglang/srt/models/granite.py +++ b/python/sglang/srt/models/granite.py @@ -119,6 +119,14 @@ def forward(self, x): x, _ = self.down_proj(x) return x + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_GATE_UP_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING) + class GraniteAttention(nn.Module): def __init__( @@ -220,6 +228,14 @@ def forward( output, _ = self.o_proj(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_QKV_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING) + class GraniteDecoderLayer(nn.Module): def __init__( @@ -420,6 +436,13 @@ def get_num_params(self): return len(params_dict) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) (".qkv_proj", ".q_proj", "q"), @@ -471,6 +494,33 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): weight_loader = getattr(param, "weight_loader", default_weight_loader) weight_loader(param, loaded_weight) + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import AutoWeightsLoader + + skip_prefixes = [] + if self.config.tie_word_embeddings: + skip_prefixes.append("lm_head.") + + loader = AutoWeightsLoader( + self, + skip_prefixes=skip_prefixes, + skip_substrs=["projector", "model.vision_tower"], + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + loaded = loader.load_weights(weights) + + if self.config.tie_word_embeddings: + params_dict = dict(self.named_parameters()) + if "lm_head.weight" in params_dict: + embed = dict(self.model.named_parameters()).get("embed_tokens.weight") + if embed is not None: + lm_head = params_dict["lm_head.weight"] + wl = getattr(lm_head, "weight_loader", default_weight_loader) + wl(lm_head, embed.data) + loaded.add("lm_head.weight") + + return loaded + def get_weights_by_name( self, name: str, truncate_size: int = 100, tp_size: int = 1 ) -> Optional[torch.Tensor]: diff --git a/python/sglang/srt/models/internlm2.py b/python/sglang/srt/models/internlm2.py index d376d96c03a2..8c119eeea363 100644 --- a/python/sglang/srt/models/internlm2.py +++ b/python/sglang/srt/models/internlm2.py @@ -80,10 +80,25 @@ def forward(self, x): x, _ = self.w2(x) return x + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + StackedParamsDispatch, + load_with_stacked_dispatch, + ) + + mapping = StackedParamsDispatch( + mappings=( + ("gate_up_proj", "w1", 0), + ("gate_up_proj", "w3", 1), + ) + ) + return load_with_stacked_dispatch(self, weights, mapping) + class InternLM2Attention(nn.Module): def __init__( self, + config: PretrainedConfig, hidden_size: int, num_heads: int, num_kv_heads: int, @@ -95,6 +110,7 @@ def __init__( prefix: str = "", ) -> None: super().__init__() + self.config = config self.hidden_size = hidden_size tp_size = get_parallel().tp_size self.total_num_heads = num_heads @@ -164,6 +180,34 @@ def forward( output, _ = self.wo(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + loaded: set[str] = set() + params_dict = dict(self.named_parameters()) + config = self.config + for name, tensor in weights: + if "wqkv" in name and name in params_dict: + param = params_dict[name] + kv_groups = config.num_attention_heads // config.num_key_value_heads + head_dim = config.hidden_size // config.num_attention_heads + loaded_weight = tensor.view( + -1, 2 + kv_groups, head_dim, tensor.shape[-1] + ) + wq, wk, wv = torch.split(loaded_weight, [kv_groups, 1, 1], dim=1) + wq = wq.reshape(-1, wq.shape[-1]) + wk = wk.reshape(-1, wk.shape[-1]) + wv = wv.reshape(-1, wv.shape[-1]) + weight_loader = param.weight_loader + weight_loader(param, wq, "q") + weight_loader(param, wk, "k") + weight_loader(param, wv, "v") + loaded.add(name) + continue + if name in params_dict: + wl = getattr(params_dict[name], "weight_loader", default_weight_loader) + wl(params_dict[name], tensor) + loaded.add(name) + return loaded + class InternLMDecoderLayer(nn.Module): def __init__( @@ -179,6 +223,7 @@ def __init__( rope_scaling = getattr(config, "rope_scaling", None) max_position_embeddings = getattr(config, "max_position_embeddings", 8192) self.attention = InternLM2Attention( + config=config, hidden_size=self.hidden_size, num_heads=config.num_attention_heads, num_kv_heads=config.num_key_value_heads, @@ -309,6 +354,13 @@ def forward( ) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) ("gate_up_proj", "w1", 0), @@ -355,5 +407,14 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): ) weight_loader(param, loaded_weight) + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import AutoWeightsLoader + + loader = AutoWeightsLoader( + self, + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + return loader.load_weights(weights) + EntryClass = InternLM2ForCausalLM diff --git a/python/sglang/srt/models/olmo.py b/python/sglang/srt/models/olmo.py index e1ff3e5d2989..429b731c16cd 100644 --- a/python/sglang/srt/models/olmo.py +++ b/python/sglang/srt/models/olmo.py @@ -124,6 +124,14 @@ def forward( output, _ = self.o_proj(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_QKV_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING) + class OlmoMLP(nn.Module): """ @@ -173,6 +181,14 @@ def forward( x, _ = self.down_proj(x) return x + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_GATE_UP_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING) + class OlmoDecoderLayer(nn.Module): """ @@ -340,6 +356,13 @@ def forward( ) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) ("qkv_proj", "q_proj", "q"), @@ -377,5 +400,31 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): weight_loader = getattr(param, "weight_loader", default_weight_loader) weight_loader(param, loaded_weight) + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import AutoWeightsLoader + + skip_prefixes = [] + if self.config.tie_word_embeddings: + skip_prefixes.append("lm_head.") + + loader = AutoWeightsLoader( + self, + skip_prefixes=skip_prefixes, + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + loaded = loader.load_weights(weights) + + if self.config.tie_word_embeddings: + params_dict = dict(self.named_parameters(remove_duplicate=False)) + if "lm_head.weight" in params_dict: + embed = dict(self.model.named_parameters()).get("embed_tokens.weight") + if embed is not None: + lm_head = params_dict["lm_head.weight"] + wl = getattr(lm_head, "weight_loader", default_weight_loader) + wl(lm_head, embed.data) + loaded.add("lm_head.weight") + + return loaded + EntryClass = OlmoForCausalLM diff --git a/python/sglang/srt/models/olmo2.py b/python/sglang/srt/models/olmo2.py index f606efaf0e14..f3b1bbc29f8e 100644 --- a/python/sglang/srt/models/olmo2.py +++ b/python/sglang/srt/models/olmo2.py @@ -212,6 +212,14 @@ def forward( output, _ = self.o_proj(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_QKV_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING) + class Olmo2MLP(nn.Module): """ @@ -261,6 +269,14 @@ def forward( x, _ = self.down_proj(x) return x + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_GATE_UP_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_GATE_UP_MAPPING) + class Olmo2DecoderLayer(nn.Module): """ @@ -441,6 +457,13 @@ def forward( ) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) ("qkv_proj", "q_proj", "q"), @@ -481,5 +504,31 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): weight_loader = getattr(param, "weight_loader", default_weight_loader) weight_loader(param, loaded_weight) + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import AutoWeightsLoader + + skip_prefixes = [] + if self.config.tie_word_embeddings: + skip_prefixes.append("lm_head.") + + loader = AutoWeightsLoader( + self, + skip_prefixes=skip_prefixes, + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + loaded = loader.load_weights(weights) + + if self.config.tie_word_embeddings: + params_dict = dict(self.named_parameters(remove_duplicate=False)) + if "lm_head.weight" in params_dict: + embed = dict(self.model.named_parameters()).get("embed_tokens.weight") + if embed is not None: + lm_head = params_dict["lm_head.weight"] + wl = getattr(lm_head, "weight_loader", default_weight_loader) + wl(lm_head, embed.data) + loaded.add("lm_head.weight") + + return loaded + EntryClass = Olmo2ForCausalLM diff --git a/python/sglang/srt/models/qwen3.py b/python/sglang/srt/models/qwen3.py index 743232dc3817..01d929e94c85 100644 --- a/python/sglang/srt/models/qwen3.py +++ b/python/sglang/srt/models/qwen3.py @@ -310,6 +310,32 @@ def forward( output, _ = self.o_proj(attn_output) return output + def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + STANDARD_QKV_MAPPING, + load_with_stacked_dispatch, + ) + + return load_with_stacked_dispatch(self, weights, STANDARD_QKV_MAPPING) + + +def _prepare_qwen3_checkpoint_weights( + weights: Iterable[Tuple[str, torch.Tensor]], + params_dict: dict[str, torch.nn.Parameter], +) -> Iterable[Tuple[str, torch.Tensor]]: + for name, loaded_weight in weights: + if not name.startswith("model.") and ( + name.startswith("layers.") + or name.startswith("embed_tokens.") + or name.startswith("norm.") + ): + name = add_prefix(name, "model") + if "scale" in name: + name = maybe_remap_kv_scale_name(name, params_dict) + if name is None: + continue + yield name, loaded_weight + class Qwen3DecoderLayer(nn.Module): def __init__( @@ -628,6 +654,13 @@ def end_layer(self): return self.model.end_layer def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) ("qkv_proj", "q_proj", "q"), @@ -703,6 +736,43 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): else: logger.warning(f"Parameter {name} not found in params_dict") + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import ( + AutoWeightsLoader, + filter_pp_weights, + ) + + params_dict = dict(self.named_parameters()) + weights = _prepare_qwen3_checkpoint_weights(weights, params_dict) + + if hasattr(self.model, "start_layer"): + weights = filter_pp_weights( + weights, self.model.start_layer, self.model.end_layer + ) + + skip_prefixes = [] + if self.config.tie_word_embeddings: + skip_prefixes.append("lm_head.") + + loader = AutoWeightsLoader( + self, + skip_prefixes=skip_prefixes, + skip_substrs=["projector", "model.vision_tower"], + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + loaded = loader.load_weights(weights) + + if self.config.tie_word_embeddings and self.pp_group.is_last_rank: + if "lm_head.weight" in params_dict: + embed = dict(self.model.named_parameters()).get("embed_tokens.weight") + if embed is not None: + lm_head = params_dict["lm_head.weight"] + wl = getattr(lm_head, "weight_loader", default_weight_loader) + wl(lm_head, embed.data) + loaded.add("lm_head.weight") + + return loaded + def get_embed_and_head(self): return self.model.embed_tokens.weight, self.lm_head.weight diff --git a/python/sglang/srt/models/qwen3_classification.py b/python/sglang/srt/models/qwen3_classification.py index 352f7c7d8434..9adfff0d6d75 100644 --- a/python/sglang/srt/models/qwen3_classification.py +++ b/python/sglang/srt/models/qwen3_classification.py @@ -28,7 +28,7 @@ from sglang.srt.layers.quantization.base_config import QuantizationConfig 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 import Qwen3Model +from sglang.srt.models.qwen3 import Qwen3Model, _prepare_qwen3_checkpoint_weights from sglang.srt.utils import add_prefix logger = logging.getLogger(__name__) @@ -75,6 +75,13 @@ def forward( ) def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): + from sglang.srt.environ import envs + + if envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.get(): + return self._load_weights_v2(weights) + return self._legacy_load_weights(weights) + + def _legacy_load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): stacked_params_mapping = [ # (param_name, shard_name, shard_id) ("qkv_proj", "q_proj", "q"), @@ -124,6 +131,20 @@ def load_weights(self, weights: Iterable[Tuple[str, torch.Tensor]]): else: logger.warning(f"Parameter {name} not found in params_dict") + def _load_weights_v2(self, weights: Iterable[Tuple[str, torch.Tensor]]) -> set[str]: + from sglang.srt.model_loader.auto_loader import AutoWeightsLoader + + params_dict = dict(self.named_parameters()) + weights = _prepare_qwen3_checkpoint_weights(weights, params_dict) + + loader = AutoWeightsLoader( + self, + skip_prefixes=["lm_head."], + skip_substrs=["projector"], + ignore_unexpected_suffixes=[".bias", ".kv_scale"], + ) + return loader.load_weights(weights) + class Qwen3ForSequenceClassification(Qwen3ForPooledOutput): def __init__( diff --git a/test/manual/test_weight_loader_v2_equiv.py b/test/manual/test_weight_loader_v2_equiv.py index 595b4339f996..bc39a231e2b7 100644 --- a/test/manual/test_weight_loader_v2_equiv.py +++ b/test/manual/test_weight_loader_v2_equiv.py @@ -27,6 +27,7 @@ from sglang.test.test_utils import publish_build_topology MODEL = "Qwen/Qwen2-0.5B" +QWEN3_MODEL = "Qwen/Qwen3-0.6B" def _init_model_parallel() -> None: @@ -52,6 +53,10 @@ def _init_model_parallel() -> None: def _load_qwen2_native(v2: bool) -> torch.nn.Module: + return _load_native_model(MODEL, v2) + + +def _load_native_model(model_path: str, v2: bool) -> torch.nn.Module: from sglang.srt.configs.device_config import DeviceConfig from sglang.srt.configs.load_config import LoadConfig from sglang.srt.configs.model_config import ModelConfig @@ -60,7 +65,7 @@ def _load_qwen2_native(v2: bool) -> torch.nn.Module: from sglang.srt.utils import get_device server_args = ServerArgs( - model_path=MODEL, + model_path=model_path, dtype=torch.float16, trust_remote_code=True, ) @@ -108,6 +113,28 @@ def test_qwen2_v1_v2_state_dict_identical(self): msg=name, ) + @unittest.skipIf(not torch.cuda.is_available(), "needs GPU") + def test_qwen3_v1_v2_state_dict_identical(self): + model_v1 = _load_native_model(QWEN3_MODEL, v2=False) + state_v1 = _state_dict_cpu(model_v1) + del model_v1 + torch.cuda.empty_cache() + + model_v2 = _load_native_model(QWEN3_MODEL, v2=True) + state_v2 = _state_dict_cpu(model_v2) + del model_v2 + torch.cuda.empty_cache() + + self.assertEqual(set(state_v1.keys()), set(state_v2.keys())) + for name in sorted(state_v1.keys()): + torch.testing.assert_close( + state_v1[name], + state_v2[name], + rtol=0, + atol=0, + msg=name, + ) + if __name__ == "__main__": unittest.main() diff --git a/test/registered/model_loading/test_weight_loader_v2_e2e.py b/test/registered/model_loading/test_weight_loader_v2_e2e.py index 005efe2ca005..82d95c728c2e 100644 --- a/test/registered/model_loading/test_weight_loader_v2_e2e.py +++ b/test/registered/model_loading/test_weight_loader_v2_e2e.py @@ -25,6 +25,7 @@ from sglang.test.test_utils import CustomTestCase MODEL = "Qwen/Qwen2-0.5B" +QWEN3_MODEL = "Qwen/Qwen3-0.6B" SHORT_PROMPT = "The capital of the United Kingdom is" @@ -65,6 +66,28 @@ def test_qwen2_native_v1_v2_generation_match(self): debug_text="qwen2 native v1 vs v2 weight loader", ) + def test_qwen3_native_v1_v2_generation_match(self): + prompts = [SHORT_PROMPT] + max_new_tokens = 32 + kwargs = self._runner_kwargs() + + with envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.override(False): + with SRTRunner(QWEN3_MODEL, **kwargs) as runner_v1: + out_v1 = runner_v1.forward(prompts, max_new_tokens=max_new_tokens) + + with envs.SGLANG_ENABLE_WEIGHT_LOADER_V2.override(True): + with SRTRunner(QWEN3_MODEL, **kwargs) as runner_v2: + out_v2 = runner_v2.forward(prompts, max_new_tokens=max_new_tokens) + + check_close_model_outputs( + hf_outputs=out_v1, + srt_outputs=out_v2, + prefill_tolerance=1e-6, + decode_tolerance=1e-6, + rouge_l_tolerance=1.0, + debug_text="qwen3 native v1 vs v2 weight loader", + ) + def test_transformers_impl_loads_and_generates(self): prompts = [SHORT_PROMPT] max_new_tokens = 16 diff --git a/test/registered/unit/model_loader/test_stacked_params_dispatch.py b/test/registered/unit/model_loader/test_stacked_params_dispatch.py new file mode 100644 index 000000000000..bd1d95848343 --- /dev/null +++ b/test/registered/unit/model_loader/test_stacked_params_dispatch.py @@ -0,0 +1,126 @@ +# Copyright 2023-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. +# ============================================================================== + +from sglang.test.ci.ci_register import register_cpu_ci + +register_cpu_ci(est_time=6, suite="base-b-test-cpu") + +import unittest + +import torch +from torch import nn + +from sglang.srt.model_loader.auto_loader import ( + STANDARD_GATE_UP_MAPPING, + STANDARD_QKV_MAPPING, + StackedParamsDispatch, + load_with_stacked_dispatch, +) + + +class _ParamWithLoader(nn.Parameter): + def __new__(cls, data, weight_loader=None): + param = super().__new__(cls, data) + param.weight_loader = weight_loader or (lambda p, t, *a: p.copy_(t)) + return param + + +class TestStackedParamsDispatch(unittest.TestCase): + def test_try_load_qkv_shard(self): + qkv = _ParamWithLoader(torch.zeros(6, 4)) + params = {"qkv_proj.weight": qkv} + calls = [] + + def wl(param, tensor, shard_id): + calls.append(shard_id) + + qkv.weight_loader = wl + + target = STANDARD_QKV_MAPPING.try_load( + "q_proj.weight", torch.ones(2, 4), params + ) + self.assertEqual(target, "qkv_proj.weight") + self.assertEqual(calls, ["q"]) + + def test_try_load_returns_target_when_param_missing(self): + target = STANDARD_QKV_MAPPING.try_load("q_proj.weight", torch.ones(2, 4), {}) + self.assertEqual(target, "qkv_proj.weight") + + def test_try_load_no_match(self): + self.assertIsNone( + STANDARD_QKV_MAPPING.try_load("o_proj.weight", torch.ones(2, 4), {}) + ) + + def test_gate_up_mapping(self): + gate_up = _ParamWithLoader(torch.zeros(8, 4)) + params = {"gate_up_proj.weight": gate_up} + shard_ids = [] + + def wl(param, tensor, shard_id): + shard_ids.append(shard_id) + + gate_up.weight_loader = wl + + target = STANDARD_GATE_UP_MAPPING.try_load( + "up_proj.weight", torch.ones(4, 4), params + ) + self.assertEqual(target, "gate_up_proj.weight") + self.assertEqual(shard_ids, [1]) + + def test_load_with_stacked_dispatch_direct_param(self): + linear = nn.Linear(3, 2, bias=False) + linear.weight = nn.Parameter(torch.zeros(3, 2), requires_grad=False) + linear.weight.weight_loader = lambda p, t: p.data.copy_(t) + module = nn.Module() + module.down_proj = linear + + loaded = load_with_stacked_dispatch( + module, + [("down_proj.weight", torch.ones(3, 2))], + StackedParamsDispatch(mappings=()), + ) + self.assertIn("down_proj.weight", loaded) + self.assertTrue(torch.allclose(linear.weight.data, torch.ones(3, 2))) + + def test_load_with_stacked_dispatch_stacked_then_direct(self): + qkv_linear = nn.Linear(2, 6, bias=False) + o_linear = nn.Linear(2, 3, bias=False) + qkv_linear.weight = nn.Parameter(torch.zeros(6, 2), requires_grad=False) + o_linear.weight = nn.Parameter(torch.zeros(3, 2), requires_grad=False) + qkv_calls = [] + + def qkv_wl(param, tensor, shard_id): + qkv_calls.append(shard_id) + + qkv_linear.weight.weight_loader = qkv_wl + o_linear.weight.weight_loader = lambda p, t: p.data.copy_(t) + + module = nn.Module() + module.qkv_proj = qkv_linear + module.o_proj = o_linear + + loaded = load_with_stacked_dispatch( + module, + [ + ("q_proj.weight", torch.ones(2, 2)), + ("o_proj.weight", torch.full((3, 2), 2.0)), + ], + STANDARD_QKV_MAPPING, + ) + self.assertEqual(loaded, {"qkv_proj.weight", "o_proj.weight"}) + self.assertEqual(qkv_calls, ["q"]) + + +if __name__ == "__main__": + unittest.main()