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.
Files changed (54) hide show
  1. logogram/__init__.py +6 -0
  2. logogram/__main__.py +5 -0
  3. logogram/analysis.py +419 -0
  4. logogram/atp.py +120 -0
  5. logogram/backends/__init__.py +5 -0
  6. logogram/backends/base.py +202 -0
  7. logogram/backends/hub.py +375 -0
  8. logogram/backends/saes.py +277 -0
  9. logogram/backends/transformer_lens.py +872 -0
  10. logogram/cli.py +496 -0
  11. logogram/compare.py +177 -0
  12. logogram/datasets.py +159 -0
  13. logogram/direct.py +193 -0
  14. logogram/engine.py +550 -0
  15. logogram/examples/ioi-gpt2/.gitignore +3 -0
  16. logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
  17. logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
  18. logogram/examples/ioi-gpt2/project.json +6 -0
  19. logogram/exports.py +33 -0
  20. logogram/features.py +368 -0
  21. logogram/fileio.py +63 -0
  22. logogram/ioi.py +220 -0
  23. logogram/paths.py +204 -0
  24. logogram/project.py +444 -0
  25. logogram/prompts.py +204 -0
  26. logogram/research.py +84 -0
  27. logogram/results.py +240 -0
  28. logogram/runner.py +396 -0
  29. logogram/runs.py +98 -0
  30. logogram/sae.py +161 -0
  31. logogram/schema.py +302 -0
  32. logogram/server/__init__.py +1 -0
  33. logogram/server/app.py +1083 -0
  34. logogram/server/models.py +426 -0
  35. logogram/server/security.py +212 -0
  36. logogram/server/state.py +585 -0
  37. logogram/sites.py +249 -0
  38. logogram/spec.py +518 -0
  39. logogram/stats.py +171 -0
  40. logogram/steering.py +258 -0
  41. logogram/system.py +379 -0
  42. logogram/updates.py +194 -0
  43. logogram/verify.py +39 -0
  44. logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
  45. logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
  46. logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
  47. logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
  48. logogram/web_dist/favicon.svg +1 -0
  49. logogram/web_dist/index.html +15 -0
  50. logogram-0.1.0.dist-info/METADATA +550 -0
  51. logogram-0.1.0.dist-info/RECORD +54 -0
  52. logogram-0.1.0.dist-info/WHEEL +4 -0
  53. logogram-0.1.0.dist-info/entry_points.txt +2 -0
  54. 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