diff --git a/vllm/model_executor/models/mistral_large_3.py b/vllm/model_executor/models/mistral_large_3.py index ff7e9b60c1d..603ce5c0f01 100644 --- a/vllm/model_executor/models/mistral_large_3.py +++ b/vllm/model_executor/models/mistral_large_3.py @@ -2,62 +2,91 @@ # SPDX-FileCopyrightText: Copyright contributors to the vLLM project from collections.abc import Iterable -import regex as re +import regex import torch from vllm.model_executor.models.deepseek_v2 import DeepseekV3ForCausalLM +from vllm.model_executor.models.utils import AutoWeightsLoader, WeightsMapper class MistralLarge3ForCausalLM(DeepseekV3ForCausalLM): - # fmt: off - remapping = { - r"layers\.(\d+)\.attention_norm\.weight": r"model.layers.\1.input_layernorm.weight", # noqa: E501 - r"layers\.(\d+)\.attention\.wq_a\.(\w+)": r"model.layers.\1.self_attn.q_a_proj.\2", # noqa: E501 - r"layers\.(\d+)\.attention\.q_a_norm\.weight": r"model.layers.\1.self_attn.q_a_layernorm.weight", # noqa: E501 - r"layers\.(\d+)\.attention\.wq_b\.(\w+)": r"model.layers.\1.self_attn.q_b_proj.\2", # noqa: E501 - r"layers\.(\d+)\.attention\.wkv_a_with_mqa\.(\w+)": r"model.layers.\1.self_attn.kv_a_proj_with_mqa.\2", # noqa: E501 - r"layers\.(\d+)\.attention\.kv_a_norm\.weight": r"model.layers.\1.self_attn.kv_a_layernorm.weight", # noqa: E501 - r"layers\.(\d+)\.attention\.wkv_b\.(\w+)": r"model.layers.\1.self_attn.kv_b_proj.\2", # noqa: E501 - r"layers\.(\d+)\.attention\.wo\.(\w+)": r"model.layers.\1.self_attn.o_proj.\2", # noqa: E501 - r"layers\.(\d+)\.ffn_norm\.weight": r"model.layers.\1.post_attention_layernorm.weight", # noqa: E501 - r"layers\.(\d+)\.feed_forward\.w1\.(\w+)": r"model.layers.\1.mlp.gate_proj.\2", # noqa: E501 - r"layers\.(\d+)\.feed_forward\.w2\.(\w+)": r"model.layers.\1.mlp.down_proj.\2", # noqa: E501 - r"layers\.(\d+)\.feed_forward\.w3\.(\w+)": r"model.layers.\1.mlp.up_proj.\2", # noqa: E501 - r"layers\.(\d+)\.gate\.weight": r"model.layers.\1.mlp.gate.weight", # noqa: E501 - r"layers\.(\d+)\.shared_experts\.w1\.(\w+)": r"model.layers.\1.mlp.shared_experts.gate_proj.\2", # noqa: E501 - r"layers\.(\d+)\.shared_experts\.w2\.(\w+)": r"model.layers.\1.mlp.shared_experts.down_proj.\2", # noqa: E501 - r"layers\.(\d+)\.shared_experts\.w3\.(\w+)": r"model.layers.\1.mlp.shared_experts.up_proj.\2", # noqa: E501 - r"layers\.(\d+)\.experts\.(\d+)\.w1\.(\w+)": r"model.layers.\1.mlp.experts.\2.gate_proj.\3", # noqa: E501 - r"layers\.(\d+)\.experts\.(\d+)\.w2\.(\w+)": r"model.layers.\1.mlp.experts.\2.down_proj.\3", # noqa: E501 - r"layers\.(\d+)\.experts\.(\d+)\.w3\.(\w+)": r"model.layers.\1.mlp.experts.\2.up_proj.\3", # noqa: E501 - r"norm\.weight": "model.norm.weight", # noqa: E501 - r"tok_embeddings\.weight": "model.embed_tokens.weight", # noqa: E501 - r"output\.weight": "lm_head.weight", # noqa: E501 - } - # fmt: on + # WeightsMapper applies all matching patterns sequentially (no break on first + # match). This is safe here because every pattern is anchored at both ends + # (\A...\Z) and after substitution the resulting key always starts with + # "model." or "lm_head.", so no later pattern can accidentally match again. + hf_to_vllm_mapper = WeightsMapper( + orig_to_new_regex={ # noqa: B950 + regex.compile( + r"\Alayers\.(\d+)\.attention_norm\.weight\Z" + ): r"model.layers.\1.input_layernorm.weight", + regex.compile( + r"\Alayers\.(\d+)\.attention\.wq_a\.(\w+)\Z" + ): r"model.layers.\1.self_attn.q_a_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.attention\.q_a_norm\.weight\Z" + ): r"model.layers.\1.self_attn.q_a_layernorm.weight", + regex.compile( + r"\Alayers\.(\d+)\.attention\.wq_b\.(\w+)\Z" + ): r"model.layers.\1.self_attn.q_b_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.attention\.wkv_a_with_mqa\.(\w+)\Z" + ): r"model.layers.\1.self_attn.kv_a_proj_with_mqa.\2", + regex.compile( + r"\Alayers\.(\d+)\.attention\.kv_a_norm\.weight\Z" + ): r"model.layers.\1.self_attn.kv_a_layernorm.weight", + regex.compile( + r"\Alayers\.(\d+)\.attention\.wkv_b\.(\w+)\Z" + ): r"model.layers.\1.self_attn.kv_b_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.attention\.wo\.(\w+)\Z" + ): r"model.layers.\1.self_attn.o_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.ffn_norm\.weight\Z" + ): r"model.layers.\1.post_attention_layernorm.weight", + regex.compile( + r"\Alayers\.(\d+)\.feed_forward\.w1\.(\w+)\Z" + ): r"model.layers.\1.mlp.gate_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.feed_forward\.w2\.(\w+)\Z" + ): r"model.layers.\1.mlp.down_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.feed_forward\.w3\.(\w+)\Z" + ): r"model.layers.\1.mlp.up_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.gate\.weight\Z" + ): r"model.layers.\1.mlp.gate.weight", + regex.compile( + r"\Alayers\.(\d+)\.shared_experts\.w1\.(\w+)\Z" + ): r"model.layers.\1.mlp.shared_experts.gate_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.shared_experts\.w2\.(\w+)\Z" + ): r"model.layers.\1.mlp.shared_experts.down_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.shared_experts\.w3\.(\w+)\Z" + ): r"model.layers.\1.mlp.shared_experts.up_proj.\2", + regex.compile( + r"\Alayers\.(\d+)\.experts\.(\d+)\.w1\.(\w+)\Z" + ): r"model.layers.\1.mlp.experts.\2.gate_proj.\3", + regex.compile( + r"\Alayers\.(\d+)\.experts\.(\d+)\.w2\.(\w+)\Z" + ): r"model.layers.\1.mlp.experts.\2.down_proj.\3", + regex.compile( + r"\Alayers\.(\d+)\.experts\.(\d+)\.w3\.(\w+)\Z" + ): r"model.layers.\1.mlp.experts.\2.up_proj.\3", + regex.compile(r"\Anorm\.weight\Z"): "model.norm.weight", + regex.compile(r"\Atok_embeddings\.weight\Z"): "model.embed_tokens.weight", + regex.compile(r"\Aoutput\.weight\Z"): "lm_head.weight", + }, + orig_to_new_suffix={ + ".qscale_act": ".input_scale", + ".qscale_weight": ".weight_scale", + }, + ) + # Bypass super().load_weights() and construct AutoWeightsLoader(self) + # directly (same pattern as Qwen2ForCausalLM). Any logic in the parent + # class's load_weights is a thin wrapper around AutoWeightsLoader, and + # we must apply hf_to_vllm_mapper before the loader walks the tree. def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]: - return super().load_weights(map(self._remap_mistral_to_ds, weights)) - - def _remap_mistral_to_ds( - self, weight: tuple[str, torch.Tensor] - ) -> tuple[str, torch.Tensor]: - """Remap Mistral parameters to DeepseekV2 parameters.""" - name, loaded_weight = weight - - for k, v in self.remapping.items(): - match = re.fullmatch(k, name) - if match: - name = re.sub(k, v, name) - break - else: - raise ValueError(f"Cannot remap {name}") - - # Remapping scale names. We could do this in the regex above but it - # would triple the number of lines for most layers. - if name.endswith(".qscale_act"): - name = re.sub(r"\.qscale_act$", ".input_scale", name) - elif name.endswith(".qscale_weight"): - name = re.sub(r"\.qscale_weight$", ".weight_scale", name) - - return name, loaded_weight + loader = AutoWeightsLoader(self) + return loader.load_weights(weights, mapper=self.hf_to_vllm_mapper) diff --git a/vllm/model_executor/models/mistral_large_3_eagle.py b/vllm/model_executor/models/mistral_large_3_eagle.py index bde5bc9451f..8ace01205d0 100644 --- a/vllm/model_executor/models/mistral_large_3_eagle.py +++ b/vllm/model_executor/models/mistral_large_3_eagle.py @@ -5,6 +5,7 @@ import copy from collections.abc import Iterable from functools import partial +import regex import torch import torch.nn as nn @@ -22,7 +23,7 @@ from vllm.model_executor.models.deepseek_v2 import ( from vllm.model_executor.models.mistral_large_3 import MistralLarge3ForCausalLM from .interfaces import SupportsMultiModal -from .utils import make_empty_intermediate_tensors_factory, maybe_prefix +from .utils import WeightsMapper, make_empty_intermediate_tensors_factory, maybe_prefix logger = init_logger(__name__) @@ -107,11 +108,13 @@ class EagleMistralLarge3Model(DeepseekV2Model): class EagleMistralLarge3ForCausalLM(MistralLarge3ForCausalLM): - remapping = MistralLarge3ForCausalLM.remapping | { - r"eagle_linear\.weight": r"model.fc.weight", - r"eagle_linear\.qscale_act": r"model.fc.input_scale", - r"eagle_linear\.qscale_weight": r"model.fc.weight_scale", - } + hf_to_vllm_mapper = MistralLarge3ForCausalLM.hf_to_vllm_mapper | WeightsMapper( + orig_to_new_regex={ + regex.compile(r"\Aeagle_linear\.weight\Z"): r"model.fc.weight", + regex.compile(r"\Aeagle_linear\.qscale_act\Z"): r"model.fc.input_scale", + regex.compile(r"\Aeagle_linear\.qscale_weight\Z"): r"model.fc.weight_scale", + }, + ) def __init__(self, *, vllm_config: VllmConfig, prefix: str = ""): target_layer_num = vllm_config.model_config.get_num_layers(