File size: 43,317 Bytes
b025706 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 | # 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
|