logogram 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- logogram/__init__.py +6 -0
- logogram/__main__.py +5 -0
- logogram/analysis.py +419 -0
- logogram/atp.py +120 -0
- logogram/backends/__init__.py +5 -0
- logogram/backends/base.py +202 -0
- logogram/backends/hub.py +375 -0
- logogram/backends/saes.py +277 -0
- logogram/backends/transformer_lens.py +872 -0
- logogram/cli.py +496 -0
- logogram/compare.py +177 -0
- logogram/datasets.py +159 -0
- logogram/direct.py +193 -0
- logogram/engine.py +550 -0
- logogram/examples/ioi-gpt2/.gitignore +3 -0
- logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
- logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
- logogram/examples/ioi-gpt2/project.json +6 -0
- logogram/exports.py +33 -0
- logogram/features.py +368 -0
- logogram/fileio.py +63 -0
- logogram/ioi.py +220 -0
- logogram/paths.py +204 -0
- logogram/project.py +444 -0
- logogram/prompts.py +204 -0
- logogram/research.py +84 -0
- logogram/results.py +240 -0
- logogram/runner.py +396 -0
- logogram/runs.py +98 -0
- logogram/sae.py +161 -0
- logogram/schema.py +302 -0
- logogram/server/__init__.py +1 -0
- logogram/server/app.py +1083 -0
- logogram/server/models.py +426 -0
- logogram/server/security.py +212 -0
- logogram/server/state.py +585 -0
- logogram/sites.py +249 -0
- logogram/spec.py +518 -0
- logogram/stats.py +171 -0
- logogram/steering.py +258 -0
- logogram/system.py +379 -0
- logogram/updates.py +194 -0
- logogram/verify.py +39 -0
- logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
- logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
- logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
- logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
- logogram/web_dist/favicon.svg +1 -0
- logogram/web_dist/index.html +15 -0
- logogram-0.1.0.dist-info/METADATA +550 -0
- logogram-0.1.0.dist-info/RECORD +54 -0
- logogram-0.1.0.dist-info/WHEEL +4 -0
- logogram-0.1.0.dist-info/entry_points.txt +2 -0
- logogram-0.1.0.dist-info/licenses/LICENSE +21 -0
|
@@ -0,0 +1,202 @@
|
|
|
1
|
+
"""The interface between experiments and a model library.
|
|
2
|
+
|
|
3
|
+
Experiments speak only in abstract sites (``resid_pre``, ``head`` ...) and tensors. A backend
|
|
4
|
+
maps those sites onto its own hooks. TransformerLens is the only backend in v0.1; others (for
|
|
5
|
+
example remote execution) can be added by implementing :class:`ModelBackend`.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import threading
|
|
11
|
+
from abc import ABC, abstractmethod
|
|
12
|
+
from dataclasses import asdict, dataclass, field
|
|
13
|
+
from typing import Any
|
|
14
|
+
|
|
15
|
+
import torch
|
|
16
|
+
|
|
17
|
+
STREAM_KINDS = ("resid_pre", "resid_mid", "resid_post", "attn_out", "mlp_out")
|
|
18
|
+
ALL_KINDS = (*STREAM_KINDS, "head")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True)
|
|
22
|
+
class ModelInfo:
|
|
23
|
+
id: str
|
|
24
|
+
revision: str | None
|
|
25
|
+
architecture: str
|
|
26
|
+
n_layers: int
|
|
27
|
+
n_heads: int
|
|
28
|
+
d_model: int
|
|
29
|
+
d_head: int
|
|
30
|
+
d_mlp: int | None
|
|
31
|
+
d_vocab: int
|
|
32
|
+
n_ctx: int
|
|
33
|
+
n_params: int | None
|
|
34
|
+
dtype: str
|
|
35
|
+
device: str # "cpu", "cuda" or "mps"
|
|
36
|
+
device_name: str
|
|
37
|
+
process_weights: bool
|
|
38
|
+
site_kinds: tuple[str, ...]
|
|
39
|
+
backend: str
|
|
40
|
+
backend_version: str
|
|
41
|
+
extra: dict[str, Any] = field(default_factory=dict)
|
|
42
|
+
|
|
43
|
+
def to_dict(self) -> dict[str, Any]:
|
|
44
|
+
data = asdict(self)
|
|
45
|
+
data["site_kinds"] = list(self.site_kinds)
|
|
46
|
+
return data
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
@dataclass
|
|
50
|
+
class Tokenized:
|
|
51
|
+
ids: list[int]
|
|
52
|
+
tokens: list[str]
|
|
53
|
+
offsets: list[tuple[int, int]] # character span of each token; (0, 0) for added BOS
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
@dataclass
|
|
57
|
+
class Patch:
|
|
58
|
+
"""Replace one site's activation, row by row, during a forward pass.
|
|
59
|
+
|
|
60
|
+
``heads`` gives each row's head for head sites. ``positions`` gives each row's single
|
|
61
|
+
position, or is ``None`` to replace every position. ``values`` is ``[B, pos, d]`` when every
|
|
62
|
+
position is replaced and ``[B, d]`` otherwise (``d`` is ``d_head`` for heads).
|
|
63
|
+
"""
|
|
64
|
+
|
|
65
|
+
kind: str
|
|
66
|
+
layer: int
|
|
67
|
+
values: torch.Tensor
|
|
68
|
+
heads: torch.Tensor | None = None
|
|
69
|
+
positions: torch.Tensor | None = None
|
|
70
|
+
|
|
71
|
+
|
|
72
|
+
class BackendError(RuntimeError):
|
|
73
|
+
"""A model can't be loaded or run. The message says what went wrong and how to fix it."""
|
|
74
|
+
|
|
75
|
+
|
|
76
|
+
class Cancelled(Exception):
|
|
77
|
+
"""The user cancelled the job (a run between batches, or a model load between steps)."""
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
class ModelBackend(ABC):
|
|
81
|
+
info: ModelInfo
|
|
82
|
+
|
|
83
|
+
def __init__(self) -> None:
|
|
84
|
+
# All model use goes through this lock, one forward pass at a time, so interactive
|
|
85
|
+
# requests can interleave with a running sweep between batches.
|
|
86
|
+
self.lock = threading.RLock()
|
|
87
|
+
|
|
88
|
+
@abstractmethod
|
|
89
|
+
def tokenize(self, text: str, prepend_bos: bool) -> Tokenized: ...
|
|
90
|
+
|
|
91
|
+
@abstractmethod
|
|
92
|
+
def single_token_id(self, text: str) -> int | None:
|
|
93
|
+
"""The token id if ``text`` is exactly one token, else None."""
|
|
94
|
+
|
|
95
|
+
@abstractmethod
|
|
96
|
+
def token_str(self, token_id: int) -> str: ...
|
|
97
|
+
|
|
98
|
+
@abstractmethod
|
|
99
|
+
def final_logits(self, tokens: torch.Tensor, patch: Patch | None = None) -> torch.Tensor:
|
|
100
|
+
"""Logits at the last position, ``[B, vocab]`` in float32."""
|
|
101
|
+
|
|
102
|
+
@abstractmethod
|
|
103
|
+
def capture(
|
|
104
|
+
self, tokens: torch.Tensor, sites: list[tuple[str, int]]
|
|
105
|
+
) -> dict[tuple[str, int], torch.Tensor]:
|
|
106
|
+
"""Activations for ``(kind, layer)`` sites: ``[B, pos, d]``, or ``[B, pos, H, d_head]``."""
|
|
107
|
+
|
|
108
|
+
@abstractmethod
|
|
109
|
+
def attention_pattern(self, tokens: torch.Tensor, layer: int) -> torch.Tensor:
|
|
110
|
+
"""Attention probabilities ``[B, H, query, key]`` in float32."""
|
|
111
|
+
|
|
112
|
+
def edit_logits(
|
|
113
|
+
self,
|
|
114
|
+
tokens: torch.Tensor,
|
|
115
|
+
kind: str,
|
|
116
|
+
layer: int,
|
|
117
|
+
edit: Any,
|
|
118
|
+
) -> torch.Tensor:
|
|
119
|
+
"""Logits at the last position, ``[B, vocab]``, with the activation at ``(kind, layer)``
|
|
120
|
+
replaced by ``edit(activation)``: ``[B, pos, d]`` in, the same shape out. For
|
|
121
|
+
residual-stream sites the edit changes the stream itself, as patching does."""
|
|
122
|
+
raise BackendError("Editing activations isn't supported by this model backend.")
|
|
123
|
+
|
|
124
|
+
def path_patch(
|
|
125
|
+
self,
|
|
126
|
+
tokens: torch.Tensor,
|
|
127
|
+
sender: Patch,
|
|
128
|
+
frozen_heads: dict[int, torch.Tensor],
|
|
129
|
+
frozen_mlps: dict[int, torch.Tensor] | None,
|
|
130
|
+
receivers: list[tuple[str, int, int, str]],
|
|
131
|
+
) -> torch.Tensor:
|
|
132
|
+
"""Logits at the last position, ``[B, vocab]``, after patching the sender's effect into
|
|
133
|
+
the receivers' inputs only.
|
|
134
|
+
|
|
135
|
+
First pass: the sender is patched, every attention head's output ``z`` is held at
|
|
136
|
+
``frozen_heads`` (the receiver run's own values, ``[B, pos, H, d_head]`` per layer) except
|
|
137
|
+
the sender's, and so are MLP outputs when ``frozen_mlps`` is given; the receivers' inputs
|
|
138
|
+
are recorded. Second pass: only those inputs are patched in. Receivers are
|
|
139
|
+
``("head", layer, head, "q" | "k" | "v")`` or ``("logits", -1, -1, "")``.
|
|
140
|
+
"""
|
|
141
|
+
raise BackendError("Path patching isn't supported by this model backend.")
|
|
142
|
+
|
|
143
|
+
def gradients(
|
|
144
|
+
self,
|
|
145
|
+
tokens: torch.Tensor,
|
|
146
|
+
answers: torch.Tensor,
|
|
147
|
+
distractors: torch.Tensor,
|
|
148
|
+
sites: list[tuple[str, int]],
|
|
149
|
+
) -> tuple[dict[tuple[str, int], torch.Tensor], dict[tuple[str, int], torch.Tensor]]:
|
|
150
|
+
"""Activations at ``(kind, layer)`` sites and the gradient of logit(answer) -
|
|
151
|
+
logit(distractor) at the last position with respect to each, in one forward and backward
|
|
152
|
+
pass. Shapes as in :meth:`capture`. The gradient of a residual-stream site is with respect
|
|
153
|
+
to the residual stream itself, not only to what the next component reads."""
|
|
154
|
+
raise BackendError("Gradients aren't supported by this model backend.")
|
|
155
|
+
|
|
156
|
+
def direct_effects(
|
|
157
|
+
self,
|
|
158
|
+
tokens: torch.Tensor,
|
|
159
|
+
answers: torch.Tensor,
|
|
160
|
+
distractors: torch.Tensor,
|
|
161
|
+
heads: bool,
|
|
162
|
+
) -> dict[str, torch.Tensor]:
|
|
163
|
+
"""Direct contributions to logit(answer) - logit(distractor) at the last position.
|
|
164
|
+
|
|
165
|
+
Returns float64 tensors on the CPU, one row per prompt: ``embed`` ``[B]``, ``attn_out``
|
|
166
|
+
and ``mlp_out`` ``[B, layers]``, ``head`` ``[B, layers, heads]`` when ``heads`` is true,
|
|
167
|
+
``logit_diff`` ``[B]`` and ``remainder`` ``[B]`` (what biases add: the logit difference
|
|
168
|
+
minus every component's term).
|
|
169
|
+
"""
|
|
170
|
+
raise BackendError("Direct logit attribution isn't supported by this model backend.")
|
|
171
|
+
|
|
172
|
+
def layer_logits(self, tokens: torch.Tensor, position: int, row: int) -> torch.Tensor:
|
|
173
|
+
"""Final-norm logit lens at resid_post, ``[layer, vocab]`` for one row, on CPU.
|
|
174
|
+
|
|
175
|
+
Preserve the supplied batch composition. Backends must opt in rather than assuming
|
|
176
|
+
that every model has the same normalization and vocabulary projection.
|
|
177
|
+
"""
|
|
178
|
+
raise BackendError("Per-layer predictions aren't supported by this model backend.")
|
|
179
|
+
|
|
180
|
+
@property
|
|
181
|
+
def device(self) -> torch.device:
|
|
182
|
+
return torch.device(self.info.device)
|
|
183
|
+
|
|
184
|
+
def memory_in_use(self) -> int | None:
|
|
185
|
+
"""Bytes in use: accelerator memory on a GPU, this process's resident memory on a CPU."""
|
|
186
|
+
dev = self.info.device
|
|
187
|
+
if dev == "cuda" and torch.cuda.is_available():
|
|
188
|
+
return int(torch.cuda.memory_allocated())
|
|
189
|
+
if dev == "mps" and hasattr(torch, "mps"):
|
|
190
|
+
try:
|
|
191
|
+
return int(torch.mps.current_allocated_memory())
|
|
192
|
+
except Exception:
|
|
193
|
+
return None
|
|
194
|
+
try:
|
|
195
|
+
import psutil
|
|
196
|
+
|
|
197
|
+
return int(psutil.Process().memory_info().rss)
|
|
198
|
+
except Exception:
|
|
199
|
+
return None
|
|
200
|
+
|
|
201
|
+
def close(self) -> None: # noqa: B027 - optional hook
|
|
202
|
+
"""Release memory held by the model."""
|
logogram/backends/hub.py
ADDED
|
@@ -0,0 +1,375 @@
|
|
|
1
|
+
"""Hugging Face Hub access: resolve an exact revision, download with progress, read metadata.
|
|
2
|
+
|
|
3
|
+
This is the only module that talks to the network, and only when the user starts it (loading a
|
|
4
|
+
model or asking for a memory estimate). Weights are loaded only from safetensors files, which
|
|
5
|
+
cannot execute code; pickle-based ``.bin`` checkpoints are refused.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import fnmatch
|
|
11
|
+
import json
|
|
12
|
+
import logging
|
|
13
|
+
import sys
|
|
14
|
+
import threading
|
|
15
|
+
from collections.abc import Callable
|
|
16
|
+
from dataclasses import dataclass
|
|
17
|
+
from pathlib import Path
|
|
18
|
+
from typing import Any, TypeVar
|
|
19
|
+
|
|
20
|
+
from logogram.backends.base import BackendError, Cancelled
|
|
21
|
+
|
|
22
|
+
log = logging.getLogger(__name__)
|
|
23
|
+
|
|
24
|
+
# Top-level files needed to build the model and tokenizer. Weights are added separately.
|
|
25
|
+
_SUPPORT_PATTERNS = ("*.json", "*.txt", "*.model", "*.tiktoken")
|
|
26
|
+
_SKIP_FILES = {"README.md", ".gitattributes"}
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass
|
|
30
|
+
class RepoFiles:
|
|
31
|
+
revision: str
|
|
32
|
+
files: list[tuple[str, int]] # (path, size in bytes) to download
|
|
33
|
+
n_params: int | None
|
|
34
|
+
gated: bool
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
ProgressFn = Callable[[int, int, str], None] # (bytes done, bytes total, current file)
|
|
38
|
+
T = TypeVar("T")
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
def _is_offline_error(exc: Exception) -> bool:
|
|
42
|
+
from huggingface_hub.errors import OfflineModeIsEnabled
|
|
43
|
+
|
|
44
|
+
if isinstance(exc, OfflineModeIsEnabled):
|
|
45
|
+
return True
|
|
46
|
+
name = type(exc).__name__
|
|
47
|
+
return (
|
|
48
|
+
any(k in name for k in ("Connect", "Timeout", "NameResolution"))
|
|
49
|
+
or "connect" in str(exc).lower()
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def _friendly_hub_error(model_id: str, exc: Exception) -> BackendError:
|
|
54
|
+
from huggingface_hub.errors import (
|
|
55
|
+
GatedRepoError,
|
|
56
|
+
RepositoryNotFoundError,
|
|
57
|
+
RevisionNotFoundError,
|
|
58
|
+
)
|
|
59
|
+
|
|
60
|
+
if isinstance(exc, GatedRepoError):
|
|
61
|
+
return BackendError(
|
|
62
|
+
f"{model_id} is gated. Accept its license on huggingface.co, then log in on this "
|
|
63
|
+
"machine with `hf auth login` and try again."
|
|
64
|
+
)
|
|
65
|
+
if isinstance(exc, RepositoryNotFoundError):
|
|
66
|
+
return BackendError(
|
|
67
|
+
f"No model called {model_id} was found on Hugging Face. Check the spelling (ids look "
|
|
68
|
+
"like owner/name). Private models need `hf auth login` first."
|
|
69
|
+
)
|
|
70
|
+
if isinstance(exc, RevisionNotFoundError):
|
|
71
|
+
return BackendError(
|
|
72
|
+
f"That revision of {model_id} doesn't exist. Check the commit or branch."
|
|
73
|
+
)
|
|
74
|
+
return BackendError(f"Couldn't reach Hugging Face for {model_id}: {exc}")
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def _select_files(siblings: list[tuple[str, int]]) -> list[tuple[str, int]]:
|
|
78
|
+
top_level = [(f, s) for f, s in siblings if "/" not in f and f not in _SKIP_FILES]
|
|
79
|
+
weights = [(f, s) for f, s in top_level if f.endswith(".safetensors")]
|
|
80
|
+
if not weights:
|
|
81
|
+
has_pickle = any(f.endswith((".bin", ".pt", ".pth", ".ckpt")) for f, _ in top_level)
|
|
82
|
+
if has_pickle:
|
|
83
|
+
raise BackendError(
|
|
84
|
+
"This model only publishes pickle weights (.bin). Logogram loads safetensors "
|
|
85
|
+
"weights only, because pickle files can run code when opened. Choose a model "
|
|
86
|
+
"with model.safetensors, or convert the weights yourself."
|
|
87
|
+
)
|
|
88
|
+
raise BackendError("This repository has no safetensors weights at its top level.")
|
|
89
|
+
support = [
|
|
90
|
+
(f, s)
|
|
91
|
+
for f, s in top_level
|
|
92
|
+
if any(fnmatch.fnmatch(f, p) for p in _SUPPORT_PATTERNS)
|
|
93
|
+
and not f.startswith(("onnx", "flax", "tf_"))
|
|
94
|
+
]
|
|
95
|
+
return sorted(support) + sorted(weights)
|
|
96
|
+
|
|
97
|
+
|
|
98
|
+
def resolve(model_id: str, revision: str | None) -> RepoFiles:
|
|
99
|
+
"""Ask the Hub for the exact commit and the files to fetch. Falls back to the local cache."""
|
|
100
|
+
from huggingface_hub import HfApi
|
|
101
|
+
|
|
102
|
+
try:
|
|
103
|
+
info = HfApi().model_info(model_id, revision=revision, files_metadata=True)
|
|
104
|
+
except Exception as exc: # noqa: BLE001 - classified below
|
|
105
|
+
if _is_offline_error(exc):
|
|
106
|
+
cached = cached_snapshot(model_id, revision)
|
|
107
|
+
if cached is not None:
|
|
108
|
+
sha, path = cached
|
|
109
|
+
files = [(p.name, p.stat().st_size) for p in path.iterdir() if p.is_file()]
|
|
110
|
+
return RepoFiles(
|
|
111
|
+
revision=sha, files=_select_files(files), n_params=None, gated=False
|
|
112
|
+
)
|
|
113
|
+
raise BackendError(
|
|
114
|
+
f"Couldn't reach Hugging Face, and {model_id} isn't in your local cache. "
|
|
115
|
+
"Connect to the internet to download it once."
|
|
116
|
+
) from exc
|
|
117
|
+
raise _friendly_hub_error(model_id, exc) from exc
|
|
118
|
+
siblings = [(s.rfilename, int(s.size or 0)) for s in (info.siblings or [])]
|
|
119
|
+
n_params = None
|
|
120
|
+
if getattr(info, "safetensors", None) is not None:
|
|
121
|
+
n_params = int(info.safetensors.total)
|
|
122
|
+
return RepoFiles(
|
|
123
|
+
revision=info.sha,
|
|
124
|
+
files=_select_files(siblings),
|
|
125
|
+
n_params=n_params,
|
|
126
|
+
gated=bool(info.gated),
|
|
127
|
+
)
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def cached_snapshot(model_id: str, revision: str | None) -> tuple[str, Path] | None:
|
|
131
|
+
"""The (commit, folder) of a cached snapshot containing config.json, if any."""
|
|
132
|
+
from huggingface_hub import try_to_load_from_cache
|
|
133
|
+
|
|
134
|
+
found = try_to_load_from_cache(model_id, "config.json", revision=revision or "main")
|
|
135
|
+
if isinstance(found, str):
|
|
136
|
+
path = Path(found).parent
|
|
137
|
+
return path.name, path
|
|
138
|
+
return None
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def validate_local_weights(folder: Path) -> None:
|
|
142
|
+
"""Refuse pickle-only caches and unsafe or incomplete safetensors shard indexes."""
|
|
143
|
+
if (folder / "adapter_config.json").exists():
|
|
144
|
+
raise BackendError(
|
|
145
|
+
"Adapter checkpoints are not supported. Choose a complete base or merged model published as safetensors, without a separate adapter configuration."
|
|
146
|
+
)
|
|
147
|
+
single = folder / "model.safetensors"
|
|
148
|
+
index = folder / "model.safetensors.index.json"
|
|
149
|
+
if single.is_file():
|
|
150
|
+
return
|
|
151
|
+
if not index.is_file():
|
|
152
|
+
raise BackendError(
|
|
153
|
+
"No model.safetensors weights are available. Download a model with safetensors weights."
|
|
154
|
+
)
|
|
155
|
+
try:
|
|
156
|
+
data = json.loads(index.read_text(encoding="utf-8"))
|
|
157
|
+
mapping = data["weight_map"]
|
|
158
|
+
if not isinstance(mapping, dict) or not mapping:
|
|
159
|
+
raise ValueError("Empty weight map")
|
|
160
|
+
for name in mapping.values():
|
|
161
|
+
if (
|
|
162
|
+
not isinstance(name, str)
|
|
163
|
+
or "/" in name
|
|
164
|
+
or "\\" in name
|
|
165
|
+
or ":" in name
|
|
166
|
+
or not name.endswith(".safetensors")
|
|
167
|
+
):
|
|
168
|
+
raise ValueError("Unsafe shard name")
|
|
169
|
+
if not (folder / name).is_file():
|
|
170
|
+
raise ValueError("Missing shard")
|
|
171
|
+
except (OSError, ValueError, KeyError, TypeError) as exc:
|
|
172
|
+
raise BackendError(
|
|
173
|
+
"The safetensors index is invalid or has missing shards. Download the complete safetensors checkpoint again."
|
|
174
|
+
) from exc
|
|
175
|
+
|
|
176
|
+
|
|
177
|
+
def _make_bar_class(on_value: Callable[[int], None], size: int) -> type:
|
|
178
|
+
"""A silent progress bar that reports how much of one file has arrived.
|
|
179
|
+
|
|
180
|
+
Plain HTTP downloads report bytes as they arrive. Xet downloads report network bytes often
|
|
181
|
+
and bytes written to disk rarely; the larger of the two, capped at the file size, gives
|
|
182
|
+
smooth progress that still ends exactly at the file size.
|
|
183
|
+
"""
|
|
184
|
+
from tqdm.std import tqdm
|
|
185
|
+
|
|
186
|
+
class _Bar(tqdm): # type: ignore[misc]
|
|
187
|
+
monitor_interval = 0 # no display, so no tqdm monitor thread per file
|
|
188
|
+
|
|
189
|
+
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
|
190
|
+
kwargs["disable"] = True
|
|
191
|
+
super().__init__(*args, **kwargs)
|
|
192
|
+
self._written = 0
|
|
193
|
+
self._received = 0
|
|
194
|
+
|
|
195
|
+
def _report(self) -> None:
|
|
196
|
+
on_value(min(size, max(self._written, self._received)))
|
|
197
|
+
|
|
198
|
+
def update(self, n: float | None = 1) -> bool | None:
|
|
199
|
+
if n:
|
|
200
|
+
self._written += int(n)
|
|
201
|
+
self._report()
|
|
202
|
+
return None
|
|
203
|
+
|
|
204
|
+
def update_transfer(self, n: float | None = 1) -> None:
|
|
205
|
+
if n:
|
|
206
|
+
self._received += int(n)
|
|
207
|
+
self._report()
|
|
208
|
+
|
|
209
|
+
def set_transfer_postfix_str(self, *args: Any, **kwargs: Any) -> None:
|
|
210
|
+
return None
|
|
211
|
+
|
|
212
|
+
def __getattr__(self, name: str) -> Any:
|
|
213
|
+
# Display hooks that newer huggingface_hub versions may call: nothing to display.
|
|
214
|
+
if name.startswith(("set_", "update_")):
|
|
215
|
+
return lambda *args, **kwargs: None
|
|
216
|
+
raise AttributeError(name)
|
|
217
|
+
|
|
218
|
+
return _Bar
|
|
219
|
+
|
|
220
|
+
|
|
221
|
+
def download(
|
|
222
|
+
model_id: str,
|
|
223
|
+
repo: RepoFiles,
|
|
224
|
+
progress: ProgressFn | None = None,
|
|
225
|
+
*,
|
|
226
|
+
cancel: threading.Event | None = None,
|
|
227
|
+
) -> Path:
|
|
228
|
+
"""Download the selected files at the pinned revision and return the snapshot folder.
|
|
229
|
+
|
|
230
|
+
Setting ``cancel`` raises ``Cancelled`` within a fraction of a second and stops the transfer;
|
|
231
|
+
a partial file stays in the cache, where the next attempt resumes it.
|
|
232
|
+
"""
|
|
233
|
+
from huggingface_hub import hf_hub_download, try_to_load_from_cache
|
|
234
|
+
|
|
235
|
+
total = sum(size for _, size in repo.files)
|
|
236
|
+
finished = 0 # bytes of files already complete
|
|
237
|
+
folder: Path | None = None
|
|
238
|
+
for filename, size in repo.files:
|
|
239
|
+
if cancel is not None and cancel.is_set():
|
|
240
|
+
raise Cancelled()
|
|
241
|
+
cached = try_to_load_from_cache(model_id, filename, revision=repo.revision)
|
|
242
|
+
if isinstance(cached, str):
|
|
243
|
+
finished += size
|
|
244
|
+
folder = Path(cached).parent
|
|
245
|
+
if progress:
|
|
246
|
+
progress(finished, total, filename)
|
|
247
|
+
continue
|
|
248
|
+
|
|
249
|
+
def on_value(value: int, _name: str = filename, _base: int = finished) -> None:
|
|
250
|
+
if cancel is not None and cancel.is_set():
|
|
251
|
+
# Raising stops a plain HTTP transfer. Xet transfers are aborted from outside
|
|
252
|
+
# (see _abort_xet); an exception raised in their callback would only be printed.
|
|
253
|
+
if not _called_from_xet():
|
|
254
|
+
raise Cancelled()
|
|
255
|
+
return
|
|
256
|
+
if progress:
|
|
257
|
+
progress(_base + value, total, _name)
|
|
258
|
+
|
|
259
|
+
def fetch(_name: str = filename, _size: int = size, _on: Any = on_value) -> str:
|
|
260
|
+
return hf_hub_download(
|
|
261
|
+
model_id,
|
|
262
|
+
_name,
|
|
263
|
+
revision=repo.revision,
|
|
264
|
+
tqdm_class=_make_bar_class(_on, _size),
|
|
265
|
+
)
|
|
266
|
+
|
|
267
|
+
try:
|
|
268
|
+
path = _cancellable(fetch, cancel)
|
|
269
|
+
except Cancelled:
|
|
270
|
+
raise
|
|
271
|
+
except Exception as exc: # noqa: BLE001
|
|
272
|
+
raise _friendly_hub_error(model_id, exc) from exc
|
|
273
|
+
finished += size
|
|
274
|
+
folder = Path(path).parent
|
|
275
|
+
if progress:
|
|
276
|
+
progress(finished, total, filename)
|
|
277
|
+
if folder is None:
|
|
278
|
+
raise BackendError(f"Nothing to download for {model_id}.")
|
|
279
|
+
return folder
|
|
280
|
+
|
|
281
|
+
|
|
282
|
+
def _cancellable(fn: Callable[[], T], cancel: threading.Event | None) -> T:
|
|
283
|
+
"""Run ``fn`` on a helper thread, so waiting for it can stop as soon as ``cancel`` is set."""
|
|
284
|
+
if cancel is None:
|
|
285
|
+
return fn()
|
|
286
|
+
box: dict[str, Any] = {}
|
|
287
|
+
done = threading.Event()
|
|
288
|
+
|
|
289
|
+
def target() -> None:
|
|
290
|
+
try:
|
|
291
|
+
box["value"] = fn()
|
|
292
|
+
except BaseException as exc: # noqa: BLE001 - handed to the waiting thread
|
|
293
|
+
box["error"] = exc
|
|
294
|
+
finally:
|
|
295
|
+
done.set()
|
|
296
|
+
|
|
297
|
+
threading.Thread(target=target, name="logogram-download", daemon=True).start()
|
|
298
|
+
while not done.wait(0.1):
|
|
299
|
+
if cancel.is_set():
|
|
300
|
+
_abort_xet()
|
|
301
|
+
raise Cancelled()
|
|
302
|
+
if "error" in box:
|
|
303
|
+
raise box["error"]
|
|
304
|
+
return box["value"]
|
|
305
|
+
|
|
306
|
+
|
|
307
|
+
def _called_from_xet() -> bool:
|
|
308
|
+
frame = sys._getframe(1)
|
|
309
|
+
while frame is not None:
|
|
310
|
+
if "xet" in frame.f_globals.get("__name__", ""):
|
|
311
|
+
return True
|
|
312
|
+
frame = frame.f_back
|
|
313
|
+
return False
|
|
314
|
+
|
|
315
|
+
|
|
316
|
+
def _abort_xet() -> None:
|
|
317
|
+
"""Stop Xet transfers in flight, as huggingface_hub does on Ctrl+C. Without this (if the
|
|
318
|
+
private helper moves), a cancelled Xet download finishes in the background, into the cache."""
|
|
319
|
+
try:
|
|
320
|
+
from huggingface_hub.utils._xet import abort_xet_session
|
|
321
|
+
|
|
322
|
+
abort_xet_session()
|
|
323
|
+
except Exception: # noqa: BLE001
|
|
324
|
+
log.debug("couldn't abort Xet transfers", exc_info=True)
|
|
325
|
+
|
|
326
|
+
|
|
327
|
+
@dataclass
|
|
328
|
+
class ArchitectureSummary:
|
|
329
|
+
n_layers: int
|
|
330
|
+
n_heads: int
|
|
331
|
+
d_model: int
|
|
332
|
+
d_mlp: int
|
|
333
|
+
d_vocab: int
|
|
334
|
+
n_ctx: int
|
|
335
|
+
architecture: str
|
|
336
|
+
|
|
337
|
+
|
|
338
|
+
def read_architecture(config: dict[str, Any]) -> ArchitectureSummary:
|
|
339
|
+
"""Read the shape of a transformer from a Hugging Face config.json (best effort)."""
|
|
340
|
+
|
|
341
|
+
def first(*keys: str, default: int | None = None) -> int:
|
|
342
|
+
for key in keys:
|
|
343
|
+
value = config.get(key)
|
|
344
|
+
if isinstance(value, int) and value > 0:
|
|
345
|
+
return value
|
|
346
|
+
text = config.get("text_config")
|
|
347
|
+
if isinstance(text, dict):
|
|
348
|
+
for key in keys:
|
|
349
|
+
value = text.get(key)
|
|
350
|
+
if isinstance(value, int) and value > 0:
|
|
351
|
+
return value
|
|
352
|
+
if default is None:
|
|
353
|
+
raise BackendError(f"config.json has none of {', '.join(keys)}.")
|
|
354
|
+
return default
|
|
355
|
+
|
|
356
|
+
d_model = first("hidden_size", "n_embd", "d_model")
|
|
357
|
+
return ArchitectureSummary(
|
|
358
|
+
n_layers=first("num_hidden_layers", "n_layer", "num_layers"),
|
|
359
|
+
n_heads=first("num_attention_heads", "n_head", "num_heads"),
|
|
360
|
+
d_model=d_model,
|
|
361
|
+
d_mlp=first("intermediate_size", "n_inner", "ffn_dim", default=4 * d_model),
|
|
362
|
+
d_vocab=first("vocab_size"),
|
|
363
|
+
n_ctx=first("max_position_embeddings", "n_positions", "n_ctx", default=2048),
|
|
364
|
+
architecture=(config.get("architectures") or ["unknown"])[0],
|
|
365
|
+
)
|
|
366
|
+
|
|
367
|
+
|
|
368
|
+
def fetch_config(model_id: str, revision: str) -> dict[str, Any]:
|
|
369
|
+
from huggingface_hub import hf_hub_download
|
|
370
|
+
|
|
371
|
+
try:
|
|
372
|
+
path = hf_hub_download(model_id, "config.json", revision=revision)
|
|
373
|
+
except Exception as exc: # noqa: BLE001
|
|
374
|
+
raise _friendly_hub_error(model_id, exc) from exc
|
|
375
|
+
return json.loads(Path(path).read_text(encoding="utf-8"))
|