ReferenceFull Library ReferenceNemo AutomodelNemo AutomodelComponentsModelsMistral3nemo_automodel.components.models.mistral3.state_dict_adapter

nemo_automodel.components.models.mistral3.state_dict_adapter

View as Markdown

State-dict adapter for per-tensor FP8 Mistral3 checkpoints.

Plugs into the standard nemo_automodel checkpoint flow (nemo_automodel/components/checkpoint/checkpointing.py ~lines 510, 556) and handles FP8 dequantization during load/save for text-only and VLM models:

  • Checkpoint Linear weights are stored as per-tensor FP8 with a scalar weight_scale_inv sibling (and an unused activation_scale sibling). The adapter pairs each weight with its scale on load, dequantizes through fp32 (w_bf16 = (w_fp8.float() * scale.float()).bfloat16()), and drops the scale keys. Vision tower + multi_modal_projector + lm_head are BF16 on disk and pass through unchanged.

The live HF VLM module keeps the body under model.* while the checkpoint stores text weights under language_model.model.* and top-level VLM components as vision_tower.* / multi_modal_projector.*. The LM head is also nested on disk as language_model.lm_head.weight while the runtime module exposes it as lm_head.weight.

Structurally modelled after nemo_automodel/components/models/deepseek_v3/state_dict_adapter.py.

Module Contents

Classes

NameDescription
Mistral3FP8StateDictAdapterPer-tensor FP8 dequant adapter for the Mistral3 model family.

Functions

NameDescription
_config_attr-
_dequantize_from_fp8Dequantize a single FP8 weight using its per-tensor scalar scale.
_dequantize_from_fp8_intoDequantize a per-tensor FP8 weight directly into its final model tensor.
_finish_fp8_loadsInstall one load part’s FP8 tensors into final model storage.
_identity-
_is_fp8_weight_keyReturn True iff model_key names an FP8 Linear weight.
_is_mistral3p5_128b_config-
_uses_identity_vlm_layoutReturn True for FP8 VLM checkpoints whose disk keys already match HF.
_vlm_full_hf_to_nativeMap checkpoint VLM names back to runtime parameter names.
_vlm_full_native_to_hfMap runtime VLM parameter names to checkpoint names.

Data

_CAUSAL_LM_DECODER_LAYER_KEY

_CHECKPOINT_LAYERS_PER_PART

_HF_LM_HEAD_KEY

_MISTRAL3P5_128B_NUM_HIDDEN_LAYERS

_MODEL_LM_HEAD_KEY

_NON_QUANTIZED_SUFFIXES

_VLM_DECODER_LAYER_KEY

logger

API

class nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter(
native_to_hf: typing.Callable[[str], str] = _identity,
hf_to_native: typing.Callable[[str], str] = _identity,
layout_name: str = 'vlm_full',
not_fp8_prefixes: tuple[str, ...] = (),
num_hidden_layers: int | None = None
)

Bases: StateDictAdapter

Per-tensor FP8 dequant adapter for the Mistral3 model family.

Text-only causal-LM checkpoint keys already match the model state dict. VLM checkpoints additionally select the appropriate body-key layout and exclude their BF16 vision and projector modules from FP8 conversion.

