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,872 @@
|
|
|
1
|
+
"""TransformerLens backend, built on ``TransformerBridge`` (TransformerLens 4).
|
|
2
|
+
|
|
3
|
+
Abstract sites map to TransformerLens hook points here and nowhere else.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
from __future__ import annotations
|
|
7
|
+
|
|
8
|
+
import logging
|
|
9
|
+
import threading
|
|
10
|
+
import warnings
|
|
11
|
+
from collections.abc import Callable
|
|
12
|
+
from dataclasses import asdict, dataclass
|
|
13
|
+
from importlib.metadata import version
|
|
14
|
+
from pathlib import Path
|
|
15
|
+
from typing import Any
|
|
16
|
+
|
|
17
|
+
import torch
|
|
18
|
+
|
|
19
|
+
from logogram.backends.base import (
|
|
20
|
+
ALL_KINDS,
|
|
21
|
+
BackendError,
|
|
22
|
+
Cancelled,
|
|
23
|
+
ModelBackend,
|
|
24
|
+
ModelInfo,
|
|
25
|
+
Patch,
|
|
26
|
+
Tokenized,
|
|
27
|
+
)
|
|
28
|
+
|
|
29
|
+
log = logging.getLogger(__name__)
|
|
30
|
+
|
|
31
|
+
HOOKS: dict[str, str] = {
|
|
32
|
+
"resid_pre": "blocks.{layer}.hook_resid_pre",
|
|
33
|
+
"resid_mid": "blocks.{layer}.hook_resid_mid",
|
|
34
|
+
"resid_post": "blocks.{layer}.hook_resid_post",
|
|
35
|
+
"attn_out": "blocks.{layer}.hook_attn_out",
|
|
36
|
+
"mlp_out": "blocks.{layer}.hook_mlp_out",
|
|
37
|
+
"head": "blocks.{layer}.attn.hook_z",
|
|
38
|
+
}
|
|
39
|
+
PATTERN_HOOK = "blocks.{layer}.attn.hook_pattern"
|
|
40
|
+
|
|
41
|
+
DTYPES = {"float32": torch.float32, "float16": torch.float16, "bfloat16": torch.bfloat16}
|
|
42
|
+
|
|
43
|
+
# Load-time checks compare numbers, so a short fixed input is enough. The tolerance is relative to
|
|
44
|
+
# the largest value compared, and allows for the rounding of each precision.
|
|
45
|
+
CHECK_TOLERANCE = {"float32": 1e-4, "float16": 2e-2, "bfloat16": 5e-2}
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def hook_name(kind: str, layer: int) -> str:
|
|
49
|
+
return HOOKS[kind].format(layer=layer)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _resid_mid_hooks(patch: Patch) -> list[tuple[str, Callable[..., torch.Tensor]]]:
|
|
53
|
+
"""Patch the residual stream between attention and MLP.
|
|
54
|
+
|
|
55
|
+
TransformerLens exposes ``hook_resid_mid`` as the input of the MLP's LayerNorm, so editing it
|
|
56
|
+
would change only what the MLP reads while the residual stream kept its old value. Instead,
|
|
57
|
+
since resid_mid = resid_pre + attn_out, set attn_out to (target - resid_pre) at the patched
|
|
58
|
+
positions, using this run's own resid_pre.
|
|
59
|
+
"""
|
|
60
|
+
seen: dict[str, torch.Tensor] = {}
|
|
61
|
+
|
|
62
|
+
def keep_resid_pre(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
63
|
+
seen["resid_pre"] = act
|
|
64
|
+
return act
|
|
65
|
+
|
|
66
|
+
def set_attn_out(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
67
|
+
pre = seen["resid_pre"]
|
|
68
|
+
act = act.clone()
|
|
69
|
+
values = patch.values.to(dtype=act.dtype, device=act.device)
|
|
70
|
+
if patch.positions is None:
|
|
71
|
+
act[...] = values - pre.to(act.dtype)
|
|
72
|
+
else:
|
|
73
|
+
rows = torch.arange(act.shape[0], device=act.device)
|
|
74
|
+
pos = patch.positions.to(act.device)
|
|
75
|
+
act[rows, pos] = values - pre[rows, pos].to(act.dtype)
|
|
76
|
+
return act
|
|
77
|
+
|
|
78
|
+
return [
|
|
79
|
+
(hook_name("resid_pre", patch.layer), keep_resid_pre),
|
|
80
|
+
(hook_name("attn_out", patch.layer), set_attn_out),
|
|
81
|
+
]
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
def _patch_hook(patch: Patch) -> Callable[..., torch.Tensor]:
|
|
85
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
86
|
+
act = act.clone()
|
|
87
|
+
values = patch.values.to(dtype=act.dtype, device=act.device)
|
|
88
|
+
rows = torch.arange(act.shape[0], device=act.device)
|
|
89
|
+
if patch.kind == "head":
|
|
90
|
+
heads = patch.heads.to(act.device) # type: ignore[union-attr]
|
|
91
|
+
if patch.positions is None:
|
|
92
|
+
act[rows, :, heads] = values # [B, pos, d_head]
|
|
93
|
+
else:
|
|
94
|
+
act[rows, patch.positions.to(act.device), heads] = values # [B, d_head]
|
|
95
|
+
elif patch.positions is None:
|
|
96
|
+
act[...] = values # [B, pos, d_model]
|
|
97
|
+
else:
|
|
98
|
+
act[rows, patch.positions.to(act.device)] = values # [B, d_model]
|
|
99
|
+
return act
|
|
100
|
+
|
|
101
|
+
return fn
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
@dataclass
|
|
105
|
+
class ModelChecks:
|
|
106
|
+
"""What Logogram measured about a model when it loaded (see :func:`check_model`)."""
|
|
107
|
+
|
|
108
|
+
tolerance: float
|
|
109
|
+
# Largest change in the log-probabilities against the original model (relative to their range).
|
|
110
|
+
function: float
|
|
111
|
+
# "sequential" (attention, then MLP, each added to the residual stream), "parallel" (both read
|
|
112
|
+
# the same residual and are added together) or "components" (neither could be verified).
|
|
113
|
+
structure: str
|
|
114
|
+
# Largest error in the residual additions of that structure, if one was verified.
|
|
115
|
+
residual: float | None
|
|
116
|
+
# Error of the final-norm logit lens at the last layer against the model's output, if defined.
|
|
117
|
+
lens: float | None
|
|
118
|
+
# Error of each head's output (z through its slice of W_O) summed with b_O against the attention
|
|
119
|
+
# output: small when heads add up to what attention writes, large when the model normalizes
|
|
120
|
+
# after combining them (Gemma 2, OLMo 2). None if it couldn't be measured.
|
|
121
|
+
heads: float | None = None
|
|
122
|
+
|
|
123
|
+
def to_dict(self) -> dict[str, Any]:
|
|
124
|
+
return asdict(self)
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _probe_tokens(d_vocab: int, device: str) -> torch.Tensor:
|
|
128
|
+
rows = [[(7 * i + 3) % d_vocab for i in range(8)], [(5 * i + 11) % d_vocab for i in range(8)]]
|
|
129
|
+
return torch.tensor(rows, dtype=torch.long, device=device)
|
|
130
|
+
|
|
131
|
+
|
|
132
|
+
def _relative(a: torch.Tensor, b: torch.Tensor) -> float:
|
|
133
|
+
a, b = a.float(), b.float()
|
|
134
|
+
return float((a - b).abs().max() / b.abs().max().clamp_min(1e-6))
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
def _log_prob_error(a: torch.Tensor, b: torch.Tensor) -> float:
|
|
138
|
+
return _relative(torch.log_softmax(a.float(), dim=-1), torch.log_softmax(b.float(), dim=-1))
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _soft_cap(cfg: Any) -> float | None:
|
|
142
|
+
cap = float(getattr(cfg, "output_logits_soft_cap", 0) or 0)
|
|
143
|
+
return cap if cap > 0 else None
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def final_projection(bridge: Any, residual: torch.Tensor) -> torch.Tensor:
|
|
147
|
+
"""What the model does after its last layer: final normalization, unembedding, and the logit
|
|
148
|
+
soft-capping some models (Gemma 2) apply outside both."""
|
|
149
|
+
logits = bridge.unembed(bridge.ln_final(residual))
|
|
150
|
+
cap = _soft_cap(bridge.cfg)
|
|
151
|
+
if cap is not None:
|
|
152
|
+
logits = cap * torch.tanh(logits / cap)
|
|
153
|
+
return logits
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def check_model(
|
|
157
|
+
bridge: Any,
|
|
158
|
+
probe: torch.Tensor,
|
|
159
|
+
reference: torch.Tensor,
|
|
160
|
+
n_layers: int,
|
|
161
|
+
kinds: tuple[str, ...],
|
|
162
|
+
dtype: str,
|
|
163
|
+
) -> ModelChecks:
|
|
164
|
+
"""Check, on a short input, what Logogram's measurements assume about a loaded model.
|
|
165
|
+
|
|
166
|
+
TransformerLens supports many architectures, and processes some weights; rather than trusting
|
|
167
|
+
a list, this verifies the model in front of it: that TransformerLens's version predicts what the
|
|
168
|
+
original model predicts, how each layer adds attention and MLP into the residual stream, and
|
|
169
|
+
whether the final normalization and unembedding reproduce the output (the logit lens).
|
|
170
|
+
"""
|
|
171
|
+
tolerance = CHECK_TOLERANCE[dtype]
|
|
172
|
+
stream = [
|
|
173
|
+
k for k in ("resid_pre", "resid_mid", "resid_post", "attn_out", "mlp_out") if k in kinds
|
|
174
|
+
]
|
|
175
|
+
names = [hook_name(k, layer) for layer in range(n_layers) for k in stream]
|
|
176
|
+
if "head" in kinds:
|
|
177
|
+
names += [hook_name("head", layer) for layer in range(n_layers)]
|
|
178
|
+
with torch.no_grad():
|
|
179
|
+
logits, cache = bridge.run_with_cache(probe, names_filter=names)
|
|
180
|
+
|
|
181
|
+
def act(kind: str, layer: int) -> torch.Tensor:
|
|
182
|
+
return cache[hook_name(kind, layer)].float()
|
|
183
|
+
|
|
184
|
+
structure, residual = "components", None
|
|
185
|
+
if {"resid_pre", "resid_post", "attn_out", "mlp_out"} <= set(stream):
|
|
186
|
+
chain = max(
|
|
187
|
+
(_relative(act("resid_pre", i + 1), act("resid_post", i)) for i in range(n_layers - 1)),
|
|
188
|
+
default=0.0,
|
|
189
|
+
)
|
|
190
|
+
parallel = max(
|
|
191
|
+
_relative(
|
|
192
|
+
act("resid_pre", i) + act("attn_out", i) + act("mlp_out", i), act("resid_post", i)
|
|
193
|
+
)
|
|
194
|
+
for i in range(n_layers)
|
|
195
|
+
)
|
|
196
|
+
sequential = None
|
|
197
|
+
if "resid_mid" in stream:
|
|
198
|
+
sequential = max(
|
|
199
|
+
max(
|
|
200
|
+
_relative(act("resid_pre", i) + act("attn_out", i), act("resid_mid", i)),
|
|
201
|
+
_relative(act("resid_mid", i) + act("mlp_out", i), act("resid_post", i)),
|
|
202
|
+
)
|
|
203
|
+
for i in range(n_layers)
|
|
204
|
+
)
|
|
205
|
+
if sequential is not None and max(sequential, chain) <= tolerance:
|
|
206
|
+
structure, residual = "sequential", max(sequential, chain)
|
|
207
|
+
elif max(parallel, chain) <= tolerance:
|
|
208
|
+
structure, residual = "parallel", max(parallel, chain)
|
|
209
|
+
|
|
210
|
+
heads = None
|
|
211
|
+
if "head" in kinds and "attn_out" in stream:
|
|
212
|
+
try:
|
|
213
|
+
errors = []
|
|
214
|
+
for i in range(n_layers):
|
|
215
|
+
attention = bridge.blocks[i].attn
|
|
216
|
+
z = cache[hook_name("head", i)].float()
|
|
217
|
+
with torch.no_grad():
|
|
218
|
+
out = torch.einsum("bphd,hdm->bpm", z, attention.W_O.float())
|
|
219
|
+
if getattr(attention, "b_O", None) is not None:
|
|
220
|
+
out = out + attention.b_O.float()
|
|
221
|
+
errors.append(_relative(out, act("attn_out", i)))
|
|
222
|
+
heads = max(errors)
|
|
223
|
+
except Exception: # noqa: BLE001 - no per-head output weights to check with
|
|
224
|
+
heads = None
|
|
225
|
+
|
|
226
|
+
lens = None
|
|
227
|
+
if "resid_post" in stream:
|
|
228
|
+
try:
|
|
229
|
+
with torch.no_grad():
|
|
230
|
+
final = cache[hook_name("resid_post", n_layers - 1)]
|
|
231
|
+
lens = _log_prob_error(final_projection(bridge, final), logits)
|
|
232
|
+
except Exception: # noqa: BLE001 - no usable final norm or unembedding
|
|
233
|
+
lens = None
|
|
234
|
+
return ModelChecks(
|
|
235
|
+
tolerance=tolerance,
|
|
236
|
+
function=_log_prob_error(logits, reference),
|
|
237
|
+
structure=structure,
|
|
238
|
+
residual=residual,
|
|
239
|
+
lens=lens,
|
|
240
|
+
heads=heads,
|
|
241
|
+
)
|
|
242
|
+
|
|
243
|
+
|
|
244
|
+
class TransformerLensBackend(ModelBackend):
|
|
245
|
+
def __init__(self, bridge: Any, info: ModelInfo) -> None:
|
|
246
|
+
super().__init__()
|
|
247
|
+
self.bridge = bridge
|
|
248
|
+
self.info = info
|
|
249
|
+
self.tokenizer = bridge.tokenizer
|
|
250
|
+
self._logits_to_keep = self._supports_logits_to_keep()
|
|
251
|
+
|
|
252
|
+
# -- construction ------------------------------------------------------------------------
|
|
253
|
+
|
|
254
|
+
@classmethod
|
|
255
|
+
def from_bridge(
|
|
256
|
+
cls,
|
|
257
|
+
bridge: Any,
|
|
258
|
+
*,
|
|
259
|
+
model_id: str,
|
|
260
|
+
revision: str | None,
|
|
261
|
+
dtype: str,
|
|
262
|
+
process_weights: bool,
|
|
263
|
+
n_params: int | None = None,
|
|
264
|
+
) -> TransformerLensBackend:
|
|
265
|
+
bridge.eval()
|
|
266
|
+
cfg = bridge.cfg
|
|
267
|
+
probe = _probe_tokens(int(cfg.d_vocab), str(cfg.device))
|
|
268
|
+
with torch.no_grad():
|
|
269
|
+
# The model's own predictions, before TransformerLens touches its weights.
|
|
270
|
+
reference = bridge.original_model(probe).logits.float()
|
|
271
|
+
with warnings.catch_warnings():
|
|
272
|
+
warnings.simplefilter("ignore")
|
|
273
|
+
bridge.enable_compatibility_mode(
|
|
274
|
+
disable_warnings=True, no_processing=not process_weights
|
|
275
|
+
)
|
|
276
|
+
device = torch.device(str(cfg.device)).type
|
|
277
|
+
kinds = tuple(k for k in ALL_KINDS if hook_name(k, 0) in bridge.hook_dict)
|
|
278
|
+
if "head" not in kinds or PATTERN_HOOK.format(layer=0) not in bridge.hook_dict:
|
|
279
|
+
raise BackendError(
|
|
280
|
+
f"{model_id} doesn't expose per-head attention hooks in TransformerLens, so "
|
|
281
|
+
"Logogram can't run head experiments on it."
|
|
282
|
+
)
|
|
283
|
+
checks = check_model(bridge, probe, reference, int(cfg.n_layers), kinds, dtype)
|
|
284
|
+
if checks.function > checks.tolerance:
|
|
285
|
+
if process_weights:
|
|
286
|
+
raise BackendError(
|
|
287
|
+
f"Processing the weights of {model_id} changed its predictions "
|
|
288
|
+
f"(log-probabilities moved by up to {checks.function:.1%} of their range), so "
|
|
289
|
+
"results wouldn't describe the original model. Load it with weight "
|
|
290
|
+
"processing off."
|
|
291
|
+
)
|
|
292
|
+
raise BackendError(
|
|
293
|
+
f"TransformerLens's version of {model_id} doesn't reproduce the model's own "
|
|
294
|
+
f"predictions (log-probabilities differ by up to {checks.function:.1%} of their "
|
|
295
|
+
"range), so Logogram can't measure it reliably."
|
|
296
|
+
)
|
|
297
|
+
if not checks.structure.startswith("sequential"):
|
|
298
|
+
# resid_mid is patched as resid_pre + attn_out (see _resid_mid_hooks), which is only
|
|
299
|
+
# the residual stream between attention and MLP when the layer adds them in turn.
|
|
300
|
+
kinds = tuple(k for k in kinds if k != "resid_mid")
|
|
301
|
+
if n_params is None:
|
|
302
|
+
try:
|
|
303
|
+
n_params = int(bridge.n_params_total)
|
|
304
|
+
except Exception: # noqa: BLE001
|
|
305
|
+
n_params = None
|
|
306
|
+
model_type = getattr(bridge.original_model.config, "model_type", None)
|
|
307
|
+
structure = checks.structure
|
|
308
|
+
if structure == "sequential" and model_type == "gpt2":
|
|
309
|
+
structure = "sequential_pre_norm" # GPT-2's layout, normalization included, is known
|
|
310
|
+
info = ModelInfo(
|
|
311
|
+
id=model_id,
|
|
312
|
+
revision=revision,
|
|
313
|
+
architecture=str(getattr(cfg, "architecture", None) or "unknown"),
|
|
314
|
+
n_layers=int(cfg.n_layers),
|
|
315
|
+
n_heads=int(cfg.n_heads),
|
|
316
|
+
d_model=int(cfg.d_model),
|
|
317
|
+
d_head=int(cfg.d_head),
|
|
318
|
+
d_mlp=int(cfg.d_mlp) if getattr(cfg, "d_mlp", None) else None,
|
|
319
|
+
d_vocab=int(cfg.d_vocab),
|
|
320
|
+
n_ctx=int(cfg.n_ctx),
|
|
321
|
+
n_params=n_params,
|
|
322
|
+
dtype=dtype,
|
|
323
|
+
device=device,
|
|
324
|
+
device_name=_device_name(device),
|
|
325
|
+
process_weights=process_weights,
|
|
326
|
+
site_kinds=kinds,
|
|
327
|
+
backend="transformer_lens",
|
|
328
|
+
backend_version=version("transformer-lens"),
|
|
329
|
+
extra={
|
|
330
|
+
"block_structure": structure,
|
|
331
|
+
"normalization": str(getattr(cfg, "normalization_type", "unknown")),
|
|
332
|
+
"activation": str(getattr(cfg, "act_fn", "unknown")),
|
|
333
|
+
"prediction_method": "final_norm_logit_lens"
|
|
334
|
+
if checks.lens is not None and checks.lens <= checks.tolerance
|
|
335
|
+
else None,
|
|
336
|
+
"model_type": str(model_type or "unknown"),
|
|
337
|
+
"n_key_value_heads": int(getattr(cfg, "n_key_value_heads", None) or cfg.n_heads),
|
|
338
|
+
"bos": getattr(bridge.tokenizer, "bos_token_id", None) is not None,
|
|
339
|
+
# TransformerLens marks "no soft-cap" with a value of zero or below.
|
|
340
|
+
"logit_soft_cap": _soft_cap(cfg),
|
|
341
|
+
"checks": checks.to_dict(),
|
|
342
|
+
},
|
|
343
|
+
)
|
|
344
|
+
return cls(bridge, info)
|
|
345
|
+
|
|
346
|
+
def _supports_logits_to_keep(self) -> bool:
|
|
347
|
+
probe = torch.tensor([[self._bos_or_zero(), self._bos_or_zero()]], device=self.device)
|
|
348
|
+
try:
|
|
349
|
+
with torch.no_grad():
|
|
350
|
+
out = self.bridge(probe, return_type="logits", logits_to_keep=1)
|
|
351
|
+
return tuple(out.shape[:2]) == (1, 1)
|
|
352
|
+
except Exception: # noqa: BLE001 - architecture without logits_to_keep
|
|
353
|
+
return False
|
|
354
|
+
|
|
355
|
+
def _bos_or_zero(self) -> int:
|
|
356
|
+
bos = getattr(self.tokenizer, "bos_token_id", None)
|
|
357
|
+
return int(bos) if bos is not None else 0
|
|
358
|
+
|
|
359
|
+
# -- tokens ------------------------------------------------------------------------------
|
|
360
|
+
|
|
361
|
+
def tokenize(self, text: str, prepend_bos: bool) -> Tokenized:
|
|
362
|
+
enc = self.tokenizer(text, add_special_tokens=False, return_offsets_mapping=True)
|
|
363
|
+
ids = list(enc["input_ids"])
|
|
364
|
+
offsets = [tuple(o) for o in enc["offset_mapping"]]
|
|
365
|
+
if prepend_bos:
|
|
366
|
+
bos = getattr(self.tokenizer, "bos_token_id", None)
|
|
367
|
+
if bos is None:
|
|
368
|
+
raise BackendError(
|
|
369
|
+
"This model's tokenizer has no beginning-of-sequence token. Set "
|
|
370
|
+
"tokenization.prepend_bos to false in the spec."
|
|
371
|
+
)
|
|
372
|
+
ids = [int(bos), *ids]
|
|
373
|
+
offsets = [(0, 0), *offsets]
|
|
374
|
+
tokens = [self.token_str(i) for i in ids]
|
|
375
|
+
return Tokenized(ids=ids, tokens=tokens, offsets=offsets) # type: ignore[arg-type]
|
|
376
|
+
|
|
377
|
+
def single_token_id(self, text: str) -> int | None:
|
|
378
|
+
ids = self.tokenizer(text, add_special_tokens=False)["input_ids"]
|
|
379
|
+
return int(ids[0]) if len(ids) == 1 else None
|
|
380
|
+
|
|
381
|
+
def token_str(self, token_id: int) -> str:
|
|
382
|
+
return self.tokenizer.decode([int(token_id)], clean_up_tokenization_spaces=False)
|
|
383
|
+
|
|
384
|
+
# -- forward passes ----------------------------------------------------------------------
|
|
385
|
+
|
|
386
|
+
def final_logits(self, tokens: torch.Tensor, patch: Patch | None = None) -> torch.Tensor:
|
|
387
|
+
kwargs: dict[str, Any] = {"return_type": "logits"}
|
|
388
|
+
if self._logits_to_keep:
|
|
389
|
+
kwargs["logits_to_keep"] = 1
|
|
390
|
+
with self.lock, torch.no_grad():
|
|
391
|
+
bridge = self._bridge()
|
|
392
|
+
tokens = tokens.to(self.device)
|
|
393
|
+
if patch is None:
|
|
394
|
+
logits = bridge(tokens, **kwargs)
|
|
395
|
+
else:
|
|
396
|
+
if patch.kind == "resid_mid":
|
|
397
|
+
hooks = _resid_mid_hooks(patch)
|
|
398
|
+
else:
|
|
399
|
+
hooks = [(hook_name(patch.kind, patch.layer), _patch_hook(patch))]
|
|
400
|
+
logits = bridge.run_with_hooks(tokens, fwd_hooks=hooks, **kwargs)
|
|
401
|
+
return logits[:, -1, :].float()
|
|
402
|
+
|
|
403
|
+
def capture(
|
|
404
|
+
self, tokens: torch.Tensor, sites: list[tuple[str, int]]
|
|
405
|
+
) -> dict[tuple[str, int], torch.Tensor]:
|
|
406
|
+
names = {hook_name(kind, layer): (kind, layer) for kind, layer in sites}
|
|
407
|
+
last = max(layer for _, layer in sites)
|
|
408
|
+
kwargs: dict[str, Any] = {}
|
|
409
|
+
if last + 1 < self.info.n_layers:
|
|
410
|
+
kwargs["stop_at_layer"] = last + 1
|
|
411
|
+
with self.lock, torch.no_grad():
|
|
412
|
+
_, cache = self._bridge().run_with_cache(
|
|
413
|
+
tokens.to(self.device), names_filter=list(names), **kwargs
|
|
414
|
+
)
|
|
415
|
+
return {site: cache[name].detach() for name, site in names.items()}
|
|
416
|
+
|
|
417
|
+
def attention_pattern(self, tokens: torch.Tensor, layer: int) -> torch.Tensor:
|
|
418
|
+
name = PATTERN_HOOK.format(layer=layer)
|
|
419
|
+
kwargs: dict[str, Any] = {}
|
|
420
|
+
if layer + 1 < self.info.n_layers:
|
|
421
|
+
kwargs["stop_at_layer"] = layer + 1
|
|
422
|
+
with self.lock, torch.no_grad():
|
|
423
|
+
_, cache = self._bridge().run_with_cache(
|
|
424
|
+
tokens.to(self.device), names_filter=[name], **kwargs
|
|
425
|
+
)
|
|
426
|
+
return cache[name].detach().float()
|
|
427
|
+
|
|
428
|
+
def _bridge(self) -> Any:
|
|
429
|
+
if self.bridge is None:
|
|
430
|
+
raise BackendError("The model was unloaded. Load it again to continue.")
|
|
431
|
+
return self.bridge
|
|
432
|
+
|
|
433
|
+
def edit_logits(
|
|
434
|
+
self,
|
|
435
|
+
tokens: torch.Tensor,
|
|
436
|
+
kind: str,
|
|
437
|
+
layer: int,
|
|
438
|
+
edit: Any,
|
|
439
|
+
) -> torch.Tensor:
|
|
440
|
+
if kind not in ("resid_pre", "resid_post", "attn_out", "mlp_out"):
|
|
441
|
+
raise BackendError(f"Editing {kind} activations isn't supported.")
|
|
442
|
+
|
|
443
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
444
|
+
return edit(act).to(dtype=act.dtype)
|
|
445
|
+
|
|
446
|
+
kwargs: dict[str, Any] = {"return_type": "logits"}
|
|
447
|
+
if self._logits_to_keep:
|
|
448
|
+
kwargs["logits_to_keep"] = 1
|
|
449
|
+
with self.lock, torch.no_grad():
|
|
450
|
+
logits = self._bridge().run_with_hooks(
|
|
451
|
+
tokens.to(self.device), fwd_hooks=[(hook_name(kind, layer), fn)], **kwargs
|
|
452
|
+
)
|
|
453
|
+
return logits[:, -1, :].float()
|
|
454
|
+
|
|
455
|
+
def path_patch(
|
|
456
|
+
self,
|
|
457
|
+
tokens: torch.Tensor,
|
|
458
|
+
sender: Patch,
|
|
459
|
+
frozen_heads: dict[int, torch.Tensor],
|
|
460
|
+
frozen_mlps: dict[int, torch.Tensor] | None,
|
|
461
|
+
receivers: list[tuple[str, int, int, str]],
|
|
462
|
+
) -> torch.Tensor:
|
|
463
|
+
n = self.info.n_layers
|
|
464
|
+
tokens = tokens.to(self.device)
|
|
465
|
+
|
|
466
|
+
def hold_heads(layer: int) -> Callable[..., torch.Tensor]:
|
|
467
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
468
|
+
out = frozen_heads[layer].to(device=act.device, dtype=act.dtype).clone()
|
|
469
|
+
if sender.kind == "head" and sender.layer == layer:
|
|
470
|
+
rows = torch.arange(act.shape[0], device=act.device)
|
|
471
|
+
heads = sender.heads.to(act.device) # type: ignore[union-attr]
|
|
472
|
+
values = sender.values.to(device=act.device, dtype=act.dtype)
|
|
473
|
+
if sender.positions is None:
|
|
474
|
+
out[rows, :, heads] = values
|
|
475
|
+
else:
|
|
476
|
+
out[rows, sender.positions.to(act.device), heads] = values
|
|
477
|
+
return out
|
|
478
|
+
|
|
479
|
+
return fn
|
|
480
|
+
|
|
481
|
+
def hold(value: torch.Tensor) -> Callable[..., torch.Tensor]:
|
|
482
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
483
|
+
return value.to(device=act.device, dtype=act.dtype)
|
|
484
|
+
|
|
485
|
+
return fn
|
|
486
|
+
|
|
487
|
+
first: list[tuple[str, Callable[..., torch.Tensor]]] = [
|
|
488
|
+
(hook_name("head", layer), hold_heads(layer)) for layer in range(n)
|
|
489
|
+
]
|
|
490
|
+
if sender.kind == "resid_mid":
|
|
491
|
+
first += _resid_mid_hooks(sender)
|
|
492
|
+
elif sender.kind != "head":
|
|
493
|
+
first.append((hook_name(sender.kind, sender.layer), _patch_hook(sender)))
|
|
494
|
+
for layer, value in (frozen_mlps or {}).items():
|
|
495
|
+
if not (sender.kind == "mlp_out" and sender.layer == layer):
|
|
496
|
+
first.append((hook_name("mlp_out", layer), hold(value)))
|
|
497
|
+
|
|
498
|
+
# What each receiver reads, recorded in the first pass and patched in the second.
|
|
499
|
+
reads: dict[str, list[int]] = {}
|
|
500
|
+
for kind, layer, head, part in receivers:
|
|
501
|
+
if kind == "logits":
|
|
502
|
+
reads.setdefault(hook_name("resid_post", n - 1), [])
|
|
503
|
+
else:
|
|
504
|
+
reads.setdefault(f"blocks.{layer}.attn.hook_{part}", []).append(head)
|
|
505
|
+
recorded: dict[str, torch.Tensor] = {}
|
|
506
|
+
|
|
507
|
+
def record(name: str) -> Callable[..., torch.Tensor]:
|
|
508
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
509
|
+
recorded[name] = act.detach().clone()
|
|
510
|
+
return act
|
|
511
|
+
|
|
512
|
+
return fn
|
|
513
|
+
|
|
514
|
+
def replay(name: str, heads: list[int]) -> Callable[..., torch.Tensor]:
|
|
515
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
516
|
+
value = recorded[name].to(dtype=act.dtype)
|
|
517
|
+
if not heads: # the logits read the whole residual stream
|
|
518
|
+
return value
|
|
519
|
+
act = act.clone()
|
|
520
|
+
act[:, :, heads] = value[:, :, heads]
|
|
521
|
+
return act
|
|
522
|
+
|
|
523
|
+
return fn
|
|
524
|
+
|
|
525
|
+
logits_receiver = any(kind == "logits" for kind, *_ in receivers)
|
|
526
|
+
kwargs: dict[str, Any] = {"return_type": None}
|
|
527
|
+
if not logits_receiver:
|
|
528
|
+
last = max(layer for kind, layer, *_ in receivers if kind == "head")
|
|
529
|
+
if last + 1 < n:
|
|
530
|
+
kwargs["stop_at_layer"] = last + 1
|
|
531
|
+
with self.lock, torch.no_grad():
|
|
532
|
+
bridge = self._bridge()
|
|
533
|
+
bridge.run_with_hooks(
|
|
534
|
+
tokens, fwd_hooks=first + [(name, record(name)) for name in reads], **kwargs
|
|
535
|
+
)
|
|
536
|
+
second = {"return_type": "logits"}
|
|
537
|
+
if self._logits_to_keep:
|
|
538
|
+
second["logits_to_keep"] = 1
|
|
539
|
+
logits = bridge.run_with_hooks(
|
|
540
|
+
tokens,
|
|
541
|
+
fwd_hooks=[(name, replay(name, heads)) for name, heads in reads.items()],
|
|
542
|
+
**second,
|
|
543
|
+
)
|
|
544
|
+
return logits[:, -1, :].float()
|
|
545
|
+
|
|
546
|
+
def gradients(
|
|
547
|
+
self,
|
|
548
|
+
tokens: torch.Tensor,
|
|
549
|
+
answers: torch.Tensor,
|
|
550
|
+
distractors: torch.Tensor,
|
|
551
|
+
sites: list[tuple[str, int]],
|
|
552
|
+
) -> tuple[dict[tuple[str, int], torch.Tensor], dict[tuple[str, int], torch.Tensor]]:
|
|
553
|
+
# A zero tensor added at each hook point: the gradient with respect to it is the gradient
|
|
554
|
+
# with respect to the activation there, through every later use, without cutting the graph.
|
|
555
|
+
# hook_resid_mid only feeds the MLP's normalization (see _resid_mid_hooks); the residual
|
|
556
|
+
# stream between attention and MLP is resid_pre + attn_out, so its gradient is taken at
|
|
557
|
+
# attn_out, which is added to it.
|
|
558
|
+
kept: dict[tuple[str, int], torch.Tensor] = {}
|
|
559
|
+
zeros: dict[tuple[str, int], torch.Tensor] = {}
|
|
560
|
+
|
|
561
|
+
def value(site: tuple[str, int]) -> Callable[..., torch.Tensor]:
|
|
562
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
563
|
+
kept[site] = act.detach()
|
|
564
|
+
return act
|
|
565
|
+
|
|
566
|
+
return fn
|
|
567
|
+
|
|
568
|
+
def probe(site: tuple[str, int]) -> Callable[..., torch.Tensor]:
|
|
569
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
570
|
+
zero = torch.zeros_like(act, requires_grad=True)
|
|
571
|
+
zeros[site] = zero
|
|
572
|
+
if site[0] != "resid_mid":
|
|
573
|
+
kept[site] = act.detach()
|
|
574
|
+
return act + zero
|
|
575
|
+
|
|
576
|
+
return fn
|
|
577
|
+
|
|
578
|
+
hooks: list[tuple[str, Callable[..., torch.Tensor]]] = []
|
|
579
|
+
for kind, layer in dict.fromkeys(sites):
|
|
580
|
+
if kind == "resid_mid":
|
|
581
|
+
hooks.append((hook_name("resid_mid", layer), value((kind, layer))))
|
|
582
|
+
hooks.append((hook_name("attn_out", layer), probe((kind, layer))))
|
|
583
|
+
else:
|
|
584
|
+
hooks.append((hook_name(kind, layer), probe((kind, layer))))
|
|
585
|
+
kwargs: dict[str, Any] = {"return_type": "logits"}
|
|
586
|
+
if self._logits_to_keep:
|
|
587
|
+
kwargs["logits_to_keep"] = 1
|
|
588
|
+
with self.lock, torch.enable_grad():
|
|
589
|
+
bridge = self._bridge()
|
|
590
|
+
logits = bridge.run_with_hooks(tokens.to(self.device), fwd_hooks=hooks, **kwargs)
|
|
591
|
+
last = logits[:, -1, :].float()
|
|
592
|
+
rows = torch.arange(last.shape[0], device=last.device)
|
|
593
|
+
ld = last[rows, answers.to(last.device)] - last[rows, distractors.to(last.device)]
|
|
594
|
+
order = list(zeros)
|
|
595
|
+
grads = torch.autograd.grad(ld.sum(), [zeros[s] for s in order])
|
|
596
|
+
return kept, {site: g.detach() for site, g in zip(order, grads, strict=True)}
|
|
597
|
+
|
|
598
|
+
def direct_effects(
|
|
599
|
+
self,
|
|
600
|
+
tokens: torch.Tensor,
|
|
601
|
+
answers: torch.Tensor,
|
|
602
|
+
distractors: torch.Tensor,
|
|
603
|
+
heads: bool,
|
|
604
|
+
) -> dict[str, torch.Tensor]:
|
|
605
|
+
n = self.info.n_layers
|
|
606
|
+
kept: dict[str, torch.Tensor] = {}
|
|
607
|
+
|
|
608
|
+
def keep(name: str) -> Callable[..., torch.Tensor]:
|
|
609
|
+
def fn(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
610
|
+
kept[name] = act[:, -1].detach() # only the last position is read out
|
|
611
|
+
return act
|
|
612
|
+
|
|
613
|
+
return fn
|
|
614
|
+
|
|
615
|
+
names = [hook_name("resid_pre", 0), hook_name("resid_post", n - 1)]
|
|
616
|
+
names += [hook_name(k, layer) for layer in range(n) for k in ("attn_out", "mlp_out")]
|
|
617
|
+
if heads:
|
|
618
|
+
names += [hook_name("head", layer) for layer in range(n)]
|
|
619
|
+
with self.lock:
|
|
620
|
+
bridge = self._bridge()
|
|
621
|
+
with torch.no_grad():
|
|
622
|
+
bridge.run_with_hooks(
|
|
623
|
+
tokens.to(self.device),
|
|
624
|
+
fwd_hooks=[(name, keep(name)) for name in names],
|
|
625
|
+
return_type=None,
|
|
626
|
+
)
|
|
627
|
+
|
|
628
|
+
# The direction in the residual stream that the logit difference reads, with the final
|
|
629
|
+
# normalization's scale held at its value for each prompt: the logit difference is then
|
|
630
|
+
# an affine function of the residual stream, and splits over its components.
|
|
631
|
+
def hold(scale: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
632
|
+
return scale.detach()
|
|
633
|
+
|
|
634
|
+
with torch.enable_grad(), warnings.catch_warnings():
|
|
635
|
+
# TransformerLens warns that an edited scale is recomputed from the hooked values;
|
|
636
|
+
# that is the point, and the result is checked against the measured logit difference.
|
|
637
|
+
warnings.simplefilter("ignore")
|
|
638
|
+
final = kept[hook_name("resid_post", n - 1)].clone().requires_grad_(True)
|
|
639
|
+
with bridge.hooks(fwd_hooks=[("ln_final.hook_scale", hold)]):
|
|
640
|
+
logits = final_projection(bridge, final[:, None, :])[:, 0]
|
|
641
|
+
rows = torch.arange(final.shape[0], device=final.device)
|
|
642
|
+
ld = (
|
|
643
|
+
logits[rows, answers.to(final.device)]
|
|
644
|
+
- logits[rows, distractors.to(final.device)]
|
|
645
|
+
)
|
|
646
|
+
(direction,) = torch.autograd.grad(ld.sum(), final)
|
|
647
|
+
g = direction.double()
|
|
648
|
+
|
|
649
|
+
def term(vector: torch.Tensor) -> torch.Tensor:
|
|
650
|
+
return (vector.double() * g).sum(-1)
|
|
651
|
+
|
|
652
|
+
out: dict[str, torch.Tensor] = {
|
|
653
|
+
"embed": term(kept[hook_name("resid_pre", 0)]),
|
|
654
|
+
"attn_out": torch.stack(
|
|
655
|
+
[term(kept[hook_name("attn_out", layer)]) for layer in range(n)], dim=1
|
|
656
|
+
),
|
|
657
|
+
"mlp_out": torch.stack(
|
|
658
|
+
[term(kept[hook_name("mlp_out", layer)]) for layer in range(n)], dim=1
|
|
659
|
+
),
|
|
660
|
+
"logit_diff": ld.detach().double(),
|
|
661
|
+
}
|
|
662
|
+
if heads:
|
|
663
|
+
per_layer = []
|
|
664
|
+
for layer in range(n):
|
|
665
|
+
z = kept[hook_name("head", layer)].double() # [B, H, d_head]
|
|
666
|
+
w_o = bridge.blocks[layer].attn.W_O.detach().double() # [H, d_head, d_model]
|
|
667
|
+
written = torch.einsum("bhd,hdm->bhm", z, w_o)
|
|
668
|
+
per_layer.append((written * g[:, None, :]).sum(-1))
|
|
669
|
+
out["head"] = torch.stack(per_layer, dim=1)
|
|
670
|
+
total = out["embed"] + out["attn_out"].sum(1) + out["mlp_out"].sum(1)
|
|
671
|
+
out["remainder"] = out["logit_diff"] - total
|
|
672
|
+
return {k: v.cpu() for k, v in out.items()}
|
|
673
|
+
|
|
674
|
+
def layer_logits(self, tokens: torch.Tensor, position: int, row: int) -> torch.Tensor:
|
|
675
|
+
if self.info.extra.get("prediction_method") != "final_norm_logit_lens":
|
|
676
|
+
raise BackendError(
|
|
677
|
+
"Per-layer predictions need the model's final normalization and unembedding to "
|
|
678
|
+
"reproduce its output, and for this model they didn't when it loaded. The other "
|
|
679
|
+
"analyses remain available."
|
|
680
|
+
)
|
|
681
|
+
if not 0 <= position < tokens.shape[1] or not 0 <= row < tokens.shape[0]:
|
|
682
|
+
raise ValueError("The prediction token position or batch row is out of range.")
|
|
683
|
+
with self.lock, torch.no_grad():
|
|
684
|
+
bridge = self._bridge()
|
|
685
|
+
captured: dict[int, torch.Tensor] = {}
|
|
686
|
+
|
|
687
|
+
def project(layer: int) -> Callable[..., torch.Tensor]:
|
|
688
|
+
def keep(act: torch.Tensor, hook: Any = None) -> torch.Tensor:
|
|
689
|
+
# Recompute final normalization at each layer. Cached final-layer scales
|
|
690
|
+
# would implement attribution, not the logit lens. Keep the original
|
|
691
|
+
# batch shape and dtype through both modules, including learned biases.
|
|
692
|
+
residual = act[:, position : position + 1, :]
|
|
693
|
+
logits = final_projection(bridge, residual)
|
|
694
|
+
captured[layer] = logits[row, 0].detach().float().cpu()
|
|
695
|
+
return act
|
|
696
|
+
|
|
697
|
+
return keep
|
|
698
|
+
|
|
699
|
+
hooks = [
|
|
700
|
+
(hook_name("resid_post", layer), project(layer))
|
|
701
|
+
for layer in range(self.info.n_layers)
|
|
702
|
+
]
|
|
703
|
+
bridge.run_with_hooks(tokens.to(self.device), fwd_hooks=hooks, return_type=None)
|
|
704
|
+
if len(captured) != self.info.n_layers:
|
|
705
|
+
raise BackendError("This model didn't expose every layer's residual output.")
|
|
706
|
+
return torch.stack([captured[layer] for layer in range(self.info.n_layers)])
|
|
707
|
+
|
|
708
|
+
def close(self) -> None:
|
|
709
|
+
# Waits for a forward pass in progress; callers hold the lock across multi-step work.
|
|
710
|
+
with self.lock:
|
|
711
|
+
self.bridge = None
|
|
712
|
+
if torch.cuda.is_available():
|
|
713
|
+
torch.cuda.empty_cache()
|
|
714
|
+
|
|
715
|
+
|
|
716
|
+
def architecture_support(architecture: str) -> str | None:
|
|
717
|
+
"""None if TransformerLens can load this Hugging Face architecture; otherwise, why not."""
|
|
718
|
+
try:
|
|
719
|
+
import transformer_lens.model_bridge # noqa: F401 - loads before the factory (a cycle)
|
|
720
|
+
from transformer_lens.factories.architecture_adapter_factory import (
|
|
721
|
+
SUPPORTED_ARCHITECTURES,
|
|
722
|
+
)
|
|
723
|
+
except Exception: # noqa: BLE001 - can't tell; loading the model will say
|
|
724
|
+
return None
|
|
725
|
+
if architecture == "unknown" or architecture in SUPPORTED_ARCHITECTURES:
|
|
726
|
+
return None
|
|
727
|
+
return (
|
|
728
|
+
f"TransformerLens {version('transformer-lens')} can't load {architecture} models, so "
|
|
729
|
+
"Logogram can't either. Choose a model of a supported family, such as GPT-2, Llama, "
|
|
730
|
+
"Qwen, Gemma, Pythia or OLMo."
|
|
731
|
+
)
|
|
732
|
+
|
|
733
|
+
|
|
734
|
+
def _device_name(device: str) -> str:
|
|
735
|
+
if device == "cuda" and torch.cuda.is_available():
|
|
736
|
+
return torch.cuda.get_device_name(torch.cuda.current_device())
|
|
737
|
+
if device == "mps":
|
|
738
|
+
return "Apple GPU (Metal)"
|
|
739
|
+
from logogram.system import cpu_name
|
|
740
|
+
|
|
741
|
+
return cpu_name()
|
|
742
|
+
|
|
743
|
+
|
|
744
|
+
def resolve_device(device: str) -> str:
|
|
745
|
+
if device == "auto":
|
|
746
|
+
if torch.cuda.is_available():
|
|
747
|
+
return "cuda"
|
|
748
|
+
if getattr(torch.backends, "mps", None) and torch.backends.mps.is_available():
|
|
749
|
+
return "mps"
|
|
750
|
+
return "cpu"
|
|
751
|
+
if device == "cuda" and not torch.cuda.is_available():
|
|
752
|
+
raise BackendError(
|
|
753
|
+
"The spec asks for CUDA, but PyTorch can't see a CUDA GPU here. Run `logogram doctor` "
|
|
754
|
+
"for the fix, or set model.device to auto or cpu."
|
|
755
|
+
)
|
|
756
|
+
if device == "mps" and not (
|
|
757
|
+
getattr(torch.backends, "mps", None) and torch.backends.mps.is_available()
|
|
758
|
+
):
|
|
759
|
+
raise BackendError(
|
|
760
|
+
"The spec asks for Apple's MPS, which isn't available here. Set model.device to auto "
|
|
761
|
+
"or cpu."
|
|
762
|
+
)
|
|
763
|
+
return device
|
|
764
|
+
|
|
765
|
+
|
|
766
|
+
def load_model(
|
|
767
|
+
model_id: str,
|
|
768
|
+
*,
|
|
769
|
+
revision: str | None = None,
|
|
770
|
+
dtype: str = "float32",
|
|
771
|
+
device: str = "auto",
|
|
772
|
+
process_weights: bool = True,
|
|
773
|
+
on_progress: Callable[[dict[str, Any]], None] | None = None,
|
|
774
|
+
cancel: threading.Event | None = None,
|
|
775
|
+
) -> TransformerLensBackend:
|
|
776
|
+
"""Resolve, download (with progress) and load a model from the Hugging Face Hub.
|
|
777
|
+
|
|
778
|
+
``cancel`` stops a download promptly and a load at the next step; it raises ``Cancelled``.
|
|
779
|
+
"""
|
|
780
|
+
from logogram.backends import hub
|
|
781
|
+
|
|
782
|
+
def emit(stage: str, **data: Any) -> None:
|
|
783
|
+
if cancel is not None and cancel.is_set():
|
|
784
|
+
raise Cancelled()
|
|
785
|
+
if on_progress:
|
|
786
|
+
on_progress({"stage": stage, **data})
|
|
787
|
+
|
|
788
|
+
if dtype not in DTYPES:
|
|
789
|
+
raise BackendError(f"Unknown dtype {dtype!r}; use float32, float16 or bfloat16.")
|
|
790
|
+
resolved_device = resolve_device(device)
|
|
791
|
+
emit("resolving", model_id=model_id)
|
|
792
|
+
repo = hub.resolve(model_id, revision)
|
|
793
|
+
emit("downloading", done=0, total=sum(s for _, s in repo.files), file="")
|
|
794
|
+
|
|
795
|
+
def on_download(done: int, total: int, file: str) -> None:
|
|
796
|
+
emit("downloading", done=done, total=total, file=file)
|
|
797
|
+
|
|
798
|
+
folder = hub.download(model_id, repo, on_download, cancel=cancel)
|
|
799
|
+
emit("loading", revision=repo.revision)
|
|
800
|
+
bridge = boot_local(folder, device=resolved_device, dtype=dtype)
|
|
801
|
+
emit("processing")
|
|
802
|
+
backend = TransformerLensBackend.from_bridge(
|
|
803
|
+
bridge,
|
|
804
|
+
model_id=model_id,
|
|
805
|
+
revision=repo.revision,
|
|
806
|
+
dtype=dtype,
|
|
807
|
+
process_weights=process_weights,
|
|
808
|
+
)
|
|
809
|
+
if cancel is not None and cancel.is_set():
|
|
810
|
+
backend.close()
|
|
811
|
+
raise Cancelled()
|
|
812
|
+
emit("ready")
|
|
813
|
+
return backend
|
|
814
|
+
|
|
815
|
+
|
|
816
|
+
def boot_local(folder: Path, *, device: str, dtype: str) -> Any:
|
|
817
|
+
from transformer_lens.model_bridge import TransformerBridge
|
|
818
|
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
|
819
|
+
|
|
820
|
+
from logogram.backends.hub import validate_local_weights
|
|
821
|
+
|
|
822
|
+
_quiet_transformers()
|
|
823
|
+
validate_local_weights(folder)
|
|
824
|
+
try:
|
|
825
|
+
with warnings.catch_warnings():
|
|
826
|
+
warnings.simplefilter("ignore")
|
|
827
|
+
# The bridge otherwise lets Transformers fall back to pickle weights. Load both
|
|
828
|
+
# objects explicitly and locally so neither a cache nor an index can change that.
|
|
829
|
+
model = AutoModelForCausalLM.from_pretrained(
|
|
830
|
+
str(folder),
|
|
831
|
+
use_safetensors=True,
|
|
832
|
+
local_files_only=True,
|
|
833
|
+
trust_remote_code=False,
|
|
834
|
+
torch_dtype=DTYPES[dtype],
|
|
835
|
+
attn_implementation="eager",
|
|
836
|
+
)
|
|
837
|
+
model = model.to(device)
|
|
838
|
+
tokenizer = AutoTokenizer.from_pretrained(
|
|
839
|
+
str(folder),
|
|
840
|
+
local_files_only=True,
|
|
841
|
+
trust_remote_code=False,
|
|
842
|
+
)
|
|
843
|
+
return TransformerBridge.boot_transformers(
|
|
844
|
+
str(folder),
|
|
845
|
+
device=device,
|
|
846
|
+
dtype=DTYPES[dtype],
|
|
847
|
+
hf_model=model,
|
|
848
|
+
tokenizer=tokenizer,
|
|
849
|
+
trust_remote_code=False,
|
|
850
|
+
)
|
|
851
|
+
except torch.OutOfMemoryError as exc:
|
|
852
|
+
raise BackendError(
|
|
853
|
+
"The model doesn't fit in GPU memory. Choose a smaller model, a 16-bit dtype, or "
|
|
854
|
+
"the CPU."
|
|
855
|
+
) from exc
|
|
856
|
+
except BackendError:
|
|
857
|
+
raise
|
|
858
|
+
except Exception as exc: # noqa: BLE001
|
|
859
|
+
raise BackendError(
|
|
860
|
+
f"TransformerLens couldn't load this model ({type(exc).__name__}: {exc}). Models "
|
|
861
|
+
"TransformerLens supports are listed in its documentation."
|
|
862
|
+
) from exc
|
|
863
|
+
|
|
864
|
+
|
|
865
|
+
def _quiet_transformers() -> None:
|
|
866
|
+
try:
|
|
867
|
+
import transformers
|
|
868
|
+
|
|
869
|
+
transformers.utils.logging.set_verbosity_error()
|
|
870
|
+
transformers.utils.logging.disable_progress_bar()
|
|
871
|
+
except Exception: # noqa: BLE001
|
|
872
|
+
pass
|