logogram 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- logogram/__init__.py +6 -0
- logogram/__main__.py +5 -0
- logogram/analysis.py +419 -0
- logogram/atp.py +120 -0
- logogram/backends/__init__.py +5 -0
- logogram/backends/base.py +202 -0
- logogram/backends/hub.py +375 -0
- logogram/backends/saes.py +277 -0
- logogram/backends/transformer_lens.py +872 -0
- logogram/cli.py +496 -0
- logogram/compare.py +177 -0
- logogram/datasets.py +159 -0
- logogram/direct.py +193 -0
- logogram/engine.py +550 -0
- logogram/examples/ioi-gpt2/.gitignore +3 -0
- logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
- logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
- logogram/examples/ioi-gpt2/project.json +6 -0
- logogram/exports.py +33 -0
- logogram/features.py +368 -0
- logogram/fileio.py +63 -0
- logogram/ioi.py +220 -0
- logogram/paths.py +204 -0
- logogram/project.py +444 -0
- logogram/prompts.py +204 -0
- logogram/research.py +84 -0
- logogram/results.py +240 -0
- logogram/runner.py +396 -0
- logogram/runs.py +98 -0
- logogram/sae.py +161 -0
- logogram/schema.py +302 -0
- logogram/server/__init__.py +1 -0
- logogram/server/app.py +1083 -0
- logogram/server/models.py +426 -0
- logogram/server/security.py +212 -0
- logogram/server/state.py +585 -0
- logogram/sites.py +249 -0
- logogram/spec.py +518 -0
- logogram/stats.py +171 -0
- logogram/steering.py +258 -0
- logogram/system.py +379 -0
- logogram/updates.py +194 -0
- logogram/verify.py +39 -0
- logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
- logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
- logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
- logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
- logogram/web_dist/favicon.svg +1 -0
- logogram/web_dist/index.html +15 -0
- logogram-0.1.0.dist-info/METADATA +550 -0
- logogram-0.1.0.dist-info/RECORD +54 -0
- logogram-0.1.0.dist-info/WHEEL +4 -0
- logogram-0.1.0.dist-info/entry_points.txt +2 -0
- logogram-0.1.0.dist-info/licenses/LICENSE +21 -0
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
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
|