From 77fcd5fc6376cf38e06f488d275658c5528f1319 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 03:05:10 +0000 Subject: [PATCH 01/21] add mistral large 3 codes Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../models/core/mistral_large_3/README.md | 54 ++++ examples/models/core/mistral_large_3/test.py | 37 +++ requirements.txt | 1 + tensorrt_llm/_torch/model_config.py | 25 ++ .../_torch/models/checkpoints/__init__.py | 29 +- .../models/checkpoints/hf/weight_loader.py | 2 + .../models/checkpoints/mistral/__init__.py | 0 .../checkpoints/mistral/checkpoint_loader.py | 107 +++++++ .../checkpoints/mistral/config_loader.py | 301 ++++++++++++++++++ .../checkpoints/mistral/weight_mapper.py | 132 ++++++++ 10 files changed, 683 insertions(+), 5 deletions(-) create mode 100644 examples/models/core/mistral_large_3/README.md create mode 100644 examples/models/core/mistral_large_3/test.py create mode 100644 tensorrt_llm/_torch/models/checkpoints/mistral/__init__.py create mode 100644 tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py create mode 100644 tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py create mode 100644 tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py diff --git a/examples/models/core/mistral_large_3/README.md b/examples/models/core/mistral_large_3/README.md new file mode 100644 index 000000000000..0d9394bc7aaf --- /dev/null +++ b/examples/models/core/mistral_large_3/README.md @@ -0,0 +1,54 @@ +## How to use the modules + +The following explains how to use the different modules of Mistral Large V3. + +```python +from tensorrt_llm._torch.models.modeling_deepseekv3 import DeepseekV3ForCausalLM +from tensorrt_llm._torch.models.modeling_mistral import Mistral3VLM +# from tensorrt_llm.llmapi.tokenizer import MistralTokenizer +from tensorrt_llm._torch.models.checkpoints.mistral.checkpoint_loader import MistralCheckpointLoader +from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import MistralLarge3WeightMapper +from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader +from transformers import AutoTokenizer +``` + +### Tokenizer +```python +mtok = AutoTokenizer.from_pretrained(TOKENIZER_DIR) +``` + +### Config and model instance +```python +config_loader = MistralConfigLoader() +config = config_loader.load(MODEL_DIR) + +model = Mistral3VLM(model_config=config) +assert isinstance(model.llm, DeepseekV3ForCausalLM) +``` + +### Checkpoint loading +```python +weight_mapper=MistralLarge3WeightMapper() +loader = MistralCheckpointLoader(weight_mapper=weight_mapper) + +weights_dict = loader.load_weights(MODEL_DIR) +``` + +### Weight loading +#### E2E +```python +model.load_weights(weights_dict, weight_mapper=weight_mapper) # target usage +``` +#### By module +```python +def _filter_weights(weights, prefix): + return { + name[len(prefix):]: weight + for name, weight in weights.items() if name.startswith(prefix) + } + +llm_weights = weight_mapper.rename_by_params_map( + params_map=weight_mapper.mistral_llm_mapping, + weights=_filter_weights(weights_dict, "language_model.")) +model.llm.load_weights(llm_weights, weight_mapper=weight_mapper) +``` diff --git a/examples/models/core/mistral_large_3/test.py b/examples/models/core/mistral_large_3/test.py new file mode 100644 index 000000000000..a20079cce1b1 --- /dev/null +++ b/examples/models/core/mistral_large_3/test.py @@ -0,0 +1,37 @@ +TEST_TOKENIZER = False +TEST_CONFIG_LOADER = False +TEST_CHECKPOINT_LOADER = True + +MODEL_DIR = ( + "/home/scratch.trt_llm_data/llm-models/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512" +) +if TEST_TOKENIZER: + ## Test Tokenizer + from transformers import AutoTokenizer + + mtok = AutoTokenizer.from_pretrained(MODEL_DIR) + print(f"mtok: {mtok}") + + print(mtok.encode("Hello, world!")) + print(mtok.decode([1, 22177, 1044, 4304, 1033])) + +if TEST_CONFIG_LOADER: + ## Test config loader + from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader + + config_loader = MistralConfigLoader() + config = config_loader.load(MODEL_DIR) + print(f"config: {config}") + +if TEST_CHECKPOINT_LOADER: + from tensorrt_llm._torch.models.checkpoints.mistral.checkpoint_loader import ( + MistralCheckpointLoader, + ) + from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import ( + MistralLarge3WeightMapper, + ) + + weight_mapper = MistralLarge3WeightMapper() + loader = MistralCheckpointLoader(weight_mapper=weight_mapper) + weights_dict = loader.load_weights(MODEL_DIR) + # print(f"weights_dict.keys(): {weights_dict.keys()}") diff --git a/requirements.txt b/requirements.txt index e123aafcdee3..1bc253c3f0a8 100644 --- a/requirements.txt +++ b/requirements.txt @@ -75,3 +75,4 @@ numexpr<2.14.0 # WAR for attempted use of nonexistent numpy.typing partial_json_parser apache-tvm-ffi==0.1.4 # used for reduce nvidia-cutlass-dsl host overhead torch-c-dlpack-ext==0.1.3 # used for reduce nvidia-cutlass-dsl host overhead, optional package for improved torch tensor calling perf +mistral-common diff --git a/tensorrt_llm/_torch/model_config.py b/tensorrt_llm/_torch/model_config.py index 148ec5e2e3fd..f4e4497b72a2 100644 --- a/tensorrt_llm/_torch/model_config.py +++ b/tensorrt_llm/_torch/model_config.py @@ -341,6 +341,31 @@ def load_hf_quant_config(hf_quant_config, moe_backend): 'block.*.attn.out', 'block.*.mlp.gate', 'block.*.attn.qkv', 'embedding', 'unembedding' ] + # Mistral checkpoints. + elif hf_quant_config.get("quant_method") == "compressed-tensors": + if 'NVFP4' in hf_quant_config.get("config_groups"): + quant_config.quant_algo = QuantAlgo.NVFP4 + quant_config.group_size = 16 + quant_config.exclude_modules = [ + "layers.*.attention.wq*", "layers.*.attention.wk*", + "patch_merger*", "vision_encoder*", + "vision_language_adapter*", + "language_model.model.layers.*.self_attn.q*", + "language_model.model.layers.*.self_attn.k*", + "language_model.model.layers.*.self_attn.fused_qkv*", + "language_model.model.layers.*.self_attn.kv_a_proj_with_mqa*", + "model.layers.*.self_attn.q*", + "model.layers.*.self_attn.k*", + "model.layers.*.self_attn.fused_qkv*", + "model.layers.*.self_attn.kv_a_proj_with_mqa*", + "model.layers.*.self_attn.o_proj" + ] + elif 'FP8_BLOCK' in hf_quant_config.get("config_groups"): + quant_config.quant_algo = QuantAlgo.FP8_BLOCK_SCALES + quant_config.group_size = 128 + quant_config.exclude_modules = [ + "*q_a_proj*", "*kv_a_proj_with_mqa*" + ] return quant_config, layer_quant_config diff --git a/tensorrt_llm/_torch/models/checkpoints/__init__.py b/tensorrt_llm/_torch/models/checkpoints/__init__.py index 6a7426eb5bda..590a4c7ea9b6 100644 --- a/tensorrt_llm/_torch/models/checkpoints/__init__.py +++ b/tensorrt_llm/_torch/models/checkpoints/__init__.py @@ -12,11 +12,30 @@ from .hf.qwen3_next_weight_mapper import Qwen3NextHfWeightMapper from .hf.weight_loader import HfWeightLoader from .hf.weight_mapper import HfWeightMapper +from .mistral.checkpoint_loader import (MistralCheckpointLoader, + MistralLarge3CheckpointLoader) +from .mistral.config_loader import MistralConfigLoader +from .mistral.weight_mapper import (MistralLarge3WeightMapper, + MistralWeightMapper) __all__ = [ - "HfConfigLoader", "HfWeightLoader", "HfWeightMapper", - "BaseCheckpointLoader", "HfCheckpointLoader", "NemotronHHfWeightMapper", - "Gemma3HfWeightMapper", "MixtralHfWeightMapper", "Llama4HfWeightMapper", - "Qwen2MoeHfWeightMapper", "Qwen3MoeHfWeightMapper", "Qwen2VLHfWeightMapper", - "Qwen3NextHfWeightMapper", "LlavaNextHfWeightMapper" + "HfConfigLoader", + "HfWeightLoader", + "HfWeightMapper", + "MistralConfigLoader", + "MistralWeightMapper", + "MistralCheckpointLoader", + "BaseCheckpointLoader", + "HfCheckpointLoader", + "NemotronHHfWeightMapper", + "Gemma3HfWeightMapper", + "MixtralHfWeightMapper", + "Llama4HfWeightMapper", + "Qwen2MoeHfWeightMapper", + "Qwen3MoeHfWeightMapper", + "Qwen2VLHfWeightMapper", + "Qwen3NextHfWeightMapper", + "LlavaNextHfWeightMapper", + "MistralLarge3CheckpointLoader", + "MistralLarge3WeightMapper", ] diff --git a/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py b/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py index 7c24f19ae736..fd0ba287abe4 100644 --- a/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py @@ -19,6 +19,8 @@ from tensorrt_llm.mapping import Mapping +# @register_checkpoint_weight_loader("mistral_large_3") +@register_checkpoint_weight_loader("mistral") @register_checkpoint_weight_loader("HF") class HfWeightLoader(BaseWeightLoader): """ diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/__init__.py b/tensorrt_llm/_torch/models/checkpoints/mistral/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py new file mode 100644 index 000000000000..3f7749201126 --- /dev/null +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py @@ -0,0 +1,107 @@ +from typing import Optional + +from tensorrt_llm._torch.models.checkpoints.base_config_loader import BaseConfigLoader +from tensorrt_llm._torch.models.checkpoints.base_weight_loader import BaseWeightLoader +from tensorrt_llm._torch.models.checkpoints.base_weight_mapper import BaseWeightMapper +from tensorrt_llm._torch.models.checkpoints.hf.checkpoint_loader import HfCheckpointLoader +from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader +from tensorrt_llm._torch.models.modeling_utils import register_checkpoint_loader +from tensorrt_llm.quantization.mode import QuantAlgo + +@register_checkpoint_loader("mistral") +class MistralCheckpointLoader(HfCheckpointLoader): + def __init__( + self, + *, + weight_loader: Optional[BaseWeightLoader] = None, + weight_mapper: Optional[BaseWeightMapper] = None, + config_loader: Optional[BaseConfigLoader] = None, + ): + super().__init__( + weight_loader=weight_loader, weight_mapper=weight_mapper, config_loader=config_loader + ) + self._checkpoint_format = "mistral" + self.mm_module_mapping = { + "vision_encoder": "vision_tower", + "pre_mm_projector_norm": "multi_modal_projector.norm", + "vision_language_adapter": "multi_modal_projector", + "patch_merger": "multi_modal_projector.patch_merger", + } + + def preprocess_weights(self, weights: dict) -> dict: + """ + Aggregate weights by module + """ + hf_weights = {} + + for key, value in weights.items(): + modules = key.split(".") + + if modules[0] not in self.mm_module_mapping.keys(): + hf_weights["language_model." + key] = value + + else: + modules[0] = self.mm_module_mapping[modules[0]] + hf_weights[".".join(modules)] = value + + return hf_weights + + def broadcast_per_tensor_scales(self, weights): + import math + + scales = [k for k in weights.keys() if k.endswith("qscale_weight")] + for scale in scales: + name = ".".join(scale.split(".")[:-1]) + weight_shape = weights[f"{name}.weight"].shape + broadcast = weights[scale].expand( + math.ceil(weight_shape[0] / 128), + math.ceil(weight_shape[1] / 128), + ) + weights[scale] = broadcast[:] + + def reverse_nvfp4_global_scales(self, weights): + for key in weights.keys(): + if "global_scale" in key: + weights[key] = 1.0 / weights[key] + + def load_weights(self, checkpoint_dir: str, **kwargs): + weights = super().weight_loader.load_weights(checkpoint_dir, mapping=None, **kwargs) + model_config = kwargs.get("model_config", None) + if model_config is not None: + if model_config.quant_config.quant_algo == QuantAlgo.NVFP4: + quantization_weights_map = { + "weight_packed": "weight", + "input_global_scale": "input_scale", + "weight_global_scale": "weight_scale_2", + } + elif model_config.quant_config.quant_algo == QuantAlgo.FP8_BLOCK_SCALES: + quantization_weights_map = { + "weight_scale": "weight_scale_inv", + } + params_map = self.weight_mapper.mistral_llm_mapping.copy() + params_map.update(quantization_weights_map) + weights = self.preprocess_weights(weights) + weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) + + # FIXME mimic DS fp8 till per tensor supported + self.broadcast_per_tensor_scales(weights) + self.reverse_nvfp4_global_scales(weights) + return weights + + def get_default_config_loader(self) -> MistralConfigLoader: + return MistralConfigLoader() + + +@register_checkpoint_loader("mistral_large_3") +class MistralLarge3CheckpointLoader(MistralCheckpointLoader): + def __init__( + self, + *, + weight_loader: Optional[BaseWeightLoader] = None, + weight_mapper: Optional[BaseWeightMapper] = None, + config_loader: Optional[BaseConfigLoader] = None, + ): + super().__init__( + weight_loader=weight_loader, weight_mapper=weight_mapper, config_loader=config_loader + ) + self._checkpoint_format = "mistral_large_3" diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py new file mode 100644 index 000000000000..2fa2a1706571 --- /dev/null +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py @@ -0,0 +1,301 @@ +import json +from pathlib import Path +from typing import Any, Optional + +from transformers import PretrainedConfig, WhisperConfig + +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.checkpoints.base_config_loader import BaseConfigLoader +from tensorrt_llm._torch.models.modeling_utils import register_config_loader +from tensorrt_llm.models.modeling_utils import QuantConfig + +################### +# vllm code here +# https://github.com/vllm-project/vllm/blob/48a5fff66e78985a634abac0d8d7f271da744000/vllm/transformers_utils/configs/mistral.py +################### + + +def adapt_config_dict( + config_dict: dict[str, Any], + defaults: dict[str, Any] = {}, +) -> PretrainedConfig: + config_dict = _remap_general_mistral_args(config_dict) + + if bool(config_dict.get("quantization")): + config_dict = _remap_mistral_quantization_args(config_dict) + + is_moe = bool(config_dict.get("moe")) + is_mistral_large_3 = is_moe and (config_dict["moe"].get("num_shared_experts") or 0) > 0 + if config_dict.get("model_type") == "mamba": + config_dict["architectures"] = ["Mamba2ForCausalLM"] + elif is_moe and is_mistral_large_3: + config_dict = _remap_moe_args(config_dict) + config_dict["model_type"] = "deepseek_v3" + config_dict["architectures"] = ["MistralLarge3ForCausalLM"] + + assert "llama_4_scaling" in config_dict, "MistralLarge3 expect llama4 scaling config." + llama_4_scaling_config_keys = ["original_max_position_embeddings", "beta"] + assert all( + [key in config_dict["llama_4_scaling"] for key in llama_4_scaling_config_keys] + ), f"llama_4_scaling config should define the keys: {','.join(llama_4_scaling_config_keys)}" + elif is_moe: + config_dict["architectures"] = ["MixtralForCausalLM"] + else: + config_dict["architectures"] = ["MistralForCausalLM"] + + if bool(config_dict.get("yarn")): + config_dict = _remap_mistral_yarn_args(config_dict) + + if bool(config_dict.get("llama_4_scaling")): + llama_4_scaling_config_keys = ["original_max_position_embeddings", "beta"] + assert all( + [key in config_dict["llama_4_scaling"] for key in llama_4_scaling_config_keys] + ), f"llama_4_scaling config should define the keys: {','.join(llama_4_scaling_config_keys)}" + + is_vision = (config_dict.get("multimodal") or {}).get("vision_encoder_args") or config_dict.get( + "vision_encoder" + ) + is_audio = bool( + ((config_dict.get("multimodal") or {}).get("whisper_model_args") or {}).get("encoder_args") + ) + + assert not (is_vision and is_audio), "Vision and audio are mutually exclusive" + + if is_vision: + config_dict = _remap_mistral_vision_args(config_dict) + if is_audio: + config_dict = _remap_mistral_audio_args(config_dict) + + for k, v in defaults.items(): + config_dict.setdefault(k, v) + + config = PretrainedConfig.from_dict(config_dict) + + return config + + +def _remap_mistral_vision_args(config: dict) -> dict: + if config.get("multimodal"): + vision_config = config.pop("multimodal") + else: + vision_config = config.pop("vision_encoder") + + quant_config = config.get("quantization_config") + config = { + "model_type": "pixtral", + "architectures": ["PixtralForConditionalGeneration"], + "text_config": PretrainedConfig.from_dict(config), + "vision_config": PretrainedConfig.from_dict(vision_config), + } + if quant_config: + config["quantization_config"] = quant_config + return config + + +def _remap_mistral_yarn_args(config: dict) -> dict: + yarn_config_map = { + "factor": "factor", + "original_max_position_embeddings": "original_max_position_embeddings", + "beta": "beta_fast", + "alpha": "beta_slow", + "apply_scale": "apply_yarn_scaling", + } + yarn_config = config.get("yarn") or {} + config["rope_parameters"] = { + "rope_type": "yarn", + "mscale_all_dim": 1, + } + + if rope_theta := config.pop("rope_theta", None): + config["rope_parameters"]["rope_theta"] = rope_theta + + for old_name, new_name in yarn_config_map.items(): + if old_name in yarn_config: + config["rope_parameters"][new_name] = yarn_config.pop(old_name) + + assert len(yarn_config) == 0, f"Unparsed yarn config: {yarn_config}" + + return config + + +def _remap_general_mistral_args(config: dict) -> dict: + # Mistral key -> HF key + config_mapping = { + "dim": "hidden_size", + "norm_eps": "rms_norm_eps", + "n_kv_heads": "num_key_value_heads", + "n_layers": "num_hidden_layers", + "n_heads": "num_attention_heads", + "hidden_dim": "intermediate_size", + } + # HF key -> (Mistral key, default value) + top_level_mapping_with_default = { + "model_type": ("model_type", "transformer"), + "hidden_act": ("activation", "silu"), + "tie_word_embeddings": ("tied_embeddings", False), + "max_seq_len": ("max_seq_len", config.get("max_position_embeddings", 128_000)), + "max_position_embeddings": ("max_position_embeddings", 128_000), + } + + for key, new_key in config_mapping.items(): + if key in config: + config[new_key] = config.pop(key) + + for new_key, (key, default_value) in top_level_mapping_with_default.items(): + config[new_key] = config.pop(key, default_value) + + return config + + +def _remap_mistral_quantization_args(config: dict) -> dict: + if config.get("quantization"): + quantization = config.pop("quantization", {}) + if quantization.get("qformat_weight") == "fp8_e4m3": + qscheme_act = quantization.get("qscheme_act") + assert qscheme_act in ("NO_SCALES", "TENSOR", None), ( + "Only NO_SCALES and TENSOR (default) are supported for qscheme_act" + ) + is_dynamic = qscheme_act == "NO_SCALES" + config["quantization_config"] = { + "quant_method": "fp8", + "activation_scheme": "dynamic" if is_dynamic else "static", + } + else: + raise ValueError(f"Found unknown quantization='{quantization}' in config") + + return config + + +def _remap_mistral_audio_args(config: dict) -> dict: + whisper_args = config["multimodal"].pop("whisper_model_args") + encoder_args = whisper_args["encoder_args"] + downsample_args = whisper_args["downsample_args"] + + quant_config = config.get("quantization_config") + config = { + "model_type": "whixtral", + "architectures": ["VoxtralForConditionalGeneration"], + "text_config": PretrainedConfig.from_dict(config), + "audio_config": WhisperConfig( + num_mel_bins=encoder_args["audio_encoding_args"]["num_mel_bins"], + window_size=encoder_args["audio_encoding_args"]["window_size"], + sampling_rate=encoder_args["audio_encoding_args"]["sampling_rate"], + hop_length=encoder_args["audio_encoding_args"]["hop_length"], + downsample_factor=downsample_args["downsample_factor"], + d_model=encoder_args["dim"], + encoder_layers=encoder_args["n_layers"], + encoder_ffn_dim=encoder_args["hidden_dim"], + encoder_attention_heads=encoder_args["n_heads"], + vocab_size=encoder_args["vocab_size"], + max_source_positions=encoder_args["max_source_positions"], + is_encoder_decoder=False, # Override WhisperConfig default + ), + } + if quant_config: + config["quantization_config"] = quant_config + return config + + +def _remap_moe_args(config: dict) -> dict: + moe_config_map = { + "route_every_n": "moe_layer_freq", + "first_k_dense_replace": "first_k_dense_replace", + "num_experts_per_tok": "num_experts_per_tok", + "num_experts": "n_routed_experts", + "expert_hidden_dim": "moe_intermediate_size", + "routed_scale": "routed_scaling_factor", + "num_shared_experts": "n_shared_experts", + "num_expert_groups": "n_group", + "num_expert_groups_per_tok": "topk_group", + } + moe_config = config.get("moe", {}) + for old_name, new_name in moe_config_map.items(): + if old_name in moe_config: + value = moe_config.pop(old_name) + config[new_name] = value + + config["topk_method"] = None + config["norm_topk_prob"] = True + config["scoring_func"] = "softmax" + + return config + + +###################### +# End of vllm code +###################### + + +@register_config_loader("mistral") +@register_config_loader("mistral_large_3") +class MistralConfigLoader(BaseConfigLoader): + def _load_mistral_config_dict( + self, checkpoint_dir: str, config_file_name: str + ) -> Optional[dict]: + file_path = Path(checkpoint_dir) / Path(config_file_name) + + if file_path.exists() and file_path.is_file(): + with open(file_path) as file: + return json.load(file) + return None + + # Adaptation of + # https://github.com/vllm-project/vllm/blob/48a5fff66e78985a634abac0d8d7f271da744000/vllm/transformers_utils/config.py#L175 + def _parse_mistral_config(self, checkpoint_dir: str): + config_file_name = "params.json" + + # This function loads a params.json config which + # should be used when loading models in mistral format + config_dict = self._load_mistral_config_dict(checkpoint_dir, config_file_name) + if config_dict is None: + raise ValueError( + f"Failed to load '{config_file_name}' config from '{checkpoint_dir}'. " + f"Only local checkpoints are supported for mistral format." + ) + assert isinstance(config_dict, dict) + + if (max_position_embeddings := config_dict.get("max_position_embeddings")) is None: + max_position_embeddings = 128_000 + config_dict["max_position_embeddings"] = max_position_embeddings + + pretrained_config = adapt_config_dict(config_dict) + + # Mistral configs may define sliding_window as list[int]. Convert it + # to int and add the layer_types list[str] to make it HF compatible + if (sliding_window := getattr(pretrained_config, "sliding_window", None)) and isinstance( + sliding_window, list + ): + pattern_repeats = pretrained_config.num_hidden_layers // len(sliding_window) + layer_types = sliding_window * pattern_repeats + pretrained_config.layer_types = [ + "full_attention" if layer_type is None else "sliding_attention" + for layer_type in layer_types + ] + pretrained_config.sliding_window = next(filter(None, sliding_window), None) + + return config_dict, pretrained_config + + def load(self, checkpoint_dir: str, **kwargs) -> ModelConfig: + # Re-write from ModelConfig.from_pretrained + + config_dict, pretrained_config = self._parse_mistral_config(checkpoint_dir) + + # Some checkpoints lack torch_dtype, populate with dtype + pretrained_config.torch_dtype = getattr(pretrained_config, "dtype", None) + quant_config = QuantConfig() + layer_quant_config = None + moe_backend = kwargs.get("moe_backend", "CUTLASS") + + hf_quant_config = pretrained_config.quantization_config + quant_config, layer_quant_config = ModelConfig.load_hf_quant_config( + hf_quant_config, moe_backend + ) + + model_config = ModelConfig( + pretrained_config=pretrained_config, + quant_config=quant_config, + quant_config_dict=layer_quant_config, + **kwargs, + ) + model_config._frozen = True + return model_config diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py b/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py new file mode 100644 index 000000000000..e70d09fcd0e8 --- /dev/null +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py @@ -0,0 +1,132 @@ +from torch import nn + +from tensorrt_llm._torch.models.checkpoints.hf.weight_mapper import HfWeightMapper +from tensorrt_llm._torch.models.modeling_utils import register_mapper + + +@register_mapper("mistral", "MistralForCausalLM") +@register_mapper("mistral", "PixtralForConditionalGeneration") +class MistralWeightMapper(HfWeightMapper): + def __init__(self): + super().__init__() + + self._callbacks.append(self._permute_qk) + + # TODO move to registry + self.pixtral_mapping = { + "wq": "q_proj", + "wk": "k_proj", + "wv": "v_proj", + "wo": "o_proj", + "w1": "gate_proj", + "w2": "down_proj", + "w3": "up_proj", + "w_in": "linear_1", + "w_out": "linear_2", + } + + self.mistral_llm_mapping = { + "layers": "model.layers", + "attention": "self_attn", + "qscale_act": "input_scale", + "qscale_weight": "weight_scale_inv", + "kv_fake_quantizer.qscale_act": "kv_scale", + "q_fake_quantizer.qscale_act": "attn.q_scale", + "k_fake_quantizer.qscale_act": "k_scale", + "v_fake_quantizer.qscale_act": "v_scale", + "attention_norm": "input_layernorm", + "feed_forward": "mlp", + "ffn_norm": "post_attention_layernorm", + "tok_embeddings": "model.embed_tokens", + "output": "lm_head", + "norm": "model.norm", + # For Eagle3 + "language_model.eagle_linear": "model.fc", + "language_model.layers": "layers", + "language_model.norm": "norm", + } + self.mistral_llm_mapping.update(self.pixtral_mapping) + + # Adapted from: + # https://github.com/vllm-project/vllm/blob/883b42896a9ed9791750d721fad26005b7569eba/vllm/model_executor/models/llama.py#L657 + def rename_by_params_map(self, params_map: dict[str, str], weights: dict) -> dict: + renamed_weights = {} + + for key in list(weights.keys()): + new_key = key + modules = key.split(".") + num_modules = len(modules) + for i in range(num_modules): + item = modules[i] + next_item = modules[i + 1] if i < num_modules - 1 else None + + combined_item = f"{item}.{next_item}" if next_item is not None else None + + if combined_item in params_map: + new_key = new_key.replace(combined_item, params_map[combined_item]) + elif item in params_map: + new_key = new_key.replace(item, params_map[item]) + + renamed_weights[new_key] = weights[key] + + return renamed_weights + + def _permute_qk(self, module: nn.Module, new_name: str, weights: dict): + # Adapted from: + # https://github.com/vllm-project/vllm/blob/883b42896a9ed9791750d721fad26005b7569eba/vllm/model_executor/models/llama.py#L657 + + processed_weights = {} + config = self.config.pretrained_config + + def permute(w, n_heads: int, attn_out: int): + attn_in = config.head_dim * n_heads + + return ( + w.view(n_heads, attn_in // n_heads // 2, 2, attn_out) + .transpose(1, 2) + .reshape(attn_in, attn_out) + ) + + # rotary embeds should be sliced + # If using quantized model in mistral format, + # quantization scales (qscale_weight) also need to be sliced + + if new_name in ["k_proj", "q_proj"]: + n_heads = ( + config.num_key_value_heads if new_name == "k_proj" else config.num_attention_heads + ) + + processed_weights["weight"] = permute(weights["weight"], n_heads, config.hidden_size) + + if "qscale_weight" in weights and weights["qscale_weight"].numel() > 1: + processed_weights["qscale_weight"] = permute(weights["qscale_weight"], n_heads, 1) + + return processed_weights + + return weights + + +@register_mapper("mistral_large_3") +@register_mapper("mistral_large_3", "PixtralForConditionalGeneration") +@register_mapper("mistral_large_3", "MistralLarge3ForCausalLM") +class MistralLarge3WeightMapper(MistralWeightMapper): + def __init__(self): + super().__init__() + + self.mistral_llm_mapping.update( + { + "wkv_a_with_mqa": "kv_a_proj_with_mqa", + "wkv_b": "kv_b_proj", + "wq_a": "q_a_proj", + "q_a_norm": "q_a_layernorm", + "wq_b": "q_b_proj", + "kv_a_norm": "kv_a_layernorm", + "k_fake_quantizer.qscale_act": "mla_attn.mla_attn.k_scale", + "q_fake_quantizer.qscale_act": "mla_attn.mla_attn.q_scale", + "v_fake_quantizer.qscale_act": "mla_attn.mla_attn.v_scale", + "gate": "mlp.gate", + "shared_experts": "mlp.shared_experts", + "experts": "mlp.experts", + "router_biases": "mlp.gate.e_score_correction_bias", + } + ) From e5bf6fd04ad510e79750b593ced6d1ab0902c8d6 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 03:23:46 +0000 Subject: [PATCH 02/21] fix bug of mistral checkpoint loader Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- examples/models/core/mistral_large_3/test.py | 6 +++--- .../models/checkpoints/mistral/checkpoint_loader.py | 8 +++++--- 2 files changed, 8 insertions(+), 6 deletions(-) diff --git a/examples/models/core/mistral_large_3/test.py b/examples/models/core/mistral_large_3/test.py index a20079cce1b1..535fdbafaad2 100644 --- a/examples/models/core/mistral_large_3/test.py +++ b/examples/models/core/mistral_large_3/test.py @@ -15,7 +15,7 @@ print(mtok.encode("Hello, world!")) print(mtok.decode([1, 22177, 1044, 4304, 1033])) -if TEST_CONFIG_LOADER: +if TEST_CONFIG_LOADER or True: ## Test config loader from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader @@ -33,5 +33,5 @@ weight_mapper = MistralLarge3WeightMapper() loader = MistralCheckpointLoader(weight_mapper=weight_mapper) - weights_dict = loader.load_weights(MODEL_DIR) - # print(f"weights_dict.keys(): {weights_dict.keys()}") + weights_dict = loader.load_weights(MODEL_DIR, model_config=config) + print(f"weights_dict.keys(): {weights_dict.keys()}") diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py index 3f7749201126..eefd3db18de1 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py @@ -8,6 +8,7 @@ from tensorrt_llm._torch.models.modeling_utils import register_checkpoint_loader from tensorrt_llm.quantization.mode import QuantAlgo + @register_checkpoint_loader("mistral") class MistralCheckpointLoader(HfCheckpointLoader): def __init__( @@ -65,8 +66,10 @@ def reverse_nvfp4_global_scales(self, weights): weights[key] = 1.0 / weights[key] def load_weights(self, checkpoint_dir: str, **kwargs): + model_config = kwargs.pop("model_config", None) + assert model_config is not None, "model_config is required" weights = super().weight_loader.load_weights(checkpoint_dir, mapping=None, **kwargs) - model_config = kwargs.get("model_config", None) + params_map = self.weight_mapper.mistral_llm_mapping.copy() if model_config is not None: if model_config.quant_config.quant_algo == QuantAlgo.NVFP4: quantization_weights_map = { @@ -78,8 +81,7 @@ def load_weights(self, checkpoint_dir: str, **kwargs): quantization_weights_map = { "weight_scale": "weight_scale_inv", } - params_map = self.weight_mapper.mistral_llm_mapping.copy() - params_map.update(quantization_weights_map) + params_map.update(quantization_weights_map) weights = self.preprocess_weights(weights) weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) From ef3060995a026d38fd3a566691d352af5999cd9c Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 06:19:15 +0000 Subject: [PATCH 03/21] [WIP] add mistral large 3 model definition Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- examples/llm-api/quickstart_advanced.py | 6 + .../checkpoints/mistral/config_loader.py | 1 + .../_torch/models/modeling_deepseekv3.py | 5 +- .../_torch/models/modeling_mistral.py | 165 +++++++++++++----- .../_torch/models/modeling_mistral_large3.py | 55 ++++++ .../_torch/pyexecutor/config_utils.py | 8 + tensorrt_llm/inputs/registry.py | 6 + tensorrt_llm/serve/openai_server.py | 3 +- 8 files changed, 206 insertions(+), 43 deletions(-) create mode 100644 tensorrt_llm/_torch/models/modeling_mistral_large3.py diff --git a/examples/llm-api/quickstart_advanced.py b/examples/llm-api/quickstart_advanced.py index 9b37f8c7b296..5aa7f7ce703c 100644 --- a/examples/llm-api/quickstart_advanced.py +++ b/examples/llm-api/quickstart_advanced.py @@ -23,6 +23,11 @@ def add_llm_args(parser): type=str, nargs="+", help="A single or a list of text prompts.") + parser.add_argument('--checkpoint_format', + type=str, + default=None, + choices=["HF", "mistral"], + help="Model checkpoint format.") # Build config parser.add_argument("--max_seq_len", type=int, @@ -237,6 +242,7 @@ def setup_llm(args, **kwargs): llm = LLM( model=args.model_dir, backend='pytorch', + checkpoint_format=args.checkpoint_format, disable_overlap_scheduler=args.disable_overlap_scheduler, kv_cache_config=kv_cache_config, attn_backend=args.attention_backend, diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py index 2fa2a1706571..2295fb01ff26 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py @@ -291,6 +291,7 @@ def load(self, checkpoint_dir: str, **kwargs) -> ModelConfig: hf_quant_config, moe_backend ) + kwargs.pop("trust_remote_code", None) model_config = ModelConfig( pretrained_config=pretrained_config, quant_config=quant_config, diff --git a/tensorrt_llm/_torch/models/modeling_deepseekv3.py b/tensorrt_llm/_torch/models/modeling_deepseekv3.py index 40fbaa983db1..82588c866fff 100755 --- a/tensorrt_llm/_torch/models/modeling_deepseekv3.py +++ b/tensorrt_llm/_torch/models/modeling_deepseekv3.py @@ -746,7 +746,10 @@ def __init__(self, config = model_config.pretrained_config self.top_k = top_k self.use_dp = model_config.mapping.enable_attention_dp - self.gate = DeepseekV3Gate( + gate_cls = DeepseekV3Gate + if hasattr(model_config.pretrained_config, "gate_cls"): + gate_cls = model_config.pretrained_config.gate_cls + self.gate = gate_cls( hidden_size, num_experts, top_k=top_k, diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index 9ade4dee220d..ef235a5e2805 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -1,9 +1,11 @@ +import copy import dataclasses import os from typing import Any, Dict, List, Optional, Tuple import torch import torchvision +from mistral_common.tokens.tokenizers.multimodal import ImageEncoder from torch import nn from transformers import (AutoProcessor, AutoTokenizer, Mistral3Config, MistralConfig, PretrainedConfig, PreTrainedModel) @@ -14,11 +16,15 @@ PositionalEmbeddingParams, RopeParams) from tensorrt_llm._torch.model_config import ModelConfig from tensorrt_llm._torch.models import modeling_pixtral +from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import \ + MistralWeightMapper +from tensorrt_llm._torch.models.modeling_mistral_large3 import MistralLarge3ForCausalLM, Mistral3Gate from tensorrt_llm._torch.models.modeling_multimodal_utils import ( find_input_mm_embeds, fuse_input_embeds, get_multimodal_embeddings) from tensorrt_llm._torch.models.modeling_utils import (DecoderModel, DecoderModelForCausalLM, _load_weights_impl, + filter_weights, register_auto_model) from tensorrt_llm._torch.modules.attention import Attention from tensorrt_llm._torch.modules.decoder_layer import DecoderLayer @@ -266,7 +272,9 @@ def __call__( self, inputs: TextPrompt, sampling_params: SamplingParams ) -> Tuple[List[int], Optional[ExtraProcessedInputs]]: images = inputs.get("multi_modal_data", {}).get("image") - do_rescale = self.processor.image_processor.do_rescale + mm_processor_kwargs = inputs.get("mm_processor_kwargs", {}) + do_rescale = getattr(self.processor.image_processor, "do_rescale", + False) if images is not None and isinstance(images[0], torch.Tensor): # The default multimodal input loader will normalize images to [0, 1] when the requested # format is "pt" (pytorch tensors), but not for "pil" (PIL images). @@ -276,6 +284,7 @@ def __call__( text=inputs["prompt"], images=images, do_rescale=do_rescale, + **mm_processor_kwargs, ) input_ids = processed.pop("input_ids").tolist()[0] # Remaining in `processed`: @@ -331,6 +340,7 @@ def get_mm_special_token_ids(self) -> torch.Tensor: @register_auto_model("Mistral3ForConditionalGeneration") +@register_auto_model("PixtralForConditionalGeneration") @register_input_processor( Mistral3InputProcessor, model_type="mistral3", @@ -365,34 +375,47 @@ def __init__( config = model_config.pretrained_config super().__init__(config) - self.model_config = model_config - - llm_model_config = self._get_sub_model_config(model_config, - "text_config") - # This is necessary for the auto weight mapper to figure out what it needs. - llm_model_config.pretrained_config.architectures = config.architectures - self.llm = MistralForCausalLM(llm_model_config) - - self._device = "cuda" - # NOTE: current `modelopt` does not support quantizing the vision portion. - vision_model_config = self._get_sub_model_config(model_config, - "vision_config", - quant_config=None) - self._vision_tower = modeling_pixtral.PixtralVisionModel( - vision_model_config) - self._multi_modal_projector = Mistral3MultiModalProjector(model_config) - vision_feature_layer = config.vision_feature_layer + vision_feature_layer = getattr(config, "vision_feature_layer", -1) if vision_feature_layer != -1: raise ValueError( f"Using intermediate layers ({vision_feature_layer}) in the `PixtralVisionModel` " f"is not supported. Please use `vision_feature_layer=-1`.") + self._device = "cuda" self.model_dtype = getattr(config, "torch_dtype", torch.bfloat16) - - self._image_token_ids = torch.tensor([config.image_token_index], + image_token_index = getattr( + config, "image_token_index", None) or getattr( + config.vision_config, "image_token_id", None) + self._image_token_ids = torch.tensor([image_token_index], dtype=torch.int32, device=self._device) + + model_config_cp = copy.deepcopy(model_config) + + llm_model_config = self._get_sub_model_config(model_config_cp, + "text_config") + self.model_config = model_config_cp + llm_class = MistralForCausalLM + if llm_model_config.pretrained_config.architectures[ + 0] == "MistralLarge3ForCausalLM": + llm_class = MistralLarge3ForCausalLM + + llm_model_config.pretrained_config.gate_cls = Mistral3Gate + self.llm = llm_class(llm_model_config) + self.model_config.extra_attrs.update(llm_model_config.extra_attrs) + + # NOTE: current `modelopt` does not support quantizing the vision portion. + # NOTE: attn_backend: Pixtral head size not always divisible by 128 + vision_model_config = self._get_sub_model_config(model_config_cp, + "vision_config", + attn_backend="VANILLA", + quant_config=None) + + self._vision_tower = modeling_pixtral.PixtralVisionModel( + vision_model_config) + self._multi_modal_projector = Mistral3MultiModalProjector(model_config).eval().to(self._device) self._post_config() + self.is_loaded = True # This is necessary because the executor looks at # `model.model_config.pretrained_config.vocab_size`. @@ -400,18 +423,43 @@ def _post_config(self): self.config = self.llm.config self.model_config.pretrained_config = self.llm.config - def load_weights(self, weights: Dict, *args, **kwargs): - llm_weights = _filter_weights(weights, "language_model.") - self.llm.load_weights(llm_weights, *args, **kwargs) - - vit_weights = _filter_weights(weights, "vision_tower.") - self._vision_tower.load_weights(vit_weights, *args, **kwargs) - - mm_projector_weights = _filter_weights(weights, - "multi_modal_projector.") - # `_load_weights_impl` assumes `config.hidden_size` exists, which is not the case for the - # top-level `Mistral3Config`. + def load_weights(self, weights: Dict, weight_mapper=None, *args, **kwargs): + vit_params_map = None + if weight_mapper: + if isinstance(weight_mapper, MistralWeightMapper): + vit_params_map = weight_mapper.pixtral_mapping + + llm_weights = filter_weights(weights=weights, prefix="language_model") + logger.debug(f"Loading weights for {type(self.llm)}") + self.llm.load_weights(llm_weights, + weight_mapper=weight_mapper, + *args, + **kwargs) + logger.debug(f"Successfully loaded weights for {type(self.llm)}") + + vit_weights = filter_weights(weights=weights, prefix="vision_tower") + logger.debug(f"Loading weights for {type(self._vision_tower)}") + + # FIXME rename_weights_with_regex in _load_weights_impl breaks this, fall back to manual renaming + if vit_params_map is not None: + vit_weights = weight_mapper.rename_by_params_map( + weights=vit_weights, params_map=vit_params_map) + + self._vision_tower.load_weights(vit_weights, params_map=vit_params_map) + logger.debug( + f"Successfully loaded weights for {type(self._vision_tower)}") + + logger.debug(f"Loading weights for {type(self._multi_modal_projector)}") + mm_projector_weights = filter_weights(weights=weights, + prefix="multi_modal_projector") + + if vit_params_map is not None: + mm_projector_weights = weight_mapper.rename_by_params_map( + weights=mm_projector_weights, params_map=vit_params_map) self._multi_modal_projector.load_state_dict(mm_projector_weights) + logger.debug( + f"Successfully loaded weights for {type(self._multi_modal_projector)}" + ) def infer_max_seq_len(self) -> int: return self.llm.infer_max_seq_len() @@ -423,6 +471,7 @@ def forward( input_ids: Optional[torch.LongTensor] = None, position_ids: Optional[torch.LongTensor] = None, return_context_logits: bool = False, + spec_metadata: Optional[SpecMetadata] = None, **kwargs, ) -> torch.Tensor: """Forward method.""" @@ -455,6 +504,7 @@ def forward( position_ids=position_ids, inputs_embeds=inputs_embeds, return_context_logits=return_context_logits, + spec_metadata=spec_metadata, ) @staticmethod @@ -465,16 +515,41 @@ def _get_sub_model_config( ) -> ModelConfig: # Extract the subconfig from the `transformers` config and shove it into our own # `ModelConfig` class. + assert name in [ + "text_config", "vision_config" + ], f"Expected subconfig name to be either 'text_config' or 'vision_config'. Got {name} instead." + pretrained_config = getattr(model_config.pretrained_config, name) + sub_model_config: ModelConfig[MistralConfig] = dataclasses.replace( model_config, pretrained_config=getattr(model_config.pretrained_config, name), **changes, ) + if name == "text_config": + sub_model_config._frozen = False + sub_model_config.skip_create_weights_in_init = True + if not hasattr( + sub_model_config.pretrained_config, "architectures" + ) or sub_model_config.pretrained_config.architectures is None: + sub_model_config.pretrained_config.architectures = model_config.pretrained_config.architectures + sub_model_config._frozen = True + # Make sure some fields that are not explicitly included in the sub config, but present # in the top-level config, are replicated. if (hasattr(sub_model_config.pretrained_config, "torch_dtype") and sub_model_config.pretrained_config.torch_dtype is None): - sub_model_config.pretrained_config.torch_dtype = model_config.pretrained_config.torch_dtype + sub_model_config.pretrained_config.torch_dtype = model_config.pretrained_config.torch_dtype or torch.bfloat16 + + if name == "vision_config": + pretrained_config = sub_model_config.pretrained_config + defaults = { + "head_dim": pretrained_config.hidden_size // + pretrained_config.num_attention_heads, + "hidden_act": "silu", + } + for attr, default in defaults.items(): + if not hasattr(pretrained_config, attr): + setattr(pretrained_config, attr, default) return sub_model_config @@ -572,6 +647,12 @@ def batch_pixel_values( def mm_token_ids(self): return self._image_token_ids + def load_draft_weights( + self, + weights: Dict, + weight_mapper: Optional[MistralWeightMapper] = None) -> None: + self.llm.load_draft_weights(weights, weight_mapper=weight_mapper) + # Original implementation: # https://github.com/huggingface/transformers/blob/v4.51.3/src/transformers/models/mistral3/modeling_mistral3.py#L66 @@ -586,13 +667,15 @@ def __init__(self, model_config: ModelConfig[Mistral3Config]): self.config = config hidden_size = config.vision_config.hidden_size - self._spatial_merge_size = config.spatial_merge_size + self._spatial_merge_size = getattr( + config, "spatial_merge_size", None) or getattr( + config.vision_config, "spatial_merge_size") self._patch_size = config.vision_config.patch_size self.merging_layer = Linear( in_features=hidden_size * self._spatial_merge_size**2, out_features=hidden_size, bias=False, - dtype=config.torch_dtype, + dtype=config.torch_dtype or model_config.torch_dtype, mapping=model_config.mapping, ) @@ -640,7 +723,7 @@ def __init__(self, model_config: ModelConfig[Mistral3Config]): self.model_config = model_config self.config = config - dtype = config.torch_dtype + dtype = config.torch_dtype or model_config.torch_dtype self.norm = RMSNorm( hidden_size=config.vision_config.hidden_size, # NOTE: the original implementation actually does not look at the config for this value. @@ -650,21 +733,21 @@ def __init__(self, model_config: ModelConfig[Mistral3Config]): ) self.patch_merger = Mistral3PatchMerger(model_config) # We have hidden_size * the number of vision feature layers - num_feature_layers = 1 if isinstance(config.vision_feature_layer, - int) else len( - config.vision_feature_layer) + vision_feature_layer = getattr(config, "vision_feature_layer", -1) + num_feature_layers = 1 if isinstance(vision_feature_layer, + int) else len(vision_feature_layer) self.linear_1 = Linear( in_features=config.vision_config.hidden_size * num_feature_layers, out_features=config.text_config.hidden_size, - bias=config.multimodal_projector_bias, + bias=getattr(config, "multimodal_projector_bias", None), dtype=dtype, mapping=model_config.mapping, ) - self.act = ACT2FN[config.projector_hidden_act] + self.act = ACT2FN[getattr(config, "projector_hidden_act", "gelu")] self.linear_2 = Linear( in_features=config.text_config.hidden_size, out_features=config.text_config.hidden_size, - bias=config.multimodal_projector_bias, + bias=getattr(config, "multimodal_projector_bias", None), dtype=dtype, mapping=model_config.mapping, ) diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py new file mode 100644 index 000000000000..d7b4bd34fb92 --- /dev/null +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -0,0 +1,55 @@ +import torch + +from ..models.modeling_deepseekv3 import DeepseekV3ForCausalLM +from .modeling_utils import register_auto_model +from ..model_config import ModelConfig + +from torch import nn +from typing import Dict, Optional, List + +from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import ( + MistralLarge3WeightMapper, +) +from tensorrt_llm._torch.modules.fused_moe import RenormalizeNaiveMoeRoutingMethod + + +class Mistral3Gate(nn.Module): + + def __init__( + self, + hidden_size: int, + num_experts: int, + top_k: int, + dtype: Optional[torch.dtype] = None, + **kwargs, + ): + super().__init__() + self.weight = nn.Parameter(torch.empty((num_experts, hidden_size), + dtype=dtype), + requires_grad=False) + self.top_k = top_k + self.dtype = dtype + self.routing_method = RenormalizeNaiveMoeRoutingMethod(top_k=self.top_k) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + logits: torch.Tensor = torch.ops.trtllm.cublas_mm(hidden_states, + self.weight.t(), + bias=None, + out_dtype=self.dtype) + return logits + + def load_weights(self, weights: List[Dict]): + assert len(weights) == 1 + + self.weight.copy_(weights[0]["weight"][:]) + +@register_auto_model("MistralLarge3ForCausalLM") +class MistralLarge3ForCausalLM(DeepseekV3ForCausalLM): + def __init__(self, model_config: ModelConfig): + super().__init__(model_config) + + def forward(self, *args, **kwargs): + return super().forward(*args, **kwargs) + + def load_weights(self, weights: Dict, *args, **kwargs): + super().load_weights(llm_weights) diff --git a/tensorrt_llm/_torch/pyexecutor/config_utils.py b/tensorrt_llm/_torch/pyexecutor/config_utils.py index 6013d51fa298..f2c0e51d3936 100644 --- a/tensorrt_llm/_torch/pyexecutor/config_utils.py +++ b/tensorrt_llm/_torch/pyexecutor/config_utils.py @@ -38,14 +38,22 @@ def __getitem__(self, key): def load_pretrained_config(model_name_or_path: str, trust_remote_code: bool = False, + checkpoint_format: str = None, **kwargs) -> transformers.PretrainedConfig: config_dict, _ = transformers.PretrainedConfig.get_config_dict( model_name_or_path, **kwargs) model_type = config_dict.get("model_type") + if model_type in _CONFIG_REGISTRY: config_class = _CONFIG_REGISTRY[model_type] model_config = config_class.from_pretrained(model_name_or_path, **kwargs) + elif checkpoint_format == "mistral_large_3": + from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import \ + MistralConfigLoader + model_config = getattr( + MistralConfigLoader().load(model_name_or_path).pretrained_config, + "text_config") else: model_config = transformers.AutoConfig.from_pretrained( model_name_or_path, trust_remote_code=trust_remote_code) diff --git a/tensorrt_llm/inputs/registry.py b/tensorrt_llm/inputs/registry.py index 7737600e6f19..54902a5ba378 100644 --- a/tensorrt_llm/inputs/registry.py +++ b/tensorrt_llm/inputs/registry.py @@ -600,6 +600,12 @@ def create_input_processor( logger.debug( f"Unable to load HF config from {model_path_or_dir}: {e}. Falling back." ) + elif checkpoint_format in ("mistral", "mistral_large_3"): + logger.debug(f"Detected checkpoint_format={checkpoint_format}.") + from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import \ + MistralConfigLoader + model_config = MistralConfigLoader().load(model_path_or_dir) + config = model_config.pretrained_config else: logger.debug( f"checkpoint_format={checkpoint_format}; skipping HF config load.") diff --git a/tensorrt_llm/serve/openai_server.py b/tensorrt_llm/serve/openai_server.py index e64c5d20df69..d7750f3c7577 100644 --- a/tensorrt_llm/serve/openai_server.py +++ b/tensorrt_llm/serve/openai_server.py @@ -109,7 +109,8 @@ def __init__(self, from tensorrt_llm._torch.pyexecutor.config_utils import \ load_pretrained_config self.model_config = load_pretrained_config(hf_tokenizer_path, - trust_remote_code=trust_remote_code) + trust_remote_code=trust_remote_code, + checkpoint_format=getattr(self.llm.args, "checkpoint_format", None)) except Exception: logger.debug("Failed to load AutoConfig for %s", hf_tokenizer_path) self.model_config = None From 31cf365466d86df4704d4778a4d5d878270317a1 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 07:12:31 +0000 Subject: [PATCH 04/21] [WIP] fix bugs of mistral large 3 ckpt loader Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../checkpoints/mistral/checkpoint_loader.py | 19 +----------------- .../_torch/models/modeling_mistral_large3.py | 20 ++++++++++++++++++- 2 files changed, 20 insertions(+), 19 deletions(-) diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py index eefd3db18de1..8d245ae43798 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py @@ -66,25 +66,8 @@ def reverse_nvfp4_global_scales(self, weights): weights[key] = 1.0 / weights[key] def load_weights(self, checkpoint_dir: str, **kwargs): - model_config = kwargs.pop("model_config", None) - assert model_config is not None, "model_config is required" - weights = super().weight_loader.load_weights(checkpoint_dir, mapping=None, **kwargs) - params_map = self.weight_mapper.mistral_llm_mapping.copy() - if model_config is not None: - if model_config.quant_config.quant_algo == QuantAlgo.NVFP4: - quantization_weights_map = { - "weight_packed": "weight", - "input_global_scale": "input_scale", - "weight_global_scale": "weight_scale_2", - } - elif model_config.quant_config.quant_algo == QuantAlgo.FP8_BLOCK_SCALES: - quantization_weights_map = { - "weight_scale": "weight_scale_inv", - } - params_map.update(quantization_weights_map) + weights = super().weight_loader.load_weights(checkpoint_dir, **kwargs) weights = self.preprocess_weights(weights) - weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) - # FIXME mimic DS fp8 till per tensor supported self.broadcast_per_tensor_scales(weights) self.reverse_nvfp4_global_scales(weights) diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index d7b4bd34fb92..b5daebe76c59 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -47,9 +47,27 @@ def load_weights(self, weights: List[Dict]): class MistralLarge3ForCausalLM(DeepseekV3ForCausalLM): def __init__(self, model_config: ModelConfig): super().__init__(model_config) + self.weight_mapper = MistralLarge3WeightMapper() def forward(self, *args, **kwargs): return super().forward(*args, **kwargs) def load_weights(self, weights: Dict, *args, **kwargs): - super().load_weights(llm_weights) + assert self.model_config is not None, "self.model_config is required" + params_map = self.weight_mapper.mistral_llm_mapping.copy() + if self.model_config is not None: + if self.model_config.quant_config.quant_algo == QuantAlgo.NVFP4: + quantization_weights_map = { + "weight_packed": "weight", + "input_global_scale": "input_scale", + "weight_global_scale": "weight_scale_2", + } + elif self.model_config.quant_config.quant_algo == QuantAlgo.FP8_BLOCK_SCALES: + quantization_weights_map = { + "weight_scale": "weight_scale_inv", + } + params_map.update(quantization_weights_map) + weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) + + super().load_weights(weights) + From fedb33ed4af5655abd45dbffde3dc23c8cb1c6bc Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Sun, 7 Dec 2025 23:29:01 -0800 Subject: [PATCH 05/21] [WIP] fix bug of mistral large3 llm part Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../_torch/models/checkpoints/mistral/checkpoint_loader.py | 1 - tensorrt_llm/_torch/models/modeling_mistral.py | 2 +- tensorrt_llm/_torch/models/modeling_mistral_large3.py | 1 + 3 files changed, 2 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py index 8d245ae43798..e4e29a561b83 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py @@ -6,7 +6,6 @@ from tensorrt_llm._torch.models.checkpoints.hf.checkpoint_loader import HfCheckpointLoader from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader from tensorrt_llm._torch.models.modeling_utils import register_checkpoint_loader -from tensorrt_llm.quantization.mode import QuantAlgo @register_checkpoint_loader("mistral") diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index ef235a5e2805..6c93be0f88bb 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -445,7 +445,7 @@ def load_weights(self, weights: Dict, weight_mapper=None, *args, **kwargs): vit_weights = weight_mapper.rename_by_params_map( weights=vit_weights, params_map=vit_params_map) - self._vision_tower.load_weights(vit_weights, params_map=vit_params_map) + self._vision_tower.load_weights(vit_weights) logger.debug( f"Successfully loaded weights for {type(self._vision_tower)}") diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index b5daebe76c59..ba811aa45b77 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -11,6 +11,7 @@ MistralLarge3WeightMapper, ) from tensorrt_llm._torch.modules.fused_moe import RenormalizeNaiveMoeRoutingMethod +from tensorrt_llm.quantization.mode import QuantAlgo class Mistral3Gate(nn.Module): From c7fd6a077211c31b51ef8d861e0d2d274eb9cc92 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Sun, 7 Dec 2025 23:47:28 -0800 Subject: [PATCH 06/21] [WIP] update document Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../models/core/mistral_large_3/README.md | 93 ++++++++++--------- 1 file changed, 47 insertions(+), 46 deletions(-) diff --git a/examples/models/core/mistral_large_3/README.md b/examples/models/core/mistral_large_3/README.md index 0d9394bc7aaf..7f8e1269832d 100644 --- a/examples/models/core/mistral_large_3/README.md +++ b/examples/models/core/mistral_large_3/README.md @@ -1,54 +1,55 @@ -## How to use the modules - -The following explains how to use the different modules of Mistral Large V3. - -```python -from tensorrt_llm._torch.models.modeling_deepseekv3 import DeepseekV3ForCausalLM -from tensorrt_llm._torch.models.modeling_mistral import Mistral3VLM -# from tensorrt_llm.llmapi.tokenizer import MistralTokenizer -from tensorrt_llm._torch.models.checkpoints.mistral.checkpoint_loader import MistralCheckpointLoader -from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import MistralLarge3WeightMapper -from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader -from transformers import AutoTokenizer -``` - -### Tokenizer -```python -mtok = AutoTokenizer.from_pretrained(TOKENIZER_DIR) -``` +# Mistral Large V3 -### Config and model instance -```python -config_loader = MistralConfigLoader() -config = config_loader.load(MODEL_DIR) +* Setup the model path -model = Mistral3VLM(model_config=config) -assert isinstance(model.llm, DeepseekV3ForCausalLM) +```bash +export mistral_large_3_model_path= ``` -### Checkpoint loading -```python -weight_mapper=MistralLarge3WeightMapper() -loader = MistralCheckpointLoader(weight_mapper=weight_mapper) +## LLM-only run -weights_dict = loader.load_weights(MODEL_DIR) -``` +* Run the Mistral Large V3 by `quickstart_advanced.py` -### Weight loading -#### E2E -```python -model.load_weights(weights_dict, weight_mapper=weight_mapper) # target usage +```bash +mpirun -n 1 --allow-run-as-root --oversubscribe python3 examples/llm-api/quickstart_advanced.py \ + --model_dir ${mistral_large_3_model_path} \ + --tp_size 4 \ + --moe_ep_size 4 \ + --max_tokens 100 \ + --checkpoint_format mistral \ + --kv_cache_fraction 0.25 \ + --moe_backend TRTLLM # optional ``` -#### By module -```python -def _filter_weights(weights, prefix): - return { - name[len(prefix):]: weight - for name, weight in weights.items() if name.startswith(prefix) - } - -llm_weights = weight_mapper.rename_by_params_map( - params_map=weight_mapper.mistral_llm_mapping, - weights=_filter_weights(weights_dict, "language_model.")) -model.llm.load_weights(llm_weights, weight_mapper=weight_mapper) + +* Launch the trtllm-serve and send a request + +```bash +echo " +backend: pytorch +tensor_parallel_size: 4 +moe_expert_parallel_size: 4 +enable_attention_dp: false +kv_cache_config: + free_gpu_memory_fraction: 0.25 + enable_block_reuse: true +checkpoint_format: mistral +" > serve.yml +mpirun -n 1 --allow-run-as-root --oversubscribe python3 -m tensorrt_llm.commands.serve serve \ + ${mistral_large_3_model_path} \ + --host localhost --port 8001 --backend pytorch \ + --extra_llm_api_options serve.yml \ + --tokenizer ${mistral_large_3_model_path} \ + 2>&1 | tee serve_debug.log & + +curl http://localhost:8001/v1/completions \ + -H "Content-Type: application/json" \ + -d '{ + "model": "${mistral_large_3_model_path}", + "prompt": "The capital of France is", + "max_tokens": 16, + "top_k": 16 + }' + +# The result would be like +{"id":"cmpl-7e342c1d722d4226a1bf3ed35d762c35","object":"text_completion","created":1764061351,"model":"${mistral_large_3_model_path}","choices":[{"index":0,"text":"The capital of France is **Paris**.\n\nParis is the largest city in France and","token_ids":null,"logprobs":null,"context_logits":null,"finish_reason":"length","stop_reason":null,"disaggregated_params":null,"avg_decoded_tokens_per_iter":1.0}],"usage":{"prompt_tokens":7,"total_tokens":23,"completion_tokens":16,"prompt_tokens_details":{"cached_tokens":1}},"prompt_token_ids":null} ``` From dcee6ebb490ccf1e7ffd4d6ec3f06da9dd8f5673 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 00:04:08 -0800 Subject: [PATCH 07/21] [WIP] remove debug file Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- examples/models/core/mistral_large_3/test.py | 37 -------------------- 1 file changed, 37 deletions(-) delete mode 100644 examples/models/core/mistral_large_3/test.py diff --git a/examples/models/core/mistral_large_3/test.py b/examples/models/core/mistral_large_3/test.py deleted file mode 100644 index 535fdbafaad2..000000000000 --- a/examples/models/core/mistral_large_3/test.py +++ /dev/null @@ -1,37 +0,0 @@ -TEST_TOKENIZER = False -TEST_CONFIG_LOADER = False -TEST_CHECKPOINT_LOADER = True - -MODEL_DIR = ( - "/home/scratch.trt_llm_data/llm-models/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512" -) -if TEST_TOKENIZER: - ## Test Tokenizer - from transformers import AutoTokenizer - - mtok = AutoTokenizer.from_pretrained(MODEL_DIR) - print(f"mtok: {mtok}") - - print(mtok.encode("Hello, world!")) - print(mtok.decode([1, 22177, 1044, 4304, 1033])) - -if TEST_CONFIG_LOADER or True: - ## Test config loader - from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import MistralConfigLoader - - config_loader = MistralConfigLoader() - config = config_loader.load(MODEL_DIR) - print(f"config: {config}") - -if TEST_CHECKPOINT_LOADER: - from tensorrt_llm._torch.models.checkpoints.mistral.checkpoint_loader import ( - MistralCheckpointLoader, - ) - from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import ( - MistralLarge3WeightMapper, - ) - - weight_mapper = MistralLarge3WeightMapper() - loader = MistralCheckpointLoader(weight_mapper=weight_mapper) - weights_dict = loader.load_weights(MODEL_DIR, model_config=config) - print(f"weights_dict.keys(): {weights_dict.keys()}") From 793be032cf201c085fcb3e60ab7066b2558274e1 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 00:06:05 -0800 Subject: [PATCH 08/21] [WIP] add checkpoint_format into eval.py Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- tensorrt_llm/commands/eval.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/commands/eval.py b/tensorrt_llm/commands/eval.py index d849a7c91a42..44331780e819 100644 --- a/tensorrt_llm/commands/eval.py +++ b/tensorrt_llm/commands/eval.py @@ -108,6 +108,10 @@ is_flag=True, default=False, help="Flag for disabling KV cache reuse.") +@click.option("--checkpoint_format", + type=click.Choice(["hf", "mistral"]), + default=None, + help="Checkpoint format.") @click.pass_context def main(ctx, model: str, tokenizer: Optional[str], log_level: str, backend: str, max_beam_width: int, max_batch_size: int, @@ -115,7 +119,7 @@ def main(ctx, model: str, tokenizer: Optional[str], log_level: str, ep_size: Optional[int], gpus_per_node: Optional[int], kv_cache_free_gpu_memory_fraction: float, trust_remote_code: bool, revision: Optional[str], extra_llm_api_options: Optional[str], - disable_kv_cache_reuse: bool): + disable_kv_cache_reuse: bool, checkpoint_format: Optional[str]): logger.set_level(log_level) kv_cache_config = KvCacheConfig( @@ -132,6 +136,7 @@ def main(ctx, model: str, tokenizer: Optional[str], log_level: str, "trust_remote_code": trust_remote_code, "revision": revision, "kv_cache_config": kv_cache_config, + "checkpoint_format": checkpoint_format, } if extra_llm_api_options is not None: From 32f26ebbe1067c68e14227ec9c94b0cb1eb7aa90 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 00:22:36 -0800 Subject: [PATCH 09/21] [WIP] refine codes Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../checkpoints/mistral/checkpoint_loader.py | 14 ++++----- .../checkpoints/mistral/config_loader.py | 4 +-- .../_torch/models/modeling_mistral.py | 30 +++++++++---------- .../_torch/models/modeling_mistral_large3.py | 4 +-- 4 files changed, 25 insertions(+), 27 deletions(-) diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py index e4e29a561b83..dede043e6aee 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py @@ -1,5 +1,3 @@ -from typing import Optional - from tensorrt_llm._torch.models.checkpoints.base_config_loader import BaseConfigLoader from tensorrt_llm._torch.models.checkpoints.base_weight_loader import BaseWeightLoader from tensorrt_llm._torch.models.checkpoints.base_weight_mapper import BaseWeightMapper @@ -13,9 +11,9 @@ class MistralCheckpointLoader(HfCheckpointLoader): def __init__( self, *, - weight_loader: Optional[BaseWeightLoader] = None, - weight_mapper: Optional[BaseWeightMapper] = None, - config_loader: Optional[BaseConfigLoader] = None, + weight_loader: BaseWeightLoader | None = None, + weight_mapper: BaseWeightMapper | None = None, + config_loader: BaseConfigLoader | None = None, ): super().__init__( weight_loader=weight_loader, weight_mapper=weight_mapper, config_loader=config_loader @@ -81,9 +79,9 @@ class MistralLarge3CheckpointLoader(MistralCheckpointLoader): def __init__( self, *, - weight_loader: Optional[BaseWeightLoader] = None, - weight_mapper: Optional[BaseWeightMapper] = None, - config_loader: Optional[BaseConfigLoader] = None, + weight_loader: BaseWeightLoader | None = None, + weight_mapper: BaseWeightMapper | None = None, + config_loader: BaseConfigLoader | None = None, ): super().__init__( weight_loader=weight_loader, weight_mapper=weight_mapper, config_loader=config_loader diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py index 2295fb01ff26..8973bf8eb16d 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py @@ -1,6 +1,6 @@ import json from pathlib import Path -from typing import Any, Optional +from typing import Any from transformers import PretrainedConfig, WhisperConfig @@ -231,7 +231,7 @@ def _remap_moe_args(config: dict) -> dict: class MistralConfigLoader(BaseConfigLoader): def _load_mistral_config_dict( self, checkpoint_dir: str, config_file_name: str - ) -> Optional[dict]: + ) -> dict | None: file_path = Path(checkpoint_dir) / Path(config_file_name) if file_path.exists() and file_path.is_file(): diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index 6c93be0f88bb..dfddc5599985 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -1,7 +1,7 @@ import copy import dataclasses import os -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Tuple import torch import torchvision @@ -58,7 +58,7 @@ class MistralAttention(Attention): def __init__( self, model_config: ModelConfig[MistralConfig], - layer_idx: Optional[int] = None, + layer_idx: int | None = None, ): config = model_config.pretrained_config super().__init__( @@ -117,8 +117,8 @@ def forward( position_ids: torch.IntTensor, hidden_states: torch.Tensor, attn_metadata: AttentionMetadata, - residual: Optional[torch.Tensor] = None, - spec_metadata: Optional[SpecMetadata] = None, + residual: torch.Tensor | None = None, + spec_metadata: SpecMetadata | None = None, **kwargs, ) -> torch.Tensor: if residual is None: @@ -175,11 +175,11 @@ def __init__(self, model_config: ModelConfig[MistralConfig]): def forward( self, attn_metadata: AttentionMetadata, - input_ids: Optional[torch.IntTensor] = None, - position_ids: Optional[torch.IntTensor] = None, - inputs_embeds: Optional[torch.FloatTensor] = None, - spec_metadata: Optional[SpecMetadata] = None, - lora_params: Optional[Any] = None, + input_ids: torch.IntTensor | None = None, + position_ids: torch.IntTensor | None = None, + inputs_embeds: torch.FloatTensor | None = None, + spec_metadata: SpecMetadata | None = None, + lora_params: Any | None = None, ) -> torch.Tensor: if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError( @@ -228,7 +228,7 @@ def __init__( self, model_path: str, config: PretrainedConfig, - tokenizer: Optional[AutoTokenizer], + tokenizer: AutoTokenizer | None, trust_remote_code: bool = False, **kwargs, ): @@ -270,7 +270,7 @@ def dtype(self) -> torch.dtype: @torch.inference_mode() def __call__( self, inputs: TextPrompt, sampling_params: SamplingParams - ) -> Tuple[List[int], Optional[ExtraProcessedInputs]]: + ) -> Tuple[List[int], ExtraProcessedInputs | None]: images = inputs.get("multi_modal_data", {}).get("image") mm_processor_kwargs = inputs.get("mm_processor_kwargs", {}) do_rescale = getattr(self.processor.image_processor, "do_rescale", @@ -468,10 +468,10 @@ def infer_max_seq_len(self) -> int: def forward( self, attn_metadata: AttentionMetadata, - input_ids: Optional[torch.LongTensor] = None, - position_ids: Optional[torch.LongTensor] = None, + input_ids: torch.LongTensor | None = None, + position_ids: torch.LongTensor | None = None, return_context_logits: bool = False, - spec_metadata: Optional[SpecMetadata] = None, + spec_metadata: SpecMetadata | None = None, **kwargs, ) -> torch.Tensor: """Forward method.""" @@ -650,7 +650,7 @@ def mm_token_ids(self): def load_draft_weights( self, weights: Dict, - weight_mapper: Optional[MistralWeightMapper] = None) -> None: + weight_mapper: MistralWeightMapper | None = None) -> None: self.llm.load_draft_weights(weights, weight_mapper=weight_mapper) diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index ba811aa45b77..b9579829b666 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -5,7 +5,7 @@ from ..model_config import ModelConfig from torch import nn -from typing import Dict, Optional, List +from typing import Dict, List from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import ( MistralLarge3WeightMapper, @@ -21,7 +21,7 @@ def __init__( hidden_size: int, num_experts: int, top_k: int, - dtype: Optional[torch.dtype] = None, + dtype: torch.dtype | None = None, **kwargs, ): super().__init__() From c7f80a9c6e030c16201bd8ed5a97c8afb02738f5 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 01:18:39 -0800 Subject: [PATCH 10/21] [WIP] add mistral large 3 into CI tests Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../_torch/models/modeling_mistral.py | 1 - .../defs/accuracy/references/gsm8k.yaml | 3 ++ .../defs/accuracy/references/mmlu.yaml | 3 ++ .../defs/accuracy/test_llm_api_pytorch.py | 52 +++++++++++++++++++ .../test-db/l0_gb200_multi_gpus.yml | 1 + 5 files changed, 59 insertions(+), 1 deletion(-) diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index dfddc5599985..7625013617f7 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -440,7 +440,6 @@ def load_weights(self, weights: Dict, weight_mapper=None, *args, **kwargs): vit_weights = filter_weights(weights=weights, prefix="vision_tower") logger.debug(f"Loading weights for {type(self._vision_tower)}") - # FIXME rename_weights_with_regex in _load_weights_impl breaks this, fall back to manual renaming if vit_params_map is not None: vit_weights = weight_mapper.rename_by_params_map( weights=vit_weights, params_map=vit_params_map) diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index c62ff5a0d898..af1faf9ec309 100644 --- a/tests/integration/defs/accuracy/references/gsm8k.yaml +++ b/tests/integration/defs/accuracy/references/gsm8k.yaml @@ -281,3 +281,6 @@ bigcode/starcoder2-7b: - accuracy: 26.5 bigcode/starcoder2-15b: - accuracy: 54.5 +mistral/Mistral-Large-3-675B: + - quant_algo: NVFP4 + accuracy: 90.83 diff --git a/tests/integration/defs/accuracy/references/mmlu.yaml b/tests/integration/defs/accuracy/references/mmlu.yaml index dd404ba8f7bb..c52278618fcc 100644 --- a/tests/integration/defs/accuracy/references/mmlu.yaml +++ b/tests/integration/defs/accuracy/references/mmlu.yaml @@ -340,3 +340,6 @@ mistralai/Mistral-Nemo-12b-Base: - accuracy: 69.66 - quant_algo: FP8 accuracy: 69.66 +mistral/Mistral-Large-3-675B: + - quant_algo: NVFP4 + accuracy: 87.54 diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 35e60e043600..c711bab32bf6 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -4800,3 +4800,55 @@ def test_auto_dtype(self): task.evaluate(llm, sampling_params=sampling_params, extra_evaluator_kwargs=extra_evaluator_kwargs) + + +class TestMistralLarge3_675B(LlmapiAccuracyTestHarness): + MODEL_NAME = "mistral/Mistral-Large-3-675B" + + @skip_pre_blackwell + @pytest.mark.skip_less_mpi_world_size(4) + @pytest.mark.parametrize( + "tp_size,pp_size,ep_size,attention_dp,cuda_graph,overlap_scheduler,moe_backend,eagle3", + [ + (4, 1, 4, False, True, True, "TRTLLM", False), + ], + ids=[ + "latency_moe_trtllm", + ], + ) + def test_nvfp4_4gpus(self, tp_size, pp_size, ep_size, attention_dp, + cuda_graph, overlap_scheduler, moe_backend, eagle3): + + if moe_backend == "TRTLLM" and (get_sm_version() == 120 + or get_sm_version() == 121): + pytest.skip( + "MOE TRTLLM backend does not support SM version 120 or 121") + + pytorch_config = dict( + disable_overlap_scheduler=not overlap_scheduler, + cuda_graph_config=CudaGraphConfig() if cuda_graph else None, + moe_config=MoeConfig(backend=moe_backend)) + + kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.4, + enable_block_reuse=not eagle3) + spec_config = None + if eagle3: + spec_config = EagleDecodingConfig( + max_draft_len=2, + speculative_model_dir= + f"{llm_models_root()}/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512-Eagle/", + eagle3_one_model=True) + with LLM( + f"{llm_models_root()}/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512-NVFP4/", + tensor_parallel_size=tp_size, + pipeline_parallel_size=pp_size, + moe_expert_parallel_size=ep_size, + **pytorch_config, + enable_attention_dp=attention_dp, + kv_cache_config=kv_cache_config, + speculative_config=spec_config) as llm: + + task = MMLU(self.MODEL_NAME) + task.evaluate(llm) + task = GSM8K(self.MODEL_NAME) + task.evaluate(llm) diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 5c5bc4132b58..67519af4bcca 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -75,3 +75,4 @@ l0_gb200_multi_gpus: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_nvfp4_4gpus[moe_backend=TRTLLM-mtp_nextn=2-tp4-fp8kv=True-attention_dp=True-cuda_graph=True-overlap_scheduler=True-torch_compile=False] - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_nvfp4_4gpus[moe_backend=TRTLLM-mtp_nextn=2-ep4-fp8kv=True-attention_dp=True-cuda_graph=True-overlap_scheduler=True-torch_compile=False] - accuracy/test_llm_api_pytorch.py::TestQwen3_235B_A22B::test_nvfp4_4gpus[latency_moe_trtllm_eagle3] TIMEOUT (90) + - accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_nvfp4_4gpus[latency_moe_trtllm] TIMEOUT (90) From f618f4d46f748c5b27c843df2d5d5863035fedbe Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 17:10:46 -0800 Subject: [PATCH 11/21] [WIP] clean codes and refine config_loader of mistral Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- tensorrt_llm/_torch/model_config.py | 25 ----------------- .../models/checkpoints/hf/weight_loader.py | 1 - .../checkpoints/mistral/config_loader.py | 27 +++++++++++++++---- .../checkpoints/mistral/weight_mapper.py | 1 - 4 files changed, 22 insertions(+), 32 deletions(-) diff --git a/tensorrt_llm/_torch/model_config.py b/tensorrt_llm/_torch/model_config.py index f4e4497b72a2..148ec5e2e3fd 100644 --- a/tensorrt_llm/_torch/model_config.py +++ b/tensorrt_llm/_torch/model_config.py @@ -341,31 +341,6 @@ def load_hf_quant_config(hf_quant_config, moe_backend): 'block.*.attn.out', 'block.*.mlp.gate', 'block.*.attn.qkv', 'embedding', 'unembedding' ] - # Mistral checkpoints. - elif hf_quant_config.get("quant_method") == "compressed-tensors": - if 'NVFP4' in hf_quant_config.get("config_groups"): - quant_config.quant_algo = QuantAlgo.NVFP4 - quant_config.group_size = 16 - quant_config.exclude_modules = [ - "layers.*.attention.wq*", "layers.*.attention.wk*", - "patch_merger*", "vision_encoder*", - "vision_language_adapter*", - "language_model.model.layers.*.self_attn.q*", - "language_model.model.layers.*.self_attn.k*", - "language_model.model.layers.*.self_attn.fused_qkv*", - "language_model.model.layers.*.self_attn.kv_a_proj_with_mqa*", - "model.layers.*.self_attn.q*", - "model.layers.*.self_attn.k*", - "model.layers.*.self_attn.fused_qkv*", - "model.layers.*.self_attn.kv_a_proj_with_mqa*", - "model.layers.*.self_attn.o_proj" - ] - elif 'FP8_BLOCK' in hf_quant_config.get("config_groups"): - quant_config.quant_algo = QuantAlgo.FP8_BLOCK_SCALES - quant_config.group_size = 128 - quant_config.exclude_modules = [ - "*q_a_proj*", "*kv_a_proj_with_mqa*" - ] return quant_config, layer_quant_config diff --git a/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py b/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py index fd0ba287abe4..3b1c3af1727d 100644 --- a/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/hf/weight_loader.py @@ -19,7 +19,6 @@ from tensorrt_llm.mapping import Mapping -# @register_checkpoint_weight_loader("mistral_large_3") @register_checkpoint_weight_loader("mistral") @register_checkpoint_weight_loader("HF") class HfWeightLoader(BaseWeightLoader): diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py index 8973bf8eb16d..45e7b43bc5b7 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py @@ -8,6 +8,7 @@ from tensorrt_llm._torch.models.checkpoints.base_config_loader import BaseConfigLoader from tensorrt_llm._torch.models.modeling_utils import register_config_loader from tensorrt_llm.models.modeling_utils import QuantConfig +from tensorrt_llm.quantization.mode import QuantAlgo ################### # vllm code here @@ -287,11 +288,27 @@ def load(self, checkpoint_dir: str, **kwargs) -> ModelConfig: moe_backend = kwargs.get("moe_backend", "CUTLASS") hf_quant_config = pretrained_config.quantization_config - quant_config, layer_quant_config = ModelConfig.load_hf_quant_config( - hf_quant_config, moe_backend - ) - - kwargs.pop("trust_remote_code", None) + if hf_quant_config.get("quant_method") == "compressed-tensors": + if 'NVFP4' in hf_quant_config.get("config_groups"): + quant_config.quant_algo = QuantAlgo.NVFP4 + quant_config.group_size = 16 + ignore_list = hf_quant_config.get("ignore", []) + quant_config.exclude_modules = [] + if "re:.*attn.*" in ignore_list: + quant_config.exclude_modules.append("model.layers.*.self_attn.*") + if "re:vision_encoder.*" in ignore_list: + quant_config.exclude_modules.append("vision_encoder*") + if "re:vision_language_adapter.*" in ignore_list: + quant_config.exclude_modules.append("vision_language_adapter*") + + elif 'FP8_BLOCK' in hf_quant_config.get("config_groups"): + quant_config.quant_algo = QuantAlgo.FP8_BLOCK_SCALES + quant_config.group_size = 128 + quant_config.exclude_modules = [ + "*q_a_proj*", "*kv_a_proj_with_mqa*" + ] + + kwargs.pop("trust_remote_code", None) # ModelConfig does not have this input parameter model_config = ModelConfig( pretrained_config=pretrained_config, quant_config=quant_config, diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py b/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py index e70d09fcd0e8..28362f1f90f6 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/weight_mapper.py @@ -12,7 +12,6 @@ def __init__(self): self._callbacks.append(self._permute_qk) - # TODO move to registry self.pixtral_mapping = { "wq": "q_proj", "wk": "k_proj", From 135fb6dd10230b9d9232e839646456849eee621b Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 18:31:41 -0800 Subject: [PATCH 12/21] [WIP] refine codes about CI and evaluation Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../models/core/mistral_large_3/README.md | 3 +- .../_torch/pyexecutor/config_utils.py | 2 +- tensorrt_llm/commands/eval.py | 7 +-- .../defs/accuracy/test_llm_api_pytorch.py | 52 +++++++++++++++++++ .../test_lists/test-db/l0_dgx_b200.yml | 1 + .../test-db/l0_gb200_multi_gpus.yml | 2 +- 6 files changed, 57 insertions(+), 10 deletions(-) diff --git a/examples/models/core/mistral_large_3/README.md b/examples/models/core/mistral_large_3/README.md index 7f8e1269832d..7f91c4f02d29 100644 --- a/examples/models/core/mistral_large_3/README.md +++ b/examples/models/core/mistral_large_3/README.md @@ -17,8 +17,7 @@ mpirun -n 1 --allow-run-as-root --oversubscribe python3 examples/llm-api/quickst --moe_ep_size 4 \ --max_tokens 100 \ --checkpoint_format mistral \ - --kv_cache_fraction 0.25 \ - --moe_backend TRTLLM # optional + --moe_backend TRTLLM ``` * Launch the trtllm-serve and send a request diff --git a/tensorrt_llm/_torch/pyexecutor/config_utils.py b/tensorrt_llm/_torch/pyexecutor/config_utils.py index f2c0e51d3936..e4fa9da6e6cf 100644 --- a/tensorrt_llm/_torch/pyexecutor/config_utils.py +++ b/tensorrt_llm/_torch/pyexecutor/config_utils.py @@ -48,7 +48,7 @@ def load_pretrained_config(model_name_or_path: str, config_class = _CONFIG_REGISTRY[model_type] model_config = config_class.from_pretrained(model_name_or_path, **kwargs) - elif checkpoint_format == "mistral_large_3": + elif checkpoint_format in ("mistral", "mistral_large_3"): from tensorrt_llm._torch.models.checkpoints.mistral.config_loader import \ MistralConfigLoader model_config = getattr( diff --git a/tensorrt_llm/commands/eval.py b/tensorrt_llm/commands/eval.py index 44331780e819..afa990ce93d5 100644 --- a/tensorrt_llm/commands/eval.py +++ b/tensorrt_llm/commands/eval.py @@ -108,10 +108,6 @@ is_flag=True, default=False, help="Flag for disabling KV cache reuse.") -@click.option("--checkpoint_format", - type=click.Choice(["hf", "mistral"]), - default=None, - help="Checkpoint format.") @click.pass_context def main(ctx, model: str, tokenizer: Optional[str], log_level: str, backend: str, max_beam_width: int, max_batch_size: int, @@ -119,7 +115,7 @@ def main(ctx, model: str, tokenizer: Optional[str], log_level: str, ep_size: Optional[int], gpus_per_node: Optional[int], kv_cache_free_gpu_memory_fraction: float, trust_remote_code: bool, revision: Optional[str], extra_llm_api_options: Optional[str], - disable_kv_cache_reuse: bool, checkpoint_format: Optional[str]): + disable_kv_cache_reuse: bool): logger.set_level(log_level) kv_cache_config = KvCacheConfig( @@ -136,7 +132,6 @@ def main(ctx, model: str, tokenizer: Optional[str], log_level: str, "trust_remote_code": trust_remote_code, "revision": revision, "kv_cache_config": kv_cache_config, - "checkpoint_format": checkpoint_format, } if extra_llm_api_options is not None: diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index c711bab32bf6..7009bebaf09f 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -4807,6 +4807,7 @@ class TestMistralLarge3_675B(LlmapiAccuracyTestHarness): @skip_pre_blackwell @pytest.mark.skip_less_mpi_world_size(4) + @pytest.mark.skip_less_device_memory(183000) @pytest.mark.parametrize( "tp_size,pp_size,ep_size,attention_dp,cuda_graph,overlap_scheduler,moe_backend,eagle3", [ @@ -4840,6 +4841,57 @@ def test_nvfp4_4gpus(self, tp_size, pp_size, ep_size, attention_dp, eagle3_one_model=True) with LLM( f"{llm_models_root()}/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512-NVFP4/", + checkpoint_format="mistral", + tensor_parallel_size=tp_size, + pipeline_parallel_size=pp_size, + moe_expert_parallel_size=ep_size, + **pytorch_config, + enable_attention_dp=attention_dp, + kv_cache_config=kv_cache_config, + speculative_config=spec_config) as llm: + + task = MMLU(self.MODEL_NAME) + task.evaluate(llm) + task = GSM8K(self.MODEL_NAME) + task.evaluate(llm) + + @skip_pre_blackwell + @pytest.mark.skip_less_mpi_world_size(8) + @pytest.mark.skip_less_device_memory(183000) + @pytest.mark.parametrize( + "tp_size,pp_size,ep_size,attention_dp,cuda_graph,overlap_scheduler,moe_backend,eagle3", + [ + (8, 1, 8, False, True, True, "DEEPGEMM", False), + ], + ids=[ + "latency_moe_deepgemm", + ], + ) + def test_fp8(self, tp_size, pp_size, ep_size, attention_dp, + cuda_graph, overlap_scheduler, moe_backend, eagle3): + + if moe_backend == "DEEPGEMM" and (get_sm_version() == 120 + or get_sm_version() == 121): + pytest.skip( + "MOE DEEPGEMM backend does not support SM version 120 or 121") + + pytorch_config = dict( + disable_overlap_scheduler=not overlap_scheduler, + cuda_graph_config=CudaGraphConfig() if cuda_graph else None, + moe_config=MoeConfig(backend=moe_backend)) + + kv_cache_config = KvCacheConfig(free_gpu_memory_fraction=0.4, + enable_block_reuse=not eagle3) + spec_config = None + if eagle3: + spec_config = EagleDecodingConfig( + max_draft_len=2, + speculative_model_dir= + f"{llm_models_root()}/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512-Eagle/", + eagle3_one_model=True) + with LLM( + f"{llm_models_root()}/Mistral-Large-3-675B/Mistral-Large-3-675B-Instruct-2512/", + checkpoint_format="mistral", tensor_parallel_size=tp_size, pipeline_parallel_size=pp_size, moe_expert_parallel_size=ep_size, diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index 04a4278ba6f0..00dcc1fa4ace 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -114,6 +114,7 @@ l0_dgx_b200: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus[baseline] TIMEOUT (180) - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus[baseline_mtp1] TIMEOUT (180) - accuracy/test_disaggregated_serving.py::TestDeepSeekV32Exp::test_auto_dtype[False] TIMEOUT (360) + - accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] TIMEOUT (90) - condition: ranges: system_gpu_count: diff --git a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml index 67519af4bcca..e8f7b6cf58c9 100644 --- a/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml +++ b/tests/integration/test_lists/test-db/l0_gb200_multi_gpus.yml @@ -75,4 +75,4 @@ l0_gb200_multi_gpus: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_nvfp4_4gpus[moe_backend=TRTLLM-mtp_nextn=2-tp4-fp8kv=True-attention_dp=True-cuda_graph=True-overlap_scheduler=True-torch_compile=False] - accuracy/test_llm_api_pytorch.py::TestDeepSeekV3Lite::test_nvfp4_4gpus[moe_backend=TRTLLM-mtp_nextn=2-ep4-fp8kv=True-attention_dp=True-cuda_graph=True-overlap_scheduler=True-torch_compile=False] - accuracy/test_llm_api_pytorch.py::TestQwen3_235B_A22B::test_nvfp4_4gpus[latency_moe_trtllm_eagle3] TIMEOUT (90) - - accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_nvfp4_4gpus[latency_moe_trtllm] TIMEOUT (90) + - accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] TIMEOUT (90) From 77f16a0df53f7b9b6d0567c3a5fde67d8777f2fc Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 18:42:25 -0800 Subject: [PATCH 13/21] [WIP] fix few bugs Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../checkpoints/mistral/checkpoint_loader.py | 20 +++---------------- tensorrt_llm/commands/eval.py | 2 +- .../defs/accuracy/references/gsm8k.yaml | 2 ++ .../defs/accuracy/references/mmlu.yaml | 2 ++ 4 files changed, 8 insertions(+), 18 deletions(-) diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py index dede043e6aee..433bde665b29 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/checkpoint_loader.py @@ -44,20 +44,7 @@ def preprocess_weights(self, weights: dict) -> dict: return hf_weights - def broadcast_per_tensor_scales(self, weights): - import math - - scales = [k for k in weights.keys() if k.endswith("qscale_weight")] - for scale in scales: - name = ".".join(scale.split(".")[:-1]) - weight_shape = weights[f"{name}.weight"].shape - broadcast = weights[scale].expand( - math.ceil(weight_shape[0] / 128), - math.ceil(weight_shape[1] / 128), - ) - weights[scale] = broadcast[:] - - def reverse_nvfp4_global_scales(self, weights): + def inverse_nvfp4_global_scales(self, weights): for key in weights.keys(): if "global_scale" in key: weights[key] = 1.0 / weights[key] @@ -65,9 +52,8 @@ def reverse_nvfp4_global_scales(self, weights): def load_weights(self, checkpoint_dir: str, **kwargs): weights = super().weight_loader.load_weights(checkpoint_dir, **kwargs) weights = self.preprocess_weights(weights) - # FIXME mimic DS fp8 till per tensor supported - self.broadcast_per_tensor_scales(weights) - self.reverse_nvfp4_global_scales(weights) + # The definition of global_scale is different in Mistral, need to inverse the scale + self.inverse_nvfp4_global_scales(weights) return weights def get_default_config_loader(self) -> MistralConfigLoader: diff --git a/tensorrt_llm/commands/eval.py b/tensorrt_llm/commands/eval.py index afa990ce93d5..d849a7c91a42 100644 --- a/tensorrt_llm/commands/eval.py +++ b/tensorrt_llm/commands/eval.py @@ -115,7 +115,7 @@ def main(ctx, model: str, tokenizer: Optional[str], log_level: str, ep_size: Optional[int], gpus_per_node: Optional[int], kv_cache_free_gpu_memory_fraction: float, trust_remote_code: bool, revision: Optional[str], extra_llm_api_options: Optional[str], - disable_kv_cache_reuse: bool): + disable_kv_cache_reuse: bool): logger.set_level(log_level) kv_cache_config = KvCacheConfig( diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index af1faf9ec309..d0cb5b88f7b7 100644 --- a/tests/integration/defs/accuracy/references/gsm8k.yaml +++ b/tests/integration/defs/accuracy/references/gsm8k.yaml @@ -284,3 +284,5 @@ bigcode/starcoder2-15b: mistral/Mistral-Large-3-675B: - quant_algo: NVFP4 accuracy: 90.83 + - quant_algo: FP8_BLOCK_SCALES + accuracy: 90.83 diff --git a/tests/integration/defs/accuracy/references/mmlu.yaml b/tests/integration/defs/accuracy/references/mmlu.yaml index c52278618fcc..41633c3e0695 100644 --- a/tests/integration/defs/accuracy/references/mmlu.yaml +++ b/tests/integration/defs/accuracy/references/mmlu.yaml @@ -343,3 +343,5 @@ mistralai/Mistral-Nemo-12b-Base: mistral/Mistral-Large-3-675B: - quant_algo: NVFP4 accuracy: 87.54 + - quant_algo: FP8_BLOCK_SCALES + accuracy: 90.83 From ef7aa71677d1cac080a1ed244ad2c8234cace699 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 18:55:11 -0800 Subject: [PATCH 14/21] [READY] fix code format Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../checkpoints/mistral/config_loader.py | 17 ++++------ .../_torch/models/modeling_deepseekv3.py | 21 ++++++------ .../_torch/models/modeling_mistral.py | 7 ++-- .../_torch/models/modeling_mistral_large3.py | 32 ++++++++----------- .../defs/accuracy/test_llm_api_pytorch.py | 6 ++-- 5 files changed, 37 insertions(+), 46 deletions(-) diff --git a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py index 45e7b43bc5b7..95e93fdc0523 100644 --- a/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py +++ b/tensorrt_llm/_torch/models/checkpoints/mistral/config_loader.py @@ -230,9 +230,7 @@ def _remap_moe_args(config: dict) -> dict: @register_config_loader("mistral") @register_config_loader("mistral_large_3") class MistralConfigLoader(BaseConfigLoader): - def _load_mistral_config_dict( - self, checkpoint_dir: str, config_file_name: str - ) -> dict | None: + def _load_mistral_config_dict(self, checkpoint_dir: str, config_file_name: str) -> dict | None: file_path = Path(checkpoint_dir) / Path(config_file_name) if file_path.exists() and file_path.is_file(): @@ -285,11 +283,10 @@ def load(self, checkpoint_dir: str, **kwargs) -> ModelConfig: pretrained_config.torch_dtype = getattr(pretrained_config, "dtype", None) quant_config = QuantConfig() layer_quant_config = None - moe_backend = kwargs.get("moe_backend", "CUTLASS") hf_quant_config = pretrained_config.quantization_config if hf_quant_config.get("quant_method") == "compressed-tensors": - if 'NVFP4' in hf_quant_config.get("config_groups"): + if "NVFP4" in hf_quant_config.get("config_groups"): quant_config.quant_algo = QuantAlgo.NVFP4 quant_config.group_size = 16 ignore_list = hf_quant_config.get("ignore", []) @@ -300,15 +297,13 @@ def load(self, checkpoint_dir: str, **kwargs) -> ModelConfig: quant_config.exclude_modules.append("vision_encoder*") if "re:vision_language_adapter.*" in ignore_list: quant_config.exclude_modules.append("vision_language_adapter*") - - elif 'FP8_BLOCK' in hf_quant_config.get("config_groups"): + + elif "FP8_BLOCK" in hf_quant_config.get("config_groups"): quant_config.quant_algo = QuantAlgo.FP8_BLOCK_SCALES quant_config.group_size = 128 - quant_config.exclude_modules = [ - "*q_a_proj*", "*kv_a_proj_with_mqa*" - ] + quant_config.exclude_modules = ["*q_a_proj*", "*kv_a_proj_with_mqa*"] - kwargs.pop("trust_remote_code", None) # ModelConfig does not have this input parameter + kwargs.pop("trust_remote_code", None) # ModelConfig does not have this input parameter model_config = ModelConfig( pretrained_config=pretrained_config, quant_config=quant_config, diff --git a/tensorrt_llm/_torch/models/modeling_deepseekv3.py b/tensorrt_llm/_torch/models/modeling_deepseekv3.py index 82588c866fff..8df4eae7066a 100755 --- a/tensorrt_llm/_torch/models/modeling_deepseekv3.py +++ b/tensorrt_llm/_torch/models/modeling_deepseekv3.py @@ -749,17 +749,16 @@ def __init__(self, gate_cls = DeepseekV3Gate if hasattr(model_config.pretrained_config, "gate_cls"): gate_cls = model_config.pretrained_config.gate_cls - self.gate = gate_cls( - hidden_size, - num_experts, - top_k=top_k, - n_group=config.n_group, - topk_group=config.topk_group, - routed_scaling_factor=config.routed_scaling_factor, - dtype=dtype, - fuse_routing_kernel=True, - apply_routing=False, - moe_backend=model_config.moe_backend) + self.gate = gate_cls(hidden_size, + num_experts, + top_k=top_k, + n_group=config.n_group, + topk_group=config.topk_group, + routed_scaling_factor=config.routed_scaling_factor, + dtype=dtype, + fuse_routing_kernel=True, + apply_routing=False, + moe_backend=model_config.moe_backend) self.experts = create_moe( num_experts=num_experts, routing_method=self.gate.routing_method, diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index 7625013617f7..b1dcbe6d15a8 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -5,7 +5,6 @@ import torch import torchvision -from mistral_common.tokens.tokenizers.multimodal import ImageEncoder from torch import nn from transformers import (AutoProcessor, AutoTokenizer, Mistral3Config, MistralConfig, PretrainedConfig, PreTrainedModel) @@ -18,7 +17,8 @@ from tensorrt_llm._torch.models import modeling_pixtral from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import \ MistralWeightMapper -from tensorrt_llm._torch.models.modeling_mistral_large3 import MistralLarge3ForCausalLM, Mistral3Gate +from tensorrt_llm._torch.models.modeling_mistral_large3 import ( + Mistral3Gate, MistralLarge3ForCausalLM) from tensorrt_llm._torch.models.modeling_multimodal_utils import ( find_input_mm_embeds, fuse_input_embeds, get_multimodal_embeddings) from tensorrt_llm._torch.models.modeling_utils import (DecoderModel, @@ -413,7 +413,8 @@ def __init__( self._vision_tower = modeling_pixtral.PixtralVisionModel( vision_model_config) - self._multi_modal_projector = Mistral3MultiModalProjector(model_config).eval().to(self._device) + self._multi_modal_projector = Mistral3MultiModalProjector( + model_config).eval().to(self._device) self._post_config() self.is_loaded = True diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index b9579829b666..15486988c8e6 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -1,21 +1,18 @@ -import torch - -from ..models.modeling_deepseekv3 import DeepseekV3ForCausalLM -from .modeling_utils import register_auto_model -from ..model_config import ModelConfig +from typing import Dict, List +import torch from torch import nn -from typing import Dict, List -from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import ( - MistralLarge3WeightMapper, -) +from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import MistralLarge3WeightMapper from tensorrt_llm._torch.modules.fused_moe import RenormalizeNaiveMoeRoutingMethod from tensorrt_llm.quantization.mode import QuantAlgo +from ..model_config import ModelConfig +from ..models.modeling_deepseekv3 import DeepseekV3ForCausalLM +from .modeling_utils import register_auto_model -class Mistral3Gate(nn.Module): +class Mistral3Gate(nn.Module): def __init__( self, hidden_size: int, @@ -25,18 +22,17 @@ def __init__( **kwargs, ): super().__init__() - self.weight = nn.Parameter(torch.empty((num_experts, hidden_size), - dtype=dtype), - requires_grad=False) + self.weight = nn.Parameter( + torch.empty((num_experts, hidden_size), dtype=dtype), requires_grad=False + ) self.top_k = top_k self.dtype = dtype self.routing_method = RenormalizeNaiveMoeRoutingMethod(top_k=self.top_k) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - logits: torch.Tensor = torch.ops.trtllm.cublas_mm(hidden_states, - self.weight.t(), - bias=None, - out_dtype=self.dtype) + logits: torch.Tensor = torch.ops.trtllm.cublas_mm( + hidden_states, self.weight.t(), bias=None, out_dtype=self.dtype + ) return logits def load_weights(self, weights: List[Dict]): @@ -44,6 +40,7 @@ def load_weights(self, weights: List[Dict]): self.weight.copy_(weights[0]["weight"][:]) + @register_auto_model("MistralLarge3ForCausalLM") class MistralLarge3ForCausalLM(DeepseekV3ForCausalLM): def __init__(self, model_config: ModelConfig): @@ -71,4 +68,3 @@ def load_weights(self, weights: Dict, *args, **kwargs): weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) super().load_weights(weights) - diff --git a/tests/integration/defs/accuracy/test_llm_api_pytorch.py b/tests/integration/defs/accuracy/test_llm_api_pytorch.py index 7009bebaf09f..872604d54848 100644 --- a/tests/integration/defs/accuracy/test_llm_api_pytorch.py +++ b/tests/integration/defs/accuracy/test_llm_api_pytorch.py @@ -4867,11 +4867,11 @@ def test_nvfp4_4gpus(self, tp_size, pp_size, ep_size, attention_dp, "latency_moe_deepgemm", ], ) - def test_fp8(self, tp_size, pp_size, ep_size, attention_dp, - cuda_graph, overlap_scheduler, moe_backend, eagle3): + def test_fp8(self, tp_size, pp_size, ep_size, attention_dp, cuda_graph, + overlap_scheduler, moe_backend, eagle3): if moe_backend == "DEEPGEMM" and (get_sm_version() == 120 - or get_sm_version() == 121): + or get_sm_version() == 121): pytest.skip( "MOE DEEPGEMM backend does not support SM version 120 or 121") From fcf15b77daf18b65abdf7fc6e3d90f4daed97934 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 19:15:35 -0800 Subject: [PATCH 15/21] [DOC] quick update for document Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- examples/models/core/mistral_large_3/README.md | 1 - 1 file changed, 1 deletion(-) diff --git a/examples/models/core/mistral_large_3/README.md b/examples/models/core/mistral_large_3/README.md index 7f91c4f02d29..dfd3fd0c2856 100644 --- a/examples/models/core/mistral_large_3/README.md +++ b/examples/models/core/mistral_large_3/README.md @@ -29,7 +29,6 @@ tensor_parallel_size: 4 moe_expert_parallel_size: 4 enable_attention_dp: false kv_cache_config: - free_gpu_memory_fraction: 0.25 enable_block_reuse: true checkpoint_format: mistral " > serve.yml From 7c0633f0ac5d1f9abf6254bd85fa9b6c5bc082a7 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 19:28:38 -0800 Subject: [PATCH 16/21] [FIX] fix issues by bot comments Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- requirements.txt | 2 +- tensorrt_llm/_torch/models/modeling_mistral.py | 4 ++-- tensorrt_llm/_torch/models/modeling_mistral_large3.py | 6 ++++-- 3 files changed, 7 insertions(+), 5 deletions(-) diff --git a/requirements.txt b/requirements.txt index 1bc253c3f0a8..8f740a9ede14 100644 --- a/requirements.txt +++ b/requirements.txt @@ -75,4 +75,4 @@ numexpr<2.14.0 # WAR for attempted use of nonexistent numpy.typing partial_json_parser apache-tvm-ffi==0.1.4 # used for reduce nvidia-cutlass-dsl host overhead torch-c-dlpack-ext==0.1.3 # used for reduce nvidia-cutlass-dsl host overhead, optional package for improved torch tensor calling perf -mistral-common +mistral-common==1.8.6 diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index b1dcbe6d15a8..f0e8906a7bce 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -432,10 +432,10 @@ def load_weights(self, weights: Dict, weight_mapper=None, *args, **kwargs): llm_weights = filter_weights(weights=weights, prefix="language_model") logger.debug(f"Loading weights for {type(self.llm)}") + all_kwargs = {'weight_mapper': weight_mapper, **kwargs} self.llm.load_weights(llm_weights, - weight_mapper=weight_mapper, *args, - **kwargs) + **all_kwargs) logger.debug(f"Successfully loaded weights for {type(self.llm)}") vit_weights = filter_weights(weights=weights, prefix="vision_tower") diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index 15486988c8e6..50760b2f4837 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -50,10 +50,11 @@ def __init__(self, model_config: ModelConfig): def forward(self, *args, **kwargs): return super().forward(*args, **kwargs) - def load_weights(self, weights: Dict, *args, **kwargs): + def load_weights(self, weights: Dict): assert self.model_config is not None, "self.model_config is required" params_map = self.weight_mapper.mistral_llm_mapping.copy() if self.model_config is not None: + quantization_weights_map: Dict[str, str] = {} if self.model_config.quant_config.quant_algo == QuantAlgo.NVFP4: quantization_weights_map = { "weight_packed": "weight", @@ -64,7 +65,8 @@ def load_weights(self, weights: Dict, *args, **kwargs): quantization_weights_map = { "weight_scale": "weight_scale_inv", } - params_map.update(quantization_weights_map) + if quantization_weights_map: + params_map.update(quantization_weights_map) weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) super().load_weights(weights) From b36858a5ab4c9c1ac6a728584f3563a4c03b620e Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 20:23:23 -0800 Subject: [PATCH 17/21] [FIX] fix bug of code format Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- tensorrt_llm/_torch/models/modeling_mistral.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index f0e8906a7bce..b41da0c02d5a 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -433,9 +433,7 @@ def load_weights(self, weights: Dict, weight_mapper=None, *args, **kwargs): llm_weights = filter_weights(weights=weights, prefix="language_model") logger.debug(f"Loading weights for {type(self.llm)}") all_kwargs = {'weight_mapper': weight_mapper, **kwargs} - self.llm.load_weights(llm_weights, - *args, - **all_kwargs) + self.llm.load_weights(llm_weights, *args, **all_kwargs) logger.debug(f"Successfully loaded weights for {type(self.llm)}") vit_weights = filter_weights(weights=weights, prefix="vision_tower") From 62fa048fa23d896c2eedf873cf674faedc8c6091 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Mon, 8 Dec 2025 20:37:11 -0800 Subject: [PATCH 18/21] [FIX] fix minor bugs Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- tensorrt_llm/_torch/models/modeling_mistral.py | 3 +-- 1 file changed, 1 insertion(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_mistral.py b/tensorrt_llm/_torch/models/modeling_mistral.py index b41da0c02d5a..2667d20d55ec 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral.py +++ b/tensorrt_llm/_torch/models/modeling_mistral.py @@ -432,8 +432,7 @@ def load_weights(self, weights: Dict, weight_mapper=None, *args, **kwargs): llm_weights = filter_weights(weights=weights, prefix="language_model") logger.debug(f"Loading weights for {type(self.llm)}") - all_kwargs = {'weight_mapper': weight_mapper, **kwargs} - self.llm.load_weights(llm_weights, *args, **all_kwargs) + self.llm.load_weights(llm_weights) logger.debug(f"Successfully loaded weights for {type(self.llm)}") vit_weights = filter_weights(weights=weights, prefix="vision_tower") From cd00efff0fb55c4d7d623da7bb519c614f69202f Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Tue, 9 Dec 2025 10:24:02 +0000 Subject: [PATCH 19/21] [FIX] minor changes based on comment Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- .../_torch/models/modeling_mistral_large3.py | 34 +++++++++---------- 1 file changed, 16 insertions(+), 18 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index 50760b2f4837..6a0beb1b29c2 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -5,12 +5,11 @@ from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import MistralLarge3WeightMapper from tensorrt_llm._torch.modules.fused_moe import RenormalizeNaiveMoeRoutingMethod +from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.modeling_deepseekv3 import DeepseekV3ForCausalLM +from tensorrt_llm._torch.models.modeling_utils import register_auto_model from tensorrt_llm.quantization.mode import QuantAlgo -from ..model_config import ModelConfig -from ..models.modeling_deepseekv3 import DeepseekV3ForCausalLM -from .modeling_utils import register_auto_model - class Mistral3Gate(nn.Module): def __init__( @@ -53,20 +52,19 @@ def forward(self, *args, **kwargs): def load_weights(self, weights: Dict): assert self.model_config is not None, "self.model_config is required" params_map = self.weight_mapper.mistral_llm_mapping.copy() - if self.model_config is not None: - quantization_weights_map: Dict[str, str] = {} - if self.model_config.quant_config.quant_algo == QuantAlgo.NVFP4: - quantization_weights_map = { - "weight_packed": "weight", - "input_global_scale": "input_scale", - "weight_global_scale": "weight_scale_2", - } - elif self.model_config.quant_config.quant_algo == QuantAlgo.FP8_BLOCK_SCALES: - quantization_weights_map = { - "weight_scale": "weight_scale_inv", - } - if quantization_weights_map: - params_map.update(quantization_weights_map) + quantization_weights_map: Dict[str, str] = {} + if self.model_config.quant_config.quant_algo == QuantAlgo.NVFP4: + quantization_weights_map = { + "weight_packed": "weight", + "input_global_scale": "input_scale", + "weight_global_scale": "weight_scale_2", + } + elif self.model_config.quant_config.quant_algo == QuantAlgo.FP8_BLOCK_SCALES: + quantization_weights_map = { + "weight_scale": "weight_scale_inv", + } + if quantization_weights_map: + params_map.update(quantization_weights_map) weights = self.weight_mapper.rename_by_params_map(weights=weights, params_map=params_map) super().load_weights(weights) From 4443c3abfffd332ed9bcedfd16a3d2c10b79bbb4 Mon Sep 17 00:00:00 2001 From: bhsueh <11360707+byshiue@users.noreply.github.com> Date: Tue, 9 Dec 2025 10:25:28 +0000 Subject: [PATCH 20/21] [FIX] minor changes based on comment Signed-off-by: bhsueh <11360707+byshiue@users.noreply.github.com> --- tensorrt_llm/_torch/models/modeling_mistral_large3.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/tensorrt_llm/_torch/models/modeling_mistral_large3.py b/tensorrt_llm/_torch/models/modeling_mistral_large3.py index 6a0beb1b29c2..c88cebdf054b 100644 --- a/tensorrt_llm/_torch/models/modeling_mistral_large3.py +++ b/tensorrt_llm/_torch/models/modeling_mistral_large3.py @@ -3,11 +3,11 @@ import torch from torch import nn -from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import MistralLarge3WeightMapper -from tensorrt_llm._torch.modules.fused_moe import RenormalizeNaiveMoeRoutingMethod from tensorrt_llm._torch.model_config import ModelConfig +from tensorrt_llm._torch.models.checkpoints.mistral.weight_mapper import MistralLarge3WeightMapper from tensorrt_llm._torch.models.modeling_deepseekv3 import DeepseekV3ForCausalLM from tensorrt_llm._torch.models.modeling_utils import register_auto_model +from tensorrt_llm._torch.modules.fused_moe import RenormalizeNaiveMoeRoutingMethod from tensorrt_llm.quantization.mode import QuantAlgo From b523e10fd209749c7809249ef997e95b6c529e87 Mon Sep 17 00:00:00 2001 From: bhsueh Date: Thu, 11 Dec 2025 06:43:27 -0800 Subject: [PATCH 21/21] [FIX] adust CI settings Signed-off-by: bhsueh --- tests/integration/defs/accuracy/references/gsm8k.yaml | 5 +---- tests/integration/defs/accuracy/references/mmlu.yaml | 5 +---- tests/integration/test_lists/test-db/l0_dgx_b200.yml | 2 +- 3 files changed, 3 insertions(+), 9 deletions(-) diff --git a/tests/integration/defs/accuracy/references/gsm8k.yaml b/tests/integration/defs/accuracy/references/gsm8k.yaml index d0cb5b88f7b7..33f7dddc6bd1 100644 --- a/tests/integration/defs/accuracy/references/gsm8k.yaml +++ b/tests/integration/defs/accuracy/references/gsm8k.yaml @@ -282,7 +282,4 @@ bigcode/starcoder2-7b: bigcode/starcoder2-15b: - accuracy: 54.5 mistral/Mistral-Large-3-675B: - - quant_algo: NVFP4 - accuracy: 90.83 - - quant_algo: FP8_BLOCK_SCALES - accuracy: 90.83 + - accuracy: 90.83 diff --git a/tests/integration/defs/accuracy/references/mmlu.yaml b/tests/integration/defs/accuracy/references/mmlu.yaml index 41633c3e0695..f728919abe06 100644 --- a/tests/integration/defs/accuracy/references/mmlu.yaml +++ b/tests/integration/defs/accuracy/references/mmlu.yaml @@ -341,7 +341,4 @@ mistralai/Mistral-Nemo-12b-Base: - quant_algo: FP8 accuracy: 69.66 mistral/Mistral-Large-3-675B: - - quant_algo: NVFP4 - accuracy: 87.54 - - quant_algo: FP8_BLOCK_SCALES - accuracy: 90.83 + - accuracy: 87.54 diff --git a/tests/integration/test_lists/test-db/l0_dgx_b200.yml b/tests/integration/test_lists/test-db/l0_dgx_b200.yml index 00dcc1fa4ace..73c218c541b6 100644 --- a/tests/integration/test_lists/test-db/l0_dgx_b200.yml +++ b/tests/integration/test_lists/test-db/l0_dgx_b200.yml @@ -114,7 +114,6 @@ l0_dgx_b200: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus[baseline] TIMEOUT (180) - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus[baseline_mtp1] TIMEOUT (180) - accuracy/test_disaggregated_serving.py::TestDeepSeekV32Exp::test_auto_dtype[False] TIMEOUT (360) - - accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] TIMEOUT (90) - condition: ranges: system_gpu_count: @@ -136,6 +135,7 @@ l0_dgx_b200: - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus[baseline_fp8kv] TIMEOUT (180) - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus[latency] TIMEOUT (180) - accuracy/test_llm_api_pytorch.py::TestDeepSeekV32::test_nvfp4_multi_gpus_chunked_prefill[baseline_fp8kv] TIMEOUT (180) + - accuracy/test_llm_api_pytorch.py::TestMistralLarge3_675B::test_fp8[latency_moe_deepgemm] TIMEOUT (90) - condition: ranges: system_gpu_count: