nemo_automodel.components.models.mistral3.state_dict_adapter
nemo_automodel.components.models.mistral3.state_dict_adapter
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_invsibling (and an unusedactivation_scalesibling). 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
Functions
Data
_MISTRAL3P5_128B_NUM_HIDDEN_LAYERS
API
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.
Build load parts for a complete causal-LM or VLM decoder.
Parameters:
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.
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
Per-tensor model → HF used by Checkpointer.save_model.
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.
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.
Convert an HF-format (possibly FP8) state dict to model-native format.
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:
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.
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
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.
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.
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.
Install one load part’s FP8 tensors into final model storage.
Return True iff model_key names an FP8 Linear weight.
Return True for FP8 VLM checkpoints whose disk keys already match HF.
Map checkpoint VLM names back to runtime parameter names.
Map runtime VLM parameter names to checkpoint names.