clef / code /models /tt_transformers /tt /load_checkpoints.py
tt-hous's picture
Add files using upload-large-folder tool
b025706 verified
Raw History Blame Contribute Delete
43.3 kB
# SPDX-FileCopyrightText: © 2025 Tenstorrent USA, Inc.
# SPDX-License-Identifier: Apache-2.0
import json
import os
import re
from pathlib import Path
import torch
from loguru import logger
from safetensors.torch import load_file as safetensors_load_file
from safetensors.torch import safe_open as safetensors_safe_open
from tqdm import tqdm
# TODO Update function for large models: For 1 layer tests we only want to load 1 checkpoint file, instead of all.
def load_hf_state_dict(ckpt_dir):
# First check if index file exists
index_path = os.path.join(ckpt_dir, "model.safetensors.index.json")
if os.path.exists(index_path):
# Multi-file case: Read the index file and load all referenced safetensor files
with open(index_path, "r") as f:
index_data = json.load(f)
# Retrieve the weight file names from the index JSON
weight_map = index_data["weight_map"]
safetensor_files = set(weight_map.values())
# Read each safetensors file mentioned in the index
loaded_weights = {}
for file in safetensor_files:
safetensor_path = os.path.join(ckpt_dir, file)
weights = safetensors_load_file(safetensor_path)
loaded_weights.update(weights) # Merge weights into a single dictionary
else:
# Single-file case: Load the single model.safetensors file
safetensor_path = os.path.join(ckpt_dir, "model.safetensors")
if not os.path.exists(safetensor_path):
raise FileNotFoundError(f"Neither model.safetensors.index.json nor model.safetensors found in {ckpt_dir}")
loaded_weights = safetensors_load_file(safetensor_path)
return loaded_weights
def load_hf_state_dict_filtered(ckpt_dir, key_prefixes, local_files_only=None):
"""
Load only the subset of HF checkpoint weights that match the given key prefixes.
Uses safetensors safe_open to avoid loading unrelated tensors into memory.
Supports local checkpoint directories or HF repo IDs.
"""
prefixes = tuple(key_prefixes)
if not prefixes:
return {}
return _load_hf_state_dict_matching(ckpt_dir, lambda key: key.startswith(prefixes), local_files_only)
_HF_LAYER_KEY = re.compile(r"^model\.layers\.(\d+)\.")
def load_hf_state_dict_for_layers(ckpt_dir, n_layers, local_files_only=None):
"""
Load an HF text checkpoint keeping only decoder layers [0, n_layers) plus every non-layer weight
(embeddings, final norm, lm_head). Reads just the shards those keys live in through safetensors
safe_open, so a one-layer unit test does not materialise the whole checkpoint.
"""
def keep(key):
m = _HF_LAYER_KEY.match(key)
return m is None or int(m.group(1)) < n_layers
return _load_hf_state_dict_matching(ckpt_dir, keep, local_files_only)
def _load_hf_state_dict_matching(ckpt_dir, key_filter, local_files_only=None):
if local_files_only is None:
local_files_only = os.getenv("CI") == "true"
ckpt_dir = str(ckpt_dir)
is_local_dir = os.path.isdir(ckpt_dir)
hf_hub_download = None
EntryNotFoundError = None
LocalEntryNotFoundError = None
if not is_local_dir:
try:
from huggingface_hub import hf_hub_download
from huggingface_hub.utils import EntryNotFoundError, LocalEntryNotFoundError
except ImportError as exc:
raise ImportError("huggingface_hub is required to resolve HF repo IDs for safetensors loading.") from exc
def resolve_file(filename, allow_missing=False):
if is_local_dir:
path = os.path.join(ckpt_dir, filename)
if os.path.exists(path):
return path
if allow_missing:
return None
raise FileNotFoundError(f"Missing safetensors file {path}")
try:
return hf_hub_download(ckpt_dir, filename=filename, local_files_only=local_files_only)
except (EntryNotFoundError, LocalEntryNotFoundError) as exc:
if allow_missing:
return None
raise FileNotFoundError(
f"Missing safetensors file {filename} for repo {ckpt_dir} (local_files_only={local_files_only})"
) from exc
loaded_weights = {}
index_path = resolve_file("model.safetensors.index.json", allow_missing=True)
if index_path is not None:
with open(index_path, "r") as f:
index_data = json.load(f)
weight_map = index_data["weight_map"]
file_to_keys = {}
for key, file in weight_map.items():
if key_filter(key):
file_to_keys.setdefault(file, []).append(key)
for file, keys in file_to_keys.items():
safetensor_path = resolve_file(file)
with safetensors_safe_open(safetensor_path, framework="pt", device="cpu") as f:
for key in keys:
loaded_weights[key] = f.get_tensor(key)
else:
safetensor_path = resolve_file("model.safetensors")
with safetensors_safe_open(safetensor_path, framework="pt", device="cpu") as f:
for key in f.keys():
if key_filter(key):
loaded_weights[key] = f.get_tensor(key)
return loaded_weights
def standardize_hf_keys(state_dict):
key_meta = "lm_head.weight"
key_hf = "model.embed_tokens.weight"
if not key_meta in state_dict and key_hf in state_dict:
state_dict[key_meta] = state_dict[key_hf]
del state_dict[key_hf]
return state_dict
def standardize_hf_keys_multimodal(state_dict):
all_keys = tuple(state_dict.keys())
new_state_dict = {}
for k in all_keys:
if "model.visual." in k:
new_state_dict[k.replace("model.visual.", "visual.")] = state_dict[k]
elif "model.vision_tower.vision_model." in k:
new_state_dict[k.replace("model.vision_tower.vision_model.", "visual.")] = state_dict[k]
elif "model.vision_tower." in k:
new_state_dict[k.replace("model.", "")] = state_dict[k]
elif "model.multi_modal_projector." in k:
new_state_dict[k.replace("model.", "")] = state_dict[k]
elif "model.vision_model." in k:
new_state_dict[k.replace("model.vision_model.", "vision_model.")] = state_dict[k]
elif "model.language_model." in k:
new_state_dict[k.replace("model.language_model.", "model.")] = state_dict[k]
else:
new_state_dict[k] = state_dict[k]
# Standardize keys used in vision parts of Qwen2.5-VL
state_dict = standardize_hf_keys(new_state_dict)
replace_whole_name = lambda pattern, repl: lambda s: re.sub(rf"(^|\.)({pattern})($|\.)", rf"\1{repl}\3", s)
output = {}
for k, v in state_dict.items():
k = replace_whole_name("qkv", "qkv_proj")(k)
k = replace_whole_name("proj", "o_proj")(k)
k = replace_whole_name("attn", "self_attn")(k)
output[k] = v
return output
def expand_fused_moe_experts(state_dict):
"""Split transformers 5.x fused Mixtral MoE expert params back to per-expert keys.
transformers 5.x replaced the per-expert ``...block_sparse_moe.experts.{i}.w{1,2,3}.weight``
tensors with 3D batched params under ``...mlp.experts.`` :
- ``gate_up_proj`` : ``[num_experts, 2*intermediate, hidden]`` (rows ``:I`` = w1/gate, ``I:`` = w3/up)
- ``down_proj`` : ``[num_experts, hidden, intermediate]`` (= w2)
and renamed the router ``block_sparse_moe.gate`` -> ``mlp.gate``. The tt Mixtral model loads the
per-expert / ``block_sparse_moe`` keys, so split them back here. Version- and model-tolerant:
a no-op unless the fused ``mlp.experts.gate_up_proj`` keys are present (i.e. Mixtral on >=5.x).
"""
fused_keys = [k for k in state_dict if k.endswith("mlp.experts.gate_up_proj")]
if not fused_keys:
return state_dict
out = dict(state_dict)
for gup_key in fused_keys:
prefix = gup_key[: -len("mlp.experts.gate_up_proj")] # e.g. "model.layers.0."
gate_up = out.pop(gup_key) # [E, 2I, H]
down = out.pop(prefix + "mlp.experts.down_proj") # [E, H, I]
num_experts = gate_up.shape[0]
inter = gate_up.shape[1] // 2
for i in range(num_experts):
base = f"{prefix}block_sparse_moe.experts.{i}."
out[base + "w1.weight"] = gate_up[i, :inter, :].contiguous() # gate -> w1, [I, H]
out[base + "w3.weight"] = gate_up[i, inter:, :].contiguous() # up -> w3, [I, H]
out[base + "w2.weight"] = down[i].contiguous() # down -> w2, [H, I]
# router gate: 5.x `...mlp.gate.weight` -> tt expects `...block_sparse_moe.gate.weight`
gate_key = prefix + "mlp.gate.weight"
if gate_key in out:
out[prefix + "block_sparse_moe.gate.weight"] = out.pop(gate_key)
return out
def convert_hf_to_meta(state_dict, head_dim, n_heads=None, n_kv_heads=None):
state_dict = expand_fused_moe_experts(state_dict)
state_dict = split_hf_keys(state_dict, n_heads, n_kv_heads)
state_dict = convert_hf_qkv_to_meta_format(state_dict, head_dim)
state_dict = map_hf_to_meta_keys(state_dict)
return state_dict
def convert_hf_to_meta_no_qkv_permute(state_dict, head_dim, n_heads=None, n_kv_heads=None):
"""Convert HF to Meta format but skip QKV weight permutation.
This keeps weights in HF format for use with HF-style RoPE.
Only key mapping is performed (q_proj -> wq, etc.).
"""
state_dict = split_hf_keys(state_dict, n_heads, n_kv_heads)
# SKIP convert_hf_qkv_to_meta_format - keep weights in HF format
state_dict = map_hf_to_meta_keys(state_dict)
return state_dict
def convert_vision_hf_to_meta(state_dict, head_dim):
state_dict = split_hf_keys(state_dict)
state_dict = map_vision_hf_to_meta_keys(state_dict, head_dim)
return state_dict
def convert_hf_qkv_to_meta_format_mllama(state_dict, head_dim):
vision_state_dict, text_state_dict, other_state_dict = map_vision_hf_to_meta_keys_split_to_submodels(state_dict)
cross_attn_text_state_dict = {k: v for k, v in text_state_dict.items() if "cross_attn" in k}
text_state_dict = {k: v for k, v in text_state_dict.items() if k not in cross_attn_text_state_dict}
text_state_dict = convert_hf_qkv_to_meta_format(text_state_dict, head_dim)
return {**vision_state_dict, **cross_attn_text_state_dict, **text_state_dict, **other_state_dict}
def convert_hf_to_meta_mllama(state_dict, head_dim, config):
state_dict = split_hf_keys(state_dict)
state_dict = convert_hf_qkv_to_meta_format_mllama(state_dict, head_dim)
state_dict = map_hf_to_meta_keys_mllama(state_dict, config)
state_dict = convert_pos_embeddings(state_dict)
state_dict = flatten_conv_linear(state_dict)
return state_dict
def convert_hf_to_meta_mllama_no_qkv_permute(state_dict, head_dim, config):
"""Convert HF to Meta format for multimodal Llama but skip QKV weight permutation.
This keeps weights in HF format for use with HF-style RoPE.
Only key mapping is performed (q_proj -> wq, etc.).
"""
state_dict = split_hf_keys(state_dict)
state_dict = map_hf_to_meta_keys_mllama(state_dict, config)
state_dict = convert_pos_embeddings(state_dict)
state_dict = flatten_conv_linear(state_dict)
return state_dict
def map_hf_to_meta_keys_vision_only(state_dict):
"""
Map Hugging Face checkpoint keys to Meta checkpoint keys.
You can use this to support other models by adding more mappings.
See replace_keys for more details on the format of replacements.
"""
replacements = [
("self_attn", "attn"),
("q_proj", "wq"),
("k_proj", "wk"),
("v_proj", "wv"),
("o_proj", "wo"),
("out_proj", "wo"),
("q_norm", "q_norm"),
("k_norm", "k_norm"),
("fc1", "c_fc"),
("fc2", "c_proj"),
("gate_proj", "w1"),
("down_proj", "w2"),
("up_proj", "w3"),
("layer_norm1", "ln_1"),
("layer_norm2", "ln_2"),
("post_layernorm", "ln_post"),
("embeddings.patch_embedding._linear", "embeddings.patch_embedding"),
("embeddings.patch_embedding", "embeddings.patch_embedding._linear"),
("embeddings.position_embedding.weight", "embeddings.position_embedding.positional_embedding"),
("patch_conv", "patch_conv._linear"),
]
return replace_keys(state_dict, replacements)
def map_vision_hf_to_meta_keys_split_to_submodels(state_dict):
vision_state_dict = dict()
text_state_dict = dict()
other_state_dict = dict()
for k, v in state_dict.items():
if k.startswith("visual") or k.startswith("vision_model") or k.startswith("vision_tower"):
selected_dict = vision_state_dict
elif k.startswith("model") or k.startswith("lm_head") or k.startswith("language_model"):
selected_dict = text_state_dict
else:
selected_dict = other_state_dict
selected_dict[k] = v
return vision_state_dict, text_state_dict, other_state_dict
def map_vision_hf_to_meta_keys(state_dict, head_dim):
vision_state_dict, text_state_dict, other_state_dict = map_vision_hf_to_meta_keys_split_to_submodels(state_dict)
text_state_dict = convert_hf_qkv_to_meta_format(text_state_dict, head_dim)
text_state_dict = map_hf_to_meta_keys(text_state_dict)
vision_state_dict = map_hf_to_meta_keys_vision_only(vision_state_dict)
return {**vision_state_dict, **text_state_dict, **other_state_dict}
def map_vision_hf_to_meta_keys_no_qkv_permute(state_dict, head_dim):
"""Map vision HF to Meta keys but skip QKV format conversion for text portion.
This keeps text weights in HF format for use with HF-style RoPE.
"""
vision_state_dict, text_state_dict, other_state_dict = map_vision_hf_to_meta_keys_split_to_submodels(state_dict)
# SKIP convert_hf_qkv_to_meta_format - keep text weights in HF format
text_state_dict = map_hf_to_meta_keys(text_state_dict)
vision_state_dict = map_hf_to_meta_keys_vision_only(vision_state_dict)
return {**vision_state_dict, **text_state_dict, **other_state_dict}
def convert_vision_hf_to_meta_no_qkv_permute(state_dict, head_dim):
"""Convert vision HF to Meta format but skip QKV weight permutation.
This keeps weights in HF format for use with HF-style RoPE.
Only key mapping is performed (q_proj -> wq, etc.).
"""
state_dict = split_hf_keys(state_dict)
state_dict = map_vision_hf_to_meta_keys_no_qkv_permute(state_dict, head_dim)
return state_dict
def load_meta_state_dict(ckpt_dir, n_layers=None, start_layer_idx=0):
checkpoints = sorted(Path(ckpt_dir).glob("*.pth"))
assert len(checkpoints) > 0, f"no checkpoint files found in {ckpt_dir}"
is_chunked = any(ckpt.stem.startswith("layers_") for ckpt in checkpoints)
if is_chunked:
checkpoints = [ckpt_name for ckpt_name in checkpoints if ckpt_name.stem.startswith("layers_")]
checkpoint = load_chunked_checkpoints(checkpoints, n_layers, start_layer_idx)
else:
checkpoint = load_sharded_checkpoints(checkpoints, n_layers)
return checkpoint
def load_chunked_checkpoints(checkpoints, n_layers, start_layer_idx):
checkpoint = {}
(f"Loading {len(checkpoints)} chunked checkpoint files")
for ckpt in tqdm(checkpoints):
if n_layers:
# Layer range is in the file name, like layers_start-end.pth
layer_range = ckpt.stem.split("_")[1]
start_layer, end_layer = map(int, layer_range.split("-"))
if start_layer > n_layers + start_layer_idx:
continue
if end_layer < start_layer_idx:
continue
loaded_ckpt = torch.load(ckpt, map_location="cpu")
checkpoint.update(loaded_ckpt)
return checkpoint
def is_param_replicated_across_shards(key: str) -> bool:
"""
Return `True` if the parameter is replicated (i.e., not sharded)
across checkpoint files and should not be concatenated.
"""
if key.startswith("vision_model."):
return any(keyword in key for keyword in ("ln", "gate", "embed", "c_proj.bias"))
else:
# for Meta checkpoint keys, key either starts with "text_model." or contains no such prefix; both cases are handled here
return any(keyword in key for keyword in ("norm", "gate"))
def load_sharded_checkpoints(checkpoints, n_layers):
checkpoint = {}
logger.info(f"Loading {len(checkpoints)} sharded checkpoint files")
for ckpt in tqdm(checkpoints):
loaded_ckpt = torch.load(ckpt, map_location="cpu")
for key, value in loaded_ckpt.items():
if "layers." in key:
layer_num = int(key.split("layers.")[1].split(".")[0])
if n_layers and layer_num >= n_layers:
continue
if key in checkpoint:
checkpoint[key] += [value]
else:
checkpoint[key] = [value]
del loaded_ckpt
# concat checkpoint values
for key, value in checkpoint.items():
if len(value) == 1 or is_param_replicated_across_shards(key):
checkpoint[key] = value[0]
else:
if key.endswith("tok_embeddings.weight") or key.endswith("output.weight"):
assert value[0].shape[1] == 8192 # FIXME: do we need this hardcoded shape?
# Concatenate along dimension 0 for llama3 token embeddings weight and lm head
checkpoint[key] = torch.cat(value, dim=0)
else:
# cat_dim is index of the smallest dimension in value[0].shape
cat_dim = torch.argmin(torch.tensor(value[0].shape))
checkpoint[key] = torch.cat(value, dim=cat_dim)
return checkpoint
def split_hf_keys(loaded_weights, n_heads=None, n_kv_heads=None):
converted_weights = {}
for key, tensor in loaded_weights.items():
if "qkv_proj" in key:
# split Q, K and V
q_key = key.replace("qkv_proj", "q_proj")
k_key = key.replace("qkv_proj", "k_proj")
v_key = key.replace("qkv_proj", "v_proj")
# Handle GQA (Grouped Query Attention) case
if n_heads is not None and n_kv_heads is not None and n_heads != n_kv_heads:
# For GQA: Q has n_heads, K and V have n_kv_heads
head_dim = tensor.shape[0] // (n_heads + 2 * n_kv_heads)
q_size = n_heads * head_dim
kv_size = n_kv_heads * head_dim
q_tensor = tensor[:q_size]
k_tensor = tensor[q_size : q_size + kv_size]
v_tensor = tensor[q_size + kv_size : q_size + 2 * kv_size]
else:
# Default case: equal split for Q, K, V
q_tensor, k_tensor, v_tensor = torch.split(tensor, tensor.shape[0] // 3, dim=0)
converted_weights[q_key] = q_tensor
converted_weights[k_key] = k_tensor
converted_weights[v_key] = v_tensor
elif "gate_up_proj" in key:
# Split Gate and Up
gate_key = key.replace("gate_up_proj", "gate_proj")
up_key = key.replace("gate_up_proj", "up_proj")
gate_tensor, up_tensor = torch.split(tensor, tensor.shape[0] // 2, dim=0)
converted_weights[gate_key] = gate_tensor
converted_weights[up_key] = up_tensor
else:
# Keep all other weights unchanged
converted_weights[key] = tensor
return converted_weights
def convert_hf_qkv_to_meta_format(loaded_weights, head_dim):
"""Convert HuggingFace QKV weights to Meta format for RoPE compatibility."""
converted_weights = {}
for key, tensor in loaded_weights.items():
if "vision_tower" in key:
# Skip conversion for vision tower weights (Mistral vision support)
converted_weights[key] = tensor
elif "q_proj.weight" in key or "k_proj.weight" in key:
# For weights: n_heads = tensor.shape[0] // head_dim
n_heads = tensor.shape[0] // head_dim
converted_weights[key] = reverse_permute(tensor, n_heads, tensor.shape[0], tensor.shape[1])
elif "q_proj.bias" in key or "k_proj.bias" in key:
# For biases: n_heads = tensor.shape[0] // head_dim
n_heads = tensor.shape[0] // head_dim
converted_weights[key] = reverse_permute(tensor, n_heads, tensor.shape[0], 1).squeeze(-1)
elif "q_norm.weight" in key or "k_norm.weight" in key:
converted_weights[key] = reverse_permute_1d(tensor)
else:
# Keep all other weights unchanged
converted_weights[key] = tensor
return converted_weights
def fuse_mlp_meta(state_dict):
key_map = {"w_gate": "w1.weight", "w_up": "w3.weight", "w_gate_up_proj": "w1_w3.weight"}
wgate_list = sorted(list(filter(lambda x: key_map["w_gate"] in x, state_dict.keys())))
wproj_list = sorted(list(filter(lambda x: key_map["w_up"] in x, state_dict.keys())))
for wgate_key, wproj_key in zip(wgate_list, wproj_list):
wgate = state_dict[wgate_key]
wproj = state_dict[wproj_key]
prefix_gate = wgate_key[: -len(key_map["w_gate"])]
fused_gate_up_proj = torch.vstack((wgate, wproj))
state_dict[f"{prefix_gate}{key_map['w_gate_up_proj']}"] = fused_gate_up_proj
del state_dict[wgate_key], state_dict[wproj_key]
return state_dict
def fuse_qkv_meta(state_dict):
# Weight keys list
wq_list = sorted(list(filter(lambda x: "wq.weight" in x, state_dict.keys())))
wk_list = sorted(list(filter(lambda x: "wk.weight" in x, state_dict.keys())))
wv_list = sorted(list(filter(lambda x: "wv.weight" in x, state_dict.keys())))
# Bias keys list
wq_bias_list = sorted(list(filter(lambda x: "wq.bias" in x, state_dict.keys())))
wk_bias_list = sorted(list(filter(lambda x: "wk.bias" in x, state_dict.keys())))
wv_bias_list = sorted(list(filter(lambda x: "wv.bias" in x, state_dict.keys())))
for wq_key, wk_key, wv_key in zip(wq_list, wk_list, wv_list):
wq = state_dict[wq_key]
wk = state_dict[wk_key]
wv = state_dict[wv_key]
prefix = wq_key[: -len("wq.weight")]
fused_qkv_weights = torch.vstack((wq, wk, wv))
state_dict[f"{prefix}wqkv.weight"] = fused_qkv_weights
del state_dict[wq_key], state_dict[wk_key], state_dict[wv_key]
# Checking for bias
if len(wq_bias_list) > 0:
for wq_bias_key, wk_bias_key, wv_bias_key in zip(wq_bias_list, wk_bias_list, wv_bias_list):
wq_bias = state_dict[wq_bias_key]
wk_bias = state_dict[wk_bias_key]
wv_bias = state_dict[wv_bias_key]
prefix = wq_bias_key[: -len("wq.bias")]
fused_qkv_bias = torch.vstack((wq_bias, wk_bias, wv_bias))
state_dict[f"{prefix}wqkv.bias"] = fused_qkv_bias
del state_dict[wq_bias_key], state_dict[wk_bias_key], state_dict[wv_bias_key]
return state_dict
def _is_hf_llama_vision(config):
return hasattr(config, "text_config") and hasattr(config.text_config, "cross_attention_layers")
def reindex_layers(state_dict, config):
"""Only for Llama-Vision models
Same functionality as in https://github.com/huggingface/transformers/blob/41980ce93e775f6c88500c51c8db7946fc6a2add/src/transformers/models/mllama/convert_mllama_weights_to_hf.py#L365-L369
"""
if not _is_hf_llama_vision(config):
return state_dict
new_state_dict = {k: v for k, v in state_dict.items()}
idx_cross_attn = len(config.text_config.cross_attention_layers) - 1
idx_self_attn = config.text_config.num_hidden_layers - len(config.text_config.cross_attention_layers) - 1
for i in range(config.text_config.num_hidden_layers - 1, -1, -1):
if i in config.text_config.cross_attention_layers:
keys = [k for k in new_state_dict if f"cross_attention_layers.{idx_cross_attn}." in k]
for key in keys:
new_key = key.replace(f"cross_attention_layers.{idx_cross_attn}.", f"layers.{i}.")
new_state_dict[new_key] = new_state_dict.pop(key)
idx_cross_attn -= 1
else:
keys = [k for k in new_state_dict if f"layers.{idx_self_attn}." in k]
for key in keys:
new_key = key.replace(f"layers.{idx_self_attn}.", f"layers.{i}.")
new_state_dict[new_key] = new_state_dict.pop(key)
idx_self_attn -= 1
return new_state_dict
def rename_layers_to_cross_attn(state_dict, config):
if not _is_hf_llama_vision(config):
return state_dict
mapping = {
"self_attn.q_proj.weight": "cross_attn.q_proj.weight",
"self_attn.k_proj.weight": "cross_attn.k_proj.weight",
"self_attn.v_proj.weight": "cross_attn.v_proj.weight",
"self_attn.o_proj.weight": "cross_attn.o_proj.weight",
"self_attn.q_proj.bias": "cross_attn.q_proj.bias",
"self_attn.k_proj.bias": "cross_attn.k_proj.bias",
"self_attn.v_proj.bias": "cross_attn.v_proj.bias",
"self_attn.o_proj.bias": "cross_attn.o_proj.bias",
"self_attn.q_norm.weight": "cross_attn.q_norm.weight",
"self_attn.k_norm.weight": "cross_attn.k_norm.weight",
}
new_state_dict = {}
for key, tensor in state_dict.items():
matched = False
for idx in config.text_config.cross_attention_layers:
if matched:
break
for self_attn, cross_attn in mapping.items():
self_pattern = f"layers.{idx}.{self_attn}"
cross_pattern = f"layers.{idx}.{cross_attn}"
if self_pattern in key:
key = key.replace(self_pattern, cross_pattern)
new_state_dict[key] = tensor
matched = True
break
if not matched:
new_state_dict[key] = tensor
return new_state_dict
def convert_meta_to_hf(state_dict, head_dim, fuse_qkv=False, fuse_mlp=False, config=None):
state_dict = reindex_layers(state_dict, config)
state_dict = convert_meta_qkv_to_hf_format(state_dict, head_dim)
if fuse_qkv:
state_dict = fuse_qkv_meta(state_dict)
if fuse_mlp:
state_dict = fuse_mlp_meta(state_dict)
state_dict = map_meta_to_hf_keys(state_dict)
state_dict = rename_layers_to_cross_attn(state_dict, config)
return state_dict
def convert_meta_to_hf_no_qkv_permute(state_dict, fuse_qkv=False, fuse_mlp=False, config=None):
state_dict = reindex_layers(state_dict, config)
if fuse_qkv:
state_dict = fuse_qkv_meta(state_dict)
if fuse_mlp:
state_dict = fuse_mlp_meta(state_dict)
state_dict = map_meta_to_hf_keys(state_dict)
state_dict = rename_layers_to_cross_attn(state_dict, config)
return state_dict
def replace_keys(state_dict, replacements):
"""
Replacements are in the form (pattern, replacement).
Patterns can use ^ to match the start of the string but are otherwise
matched as whole words. These are not regular expressions, e.g. . is not
a special character.
"""
for pattern, replacement in replacements:
pre = r"^" if pattern.startswith("^") else r"(?=^|\b)"
post = r"\." if pattern.endswith(".") else r"(?=\b|$)"
pattern = pattern[1:] if pattern.startswith("^") else pattern
pattern = pattern[:-1] if pattern.endswith(".") else pattern
pattern = pre + pattern + post
state_dict = {re.sub(pattern, replacement, k): v for k, v in state_dict.items()}
return state_dict
def map_hf_to_meta_keys_mllama(loaded_weights, config):
replacements = [
(r"^model.norm.weight", r"text_model.norm.weight"),
(r"^lm_head.weight", r"text_model.output.weight"),
(r"^model.embed_tokens", r"text_model.tok_embeddings"),
(r"^vision_model.patch_embedding", r"vision_model.conv1._linear"),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).self_attn.q_proj",
r"vision_model.\1.resblocks.\2.attn.wq",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).self_attn.k_proj",
r"vision_model.\1.resblocks.\2.attn.wk",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).self_attn.v_proj",
r"vision_model.\1.resblocks.\2.attn.wv",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).self_attn.o_proj",
r"vision_model.\1.resblocks.\2.attn.wo",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).mlp.fc1",
r"vision_model.\1.resblocks.\2.mlp.c_fc",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).mlp.fc2",
r"vision_model.\1.resblocks.\2.mlp.c_proj",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).input_layernorm",
r"vision_model.\1.resblocks.\2.ln_1",
),
(
r"^vision_model.(global_transformer|transformer).layers.(\d+).post_attention_layernorm",
r"vision_model.\1.resblocks.\2.ln_2",
),
(
r"^vision_model.global_transformer.layers.(\d+).(gate_ffn|gate_attn)",
r"vision_model.global_transformer.resblocks.\1.\2",
),
(r"^vision_model.layernorm_(pre|post).(weight|bias)", r"vision_model.ln_\1.\2"),
(r"^vision_model.gated_positional_embedding.embedding", r"vision_model.positional_embedding"),
(r"^vision_model.gated_positional_embedding.tile_embedding.weight", r"vision_model.gated_positional_embedding"),
(r"^vision_model.gated_positional_embedding.gate", r"vision_model.gated_positional_embedding_gate"),
(r"^vision_model.pre_tile_positional_embedding.embedding.weight", r"vision_model.pre_tile_pos_embed.embedding"),
(
r"^vision_model.post_tile_positional_embedding.embedding.weight",
r"vision_model.post_tile_pos_embed.embedding",
),
(r"^vision_model.pre_tile_positional_embedding.gate", r"vision_model.pre_tile_pos_embed.gate"),
(r"^vision_model.post_tile_positional_embedding.gate", r"vision_model.post_tile_pos_embed.gate"),
(r"^vision_model.", r"vision_model.vision_encoder."),
(r"^model.multi_modal_projector.", r"vision_model.vision_projection."),
(r"^multi_modal_projector.", r"vision_model.vision_projection."),
]
self_attn_replacements = {
(r"^model.layers.(\d+).mlp.gate_proj.", r"text_model.layers.\1.feed_forward.w1."),
(r"^model.layers.(\d+).mlp.down_proj.", r"text_model.layers.\1.feed_forward.w2."),
(r"^model.layers.(\d+).mlp.up_proj.", r"text_model.layers.\1.feed_forward.w3."),
(r"^model.layers.(\d+).input_layernorm.weight", r"text_model.layers.\1.attention_norm.weight"),
(r"^model.layers.(\d+).post_attention_layernorm.weight", r"text_model.layers.\1.ffn_norm.weight"),
(r"^model.layers.(\d+).self_attn.(q|k|v|o)_proj.weight", r"text_model.layers.\1.attention.w\2.weight"),
}
cross_attn_replacements = {
(r"^model.layers.(\d+).mlp.gate_proj.weight", r"text_model.cross_attention_layers.\1.feed_forward.w1.weight"),
(r"^model.layers.(\d+).mlp.down_proj.weight", r"text_model.cross_attention_layers.\1.feed_forward.w2.weight"),
(r"^model.layers.(\d+).mlp.up_proj.weight", r"text_model.cross_attention_layers.\1.feed_forward.w3.weight"),
(r"^model.layers.(\d+).input_layernorm.weight", r"text_model.cross_attention_layers.\1.attention_norm.weight"),
(
r"^model.layers.(\d+).post_attention_layernorm.weight",
r"text_model.cross_attention_layers.\1.ffn_norm.weight",
),
(r"^model.layers.(\d+).cross_attn_attn_gate", r"text_model.cross_attention_layers.\1.gate_attn"),
(r"^model.layers.(\d+).cross_attn_mlp_gate", r"text_model.cross_attention_layers.\1.gate_ffwd"),
(r"^model.layers.(\d+).cross_attn.(q|k|v|o)_proj", r"text_model.cross_attention_layers.\1.attention.w\2"),
(r"^model.layers.(\d+).cross_attn.(q|k)_norm", r"text_model.cross_attention_layers.\1.attention.\2_norm"),
}
idx_cross_attn = 0
for i in range(config.text_config.num_hidden_layers):
if i in config.text_config.cross_attention_layers:
cur_replacements = [
(
k.replace(r"layers.(\d+).", rf"layers.{i}."),
v.replace(r"cross_attention_layers.\1.", rf"cross_attention_layers.{idx_cross_attn}.").replace(
r"\2", r"\1"
),
)
for k, v in cross_attn_replacements
]
idx_cross_attn += 1
else:
cur_replacements = [
(
k.replace(r"layers.(\d+).", rf"layers.{i}."),
v.replace(r"layers.\1.", rf"layers.{i-idx_cross_attn}.").replace(r"\2", r"\1"),
)
for k, v in self_attn_replacements
]
replacements.extend(cur_replacements)
state_dict = replace_keys(loaded_weights, replacements)
state_dict["text_model.learnable_embedding.weight"] = state_dict["text_model.tok_embeddings.weight"][-8:]
state_dict["text_model.tok_embeddings.weight"] = state_dict["text_model.tok_embeddings.weight"][:-8]
return state_dict
def convert_pos_embeddings(state_dict):
do_convert = lambda key: (
("tile_pos_embed.embedding" in key) or (key == "vision_model.vision_encoder.gated_positional_embedding")
)
state_dict = {k: invert_pre_compute_positional_embedding(v) if do_convert(k) else v for k, v in state_dict.items()}
return state_dict
def invert_pre_compute_positional_embedding(precomputed_embeddings):
"""Inverts https://github.com/huggingface/transformers/blob/41980ce93e775f6c88500c51c8db7946fc6a2add/src/transformers/models/mllama/convert_mllama_weights_to_hf.py#L122-L148
Note: original embeddings can't be reconstructed since non-used parts (non-supported aspect ratios) are random numbers
"""
# TBD: remove hardcode
if tuple(precomputed_embeddings.shape) == (9, 5120):
max_aspect_ratio_id, max_num_tiles, num_patches, hidden_size = 9 - 1, 4, 1, 1280
elif tuple(precomputed_embeddings.shape) == (9, 8197120):
max_aspect_ratio_id, max_num_tiles, num_patches, hidden_size = 9 - 1, 4, 1601, 1280
else:
raise ValueError(f"Unknown embedding shape: {precomputed_embeddings.shape}")
precomputed_embeddings = precomputed_embeddings.reshape(
max_aspect_ratio_id + 1, max_num_tiles, num_patches, hidden_size
)
from transformers.models.mllama.image_processing_mllama import get_all_supported_aspect_ratios
supported_aspect_ratios = get_all_supported_aspect_ratios(max_num_tiles)
embedding = torch.zeros(max_num_tiles, max_num_tiles, num_patches, hidden_size, dtype=precomputed_embeddings.dtype)
for i, (height, width) in enumerate(supported_aspect_ratios):
aspect_ratio_id = i + 1
current_embedding = precomputed_embeddings[aspect_ratio_id, : height * width]
embedding[:height, :width] = current_embedding.reshape(height, width, num_patches, hidden_size)
return embedding
def flatten_conv_linear(state_dict):
do_flatten = lambda key: (("conv" in key) and ("_linear.weight" in key))
state_dict = {k: v.flatten(start_dim=1) if do_flatten(k) else v for k, v in state_dict.items()}
return state_dict
# HF name of each decoder-layer norm, keyed by the tt_transformers norm type. map_hf_to_meta_keys and
# map_meta_to_hf_keys below carry the same two pairs; one-layer tests that read a single norm weight
# straight from the checkpoint take the HF name from here instead of re-encoding it.
HF_LAYER_NORM_KEYS = {"attention": "input_layernorm", "ffn": "post_attention_layernorm"}
def map_hf_to_meta_keys(loaded_weights):
"""
Map Hugging Face checkpoint keys to Meta checkpoint keys.
You can use this to support other models by adding more mappings.
See replace_keys for more details on the format of replacements.
"""
replacements = [
("^emb.weight", "weight"),
("model.language_model.", ""),
("model.", ""),
("embed_tokens", "tok_embeddings"),
("lm_head", "output"),
("input_layernorm", "attention_norm"),
("post_attention_layernorm", "ffn_norm"),
("self_attn", "attention"),
("mlp", "feed_forward"),
("gate_proj", "w1"),
("down_proj", "w2"),
("up_proj", "w3"),
("q_proj", "wq"),
("k_proj", "wk"),
("v_proj", "wv"),
("o_proj", "wo"),
("q_norm", "q_norm"),
("k_norm", "k_norm"),
("patch_conv.weight", "patch_conv._linear.weight"), # Minimal addition for Mistral vision
]
return replace_keys(loaded_weights, replacements)
def map_meta_to_hf_keys(state_dict):
"""
Map Hugging Face checkpoint keys to Meta checkpoint keys.
You can use this to support other models by adding more mappings.
See replace_keys for more details on the format of replacements.
"""
tok_embeddings_layers = [layer for layer in state_dict if ("tok_embeddings" in layer) or ("emb.weight" in layer)]
learnable_embedding_layers = [layer for layer in state_dict if "learnable_embedding" in layer]
assert len(learnable_embedding_layers) <= len(tok_embeddings_layers) <= 1
if len(learnable_embedding_layers) == 1:
state_dict[tok_embeddings_layers[0]] = torch.cat(
[
state_dict[tok_embeddings_layers[0]],
state_dict.pop(learnable_embedding_layers[0]),
],
dim=0,
)
replacements = [
("layers", "model.layers"),
("attention_norm", "input_layernorm"),
("ffn_norm", "post_attention_layernorm"),
("attention", "self_attn"),
("wq", "q_proj"),
("wk", "k_proj"),
("wv", "v_proj"),
("wo", "o_proj"),
("wqkv", "qkv_proj"),
("feed_forward", "mlp"),
("w1", "gate_proj"),
("w2", "down_proj"),
("w3", "up_proj"),
("w1_w3", "gate_up_proj"),
("emb.weight", "weight"),
("tok_embeddings", "model.embed_tokens"),
("norm", "model.norm"),
("output", "lm_head"),
]
return replace_keys(state_dict, replacements)
def convert_meta_qkv_to_hf_format(loaded_weights, head_dim):
"""Convert Meta QKV weights back to HuggingFace format."""
converted_weights = {}
for key, tensor in loaded_weights.items():
if "wq.weight" in key or "wk.weight" in key:
# For weights: n_heads = tensor.shape[0] // head_dim
n_heads = tensor.shape[0] // head_dim
converted_weights[key] = permute(tensor, n_heads, tensor.shape[0], tensor.shape[1])
elif "wq.bias" in key or "wk.bias" in key:
# For biases: n_heads = tensor.shape[0] // head_dim
n_heads = tensor.shape[0] // head_dim
converted_weights[key] = permute(tensor.unsqueeze(-1), n_heads, tensor.shape[0], 1).squeeze(-1)
elif "q_norm.weight" in key or "k_norm.weight" in key:
converted_weights[key] = permute_1d(tensor)
else:
# Keep all other weights unchanged
converted_weights[key] = tensor
return converted_weights
def reverse_permute(tensor, n_heads, dim1, dim2):
return tensor.view(n_heads, 2, dim1 // n_heads // 2, dim2).transpose(1, 2).reshape(dim1, dim2)
def permute(tensor, n_heads, dim1, dim2):
return tensor.view(n_heads, dim1 // n_heads // 2, 2, dim2).transpose(1, 2).reshape(dim1, dim2)
def reverse_permute_1d(tensor):
"""Convert the last dim of a tensor from separate real and imaginary parts (r1, r2, i1, i2, ...) to interleaved rope format (r1, i1, r2, i2, ...)"""
shape = tensor.shape
dim = shape[-1]
assert dim % 2 == 0, "Last dimension must be even"
reals = tensor[..., : dim // 2]
imags = tensor[..., dim // 2 :]
interleaved = torch.stack((reals, imags), dim=-1).flatten(start_dim=len(shape) - 1)
return interleaved
def permute_1d(tensor):
"""Convert the last dim of a tensor from interleaved rope format (r1, i1, r2, i2, ...) to separate real and imaginary parts (r1, r2, i1, i2, ...)"""
shape = tensor.shape
dim = shape[-1]
assert dim % 2 == 0, "Last dimension must be even"
reshaped = tensor.reshape(*shape[:-1], dim // 2, 2)
reals = reshaped[..., 0]
imags = reshaped[..., 1]
return torch.cat((reals, imags), dim=-1)
def convert_rope_style_hf_to_meta(cos_hf: torch.Tensor, sin_hf: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
"""
Converts RoPE cos/sin tensors from Hugging Face style (half-dim duplicated)
to Meta style (pairwise duplicated / odd-even interleaved).
Args:
cos_hf: Cosine tensor in HF format [..., seq_len, head_dim]
(e.g., [c0, c1, ..., c_{d/2-1}, c0, c1, ..., c_{d/2-1}])
sin_hf: Sine tensor in HF format [..., seq_len, head_dim]
(e.g., [s0, s1, ..., s_{d/2-1}, s0, s1, ..., s_{d/2-1}])
Returns:
A tuple containing (cos_meta, sin_meta) in Meta format [..., seq_len, head_dim]
(e.g., [c0, c0, c1, c1, ..., c_{d/2-1}, c_{d/2-1}],
[s0, s0, s1, s1, ..., s_{d/2-1}, s_{d/2-1}])
"""
# Input validation (optional but good practice)
if cos_hf.shape != sin_hf.shape:
raise ValueError("cos_hf and sin_hf must have the same shape.")
if len(cos_hf.shape) < 2:
raise ValueError("Input tensors must have at least 2 dimensions (seq_len, head_dim).")
head_dim = cos_hf.shape[-1]
if head_dim % 2 != 0:
raise ValueError(f"Head dimension ({head_dim}) must be even.")
half_head_dim = head_dim // 2
# Select the first half (contains the unique frequencies)
cos_unique = cos_hf[..., :half_head_dim]
sin_unique = sin_hf[..., :half_head_dim]
# Repeat each unique frequency pairwise
cos_meta = torch.repeat_interleave(cos_unique, repeats=2, dim=-1)
sin_meta = torch.repeat_interleave(sin_unique, repeats=2, dim=-1)
return cos_meta, sin_meta
# Minimal addition for Mistral vision support
def map_vision_meta_to_hf_keys(loaded_weights):
"""
Map vision model Meta checkpoint keys to HuggingFace checkpoint keys.
Added for Mistral-Small-3.1-24B-Instruct-2503 vision support.
"""
base_mapping = [
("w1", "gate_proj"),
("w2", "down_proj"),
("w3", "up_proj"),
("wq", "q_proj"),
("wk", "k_proj"),
("wv", "v_proj"),
("wo", "o_proj"),
("_linear.weight", "weight"),
]
return replace_keys(loaded_weights, base_mapping)
# Minimal addition for Mistral vision support
def convert_vision_meta_to_hf(state_dict, head_dim):
"""
Convert vision model state dict from Meta to HuggingFace format.
Added for Mistral-Small-3.1-24B-Instruct-2503 vision support.
"""
state_dict = map_vision_meta_to_hf_keys(state_dict)
return state_dict