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
logogram/engine.py ADDED
@@ -0,0 +1,550 @@
1
+ """Run an experiment spec against a loaded model.
2
+
3
+ The sweep works layer by layer. For each layer it captures the source activations it needs (only
4
+ that layer's sites, stopping the forward pass early), then runs the receiver prompts with one
5
+ site replaced per row. Rows for many sites share a batch; batches never mix token lengths, so no
6
+ padding or attention masks are involved. After each layer the finished sites are reported, so
7
+ results can be painted while the sweep continues.
8
+
9
+ Terminology: the *receiver* is the prompt the model runs on; the *source* supplies the patched
10
+ activation (the other prompt of the pair, a mean, a donor, or zeros). The *reference* run defines
11
+ the gap used for normalization: the source prompt for patching, the corrupt prompt for ablation.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import threading
17
+ from collections.abc import Callable, Iterator
18
+ from dataclasses import dataclass, field
19
+ from typing import Any
20
+
21
+ import numpy as np
22
+ import torch
23
+
24
+ from logogram.backends.base import Cancelled, ModelBackend, Patch
25
+ from logogram.prompts import LengthGroup, PreparedPrompt, group_by_length
26
+ from logogram.sites import ResolvedSite, expand_scope, resolve_position
27
+ from logogram.spec import (
28
+ Ablation,
29
+ ActivationPatching,
30
+ AllPositions,
31
+ AttributionPatching,
32
+ DirectLogitAttribution,
33
+ FeaturesScope,
34
+ MeanBaseline,
35
+ PathPatching,
36
+ ResampleBaseline,
37
+ SitesScope,
38
+ Spec,
39
+ Steering,
40
+ ZeroBaseline,
41
+ )
42
+
43
+
44
+ class EngineError(ValueError):
45
+ pass
46
+
47
+
48
+ @dataclass
49
+ class Baselines:
50
+ clean_ld: np.ndarray
51
+ corrupt_ld: np.ndarray
52
+ clean_prob: np.ndarray
53
+ corrupt_prob: np.ndarray
54
+
55
+ def ld(self, which: str) -> np.ndarray:
56
+ return self.clean_ld if which == "clean" else self.corrupt_ld
57
+
58
+ def prob(self, which: str) -> np.ndarray:
59
+ return self.clean_prob if which == "clean" else self.corrupt_prob
60
+
61
+
62
+ @dataclass
63
+ class EngineResult:
64
+ sites: list[ResolvedSite]
65
+ layout: dict[str, Any]
66
+ prompts: list[PreparedPrompt]
67
+ baselines: Baselines
68
+ receiver: str
69
+ reference: str
70
+ patched_ld: np.ndarray # [S, n]; NaN where nothing was run patched
71
+ patched_prob: np.ndarray # [S, n]; NaN where not measured
72
+ donors: list[list[int]] | None = None
73
+ warnings: list[str] = field(default_factory=list)
74
+ # What the per-prompt values are: "intervention" (a patched forward pass), "estimate" (a
75
+ # linear estimate of one) or "attribution" (a term of a decomposition of one forward pass).
76
+ measure: str = "intervention"
77
+ # Per-prompt values [S, n] when they aren't patched - receiver (an estimate or a term), and
78
+ # what normalizes them [n] when it isn't reference - receiver.
79
+ delta: np.ndarray | None = None
80
+ gap: np.ndarray | None = None
81
+ # Method-specific numbers for the summary.
82
+ extra: dict[str, Any] = field(default_factory=dict)
83
+
84
+ @property
85
+ def receiver_ld(self) -> np.ndarray:
86
+ return self.baselines.ld(self.receiver)
87
+
88
+ @property
89
+ def reference_ld(self) -> np.ndarray:
90
+ return self.baselines.ld(self.reference)
91
+
92
+ @property
93
+ def receiver_prob(self) -> np.ndarray:
94
+ return self.baselines.prob(self.receiver)
95
+
96
+
97
+ ProgressFn = Callable[[int, int, int], None] # (rows done, rows total, current layer)
98
+ LayerFn = Callable[[int, list[int], "EngineResult"], None] # (layer, sites done, partial)
99
+
100
+
101
+ def _chunks(n: int, size: int) -> Iterator[slice]:
102
+ for start in range(0, n, size):
103
+ yield slice(start, min(start + size, n))
104
+
105
+
106
+ def _metric(
107
+ logits: torch.Tensor, answers: torch.Tensor, distractors: torch.Tensor
108
+ ) -> tuple[np.ndarray, np.ndarray]:
109
+ rows = torch.arange(logits.shape[0], device=logits.device)
110
+ answers = answers.to(logits.device)
111
+ distractors = distractors.to(logits.device)
112
+ ld = logits[rows, answers] - logits[rows, distractors]
113
+ prob = torch.log_softmax(logits, dim=-1)[rows, answers].exp()
114
+ return ld.double().cpu().numpy(), prob.double().cpu().numpy()
115
+
116
+
117
+ def _answer_tensors(prompts: list[PreparedPrompt], idx: list[int]) -> tuple[torch.Tensor, ...]:
118
+ return (
119
+ torch.tensor([prompts[i].answer_id for i in idx], dtype=torch.long),
120
+ torch.tensor([prompts[i].distractor_id for i in idx], dtype=torch.long),
121
+ )
122
+
123
+
124
+ def compute_baselines(
125
+ backend: ModelBackend,
126
+ prompts: list[PreparedPrompt],
127
+ groups: list[LengthGroup],
128
+ batch_size: int,
129
+ cancel: threading.Event | None = None,
130
+ ) -> Baselines:
131
+ n = len(prompts)
132
+ out = {k: np.zeros(n) for k in ("clean_ld", "corrupt_ld", "clean_prob", "corrupt_prob")}
133
+ for which in ("clean", "corrupt"):
134
+ for group in groups:
135
+ tokens = group.clean if which == "clean" else group.corrupt
136
+ for sl in _chunks(len(group.members), batch_size):
137
+ if cancel is not None and cancel.is_set():
138
+ raise Cancelled()
139
+ idx = group.members[sl]
140
+ logits = backend.final_logits(tokens[sl])
141
+ ld, prob = _metric(logits, *_answer_tensors(prompts, idx))
142
+ out[f"{which}_ld"][idx] = ld
143
+ out[f"{which}_prob"][idx] = prob
144
+ return Baselines(**out)
145
+
146
+
147
+ def check_gap(
148
+ spec: Spec,
149
+ baselines: Baselines,
150
+ prompts: list[PreparedPrompt],
151
+ receiver: str,
152
+ reference: str,
153
+ ) -> list[str]:
154
+ """Refuse a normalization the gap can't support, and warn when it is unreliable."""
155
+ warnings: list[str] = []
156
+ n = len(prompts)
157
+ gap = baselines.ld(reference) - baselines.ld(receiver)
158
+ mean_gap = float(gap.mean())
159
+ if spec.metric.normalization == "dataset_gap":
160
+ if abs(mean_gap) < 1e-3:
161
+ raise EngineError(
162
+ "The clean and corrupt prompts give almost the same logit difference "
163
+ f"(mean gap {mean_gap:.4f}), so a normalized effect is undefined. Run the "
164
+ "baseline check: the model may not show the behavior on these prompts."
165
+ )
166
+ if n > 1:
167
+ se = float(gap.std(ddof=1)) / np.sqrt(n)
168
+ if abs(mean_gap) < 3 * se:
169
+ warnings.append(
170
+ f"The mean clean–corrupt gap ({mean_gap:.3f}) is small next to its standard "
171
+ f"error ({se:.3f}), so normalized effects and their intervals are unreliable. "
172
+ "Use more prompts, or check the baseline."
173
+ )
174
+ if spec.metric.normalization == "prompt_gap":
175
+ zero = [p.index for p, g in zip(prompts, gap, strict=True) if abs(g) < 1e-6]
176
+ if zero:
177
+ raise EngineError(
178
+ f"Prompt(s) {', '.join(map(str, zero[:5]))} have no clean–corrupt gap, so their "
179
+ "own gap can't normalize an effect. Normalize by the dataset gap, or fix those "
180
+ "prompts."
181
+ )
182
+ small = int((np.abs(gap) < 0.1).sum())
183
+ if small:
184
+ warnings.append(
185
+ f"{small} prompt(s) have a clean–corrupt gap below 0.1, so their per-prompt "
186
+ "normalized effects are unstable."
187
+ )
188
+ return warnings
189
+
190
+
191
+ def draw_donors(
192
+ prompts: list[PreparedPrompt],
193
+ groups: list[LengthGroup],
194
+ k: int,
195
+ seed: int,
196
+ same_length: bool,
197
+ ) -> list[list[int]]:
198
+ """For each prompt, ``k`` distinct donors drawn without replacement, never the prompt itself.
199
+
200
+ Uses the raw PCG64 stream (stable across NumPy versions) and a partial Fisher-Yates shuffle.
201
+ """
202
+ bitgen = np.random.PCG64(seed)
203
+ group_of = {i: g for g in groups for i in g.members}
204
+ donors: list[list[int]] = []
205
+ for i in range(len(prompts)):
206
+ pool = group_of[i].members if same_length else list(range(len(prompts)))
207
+ candidates = [j for j in pool if j != i]
208
+ if len(candidates) < k:
209
+ where = "of the same length " if same_length else ""
210
+ raise EngineError(
211
+ f"Resampling needs {k} donor prompts {where}for prompt {i}, but only "
212
+ f"{len(candidates)} are available. Lower the donor count, use more prompts, or "
213
+ "use prompts of equal length."
214
+ )
215
+ raw = bitgen.random_raw(k)
216
+ for t in range(k):
217
+ j = t + int(raw[t] % np.uint64(len(candidates) - t))
218
+ candidates[t], candidates[j] = candidates[j], candidates[t]
219
+ donors.append(candidates[:k])
220
+ return donors
221
+
222
+
223
+ @dataclass
224
+ class _Row:
225
+ site: ResolvedSite
226
+ prompt: int # index into prompts
227
+ local: int # index within the prompt's length group
228
+ donor: int | None = None # index into prompts (resample)
229
+
230
+
231
+ class _LayerSources:
232
+ """Activations captured for one layer: per length group, and at resolved positions."""
233
+
234
+ def __init__(
235
+ self,
236
+ backend: ModelBackend,
237
+ prompts: list[PreparedPrompt],
238
+ groups: list[LengthGroup],
239
+ layer: int,
240
+ kinds: list[str],
241
+ which: str,
242
+ batch_size: int,
243
+ cancel: threading.Event | None = None,
244
+ ) -> None:
245
+ self.prompts = prompts
246
+ self.groups = groups
247
+ self.by_group: dict[str, list[torch.Tensor]] = {k: [] for k in kinds}
248
+ for group in groups:
249
+ tokens = group.clean if which == "clean" else group.corrupt
250
+ # Capture in batches, so a forward pass never holds more than batch_size prompts.
251
+ parts: dict[str, list[torch.Tensor]] = {k: [] for k in kinds}
252
+ for sl in _chunks(len(group.members), batch_size):
253
+ if cancel is not None and cancel.is_set():
254
+ raise Cancelled()
255
+ acts = backend.capture(tokens[sl], [(k, layer) for k in kinds])
256
+ for k in kinds:
257
+ parts[k].append(acts[(k, layer)])
258
+ for k in kinds:
259
+ self.by_group[k].append(torch.cat(parts[k], dim=0))
260
+ self._tables: dict[tuple[str, str], torch.Tensor] = {}
261
+
262
+ def group_tensor(self, kind: str, group_index: int) -> torch.Tensor:
263
+ return self.by_group[kind][group_index]
264
+
265
+ def table(self, kind: str, site: ResolvedSite) -> torch.Tensor:
266
+ """``[n, ...]``: each prompt's activation at its own resolved position for this site."""
267
+ key = (kind, site.site.position.model_dump_json())
268
+ if key not in self._tables:
269
+ rows = []
270
+ for gi, group in enumerate(self.groups):
271
+ acts = self.by_group[kind][gi]
272
+ for local, p in enumerate(group.members):
273
+ pos = resolve_position(site.site.position, self.prompts[p])
274
+ rows.append((p, acts[local, pos]))
275
+ rows.sort(key=lambda r: r[0])
276
+ self._tables[key] = torch.stack([r[1] for r in rows])
277
+ return self._tables[key]
278
+
279
+
280
+ def run_experiment(
281
+ spec: Spec,
282
+ backend: ModelBackend,
283
+ prompts: list[PreparedPrompt],
284
+ **kwargs: Any,
285
+ ) -> EngineResult:
286
+ """Run whichever experiment the spec describes. The app and the CLI both come through here."""
287
+ sae = kwargs.pop("sae", None)
288
+ if isinstance(spec.scope, FeaturesScope) or (
289
+ isinstance(spec.scope, SitesScope)
290
+ and any(s.kind == "sae_feature" for s in spec.scope.sites)
291
+ ):
292
+ if sae is None:
293
+ raise EngineError("This spec measures SAE features, but no SAE is loaded.")
294
+ from logogram.features import run_features
295
+
296
+ return run_features(spec, backend, prompts, sae, **kwargs)
297
+ if isinstance(spec.experiment, DirectLogitAttribution):
298
+ from logogram.direct import run_direct_effects
299
+
300
+ return run_direct_effects(spec, backend, prompts, **kwargs)
301
+ if isinstance(spec.experiment, AttributionPatching):
302
+ from logogram.atp import run_attribution_patching
303
+
304
+ return run_attribution_patching(spec, backend, prompts, **kwargs)
305
+ if isinstance(spec.experiment, Steering):
306
+ from logogram.steering import run_steering
307
+
308
+ return run_steering(spec, backend, prompts, **kwargs)
309
+ if isinstance(spec.experiment, PathPatching):
310
+ from logogram.paths import run_path_patching
311
+
312
+ return run_path_patching(spec, backend, prompts, **kwargs)
313
+ return run_engine(spec, backend, prompts, **kwargs)
314
+
315
+
316
+ def run_engine(
317
+ spec: Spec,
318
+ backend: ModelBackend,
319
+ prompts: list[PreparedPrompt],
320
+ *,
321
+ on_progress: ProgressFn | None = None,
322
+ on_layer: LayerFn | None = None,
323
+ cancel: threading.Event | None = None,
324
+ on_start: Callable[[list[ResolvedSite], dict[str, Any]], None] | None = None,
325
+ receiver_override: str | None = None,
326
+ source_override: str | None = None,
327
+ ) -> EngineResult:
328
+ """Run the sweep described by ``spec`` on already-prepared prompts.
329
+
330
+ ``receiver_override``/``source_override`` exist for sanity tests (for example patching
331
+ corrupt activations into the corrupt run, which must change nothing).
332
+ """
333
+ info = backend.info
334
+ batch_size = spec.execution.batch_size
335
+ sites, layout = expand_scope(spec, info, prompts)
336
+ if on_start is not None:
337
+ on_start(sites, layout)
338
+ groups = group_by_length(prompts)
339
+ group_of = {i: gi for gi, g in enumerate(groups) for i in g.members}
340
+ local_of = {i: li for g in groups for li, i in enumerate(g.members)}
341
+ n = len(prompts)
342
+ warnings: list[str] = []
343
+
344
+ baselines = compute_baselines(backend, prompts, groups, batch_size, cancel)
345
+
346
+ exp = spec.experiment
347
+ baseline_kind: str
348
+ if isinstance(exp, ActivationPatching):
349
+ receiver, source = (
350
+ ("corrupt", "clean") if exp.direction == "clean_to_corrupt" else ("clean", "corrupt")
351
+ )
352
+ reference = source
353
+ baseline_kind = "patch"
354
+ elif isinstance(exp, Ablation):
355
+ receiver, reference = "clean", "corrupt"
356
+ b = exp.baseline
357
+ if isinstance(b, ZeroBaseline):
358
+ source, baseline_kind = "", "zero"
359
+ elif isinstance(b, MeanBaseline):
360
+ source, baseline_kind = b.reference, "mean"
361
+ elif isinstance(b, ResampleBaseline):
362
+ source, baseline_kind = b.pool, "resample"
363
+ else: # pragma: no cover
364
+ raise EngineError(f"Unknown baseline {b!r}")
365
+ else: # pragma: no cover
366
+ raise EngineError(f"Unknown experiment {exp!r}")
367
+ receiver = receiver_override or receiver
368
+ source = source_override or source
369
+ warnings.extend(check_gap(spec, baselines, prompts, receiver, reference))
370
+
371
+ donors: list[list[int]] | None = None
372
+ k = 1
373
+ if baseline_kind == "resample":
374
+ assert isinstance(exp, Ablation) and isinstance(exp.baseline, ResampleBaseline)
375
+ k = exp.baseline.donors
376
+ same_length = any(isinstance(s.site.position, AllPositions) for s in sites)
377
+ donors = draw_donors(prompts, groups, k, exp.baseline.seed, same_length)
378
+ if baseline_kind == "mean" and any(isinstance(s.site.position, AllPositions) for s in sites):
379
+ alone = sum(1 for g in groups if len(g.members) == 1)
380
+ if alone:
381
+ warnings.append(
382
+ f"{alone} prompt(s) are the only prompt of their length, so their per-position "
383
+ "mean is their own activation."
384
+ )
385
+
386
+ patched_ld = np.zeros((len(sites), n))
387
+ patched_prob = np.zeros((len(sites), n))
388
+ # Shares the arrays being filled, so per-layer callbacks can summarize finished sites.
389
+ partial = EngineResult(
390
+ sites=sites,
391
+ layout=layout,
392
+ prompts=prompts,
393
+ baselines=baselines,
394
+ receiver=receiver,
395
+ reference=reference,
396
+ patched_ld=patched_ld,
397
+ patched_prob=patched_prob,
398
+ donors=donors,
399
+ warnings=warnings,
400
+ )
401
+ layers = sorted({s.layer for s in sites})
402
+ total_rows = len(sites) * n * k
403
+ done_rows = 0
404
+ model_dtype = {"float32": torch.float32, "float16": torch.float16, "bfloat16": torch.bfloat16}[
405
+ info.dtype
406
+ ]
407
+
408
+ for layer in layers:
409
+ layer_sites = [s for s in sites if s.layer == layer]
410
+ kinds = list(dict.fromkeys(s.kind for s in layer_sites))
411
+ sources = (
412
+ _LayerSources(backend, prompts, groups, layer, kinds, source, batch_size, cancel)
413
+ if baseline_kind != "zero"
414
+ else None
415
+ )
416
+ for gi, group in enumerate(groups):
417
+ receiver_tokens = group.clean if receiver == "clean" else group.corrupt
418
+ # Rows are grouped by (kind, all-positions?) so one hook serves the whole batch.
419
+ buckets: dict[tuple[str, bool], list[_Row]] = {}
420
+ for site in layer_sites:
421
+ all_pos = isinstance(site.site.position, AllPositions)
422
+ bucket = buckets.setdefault((site.kind, all_pos), [])
423
+ for p in group.members:
424
+ if donors is not None:
425
+ for d in donors[p]:
426
+ bucket.append(_Row(site, p, local_of[p], d))
427
+ else:
428
+ bucket.append(_Row(site, p, local_of[p]))
429
+ for (kind, all_pos), rows in buckets.items():
430
+ for sl in _chunks(len(rows), batch_size):
431
+ if cancel is not None and cancel.is_set():
432
+ raise Cancelled()
433
+ chunk = rows[sl]
434
+ patch = _build_patch(
435
+ chunk,
436
+ kind,
437
+ layer,
438
+ all_pos,
439
+ baseline_kind,
440
+ sources,
441
+ gi,
442
+ group.length,
443
+ group_of,
444
+ local_of,
445
+ prompts,
446
+ info.d_model,
447
+ info.d_head,
448
+ model_dtype,
449
+ backend.device,
450
+ )
451
+ local_idx = torch.tensor([r.local for r in chunk], dtype=torch.long)
452
+ logits = backend.final_logits(receiver_tokens[local_idx], patch)
453
+ ld, prob = _metric(logits, *_answer_tensors(prompts, [r.prompt for r in chunk]))
454
+ for r, v_ld, v_prob in zip(chunk, ld, prob, strict=True):
455
+ patched_ld[r.site.index, r.prompt] += v_ld / k
456
+ patched_prob[r.site.index, r.prompt] += v_prob / k
457
+ done_rows += len(chunk)
458
+ if on_progress is not None:
459
+ on_progress(done_rows, total_rows, layer)
460
+ if on_layer is not None:
461
+ on_layer(layer, [s.index for s in layer_sites], partial)
462
+
463
+ return partial
464
+
465
+
466
+ def _build_patch(
467
+ rows: list[_Row],
468
+ kind: str,
469
+ layer: int,
470
+ all_pos: bool,
471
+ baseline_kind: str,
472
+ sources: _LayerSources | None,
473
+ group_index: int,
474
+ length: int,
475
+ group_of: dict[int, int],
476
+ local_of: dict[int, int],
477
+ prompts: list[PreparedPrompt],
478
+ d_model: int,
479
+ d_head: int,
480
+ dtype: torch.dtype,
481
+ device: torch.device,
482
+ ) -> Patch:
483
+ B = len(rows)
484
+ is_head = kind == "head"
485
+ heads = torch.tensor([r.site.head for r in rows], dtype=torch.long) if is_head else None
486
+ positions = None
487
+ if not all_pos:
488
+ positions = torch.tensor(
489
+ [resolve_position(r.site.site.position, prompts[r.prompt]) for r in rows],
490
+ dtype=torch.long,
491
+ )
492
+ d = d_head if is_head else d_model
493
+
494
+ if baseline_kind == "zero":
495
+ shape = (B, length, d) if all_pos else (B, d)
496
+ values = torch.zeros(shape, dtype=dtype, device=device)
497
+ return Patch(kind=kind, layer=layer, values=values, heads=heads, positions=positions)
498
+
499
+ assert sources is not None
500
+ dev_heads = heads.to(device) if heads is not None else None
501
+
502
+ if baseline_kind in ("patch", "resample"):
503
+ if baseline_kind == "patch":
504
+ src_prompts = [r.prompt for r in rows]
505
+ else:
506
+ src_prompts = [r.donor for r in rows] # type: ignore[misc]
507
+ if all_pos:
508
+ # Patch pairs share a length; resample donors are drawn from the same length group.
509
+ assert all(group_of[p] == group_index for p in src_prompts)
510
+ acts = sources.group_tensor(kind, group_index)
511
+ idx = torch.tensor([local_of[p] for p in src_prompts], dtype=torch.long, device=device)
512
+ values = acts[idx, :, dev_heads] if is_head else acts[idx]
513
+ else:
514
+ table = sources.table(kind, rows[0].site)
515
+ if any(r.site.site.position != rows[0].site.site.position for r in rows):
516
+ # Rows with different positions (layer x position sweeps) gather from their own table.
517
+ values = torch.stack(
518
+ [
519
+ _table_row(sources.table(kind, r.site), p, r.site.head if is_head else None)
520
+ for r, p in zip(rows, src_prompts, strict=True)
521
+ ]
522
+ )
523
+ else:
524
+ idx = torch.tensor(src_prompts, dtype=torch.long, device=device)
525
+ values = table[idx, dev_heads] if is_head else table[idx]
526
+ return Patch(kind=kind, layer=layer, values=values, heads=heads, positions=positions)
527
+
528
+ # Mean ablation.
529
+ if all_pos:
530
+ mean = sources.group_tensor(kind, group_index).float().mean(dim=0) # [L, ...]
531
+ if is_head:
532
+ values = mean[:, dev_heads].permute(1, 0, 2) # [B, L, d_head]
533
+ else:
534
+ values = mean.unsqueeze(0).expand(B, *mean.shape)
535
+ else:
536
+ means = {}
537
+ out = []
538
+ for r in rows:
539
+ key = r.site.site.position.model_dump_json()
540
+ if key not in means:
541
+ means[key] = sources.table(kind, r.site).float().mean(dim=0)
542
+ m = means[key]
543
+ out.append(m[r.site.head] if is_head else m)
544
+ values = torch.stack(out)
545
+ return Patch(kind=kind, layer=layer, values=values, heads=heads, positions=positions)
546
+
547
+
548
+ def _table_row(table: torch.Tensor, prompt: int, head: int | None) -> torch.Tensor:
549
+ row = table[prompt]
550
+ return row[head] if head is not None else row
@@ -0,0 +1,3 @@
1
+ # Logogram caches
2
+ .logogram/
3
+ __pycache__/
@@ -0,0 +1,32 @@
1
+ {"clean": "When Kelly and Chris got to the restaurant, Chris brought a pen to", "corrupt": "When Kelly and Chris got to the restaurant, Kelly brought a pen to", "answer": " Kelly", "distractor": " Chris", "positions": {"IO": [5, 10], "S1": [15, 20], "S2": [44, 49], "end": [64, 66]}, "id": "ioi-00000", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
2
+ {"clean": "After Henry and Hannah arrived at the kitchen, Henry handed a box to", "corrupt": "After Henry and Hannah arrived at the kitchen, Hannah handed a box to", "answer": " Hannah", "distractor": " Henry", "positions": {"IO": [16, 22], "S1": [6, 11], "S2": [47, 52], "end": [66, 68]}, "id": "ioi-00001", "meta": {"template": "arrived", "pattern": "BABA", "corruption": "flip"}}
3
+ {"clean": "When Steve and Amy got to the hotel, Amy brought a ticket to", "corrupt": "When Steve and Amy got to the hotel, Steve brought a ticket to", "answer": " Steve", "distractor": " Amy", "positions": {"IO": [5, 10], "S1": [15, 18], "S2": [37, 40], "end": [58, 60]}, "id": "ioi-00002", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
4
+ {"clean": "When Owen and Noah went to the lake, Owen gave a cup to", "corrupt": "When Owen and Noah went to the lake, Noah gave a cup to", "answer": " Noah", "distractor": " Owen", "positions": {"IO": [14, 18], "S1": [5, 9], "S2": [37, 41], "end": [53, 55]}, "id": "ioi-00003", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
5
+ {"clean": "When Joshua and Zoe went to the zoo, Zoe gave a flower to", "corrupt": "When Joshua and Zoe went to the zoo, Joshua gave a flower to", "answer": " Joshua", "distractor": " Zoe", "positions": {"IO": [5, 11], "S1": [16, 19], "S2": [37, 40], "end": [55, 57]}, "id": "ioi-00004", "meta": {"template": "went", "pattern": "ABBA", "corruption": "flip"}}
6
+ {"clean": "When Victoria and Nancy went to the farm, Victoria gave a plate to", "corrupt": "When Victoria and Nancy went to the farm, Nancy gave a plate to", "answer": " Nancy", "distractor": " Victoria", "positions": {"IO": [18, 23], "S1": [5, 13], "S2": [42, 50], "end": [64, 66]}, "id": "ioi-00005", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
7
+ {"clean": "After Olivia and Sam arrived at the lake, Sam handed a hat to", "corrupt": "After Olivia and Sam arrived at the lake, Olivia handed a hat to", "answer": " Olivia", "distractor": " Sam", "positions": {"IO": [6, 12], "S1": [17, 20], "S2": [42, 45], "end": [59, 61]}, "id": "ioi-00006", "meta": {"template": "arrived", "pattern": "ABBA", "corruption": "flip"}}
8
+ {"clean": "When Jason and Joshua went to the river, Jason gave a necklace to", "corrupt": "When Jason and Joshua went to the river, Joshua gave a necklace to", "answer": " Joshua", "distractor": " Jason", "positions": {"IO": [15, 21], "S1": [5, 10], "S2": [41, 46], "end": [63, 65]}, "id": "ioi-00007", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
9
+ {"clean": "When Frank and Jake went to the office, Jake gave a card to", "corrupt": "When Frank and Jake went to the office, Frank gave a card to", "answer": " Frank", "distractor": " Jake", "positions": {"IO": [5, 10], "S1": [15, 19], "S2": [40, 44], "end": [57, 59]}, "id": "ioi-00008", "meta": {"template": "went", "pattern": "ABBA", "corruption": "flip"}}
10
+ {"clean": "When Diana and Lisa got to the office, Diana brought a plate to", "corrupt": "When Diana and Lisa got to the office, Lisa brought a plate to", "answer": " Lisa", "distractor": " Diana", "positions": {"IO": [15, 19], "S1": [5, 10], "S2": [39, 44], "end": [61, 63]}, "id": "ioi-00009", "meta": {"template": "got", "pattern": "BABA", "corruption": "flip"}}
11
+ {"clean": "When Brian and Sarah got to the market, Sarah brought a pen to", "corrupt": "When Brian and Sarah got to the market, Brian brought a pen to", "answer": " Brian", "distractor": " Sarah", "positions": {"IO": [5, 10], "S1": [15, 20], "S2": [40, 45], "end": [60, 62]}, "id": "ioi-00010", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
12
+ {"clean": "When Scott and David got to the zoo, Scott brought a hat to", "corrupt": "When Scott and David got to the zoo, David brought a hat to", "answer": " David", "distractor": " Scott", "positions": {"IO": [15, 20], "S1": [5, 10], "S2": [37, 42], "end": [57, 59]}, "id": "ioi-00011", "meta": {"template": "got", "pattern": "BABA", "corruption": "flip"}}
13
+ {"clean": "When Claire and Hugo got to the house, Hugo brought a box to", "corrupt": "When Claire and Hugo got to the house, Claire brought a box to", "answer": " Claire", "distractor": " Hugo", "positions": {"IO": [5, 11], "S1": [16, 20], "S2": [39, 43], "end": [58, 60]}, "id": "ioi-00012", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
14
+ {"clean": "After Megan and Julia arrived at the house, Megan handed a gift to", "corrupt": "After Megan and Julia arrived at the house, Julia handed a gift to", "answer": " Julia", "distractor": " Megan", "positions": {"IO": [16, 21], "S1": [6, 11], "S2": [44, 49], "end": [64, 66]}, "id": "ioi-00013", "meta": {"template": "arrived", "pattern": "BABA", "corruption": "flip"}}
15
+ {"clean": "When Lisa and Victoria went to the theater, Victoria gave a flower to", "corrupt": "When Lisa and Victoria went to the theater, Lisa gave a flower to", "answer": " Lisa", "distractor": " Victoria", "positions": {"IO": [5, 9], "S1": [14, 22], "S2": [44, 52], "end": [67, 69]}, "id": "ioi-00014", "meta": {"template": "went", "pattern": "ABBA", "corruption": "flip"}}
16
+ {"clean": "When Jake and Kelly went to the farm, Jake gave a shirt to", "corrupt": "When Jake and Kelly went to the farm, Kelly gave a shirt to", "answer": " Kelly", "distractor": " Jake", "positions": {"IO": [14, 19], "S1": [5, 9], "S2": [38, 42], "end": [56, 58]}, "id": "ioi-00015", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
17
+ {"clean": "When Sophie and Claire got to the garden, Claire brought a computer to", "corrupt": "When Sophie and Claire got to the garden, Sophie brought a computer to", "answer": " Sophie", "distractor": " Claire", "positions": {"IO": [5, 11], "S1": [16, 22], "S2": [42, 48], "end": [68, 70]}, "id": "ioi-00016", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
18
+ {"clean": "When Ivan and Fiona went to the farm, Ivan gave a box to", "corrupt": "When Ivan and Fiona went to the farm, Fiona gave a box to", "answer": " Fiona", "distractor": " Ivan", "positions": {"IO": [14, 19], "S1": [5, 9], "S2": [38, 42], "end": [54, 56]}, "id": "ioi-00017", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
19
+ {"clean": "When Tim and Brian got to the theater, Brian brought a plate to", "corrupt": "When Tim and Brian got to the theater, Tim brought a plate to", "answer": " Tim", "distractor": " Brian", "positions": {"IO": [5, 8], "S1": [13, 18], "S2": [39, 44], "end": [61, 63]}, "id": "ioi-00018", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
20
+ {"clean": "When Sarah and Carl got to the cafe, Sarah brought a flower to", "corrupt": "When Sarah and Carl got to the cafe, Carl brought a flower to", "answer": " Carl", "distractor": " Sarah", "positions": {"IO": [15, 19], "S1": [5, 10], "S2": [37, 42], "end": [60, 62]}, "id": "ioi-00019", "meta": {"template": "got", "pattern": "BABA", "corruption": "flip"}}
21
+ {"clean": "After Ivan and Eric arrived at the mall, Eric handed a bag to", "corrupt": "After Ivan and Eric arrived at the mall, Ivan handed a bag to", "answer": " Ivan", "distractor": " Eric", "positions": {"IO": [6, 10], "S1": [15, 19], "S2": [41, 45], "end": [59, 61]}, "id": "ioi-00020", "meta": {"template": "arrived", "pattern": "ABBA", "corruption": "flip"}}
22
+ {"clean": "When Emily and Carl went to the lake, Emily gave a ring to", "corrupt": "When Emily and Carl went to the lake, Carl gave a ring to", "answer": " Carl", "distractor": " Emily", "positions": {"IO": [15, 19], "S1": [5, 10], "S2": [38, 43], "end": [56, 58]}, "id": "ioi-00021", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
23
+ {"clean": "When Anthony and Bob went to the house, Bob gave a cup to", "corrupt": "When Anthony and Bob went to the house, Anthony gave a cup to", "answer": " Anthony", "distractor": " Bob", "positions": {"IO": [5, 12], "S1": [17, 20], "S2": [40, 43], "end": [55, 57]}, "id": "ioi-00022", "meta": {"template": "went", "pattern": "ABBA", "corruption": "flip"}}
24
+ {"clean": "After Alice and Charles arrived at the cafe, Alice handed a box to", "corrupt": "After Alice and Charles arrived at the cafe, Charles handed a box to", "answer": " Charles", "distractor": " Alice", "positions": {"IO": [16, 23], "S1": [6, 11], "S2": [45, 50], "end": [64, 66]}, "id": "ioi-00023", "meta": {"template": "arrived", "pattern": "BABA", "corruption": "flip"}}
25
+ {"clean": "After Jason and Ryan arrived at the club, Ryan handed a book to", "corrupt": "After Jason and Ryan arrived at the club, Jason handed a book to", "answer": " Jason", "distractor": " Ryan", "positions": {"IO": [6, 11], "S1": [16, 20], "S2": [42, 46], "end": [61, 63]}, "id": "ioi-00024", "meta": {"template": "arrived", "pattern": "ABBA", "corruption": "flip"}}
26
+ {"clean": "When Lisa and Vera went to the hospital, Lisa gave a letter to", "corrupt": "When Lisa and Vera went to the hospital, Vera gave a letter to", "answer": " Vera", "distractor": " Lisa", "positions": {"IO": [14, 18], "S1": [5, 9], "S2": [41, 45], "end": [60, 62]}, "id": "ioi-00025", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
27
+ {"clean": "When Oscar and John got to the museum, John brought a ring to", "corrupt": "When Oscar and John got to the museum, Oscar brought a ring to", "answer": " Oscar", "distractor": " John", "positions": {"IO": [5, 10], "S1": [15, 19], "S2": [39, 43], "end": [59, 61]}, "id": "ioi-00026", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
28
+ {"clean": "When Jennifer and Rachel went to the library, Jennifer gave a letter to", "corrupt": "When Jennifer and Rachel went to the library, Rachel gave a letter to", "answer": " Rachel", "distractor": " Jennifer", "positions": {"IO": [18, 24], "S1": [5, 13], "S2": [46, 54], "end": [69, 71]}, "id": "ioi-00027", "meta": {"template": "went", "pattern": "BABA", "corruption": "flip"}}
29
+ {"clean": "After Bob and David arrived at the club, David handed a letter to", "corrupt": "After Bob and David arrived at the club, Bob handed a letter to", "answer": " Bob", "distractor": " David", "positions": {"IO": [6, 9], "S1": [14, 19], "S2": [41, 46], "end": [63, 65]}, "id": "ioi-00028", "meta": {"template": "arrived", "pattern": "ABBA", "corruption": "flip"}}
30
+ {"clean": "After Nora and Joshua arrived at the farm, Nora handed a snack to", "corrupt": "After Nora and Joshua arrived at the farm, Joshua handed a snack to", "answer": " Joshua", "distractor": " Nora", "positions": {"IO": [15, 21], "S1": [6, 10], "S2": [43, 47], "end": [63, 65]}, "id": "ioi-00029", "meta": {"template": "arrived", "pattern": "BABA", "corruption": "flip"}}
31
+ {"clean": "When Hugo and Tom got to the theater, Tom brought a shirt to", "corrupt": "When Hugo and Tom got to the theater, Hugo brought a shirt to", "answer": " Hugo", "distractor": " Tom", "positions": {"IO": [5, 9], "S1": [14, 17], "S2": [38, 41], "end": [58, 60]}, "id": "ioi-00030", "meta": {"template": "got", "pattern": "ABBA", "corruption": "flip"}}
32
+ {"clean": "After Julia and Susan arrived at the airport, Julia handed a snack to", "corrupt": "After Julia and Susan arrived at the airport, Susan handed a snack to", "answer": " Susan", "distractor": " Julia", "positions": {"IO": [16, 21], "S1": [6, 11], "S2": [46, 51], "end": [67, 69]}, "id": "ioi-00031", "meta": {"template": "arrived", "pattern": "BABA", "corruption": "flip"}}
@@ -0,0 +1,42 @@
1
+ {
2
+ "logogram_spec": 1,
3
+ "name": "Which heads restore the answer?",
4
+ "notes": "Run each corrupt prompt (the subject is replaced by the indirect object, so the model should prefer the other name) and patch in one head's output from the clean prompt at every position. A normalized effect of 1 means that head alone restores the clean behavior; negative values mean the head works against the answer.",
5
+ "model": {
6
+ "id": "openai-community/gpt2",
7
+ "revision": "607a30d783dfa663caf39e06633721c8d4cfcd7e",
8
+ "dtype": "float32",
9
+ "device": "auto",
10
+ "process_weights": true
11
+ },
12
+ "dataset": {
13
+ "path": "datasets/ioi.jsonl",
14
+ "sha256": "2bd046d5d0190d42863dfa64e185f61228dd26a58458f7b6cf93140847478ac3",
15
+ "limit": null
16
+ },
17
+ "tokenization": {
18
+ "prepend_bos": true
19
+ },
20
+ "experiment": {
21
+ "kind": "activation_patching",
22
+ "direction": "clean_to_corrupt"
23
+ },
24
+ "scope": {
25
+ "kind": "heads",
26
+ "position": {
27
+ "kind": "all"
28
+ }
29
+ },
30
+ "metric": {
31
+ "kind": "logit_diff",
32
+ "normalization": "dataset_gap"
33
+ },
34
+ "statistics": {
35
+ "bootstrap": 1000,
36
+ "ci": 0.95,
37
+ "seed": 0
38
+ },
39
+ "execution": {
40
+ "batch_size": 64
41
+ }
42
+ }
@@ -0,0 +1,6 @@
1
+ {
2
+ "logogram_project": 1,
3
+ "name": "IOI example",
4
+ "description": "Indirect object identification in GPT-2 small: which attention heads carry the answer?",
5
+ "created": "2026-10-07T00:00:00Z"
6
+ }