File size: 6,397 Bytes
e65937c | 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 | import json
from pathlib import Path
from typing import Any
import torch
from huggingface_hub import hf_hub_download
from huggingface_hub.errors import EntryNotFoundError
from loguru import logger
from safetensors import safe_open
_WEIGHT_ALIASES: dict[str, list[str]] = {
"embed_tokens.weight": ["tok_embeddings.weight", "llm.embed.weight"],
"lm_head.weight": ["output.weight", "llm.unembed.weight"],
"model.norm.weight": ["llm.norm.weight", "norm.weight"],
}
def _resolve_key(name: str, weight_map: dict[str, str]) -> str | None:
"""Try exact match, then suffix match, then known aliases.
When multiple keys share a suffix, the shortest key wins (most specific).
"""
for candidate in [name, *_WEIGHT_ALIASES.get(name, [])]:
if candidate in weight_map:
return candidate
matches = [k for k in weight_map if k.endswith(candidate)]
if matches:
return min(matches, key=len)
return None
def is_config_only_dir(path: str | Path) -> bool:
"""Return True if ``path`` is a local directory with a ``config.json`` but no
weight files (``*.safetensors`` / ``*.bin``).
Used to distinguish a saved speculator *config* (from which a fresh draft is
initialized) from a full checkpoint whose weights should be loaded.
:param path: A local directory path. Hub ids and non-directories return False.
:return: True when the directory holds a config but no weights.
"""
directory = Path(path)
if not directory.is_dir():
return False
has_config = (directory / "config.json").is_file()
# Weight files, plus sharded-checkpoint index files (e.g.
# model.safetensors.index.json) -- the latter end in .json and would not match
# the *.safetensors / *.bin globs, so a shard manifest must be checked explicitly
# to avoid treating an incomplete sharded checkpoint as config-only.
has_weights = (
any(directory.glob("*.safetensors"))
or any(directory.glob("*.bin"))
or any(directory.glob("*.safetensors.index.json"))
or any(directory.glob("*.bin.index.json"))
)
return has_config and not has_weights
def list_checkpoint_keys(checkpoint_dir: str | Path) -> list[str]:
"""List all tensor keys in a checkpoint without loading weights.
Supports sharded safetensors (via index) and single safetensors formats.
:param checkpoint_dir: Path to a local checkpoint directory.
:return: List of tensor key names present in the checkpoint.
"""
checkpoint_dir = Path(checkpoint_dir)
index_path = checkpoint_dir / "model.safetensors.index.json"
if index_path.exists():
with index_path.open() as f:
return list(json.load(f)["weight_map"].keys())
single = checkpoint_dir / "model.safetensors"
if single.exists():
with safe_open(str(single), framework="pt") as f:
return list(f.keys())
raise FileNotFoundError(
f"No safetensors checkpoint found at {checkpoint_dir}. "
"Expected model.safetensors.index.json or model.safetensors."
)
def load_model_layers(
layer_names: list[str], model_path: str
) -> dict[str, torch.Tensor]:
"""
Load one or more named tensors from a HF repo using safetensors shards.
Supports both exact keys and suffix pattern matching.
:param layer_names: list of tensor names or suffix patterns to load, e.g.
["model.embed_tokens.weight", "lm_head.weight"]
:param model_path: either a local directory of huggingface model
containing model.safetensors.index
:return: dict mapping input names/patterns to loaded tensors
"""
# download the index file or build weight map for single-file models
try:
index_file = _resolve_file(model_path, "model.safetensors.index.json")
with Path(index_file).open() as f:
index = json.load(f)
weight_map: dict[str, str] = index["weight_map"]
except (FileNotFoundError, EntryNotFoundError):
logger.warning(
"`model.safetensors.index.json` file not found. "
"Checking for `model.safetensors` instead."
)
model_file = _resolve_file(model_path, "model.safetensors")
# Build virtual weight map for single-file models
with safe_open(model_file, framework="pt", device="cpu") as f:
weight_map = dict.fromkeys(f.keys(), "model.safetensors")
# Resolve names: try exact match, then suffix match, then known aliases
name_to_key = {} # Maps input name to actual checkpoint key
for name in layer_names:
key = _resolve_key(name, weight_map)
if key:
name_to_key[name] = key
else:
logger.warning(f"Tensor '{name}' not found in weight_map.")
# group requested names by shard filename
shard_to_names: dict[str, list[tuple[str, str]]] = {}
for name, key in name_to_key.items():
shard = weight_map[key]
shard_to_names.setdefault(shard, []).append((name, key))
if not shard_to_names:
raise ValueError("None of the requested tensor names were found in the index.")
# fetch each required shard and extract only the requested tensors
out: dict[str, Any] = {}
for shard_file, name_key_pairs in shard_to_names.items():
shard_path = _resolve_file(model_path, shard_file)
with safe_open(shard_path, framework="pt", device="cpu") as f:
for name, key in name_key_pairs:
out[name] = f.get_tensor(key)
return out
def _resolve_file(model_path: str, file_name: str) -> Path:
"""
If model_path is a local directory, return path/<filename> if it exists.
Otherwise treat model_path as a HF repo_id and download with hf_hub_download.
:param model_path: local directory or HF repo_id
:param file_name: filename to look for or download
:return: local path to the resolved file
"""
model_path_obj = Path(model_path)
if model_path_obj.is_dir():
logger.info("Loading from local directory: {}", model_path)
p = model_path_obj / file_name
if not p.exists():
raise FileNotFoundError(f"Expected local file missing: {p}")
return p
# Treat as repo_id on the Hub
logger.info(f"Loading from huggingface directory: {model_path}: {file_name}")
return Path(hf_hub_download(repo_id=model_path, filename=file_name))
|