_not_fp8_prefixes
= tuple(not_fp8_prefixes)
nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter._iter_checkpoint_load_parts(
model_state_dict: dict[str, torch.Tensor],
decoder_layer_key: re.Pattern[str]

Build load parts for a complete causal-LM or VLM decoder.

Parameters:

model_state_dict
dict[str, torch.Tensor]

Native model names mapped to final tensors of arbitrary model-defined rank and shape. Quantized linear weights must use BF16 or FP32 final storage; non-quantized tensors remain direct DCP destinations and are mutated in place during the load.

decoder_layer_key
re.Pattern[str]

Pattern that identifies decoder-layer names and captures the zero-based layer index in group 2. It must exclude vision-tower layers from the bounded FP8 groups.

Returns: Iterator[CheckpointLoadPart]

Dependency-complete load parts. Each quantized destination has the same shape, strides, device, and DTensor

nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter.convert_single_tensor_to_hf(
fqn: str,
tensor: typing.Any,
kwargs = {}
) -> list[tuple[str, typing.Any]]

Per-tensor model → HF used by Checkpointer.save_model.

nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter.for_causal_lm(
config: typing.Any | None = None
) -> 'Mistral3FP8StateDictAdapter'
classmethod

Text-only path for per-tensor FP8 Ministral3ForCausalLM checkpoints.

Devstral-2 stores text-model keys in the same layout exposed by the local Ministral3ForCausalLM implementation. Linear weights are FP8, while embeddings, norms, and the untied LM head remain BF16 and are excluded by _NON_QUANTIZED_SUFFIXES.

nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter.for_vlm_full(
config: typing.Any | None = None
) -> 'Mistral3FP8StateDictAdapter'
classmethod

Full-VLM path for Mistral3ForConditionalGeneration checkpoints.

Mistral3 FP8 VLM checkpoints have two observed body-key layouts. The Mistral-Medium-3.5 128B checkpoint already stores keys in the same layout as HF’s VLM state_dict() (model.language_model.* / model.vision_tower.* / model.multi_modal_projector.*). Newer Ministral/Devstral-style checkpoints store text weights under language_model.model.* and non-text component names at top level.

The LM head has one extra quirk in the nested layout: the model exposes it at the top level (lm_head.weight) while the checkpoint nests it (language_model.lm_head.weight). Tied checkpoints (Ministral-3) never serialize the head, so the head translation is a harmless no-op there; untied checkpoints (Devstral-24B) rely on it to find the head during the DCP load.

Only the language_model layer weights are FP8; vision / mm_projector / lm_head are BF16 on disk and must be passed through without a scale_inv placeholder — otherwise DCP would fail trying to fetch a non-existent _scale_inv key.

nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter.from_hf(
hf_state_dict: dict[str, typing.Any],
device_mesh: 'DeviceMesh' | None = None,
kwargs = {}
) -> dict[str, typing.Any]

Convert an HF-format (possibly FP8) state dict to model-native format.

nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter.iter_checkpoint_load_parts(
model_state_dict: dict[str, torch.Tensor],
device_mesh: 'DeviceMesh' | None = None

Load Mistral3 FP8 weights in bounded decoder-layer groups.

Each quantized model tensor has BF16 or FP32 model shape and storage. Its load part creates an FP8 destination with the same shape, device, strides, and distributed placements, plus the scalar BF16 _scale_inv destination stored by the checkpoint. After DCP fills both tensors, the part copies and scales the FP8 value directly into the final model tensor. Non-quantized tensors, including VLM vision and projector weights, use their final model storage as the DCP destination.

This path requires the complete decoder on every rank. Pipeline-parallel ranks own different layer subsets and therefore retain the existing rank-local DCP path until part scheduling can be coordinated across stages. Tied VLM checkpoints also retain that path because they omit the LM-head tensor expected by grouped loading.

Parameters:

model_state_dict
dict[str, torch.Tensor]

Native model names mapped to final parameter and persistent-buffer tensors. Each tensor has arbitrary model-defined rank and shape. Decoder tensors may be local DTensor shards; their global shapes, local shards, and placements are preserved by torch.empty_like. Non-quantized destinations alias and are populated through final model storage.

device_mesh
'DeviceMesh' | NoneDefaults to None

Optional distributed mesh. The tensors already carry their final placements, so this value is not otherwise needed.

Returns: Iterator[CheckpointLoadPart] | None

One direct-load part for tensors outside decoder layers plus bounded temporary-load parts for decoder

nemo_automodel.components.models.mistral3.state_dict_adapter.Mistral3FP8StateDictAdapter.to_hf(
state_dict: dict[str, typing.Any],
exclude_key_regex: str | None = None,
quantization: bool = False,
kwargs = {}
) -> dict[str, typing.Any]

Convert a model-native state dict to HF (on-disk) layout.

When quantization=True the weight placeholder is also cast to torch.float8_e4m3fn so the DCP storage reader fetches FP8 bytes verbatim from safetensors (a bf16 target would silently cast-on-read and lose the scale multiply — see deepseek_v3/state_dict_adapter.py:220). A scalar _scale_inv placeholder is also emitted so DCP pulls it alongside the weight.

nemo_automodel.components.models.mistral3.state_dict_adapter._config_attr(
config: typing.Any | None,
attr: str
) -> typing.Any
nemo_automodel.components.models.mistral3.state_dict_adapter._dequantize_from_fp8(
weight_fp8: torch.Tensor,
scale_inv: torch.Tensor,
target_dtype: torch.dtype = torch.bfloat16
) -> torch.Tensor

Dequantize a single FP8 weight using its per-tensor scalar scale.

Supported Mistral3 checkpoints use per-tensor quantization (weight_block_size=None), so scale_inv is a 0-d scalar and dequantization collapses to a simple multiply. The per-block formula (transformers.integrations.finegrained_fp8.Fp8Dequantize.convert, finegrained_fp8.py:867-906) is not needed here.

nemo_automodel.components.models.mistral3.state_dict_adapter._dequantize_from_fp8_into(
target: torch.Tensor,
weight_fp8: torch.Tensor,
scale_inv: torch.Tensor
) -> None

Dequantize a per-tensor FP8 weight directly into its final model tensor.

target keeps its native model shape, BF16 or FP32 dtype, device, strides, distributed placements, and storage. weight_fp8 has the same shape and distributed placements but uses the checkpoint’s FP8 dtype and temporary storage. scale_inv is the checkpoint’s scalar BF16 inverse scale.

nemo_automodel.components.models.mistral3.state_dict_adapter._finish_fp8_loads(
conversions: tuple[tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...]
) -> None

Install one load part’s FP8 tensors into final model storage.

nemo_automodel.components.models.mistral3.state_dict_adapter._identity(
k: str
) -> str
nemo_automodel.components.models.mistral3.state_dict_adapter._is_fp8_weight_key(
model_key: str,
not_fp8_prefixes: tuple[str, ...] = ()
) -> bool

Return True iff model_key names an FP8 Linear weight.

nemo_automodel.components.models.mistral3.state_dict_adapter._is_mistral3p5_128b_config(
config: typing.Any | None
) -> bool
nemo_automodel.components.models.mistral3.state_dict_adapter._uses_identity_vlm_layout(
config: typing.Any | None
) -> bool

Return True for FP8 VLM checkpoints whose disk keys already match HF.

nemo_automodel.components.models.mistral3.state_dict_adapter._vlm_full_hf_to_native(
hf_key: str
) -> str

Map checkpoint VLM names back to runtime parameter names.

nemo_automodel.components.models.mistral3.state_dict_adapter._vlm_full_native_to_hf(
model_key: str
) -> str

Map runtime VLM parameter names to checkpoint names.

nemo_automodel.components.models.mistral3.state_dict_adapter._CAUSAL_LM_DECODER_LAYER_KEY = re.compile('^(model\\.layers\\.(\\d+))\\.')
nemo_automodel.components.models.mistral3.state_dict_adapter._CHECKPOINT_LAYERS_PER_PART = 8
nemo_automodel.components.models.mistral3.state_dict_adapter._HF_LM_HEAD_KEY = 'language_model.lm_head.weight'
nemo_automodel.components.models.mistral3.state_dict_adapter._MISTRAL3P5_128B_NUM_HIDDEN_LAYERS = 88
nemo_automodel.components.models.mistral3.state_dict_adapter._MODEL_LM_HEAD_KEY = 'lm_head.weight'
nemo_automodel.components.models.mistral3.state_dict_adapter._NON_QUANTIZED_SUFFIXES = ('embed_tokens.weight', 'lm_head.weight', 'input_layernorm.weight', 'post_attent...
nemo_automodel.components.models.mistral3.state_dict_adapter._VLM_DECODER_LAYER_KEY = re.compile('^(model\\.language_model(?:\\.model)?\\.layers\\.(\\d+))\\.')
nemo_automodel.components.models.mistral3.state_dict_adapter.logger = logging.getLogger(__name__)