|
|
|
|
| import gc
|
| import json
|
| import os
|
| import re
|
| import warnings
|
| from collections import defaultdict
|
| from functools import partial
|
| from pathlib import Path
|
| from pprint import pprint
|
|
|
| import torch
|
| from lightning.fabric.utilities.load import _NotYetLoadedTensor as NotYetLoadedTensor
|
| from safetensors.torch import load_file as load_safetensors
|
| from tqdm import tqdm
|
|
|
| from litgpt.config import Config
|
| from litgpt.utils import (
|
| extend_checkpoint_dir,
|
| incremental_save,
|
| lazy_load,
|
| save_config,
|
| )
|
|
|
|
|
| def copy_weights_gpt_neox(
|
| config: Config,
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "gpt_neox.embed_in.weight": "transformer.wte.weight",
|
| "gpt_neox.layers.{}.input_layernorm.bias": "transformer.h.{}.norm_1.bias",
|
| "gpt_neox.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "gpt_neox.layers.{}.attention.query_key_value.bias": "transformer.h.{}.attn.qkv.bias",
|
| "gpt_neox.layers.{}.attention.query_key_value.weight": "transformer.h.{}.attn.qkv.weight",
|
| "gpt_neox.layers.{}.attention.dense.bias": "transformer.h.{}.attn.proj.bias",
|
| "gpt_neox.layers.{}.attention.dense.weight": "transformer.h.{}.attn.proj.weight",
|
| "gpt_neox.layers.{}.attention.rotary_emb.inv_freq": None,
|
| "gpt_neox.layers.{}.attention.bias": None,
|
| "gpt_neox.layers.{}.attention.masked_bias": None,
|
| "gpt_neox.layers.{}.post_attention_layernorm.bias": "transformer.h.{}.norm_2.bias",
|
| "gpt_neox.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "gpt_neox.layers.{}.mlp.dense_h_to_4h.bias": "transformer.h.{}.mlp.fc.bias",
|
| "gpt_neox.layers.{}.mlp.dense_h_to_4h.weight": "transformer.h.{}.mlp.fc.weight",
|
| "gpt_neox.layers.{}.mlp.dense_4h_to_h.bias": "transformer.h.{}.mlp.proj.bias",
|
| "gpt_neox.layers.{}.mlp.dense_4h_to_h.weight": "transformer.h.{}.mlp.proj.weight",
|
| "gpt_neox.final_layer_norm.bias": "transformer.ln_f.bias",
|
| "gpt_neox.final_layer_norm.weight": "transformer.ln_f.weight",
|
| "embed_out.weight": "lm_head.weight",
|
| }
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, layer_idx = layer_template(from_name)
|
| to_name = weight_map[name_template]
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(layer_idx)
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if from_name.endswith((".query_key_value.weight", ".query_key_value.bias")):
|
|
|
| param = qkv_reassemble(param, config)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_falcon(
|
| config: Config,
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "transformer.word_embeddings.weight": "transformer.wte.weight",
|
| "transformer.h.{}.self_attention.query_key_value.weight": "transformer.h.{}.attn.qkv.weight",
|
| "transformer.h.{}.self_attention.dense.weight": "transformer.h.{}.attn.proj.weight",
|
| "transformer.h.{}.mlp.dense_h_to_4h.weight": "transformer.h.{}.mlp.fc.weight",
|
| "transformer.h.{}.mlp.dense_4h_to_h.weight": "transformer.h.{}.mlp.proj.weight",
|
| "transformer.ln_f.bias": "transformer.ln_f.bias",
|
| "transformer.ln_f.weight": "transformer.ln_f.weight",
|
| "lm_head.weight": "lm_head.weight",
|
| }
|
|
|
| if "7b" in config.name:
|
| weight_map.update(
|
| {
|
| "transformer.h.{}.input_layernorm.bias": "transformer.h.{}.norm_1.bias",
|
| "transformer.h.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| }
|
| )
|
| elif "40b" in config.name or "180B" in config.name:
|
| weight_map.update(
|
| {
|
| "transformer.h.{}.ln_attn.bias": "transformer.h.{}.norm_1.bias",
|
| "transformer.h.{}.ln_attn.weight": "transformer.h.{}.norm_1.weight",
|
| "transformer.h.{}.ln_mlp.bias": "transformer.h.{}.norm_2.bias",
|
| "transformer.h.{}.ln_mlp.weight": "transformer.h.{}.norm_2.weight",
|
| }
|
| )
|
| else:
|
| raise NotImplementedError
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, layer_idx = layer_template(from_name)
|
| to_name = weight_map[name_template].format(layer_idx)
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if from_name.endswith((".query_key_value.weight", ".query_key_value.bias")):
|
|
|
| param = qkv_reassemble(param, config)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_hf_llama(
|
| config: Config,
|
| qkv_weights: dict[int, list[NotYetLoadedTensor | None]],
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "model.layers.{}.input_layernorm.bias": "transformer.h.{}.norm_1.bias",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.self_attn.rotary_emb.inv_freq": None,
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "model.layers.{}.post_attention_layernorm.bias": "transformer.h.{}.norm_2.bias",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| "model.norm.bias": "transformer.ln_f.bias",
|
| "lm_head.weight": "lm_head.weight",
|
| }
|
| if config.mlp_class_name == "LLaMAMoE":
|
| weight_map.update(
|
| {
|
| "model.layers.{}.block_sparse_moe.gate.weight": "transformer.h.{}.mlp.gate.weight",
|
| "model.layers.{}.block_sparse_moe.experts.{}.w1.weight": "transformer.h.{}.mlp.experts.{}.fc_1.weight",
|
| "model.layers.{}.block_sparse_moe.experts.{}.w3.weight": "transformer.h.{}.mlp.experts.{}.fc_2.weight",
|
| "model.layers.{}.block_sparse_moe.experts.{}.w2.weight": "transformer.h.{}.mlp.experts.{}.proj.weight",
|
| }
|
| )
|
| elif config.mlp_class_name in ("LLaMAMLP", "GemmaMLP"):
|
| weight_map.update(
|
| {
|
| "model.layers.{}.mlp.gate_proj.weight": "transformer.h.{}.mlp.fc_1.weight",
|
| "model.layers.{}.mlp.up_proj.weight": "transformer.h.{}.mlp.fc_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| }
|
| )
|
| else:
|
| raise NotImplementedError
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, *ids = layer_template(from_name, num_matches=2)
|
| to_name = weight_map[name_template]
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(ids[0], defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(*ids)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict:
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_gemma_2(
|
| qkv_weights: dict[int, list[NotYetLoadedTensor | None]],
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.mlp.gate_proj.weight": "transformer.h.{}.mlp.fc_1.weight",
|
| "model.layers.{}.mlp.up_proj.weight": "transformer.h.{}.mlp.fc_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| "model.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.post_attention_norm.weight",
|
| "model.layers.{}.pre_feedforward_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "model.layers.{}.post_feedforward_layernorm.weight": "transformer.h.{}.post_mlp_norm.weight",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| "lm_head.weight": "lm_head.weight",
|
| }
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, *ids = layer_template(from_name, num_matches=2)
|
| to_name = weight_map[name_template]
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(ids[0], defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(*ids)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict:
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_gemma_3(
|
| qkv_weights: dict[int, list[NotYetLoadedTensor | None]],
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| config: Config | None = None,
|
| ) -> None:
|
| GEMMA3_LANGUAGE_MODEL_PREFIX = (
|
| "model.language_model"
|
| if any(k.startswith("model.language_model") for k in hf_weights)
|
| else "language_model.model"
|
| )
|
|
|
| GEMMA3_VISION_MODEL_PREFIX = (
|
| "model.vision_tower" if any(k.startswith("model.vision_tower") for k in hf_weights) else "vision_tower"
|
| )
|
|
|
| GEMMA3_MM_PROJECTOR_PREFIX = (
|
| "model.multi_modal_projector"
|
| if any(k.startswith("model.multi_modal_projector") for k in hf_weights)
|
| else "multi_modal_projector"
|
| )
|
|
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.mlp.gate_proj.weight": "transformer.h.{}.mlp.fc_1.weight",
|
| "model.layers.{}.mlp.up_proj.weight": "transformer.h.{}.mlp.fc_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| "model.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.post_attention_norm.weight",
|
| "model.layers.{}.pre_feedforward_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "model.layers.{}.post_feedforward_layernorm.weight": "transformer.h.{}.post_mlp_norm.weight",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| "lm_head.weight": "lm_head.weight",
|
| "model.layers.{}.self_attn.q_norm.weight": "transformer.h.{}.attn.norm_q.weight",
|
| "model.layers.{}.self_attn.k_norm.weight": "transformer.h.{}.attn.norm_k.weight",
|
| }
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| is_multimodal = any(k.startswith(GEMMA3_LANGUAGE_MODEL_PREFIX) for k in hf_weights)
|
| if is_multimodal:
|
| warnings.warn("For Gemma3 models only the text component is supported.")
|
| new_weight_map = dict()
|
| prefix = "model"
|
| for k, v in weight_map.items():
|
| if k.startswith(prefix):
|
| k = GEMMA3_LANGUAGE_MODEL_PREFIX + k[len(prefix) :]
|
| new_weight_map[k] = v
|
| weight_map = new_weight_map
|
| for from_name, param in hf_weights.items():
|
| if from_name.startswith(GEMMA3_VISION_MODEL_PREFIX) or from_name.startswith(GEMMA3_MM_PROJECTOR_PREFIX):
|
| continue
|
| name_template, *ids = layer_template(from_name, num_matches=2)
|
| to_name = weight_map.get(name_template)
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
|
|
| if is_multimodal and to_name == "transformer.wte.weight" and config is not None:
|
| param = param[: config.vocab_size]
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(ids[0], defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
|
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(*ids)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict:
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_phi(
|
| config: Config,
|
| qkv_weights: dict,
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| if any(layer_name.startswith(("layers.", "transformer.")) for layer_name in hf_weights):
|
| raise ValueError(
|
| "You are using an outdated Phi checkpoint. Please reload it as described in 'tutorials/download_phi.md'"
|
| )
|
|
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "model.layers.{}.input_layernorm.bias": "transformer.h.{}.norm_1.bias",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.q_proj.bias": None,
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.k_proj.bias": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.bias": None,
|
| "model.layers.{}.self_attn.dense.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.self_attn.dense.bias": "transformer.h.{}.attn.proj.bias",
|
| "model.layers.{}.mlp.fc1.weight": "transformer.h.{}.mlp.fc.weight",
|
| "model.layers.{}.mlp.fc1.bias": "transformer.h.{}.mlp.fc.bias",
|
| "model.layers.{}.mlp.fc2.weight": "transformer.h.{}.mlp.proj.weight",
|
| "model.layers.{}.mlp.fc2.bias": "transformer.h.{}.mlp.proj.bias",
|
| "model.final_layernorm.weight": "transformer.ln_f.weight",
|
| "model.final_layernorm.bias": "transformer.ln_f.bias",
|
| "lm_head.weight": "lm_head.weight",
|
| "lm_head.bias": "lm_head.bias",
|
| }
|
|
|
| if config.name.startswith(("Phi-3", "phi-4", "Phi-4")):
|
| weight_map.update(
|
| {
|
| "model.layers.{}.self_attn.qkv_proj.weight": "transformer.h.{}.attn.qkv.weight",
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| }
|
| )
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, layer_idx = layer_template(from_name)
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(layer_idx, defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
| elif from_name.endswith("gate_up_proj.weight"):
|
| weight = load_param(param, f"layer {layer_idx} gate_up_proj", dtype, verbose=debug_mode)
|
| fc_1, fc_2 = weight.chunk(2, dim=0)
|
| state_dict[f"transformer.h.{layer_idx}.mlp.fc_1.weight"] = fc_1
|
| state_dict[f"transformer.h.{layer_idx}.mlp.fc_2.weight"] = fc_2
|
| continue
|
| to_name = weight_map[name_template]
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(layer_idx)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict and config.name.startswith("Phi-4"):
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_qwen_2_5(
|
| config: Config,
|
| qkv_weights: dict[int, list[NotYetLoadedTensor | None]],
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.q_proj.bias": None,
|
| "model.layers.{}.self_attn.k_proj.bias": None,
|
| "model.layers.{}.self_attn.v_proj.bias": None,
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "model.layers.{}.mlp.gate_proj.weight": "transformer.h.{}.mlp.fc_1.weight",
|
| "model.layers.{}.mlp.up_proj.weight": "transformer.h.{}.mlp.fc_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| "lm_head.weight": "lm_head.weight",
|
| }
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, *ids = layer_template(from_name, num_matches=2)
|
| to_name = weight_map[name_template]
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(ids[0], defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(*ids)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict:
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_olmo2(
|
| config: Config,
|
| qkv_weights: dict[int, list[NotYetLoadedTensor | None]],
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.self_attn.q_norm.weight": "transformer.h.{}.attn.norm_q.weight",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.k_norm.weight": "transformer.h.{}.attn.norm_k.weight",
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.self_attn.rotary_emb.inv_freq": None,
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.post_attention_norm.weight",
|
| "model.layers.{}.post_attention_layernorm.bias": "transformer.h.{}.post_attention_norm.bias",
|
| "model.layers.{}.post_feedforward_layernorm.weight": "transformer.h.{}.post_mlp_norm.weight",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| "model.norm.bias": "transformer.ln_f.bias",
|
| "lm_head.weight": "lm_head.weight",
|
| }
|
| if config.mlp_class_name in ("LLaMAMLP", "GemmaMLP"):
|
| weight_map.update(
|
| {
|
| "model.layers.{}.mlp.gate_proj.weight": "transformer.h.{}.mlp.fc_1.weight",
|
| "model.layers.{}.mlp.up_proj.weight": "transformer.h.{}.mlp.fc_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| }
|
| )
|
| else:
|
| raise NotImplementedError
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, *ids = layer_template(from_name, num_matches=2)
|
| to_name = weight_map[name_template]
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(ids[0], defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(*ids)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict:
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def copy_weights_qwen_3(
|
| config: Config,
|
| qkv_weights: dict[int, list[NotYetLoadedTensor | None]],
|
| state_dict: dict[str, torch.Tensor],
|
| hf_weights: dict[str, torch.Tensor | NotYetLoadedTensor],
|
| saver: incremental_save | None = None,
|
| dtype: torch.dtype | None = None,
|
| pbar: tqdm | None = None,
|
| progress_per_file: float | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| weight_map = {
|
| "model.embed_tokens.weight": "transformer.wte.weight",
|
| "model.layers.{}.input_layernorm.weight": "transformer.h.{}.norm_1.weight",
|
| "model.layers.{}.self_attn.q_proj.weight": None,
|
| "model.layers.{}.self_attn.k_proj.weight": None,
|
| "model.layers.{}.self_attn.v_proj.weight": None,
|
| "model.layers.{}.self_attn.q_norm.weight": "transformer.h.{}.attn.norm_q.weight",
|
| "model.layers.{}.self_attn.k_norm.weight": "transformer.h.{}.attn.norm_k.weight",
|
| "model.layers.{}.self_attn.o_proj.weight": "transformer.h.{}.attn.proj.weight",
|
| "model.layers.{}.post_attention_layernorm.weight": "transformer.h.{}.norm_2.weight",
|
| "model.norm.weight": "transformer.ln_f.weight",
|
| "lm_head.weight": "lm_head.weight",
|
| }
|
| if config.mlp_class_name == "LLaMAMoE":
|
| weight_map.update(
|
| {
|
| "model.layers.{}.mlp.experts.{}.gate_proj.weight": "transformer.h.{}.mlp.experts.{}.fc_1.weight",
|
| "model.layers.{}.mlp.experts.{}.up_proj.weight": "transformer.h.{}.mlp.experts.{}.fc_2.weight",
|
| "model.layers.{}.mlp.experts.{}.down_proj.weight": "transformer.h.{}.mlp.experts.{}.proj.weight",
|
| "model.layers.{}.mlp.gate.weight": "transformer.h.{}.mlp.gate.weight",
|
| }
|
| )
|
| elif config.mlp_class_name == "LLaMAMLP":
|
| weight_map.update(
|
| {
|
| "model.layers.{}.mlp.gate_proj.weight": "transformer.h.{}.mlp.fc_1.weight",
|
| "model.layers.{}.mlp.up_proj.weight": "transformer.h.{}.mlp.fc_2.weight",
|
| "model.layers.{}.mlp.down_proj.weight": "transformer.h.{}.mlp.proj.weight",
|
| }
|
| )
|
| else:
|
| raise NotImplementedError
|
|
|
| if progress_per_file is not None:
|
| progress_per_file = progress_per_file / max(1, len(hf_weights) + len(qkv_weights))
|
|
|
| for from_name, param in hf_weights.items():
|
| name_template, *ids = layer_template(from_name, num_matches=2)
|
| to_name = weight_map[name_template]
|
| param = load_param(param, from_name, dtype, verbose=debug_mode)
|
| if any(w in from_name for w in ("q_proj", "k_proj", "v_proj")):
|
| qkv = qkv_weights.setdefault(ids[0], defaultdict(dict))
|
| weight_name, weight_type = from_name.split(".")[-2:]
|
| qkv[weight_type][weight_name] = param
|
| if to_name is None:
|
| continue
|
| to_name = to_name.format(*ids)
|
| if saver is not None:
|
| param = saver.store_early(param)
|
| state_dict[to_name] = param
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
| if "lm_head.weight" not in state_dict:
|
| state_dict["lm_head.weight"] = state_dict["transformer.wte.weight"]
|
|
|
| for i in list(qkv_weights):
|
| for weight_type in list(qkv_weights[i]):
|
| qkv = qkv_weights[i][weight_type]
|
| if len(qkv) != 3:
|
|
|
| continue
|
| q = load_param(qkv["q_proj"], f"layer {i} q {weight_type}", dtype, verbose=debug_mode)
|
| k = load_param(qkv["k_proj"], f"layer {i} k {weight_type}", dtype, verbose=debug_mode)
|
| v = load_param(qkv["v_proj"], f"layer {i} v {weight_type}", dtype, verbose=debug_mode)
|
| qkv = torch.cat((q, k, v))
|
| state_dict[f"transformer.h.{i}.attn.qkv.{weight_type}"] = qkv
|
| del qkv_weights[i][weight_type]
|
|
|
| if progress_per_file is not None:
|
| pbar.update(progress_per_file)
|
|
|
|
|
| def qkv_reassemble(
|
| param: torch.Tensor | NotYetLoadedTensor, config: Config
|
| ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
| """Reassemble from a normal to an interleaved placement in a QKV matrix.
|
| [Q, K, V, Q, K, V, ...] --> [Q, Q, ..., K, K, ..., V, V, ...]
|
| """
|
| q_per_kv = config.n_head // config.n_query_groups
|
| qs = []
|
| ks = []
|
| vs = []
|
| for chunk in torch.chunk(param, config.n_query_groups):
|
| split = torch.split(chunk, [config.head_size * q_per_kv, config.head_size, config.head_size])
|
| qs.append(split[0])
|
| ks.append(split[1])
|
| vs.append(split[2])
|
| q = torch.cat(qs)
|
| k = torch.cat(ks)
|
| v = torch.cat(vs)
|
| return torch.cat((q, k, v))
|
|
|
|
|
| def layer_template(layer_name: str, num_matches: int = 1) -> tuple[str, int]:
|
| pattern = r"\.(\d+)\."
|
| if not (search_res := re.findall(pattern, layer_name)):
|
| return layer_name, -1
|
| layer_name_template = re.sub(pattern, ".{}.", layer_name, count=num_matches)
|
| return layer_name_template, *(int(x) for x in search_res[:num_matches])
|
|
|
|
|
| def load_param(
|
| param: torch.Tensor | NotYetLoadedTensor, name: str, dtype: torch.dtype | None, verbose: bool = False
|
| ) -> torch.Tensor:
|
| if hasattr(param, "_load_tensor"):
|
|
|
| if verbose:
|
| print(f"Loading {name!r} into RAM")
|
| param = param._load_tensor()
|
| if dtype is not None and type(dtype) is not NotYetLoadedTensor and dtype != param.dtype:
|
| if verbose:
|
| print(f"Converting {name!r} from {param.dtype} to {dtype}")
|
| param = param.to(dtype)
|
| return param
|
|
|
|
|
| @torch.inference_mode()
|
| def convert_hf_checkpoint(
|
| checkpoint_dir: Path,
|
| *,
|
| model_name: str | None = None,
|
| dtype: str | None = None,
|
| debug_mode: bool | None = False,
|
| ) -> None:
|
| """
|
| Convert a Hugging Face Transformers checkpoint into a LitGPT compatible checkpoint.
|
|
|
| Arguments:
|
| checkpoint_dir: Where to save the downloaded files.
|
| model_name: The existing config name to load. This is useful to download alternative weights of existing
|
| architectures.
|
| dtype: The data type to convert the checkpoint files to. If not specified, the weights will remain in the
|
| dtype they are downloaded in.
|
| debug_mode: Prints the individual layers being loaded instead of a progress bar, which can be useful when
|
| developing and adding new models to LitGPT.
|
| """
|
| checkpoint_dir = extend_checkpoint_dir(checkpoint_dir)
|
| pprint(locals())
|
|
|
| if model_name is None:
|
| model_name = checkpoint_dir.name
|
| if dtype is not None:
|
| dtype = getattr(torch, dtype)
|
|
|
| config = Config.from_name(model_name)
|
| save_config(config, checkpoint_dir)
|
|
|
| if "falcon" in model_name:
|
| copy_fn = partial(copy_weights_falcon, config)
|
| elif model_name.lower().startswith("gemma-2"):
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_gemma_2, qkv_weights)
|
| elif model_name.lower().startswith("gemma-3"):
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_gemma_3, qkv_weights, config=config)
|
| elif model_name.lower().startswith("phi"):
|
|
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_phi, config, qkv_weights)
|
| elif model_name.lower().startswith(("qwen2.5", "qwq")):
|
|
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_qwen_2_5, config, qkv_weights)
|
| elif model_name.lower().startswith("olmo-2-"):
|
|
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_olmo2, config, qkv_weights)
|
| elif model_name.lower().startswith("qwen3"):
|
|
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_qwen_3, config, qkv_weights)
|
| elif config.mlp_class_name in ("LLaMAMLP", "GemmaMLP", "LLaMAMoE"):
|
|
|
| qkv_weights = {}
|
| copy_fn = partial(copy_weights_hf_llama, config, qkv_weights)
|
| else:
|
| copy_fn = partial(copy_weights_gpt_neox, config)
|
|
|
|
|
| sd = {}
|
|
|
|
|
| pytorch_bin_map_json_path = checkpoint_dir / "pytorch_model.bin.index.json"
|
| model_safetensor_map_json_path = checkpoint_dir / "model.safetensors.index.json"
|
| if pytorch_bin_map_json_path.is_file():
|
| with open(pytorch_bin_map_json_path, encoding="utf-8") as json_map:
|
| bin_index = json.load(json_map)
|
| bin_files = {checkpoint_dir / bin for bin in bin_index["weight_map"].values()}
|
| elif model_safetensor_map_json_path.is_file():
|
| with open(model_safetensor_map_json_path, encoding="utf-8") as json_map:
|
| bin_index = json.load(json_map)
|
| bin_files = {checkpoint_dir / bin for bin in bin_index["weight_map"].values()}
|
| else:
|
| bin_files = set(checkpoint_dir.glob("*.bin")) | set(checkpoint_dir.glob("*.safetensors"))
|
|
|
| bin_files = {f for f in bin_files if f.name != "training_args.bin"}
|
| if not bin_files:
|
| raise ValueError(f"Expected {str(checkpoint_dir)!r} to contain .bin or .safetensors files")
|
|
|
| with incremental_save(checkpoint_dir / "lit_model.pth") as saver:
|
|
|
|
|
|
|
| if not debug_mode:
|
|
|
|
|
| total_size = max(1, sum(os.path.getsize(bin_file) for bin_file in bin_files))
|
| total_progress = 100
|
|
|
| with tqdm(
|
| total=total_progress,
|
| desc="Initializing",
|
| bar_format="{desc}{percentage:3.0f}%|{bar}| {elapsed}<{remaining}, {rate_fmt}",
|
| ) as pbar:
|
| for bin_file in sorted(bin_files):
|
| pbar.set_description(f"Loading weights: {bin_file.name}")
|
| current_file_size = os.path.getsize(bin_file)
|
| progress_per_file = (current_file_size / total_size) * total_progress
|
|
|
| hf_weights = (
|
| load_safetensors(bin_file) if bin_file.suffix == ".safetensors" else lazy_load(bin_file)
|
| )
|
| copy_fn(
|
| sd,
|
| hf_weights,
|
| saver=saver,
|
| dtype=dtype,
|
| pbar=pbar,
|
| progress_per_file=progress_per_file,
|
| debug_mode=debug_mode,
|
| )
|
| gc.collect()
|
|
|
| if pbar.n < total_progress:
|
| pbar.update(total_progress - pbar.n)
|
| pbar.close()
|
| else:
|
|
|
| for bin_file in sorted(bin_files):
|
| hf_weights = load_safetensors(bin_file) if bin_file.suffix == ".safetensors" else lazy_load(bin_file)
|
| copy_fn(sd, hf_weights, saver=saver, dtype=dtype, debug_mode=debug_mode)
|
| print(f"Saving converted checkpoint to {checkpoint_dir}")
|
| saver.save(sd)
|
|
|