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))