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/__init__.py ADDED
@@ -0,0 +1,6 @@
1
+ """Logogram: a local-first workbench for causal experiments inside language models."""
2
+
3
+ __version__ = "0.1.0"
4
+ # The release date of this version: after a few months the app suggests looking for a newer one
5
+ # (no network needed). Set it with every release.
6
+ __released__ = "2026-10-08"
logogram/__main__.py ADDED
@@ -0,0 +1,5 @@
1
+ """python -m logogram"""
2
+
3
+ from logogram.cli import main
4
+
5
+ main()
logogram/analysis.py ADDED
@@ -0,0 +1,419 @@
1
+ """Interactive analyses: the token strip, the baseline check and attention patterns."""
2
+
3
+ from __future__ import annotations
4
+
5
+ from typing import Any
6
+
7
+ import numpy as np
8
+ import torch
9
+
10
+ from logogram.backends.base import ModelBackend
11
+ from logogram.datasets import PromptRecord
12
+ from logogram.engine import compute_baselines
13
+ from logogram.prompts import PreparedPrompt, PromptIssue, group_by_length, prepare_prompt
14
+ from logogram.spec import IndexPosition, PredictionSettings
15
+
16
+
17
+ def _f(x: float) -> float | None:
18
+ x = float(x)
19
+ return x if np.isfinite(x) else None
20
+
21
+
22
+ def tokenize_pair(backend: ModelBackend, record: PromptRecord, prepend_bos: bool) -> dict[str, Any]:
23
+ """Everything the token strip shows for one prompt pair, including what's wrong with it."""
24
+ prepared, issues, clean, corrupt = prepare_prompt(backend, record, 0, prepend_bos)
25
+ n = min(len(clean.ids), len(corrupt.ids))
26
+ differs = [i for i in range(n) if clean.ids[i] != corrupt.ids[i]]
27
+ differs += list(range(n, max(len(clean.ids), len(corrupt.ids))))
28
+
29
+ def answer_info(text: str) -> dict[str, Any]:
30
+ pieces = backend.tokenize(text, prepend_bos=False)
31
+ return {
32
+ "text": text,
33
+ "tokens": pieces.tokens,
34
+ "id": pieces.ids[0] if len(pieces.ids) == 1 else None,
35
+ }
36
+
37
+ return {
38
+ "clean": {"tokens": clean.tokens, "ids": clean.ids},
39
+ "corrupt": {"tokens": corrupt.tokens, "ids": corrupt.ids},
40
+ "aligned": len(clean.ids) == len(corrupt.ids),
41
+ "differs": differs,
42
+ "labels": prepared.labels if prepared else {},
43
+ "answer": answer_info(record.answer),
44
+ "distractor": answer_info(record.distractor),
45
+ "issues": [i.to_dict() for i in issues],
46
+ }
47
+
48
+
49
+ def prepare_with_issues(
50
+ backend: ModelBackend, records: list[PromptRecord], prepend_bos: bool
51
+ ) -> tuple[list[PreparedPrompt], list[PromptIssue]]:
52
+ prepared, issues = [], []
53
+ for i, record in enumerate(records):
54
+ p, prompt_issues, _, _ = prepare_prompt(backend, record, i, prepend_bos)
55
+ issues.extend(prompt_issues)
56
+ if p is not None:
57
+ prepared.append(p)
58
+ return prepared, issues
59
+
60
+
61
+ def baseline_report(
62
+ backend: ModelBackend,
63
+ records: list[PromptRecord],
64
+ *,
65
+ prepend_bos: bool = True,
66
+ batch_size: int = 64,
67
+ top_k: int = 5,
68
+ ) -> dict[str, Any]:
69
+ """Run clean and corrupt prompts unpatched: logit differences, probabilities, top tokens."""
70
+ prepared, issues = prepare_with_issues(backend, records, prepend_bos)
71
+ if not prepared:
72
+ return {"n": 0, "prompts": [], "issues": [i.to_dict() for i in issues], "summary": None}
73
+ groups = group_by_length(prepared)
74
+ base = compute_baselines(backend, prepared, groups, batch_size)
75
+ top: dict[str, list[list[dict[str, Any]]]] = {
76
+ "clean": [[] for _ in prepared],
77
+ "corrupt": [[] for _ in prepared],
78
+ }
79
+ for which in ("clean", "corrupt"):
80
+ for group in groups:
81
+ tokens = group.clean if which == "clean" else group.corrupt
82
+ for start in range(0, len(group.members), batch_size):
83
+ idx = group.members[start : start + batch_size]
84
+ logits = backend.final_logits(tokens[start : start + batch_size])
85
+ probs = torch.softmax(logits, dim=-1)
86
+ values, ids = probs.topk(top_k, dim=-1)
87
+ for row, p in enumerate(idx):
88
+ top[which][p] = [
89
+ {"token": backend.token_str(int(t)), "id": int(t), "prob": float(v)}
90
+ for v, t in zip(values[row].tolist(), ids[row].tolist(), strict=True)
91
+ ]
92
+ rows = []
93
+ for i, p in enumerate(prepared):
94
+ rows.append(
95
+ {
96
+ "index": p.index,
97
+ "clean": p.record.clean,
98
+ "corrupt": p.record.corrupt,
99
+ "answer": p.record.answer,
100
+ "distractor": p.record.distractor,
101
+ "clean_logit_diff": _f(base.clean_ld[i]),
102
+ "corrupt_logit_diff": _f(base.corrupt_ld[i]),
103
+ "clean_answer_prob": _f(base.clean_prob[i]),
104
+ "corrupt_answer_prob": _f(base.corrupt_prob[i]),
105
+ "clean_top": top["clean"][i],
106
+ "corrupt_top": top["corrupt"][i],
107
+ }
108
+ )
109
+ gap = base.clean_ld - base.corrupt_ld
110
+ summary = {
111
+ "clean_logit_diff": _f(base.clean_ld.mean()),
112
+ "corrupt_logit_diff": _f(base.corrupt_ld.mean()),
113
+ "gap": _f(gap.mean()),
114
+ "clean_prefers_answer": int((base.clean_ld > 0).sum()),
115
+ "corrupt_prefers_answer": int((base.corrupt_ld > 0).sum()),
116
+ "clean_answer_prob": _f(base.clean_prob.mean()),
117
+ "corrupt_answer_prob": _f(base.corrupt_prob.mean()),
118
+ }
119
+ return {
120
+ "n": len(prepared),
121
+ "prompts": rows,
122
+ "issues": [i.to_dict() for i in issues],
123
+ "summary": summary,
124
+ }
125
+
126
+
127
+ def attention_report(
128
+ backend: ModelBackend,
129
+ records: list[PromptRecord],
130
+ *,
131
+ index: int,
132
+ layer: int,
133
+ head: int,
134
+ which: str = "clean",
135
+ prepend_bos: bool = True,
136
+ batch_size: int = 64,
137
+ ) -> dict[str, Any]:
138
+ """The head's attention on one prompt, and averaged over prompts of the same token length."""
139
+ info = backend.info
140
+ if not 0 <= layer < info.n_layers or not 0 <= head < info.n_heads:
141
+ raise ValueError(f"L{layer} H{head} doesn't exist in this model.")
142
+ if not 0 <= index < len(records):
143
+ raise ValueError(f"There is no prompt {index}.")
144
+ prepared, _ = prepare_with_issues(backend, records, prepend_bos)
145
+ target = next((p for p in prepared if p.index == index), None)
146
+ if target is None:
147
+ raise ValueError(f"Prompt {index} can't be used; fix it in the dataset first.")
148
+ same = [p for p in prepared if p.length == target.length]
149
+ tokens = torch.tensor(
150
+ [(p.clean if which == "clean" else p.corrupt).ids for p in same], dtype=torch.long
151
+ )
152
+ total = None
153
+ chosen = None
154
+ for start in range(0, len(same), batch_size):
155
+ pattern = backend.attention_pattern(tokens[start : start + batch_size], layer)[:, head]
156
+ pattern = pattern.double().cpu()
157
+ total = pattern.sum(0) if total is None else total + pattern.sum(0)
158
+ for row, p in enumerate(same[start : start + batch_size]):
159
+ if p.index == index:
160
+ chosen = pattern[row]
161
+ assert total is not None and chosen is not None
162
+ average = total / len(same)
163
+ tok = (target.clean if which == "clean" else target.corrupt).tokens
164
+ # The average mixes prompts, so only labels at the same index in all of them, and tokens
165
+ # identical in all of them, describe its positions.
166
+ seqs = [(p.clean if which == "clean" else p.corrupt).tokens for p in same]
167
+ common_tokens = [t if all(s[i] == t for s in seqs) else None for i, t in enumerate(tok)]
168
+ common_labels = {
169
+ label: pos
170
+ for label, pos in target.labels.items()
171
+ if all(p.labels.get(label) == pos for p in same)
172
+ }
173
+ return {
174
+ "layer": layer,
175
+ "head": head,
176
+ "which": which,
177
+ "index": index,
178
+ "tokens": tok,
179
+ "labels": target.labels,
180
+ "pattern": chosen.tolist(),
181
+ "average": average.tolist(),
182
+ "average_tokens": common_tokens,
183
+ "average_labels": common_labels,
184
+ "n_average": len(same),
185
+ "n_total": len(prepared),
186
+ "length": target.length,
187
+ }
188
+
189
+
190
+ def prediction_report(
191
+ backend: ModelBackend,
192
+ records: list[PromptRecord],
193
+ *,
194
+ index: int,
195
+ settings: PredictionSettings,
196
+ prepend_bos: bool,
197
+ batch_size: int,
198
+ ) -> dict[str, Any]:
199
+ """A descriptive logit lens, preserving the baseline's length groups and batch shape."""
200
+ if index != settings.prompt_index:
201
+ raise ValueError("The prompt index must match the prediction settings in the spec.")
202
+ if not 0 <= index < len(records):
203
+ raise ValueError(f"There is no prompt {index}.")
204
+ prepared, issues = prepare_with_issues(backend, records, prepend_bos)
205
+ target_idx = next((i for i, p in enumerate(prepared) if p.index == index), None)
206
+ if target_idx is None:
207
+ raise ValueError(f"Prompt {index} can't be used; fix it in the dataset first.")
208
+ target = prepared[target_idx]
209
+ pos = settings.position.index if isinstance(settings.position, IndexPosition) else -1
210
+ pos = target.length + pos if pos < 0 else pos
211
+ if not 0 <= pos < target.length:
212
+ raise ValueError(f"Token position {pos} is outside this prompt ({target.length} tokens).")
213
+ group = next(g for g in group_by_length(prepared) if target_idx in g.members)
214
+ within = group.members.index(target_idx)
215
+ start = (within // batch_size) * batch_size
216
+ tokens = group.clean if settings.which == "clean" else group.corrupt
217
+ batch = tokens[start : start + batch_size]
218
+ logits = backend.layer_logits(batch, pos, within - start)
219
+ if not torch.isfinite(logits).all():
220
+ raise ValueError("The prediction diagnostic produced non-finite logits. Try float32.")
221
+ probs = torch.softmax(logits, dim=-1)
222
+ # Stable ordering makes ties deterministic, including toy or degenerate models.
223
+ top = torch.argsort(logits, dim=-1, descending=True, stable=True)[:, : settings.top_k]
224
+ answer = target.answer_id
225
+ distractor = target.distractor_id
226
+ rows = []
227
+ for layer in range(backend.info.n_layers):
228
+ rows.append(
229
+ {
230
+ "layer": layer,
231
+ "top": [
232
+ {
233
+ "id": int(t),
234
+ "token": backend.token_str(int(t)),
235
+ "prob": float(probs[layer, t]),
236
+ }
237
+ for t in top[layer]
238
+ ],
239
+ "answer_prob": float(probs[layer, answer]),
240
+ "distractor_prob": float(probs[layer, distractor]),
241
+ "logit_diff": float(logits[layer, answer] - logits[layer, distractor]),
242
+ }
243
+ )
244
+ return {
245
+ "index": index,
246
+ "position": pos,
247
+ "tokens": (target.clean if settings.which == "clean" else target.corrupt).tokens,
248
+ "settings": settings.model_dump(mode="json"),
249
+ "layers": rows,
250
+ "batch_members": [prepared[i].index for i in group.members[start : start + batch_size]],
251
+ "answer": target.record.answer,
252
+ "distractor": target.record.distractor,
253
+ "issues": [issue.to_dict() for issue in issues],
254
+ }
255
+
256
+
257
+ # -- sparse autoencoder features -------------------------------------------------------------
258
+
259
+
260
+ def _first_real_token(prepend_bos: bool) -> int:
261
+ """Where real tokens start: SAEs aren't trained on the beginning-of-sequence token, whose
262
+ activations are unlike any other, so fits and splices leave it alone."""
263
+ return 1 if prepend_bos else 0
264
+
265
+
266
+ def sae_fit_report(
267
+ backend: ModelBackend,
268
+ sae: Any,
269
+ records: list[PromptRecord],
270
+ *,
271
+ prepend_bos: bool,
272
+ batch_size: int,
273
+ ) -> dict[str, Any]:
274
+ """How well the SAE fits these prompts: the variance of the activations it explains, how many
275
+ features fire per token, and what the logit difference becomes when the model runs on the
276
+ SAE's reconstruction instead of the activation (its error removed)."""
277
+ from logogram.sae import fit_on
278
+
279
+ prepared, issues = prepare_with_issues(backend, records, prepend_bos)
280
+ if not prepared:
281
+ raise ValueError("None of these prompts can be used with the loaded model.")
282
+ start = _first_real_token(prepend_bos)
283
+ key = (sae.site, sae.layer)
284
+ acts, clean, spliced = [], [], []
285
+
286
+ def splice(x: torch.Tensor) -> torch.Tensor:
287
+ out = x.float().clone()
288
+ f, stats = sae.encode(out[:, start:])
289
+ out[:, start:] = sae.decode(f, stats)
290
+ return out
291
+
292
+ for group in group_by_length(prepared):
293
+ for begin in range(0, len(group.members), batch_size):
294
+ idx = group.members[begin : begin + batch_size]
295
+ tokens = group.clean[begin : begin + batch_size]
296
+ acts.append(backend.capture(tokens, [key])[key][:, start:].float())
297
+ answers = torch.tensor([prepared[i].answer_id for i in idx])
298
+ distractors = torch.tensor([prepared[i].distractor_id for i in idx])
299
+ rows = torch.arange(len(idx))
300
+ plain = backend.final_logits(tokens)
301
+ edited = backend.edit_logits(tokens, sae.site, sae.layer, splice)
302
+ clean.append((plain[rows, answers] - plain[rows, distractors]).double())
303
+ spliced.append((edited[rows, answers] - edited[rows, distractors]).double())
304
+ flat = torch.cat([a.reshape(-1, a.shape[-1]) for a in acts])
305
+ fit = fit_on(sae, flat)
306
+ ld, ld_spliced = torch.cat(clean), torch.cat(spliced)
307
+ fit.update(
308
+ {
309
+ "logit_diff": _f(ld.mean()),
310
+ "spliced_logit_diff": _f(ld_spliced.mean()),
311
+ "n": len(prepared),
312
+ "skipped": len(issues),
313
+ }
314
+ )
315
+ sae.fit = fit
316
+ return fit
317
+
318
+
319
+ def _single(backend: ModelBackend, records: list[PromptRecord], index: int, prepend_bos: bool): # type: ignore[no-untyped-def]
320
+ prepared, _ = prepare_with_issues(backend, records, prepend_bos)
321
+ target = next((p for p in prepared if p.index == index), None)
322
+ if target is None:
323
+ raise ValueError(f"Prompt {index} can't be used; fix it in the dataset first.")
324
+ return prepared, target
325
+
326
+
327
+ def token_features_report(
328
+ backend: ModelBackend,
329
+ sae: Any,
330
+ records: list[PromptRecord],
331
+ *,
332
+ index: int,
333
+ which: str,
334
+ prepend_bos: bool,
335
+ top_k: int = 8,
336
+ ) -> dict[str, Any]:
337
+ """The features that fire most on each token of one prompt."""
338
+ _, target = _single(backend, records, index, prepend_bos)
339
+ tokenized = target.clean if which == "clean" else target.corrupt
340
+ tokens = torch.tensor([tokenized.ids], dtype=torch.long)
341
+ key = (sae.site, sae.layer)
342
+ f, _ = sae.encode(backend.capture(tokens, [key])[key][0].float())
343
+ values, ids = f.topk(min(top_k, f.shape[-1]), dim=-1)
344
+ per_token = [
345
+ [
346
+ {"feature": int(i), "activation": float(v)}
347
+ for v, i in zip(values[p].tolist(), ids[p].tolist(), strict=True)
348
+ if v > 0
349
+ ]
350
+ for p in range(f.shape[0])
351
+ ]
352
+ return {
353
+ "index": index,
354
+ "which": which,
355
+ "tokens": tokenized.tokens,
356
+ "labels": target.labels,
357
+ "features": per_token,
358
+ "active": [int((f[p] > 0).sum()) for p in range(f.shape[0])],
359
+ "first_real_token": _first_real_token(prepend_bos),
360
+ }
361
+
362
+
363
+ def feature_report(
364
+ backend: ModelBackend,
365
+ sae: Any,
366
+ records: list[PromptRecord],
367
+ *,
368
+ feature: int,
369
+ index: int,
370
+ which: str,
371
+ prepend_bos: bool,
372
+ batch_size: int,
373
+ top: int = 10,
374
+ ) -> dict[str, Any]:
375
+ """Where one feature fires: along one prompt's tokens, and on which prompts of the dataset
376
+ most strongly (all computed here, from the project's own prompts)."""
377
+ if not 0 <= feature < sae.d_sae:
378
+ raise ValueError(f"The SAE has {sae.d_sae} features; there is no feature {feature}.")
379
+ prepared, target = _single(backend, records, index, prepend_bos)
380
+ key = (sae.site, sae.layer)
381
+ start = _first_real_token(prepend_bos)
382
+ strongest: list[dict[str, Any]] = []
383
+ along: list[float] = []
384
+ for group in group_by_length(prepared):
385
+ tokens = group.clean if which == "clean" else group.corrupt
386
+ for begin in range(0, len(group.members), batch_size):
387
+ idx = group.members[begin : begin + batch_size]
388
+ f, _ = sae.encode(
389
+ backend.capture(tokens[begin : begin + batch_size], [key])[key].float()
390
+ )
391
+ column = f[..., feature].double().cpu() # [B, pos]
392
+ for row, p in enumerate(idx):
393
+ prompt = prepared[p]
394
+ values = column[row]
395
+ if prompt.index == index:
396
+ along = [float(v) for v in values]
397
+ best = int(values[start:].argmax()) + start if values.shape[0] > start else 0
398
+ seq = prompt.clean if which == "clean" else prompt.corrupt
399
+ strongest.append(
400
+ {
401
+ "index": prompt.index,
402
+ "max": float(values[best]),
403
+ "position": best,
404
+ "token": seq.tokens[best],
405
+ "text": prompt.record.clean if which == "clean" else prompt.record.corrupt,
406
+ }
407
+ )
408
+ strongest.sort(key=lambda r: (-r["max"], r["index"]))
409
+ tokenized = target.clean if which == "clean" else target.corrupt
410
+ return {
411
+ "feature": feature,
412
+ "which": which,
413
+ "index": index,
414
+ "tokens": tokenized.tokens,
415
+ "activations": along,
416
+ "top": [r for r in strongest[:top] if r["max"] > 0],
417
+ "active_prompts": sum(1 for r in strongest if r["max"] > 0),
418
+ "n": len(strongest),
419
+ }
logogram/atp.py ADDED
@@ -0,0 +1,120 @@
1
+ """Attribution patching: estimate activation patching at every site from one gradient.
2
+
3
+ Patching a site changes the logit difference by f(receiver with the source activation) - f(receiver).
4
+ To first order that is (source activation - receiver activation) · the gradient of the logit
5
+ difference with respect to the activation, at the receiver run. One forward pass on the source
6
+ prompts and one forward and backward pass on the receiver prompts give the estimate for every site
7
+ at once, so a sweep costs a few passes per batch instead of one patched run per site and prompt.
8
+
9
+ The estimate is only first order. It misses saturation (in attention patterns, normalization and
10
+ the final softmax) and can miss or even invert an effect, which is why the results say "estimated"
11
+ and offer to verify the strongest sites with real patching.
12
+ """
13
+
14
+ from __future__ import annotations
15
+
16
+ import threading
17
+ from collections.abc import Callable
18
+ from typing import Any
19
+
20
+ import numpy as np
21
+ import torch
22
+
23
+ from logogram.backends.base import Cancelled, ModelBackend
24
+ from logogram.engine import (
25
+ EngineResult,
26
+ LayerFn,
27
+ ProgressFn,
28
+ _answer_tensors,
29
+ _chunks,
30
+ check_gap,
31
+ compute_baselines,
32
+ )
33
+ from logogram.prompts import PreparedPrompt, group_by_length
34
+ from logogram.sites import ResolvedSite, expand_scope, resolve_position
35
+ from logogram.spec import AllPositions, AttributionPatching, Spec
36
+
37
+
38
+ def run_attribution_patching(
39
+ spec: Spec,
40
+ backend: ModelBackend,
41
+ prompts: list[PreparedPrompt],
42
+ *,
43
+ on_progress: ProgressFn | None = None,
44
+ on_layer: LayerFn | None = None,
45
+ cancel: threading.Event | None = None,
46
+ on_start: Callable[[list[ResolvedSite], dict[str, Any]], None] | None = None,
47
+ receiver_override: str | None = None,
48
+ source_override: str | None = None,
49
+ ) -> EngineResult:
50
+ exp = spec.experiment
51
+ assert isinstance(exp, AttributionPatching)
52
+ info = backend.info
53
+ batch_size = spec.execution.batch_size
54
+ sites, layout = expand_scope(spec, info, prompts)
55
+ if on_start is not None:
56
+ on_start(sites, layout)
57
+ groups = group_by_length(prompts)
58
+ baselines = compute_baselines(backend, prompts, groups, batch_size, cancel)
59
+ receiver, source = (
60
+ ("corrupt", "clean") if exp.direction == "clean_to_corrupt" else ("clean", "corrupt")
61
+ )
62
+ reference = source
63
+ receiver = receiver_override or receiver
64
+ source = source_override or source
65
+ warnings = check_gap(spec, baselines, prompts, receiver, reference)
66
+
67
+ n = len(prompts)
68
+ needed = list(dict.fromkeys((rs.kind, rs.layer) for rs in sites))
69
+ estimate = np.zeros((len(sites), n))
70
+ done = 0
71
+ for group in groups:
72
+ receiver_tokens = group.clean if receiver == "clean" else group.corrupt
73
+ source_tokens = group.clean if source == "clean" else group.corrupt
74
+ for sl in _chunks(len(group.members), batch_size):
75
+ if cancel is not None and cancel.is_set():
76
+ raise Cancelled()
77
+ idx = group.members[sl]
78
+ source_acts = backend.capture(source_tokens[sl], needed)
79
+ receiver_acts, grads = backend.gradients(
80
+ receiver_tokens[sl], *_answer_tensors(prompts, idx), needed
81
+ )
82
+ rows = torch.arange(len(idx))
83
+ for rs in sites:
84
+ key = (rs.kind, rs.layer)
85
+ diff = source_acts[key].double() - receiver_acts[key].double()
86
+ grad = grads[key].double()
87
+ if rs.kind == "head":
88
+ diff, grad = diff[:, :, rs.head], grad[:, :, rs.head]
89
+ per_position = (diff * grad).sum(-1) # [B, pos]
90
+ if isinstance(rs.site.position, AllPositions):
91
+ values = per_position.sum(1)
92
+ else:
93
+ positions = torch.tensor(
94
+ [resolve_position(rs.site.position, prompts[p]) for p in idx]
95
+ )
96
+ values = per_position[rows, positions]
97
+ estimate[rs.index, idx] = values.cpu().numpy()
98
+ done += len(idx)
99
+ if on_progress is not None:
100
+ on_progress(done, n, info.n_layers - 1)
101
+
102
+ receiver_ld = baselines.ld(receiver)
103
+ result = EngineResult(
104
+ sites=sites,
105
+ layout=layout,
106
+ prompts=prompts,
107
+ baselines=baselines,
108
+ receiver=receiver,
109
+ reference=reference,
110
+ # The logit difference patching would give, to first order. Probabilities aren't estimated.
111
+ patched_ld=receiver_ld[None, :] + estimate,
112
+ patched_prob=np.full((len(sites), n), np.nan),
113
+ warnings=warnings,
114
+ measure="estimate",
115
+ delta=estimate,
116
+ )
117
+ if on_layer is not None:
118
+ for layer in sorted({rs.layer for rs in sites}):
119
+ on_layer(layer, [rs.index for rs in sites if rs.layer == layer], result)
120
+ return result
@@ -0,0 +1,5 @@
1
+ """Model backends. Experiments talk to models only through ModelBackend."""
2
+
3
+ from logogram.backends.base import BackendError, ModelBackend, ModelInfo, Patch, Tokenized
4
+
5
+ __all__ = ["BackendError", "ModelBackend", "ModelInfo", "Patch", "Tokenized"]