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/exports.py
ADDED
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
"""Portable per-prompt CSV exports of a saved result table."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import csv
|
|
6
|
+
import io
|
|
7
|
+
from collections.abc import Iterator
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
|
|
10
|
+
import pyarrow.parquet as pq
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def results_csv(path: Path) -> Iterator[str]:
|
|
14
|
+
table = pq.ParquetFile(path)
|
|
15
|
+
buffer = io.StringIO(newline="")
|
|
16
|
+
writer = csv.writer(buffer, lineterminator="\n")
|
|
17
|
+
writer.writerow(table.schema_arrow.names)
|
|
18
|
+
yield buffer.getvalue()
|
|
19
|
+
for batch in table.iter_batches(batch_size=4096):
|
|
20
|
+
buffer.seek(0)
|
|
21
|
+
buffer.truncate(0)
|
|
22
|
+
for row in batch.to_pylist():
|
|
23
|
+
# Spreadsheet programs interpret these prefixes even in quoted CSV cells.
|
|
24
|
+
writer.writerow(
|
|
25
|
+
[
|
|
26
|
+
"'" + value
|
|
27
|
+
if isinstance(value, str)
|
|
28
|
+
and value.startswith(("=", "+", "-", "@", "\t", "\r", "\n"))
|
|
29
|
+
else value
|
|
30
|
+
for value in row.values()
|
|
31
|
+
]
|
|
32
|
+
)
|
|
33
|
+
yield buffer.getvalue()
|
logogram/features.py
ADDED
|
@@ -0,0 +1,368 @@
|
|
|
1
|
+
"""SAE features as sites: patch chosen features for real, or estimate every feature at once.
|
|
2
|
+
|
|
3
|
+
Patching a feature runs the receiver prompt, encodes the activation the SAE reads, and changes
|
|
4
|
+
only that feature, to its value in the source prompt (or to zero, for zero ablation): the
|
|
5
|
+
activation moves by (new - old) · the feature's decoder direction, and the SAE's error is kept as
|
|
6
|
+
it was. Attribution patching estimates, to first order, what patching each feature would do:
|
|
7
|
+
(source - receiver feature activation) · (decoder direction · the gradient of the logit
|
|
8
|
+
difference). It covers every feature from one gradient and keeps the strongest as sites.
|
|
9
|
+
|
|
10
|
+
Each run also records how well the SAE fits these prompts (variance explained), because features
|
|
11
|
+
of an SAE that doesn't fit the loaded model describe little.
|
|
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
|
+
_metric,
|
|
31
|
+
check_gap,
|
|
32
|
+
compute_baselines,
|
|
33
|
+
)
|
|
34
|
+
from logogram.prompts import PreparedPrompt, group_by_length
|
|
35
|
+
from logogram.sae import SAE, fit_on
|
|
36
|
+
from logogram.sites import ResolvedSite, ScopeError, expand_scope, resolve_position, site_label
|
|
37
|
+
from logogram.spec import (
|
|
38
|
+
Ablation,
|
|
39
|
+
ActivationPatching,
|
|
40
|
+
AllPositions,
|
|
41
|
+
AttributionPatching,
|
|
42
|
+
FeaturesScope,
|
|
43
|
+
Site,
|
|
44
|
+
Spec,
|
|
45
|
+
ZeroBaseline,
|
|
46
|
+
)
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def run_features(
|
|
50
|
+
spec: Spec, backend: ModelBackend, prompts: list[PreparedPrompt], sae: SAE, **kwargs: Any
|
|
51
|
+
) -> EngineResult:
|
|
52
|
+
exp = spec.experiment
|
|
53
|
+
if sae.params.d_in != backend.info.d_model:
|
|
54
|
+
raise ScopeError(
|
|
55
|
+
f"This SAE reads {sae.params.d_in}-dimensional activations, but the loaded model's "
|
|
56
|
+
f"are {backend.info.d_model}-dimensional: it was made for another model."
|
|
57
|
+
)
|
|
58
|
+
if isinstance(exp, AttributionPatching):
|
|
59
|
+
return _attribution(spec, backend, prompts, sae, **kwargs)
|
|
60
|
+
if isinstance(exp, ActivationPatching) or (
|
|
61
|
+
isinstance(exp, Ablation) and isinstance(exp.baseline, ZeroBaseline)
|
|
62
|
+
):
|
|
63
|
+
if isinstance(spec.scope, FeaturesScope):
|
|
64
|
+
raise ScopeError(
|
|
65
|
+
"Patching every SAE feature for real would take a run per feature. Estimate them "
|
|
66
|
+
"all with attribution patching, then verify the strongest."
|
|
67
|
+
)
|
|
68
|
+
return _patching(spec, backend, prompts, sae, **kwargs)
|
|
69
|
+
raise ScopeError(
|
|
70
|
+
"SAE features can be patched, zero-ablated, or estimated by attribution patching."
|
|
71
|
+
)
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def _directions(
|
|
75
|
+
exp: ActivationPatching | Ablation | AttributionPatching,
|
|
76
|
+
) -> tuple[str, str, str]:
|
|
77
|
+
"""(receiver, source, reference) prompts for a method."""
|
|
78
|
+
if isinstance(exp, ActivationPatching | AttributionPatching):
|
|
79
|
+
if exp.direction == "clean_to_corrupt":
|
|
80
|
+
return "corrupt", "clean", "clean"
|
|
81
|
+
return "clean", "corrupt", "corrupt"
|
|
82
|
+
return "clean", "", "corrupt" # zero ablation: no source prompt
|
|
83
|
+
|
|
84
|
+
|
|
85
|
+
def _check_sites(sites: list[ResolvedSite], sae: SAE) -> None:
|
|
86
|
+
for rs in sites:
|
|
87
|
+
if rs.kind != "sae_feature":
|
|
88
|
+
raise ScopeError(
|
|
89
|
+
"A run measures either SAE features or model components. Put "
|
|
90
|
+
f"{rs.label} in a run of its own."
|
|
91
|
+
)
|
|
92
|
+
if rs.layer != sae.layer:
|
|
93
|
+
raise ScopeError(
|
|
94
|
+
f"{rs.label} is in layer {rs.layer}, but the SAE reads layer {sae.layer}."
|
|
95
|
+
)
|
|
96
|
+
if (rs.site.feature or 0) >= sae.d_sae:
|
|
97
|
+
raise ScopeError(f"The SAE has {sae.d_sae} features, so {rs.label} doesn't exist.")
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def _fit(sae: SAE, acts: list[torch.Tensor], skip_first: bool) -> dict[str, Any]:
|
|
101
|
+
"""The SAE's fit on the receiver prompts' activations (without the first token when it is
|
|
102
|
+
the beginning-of-sequence token, whose activations SAEs usually aren't trained on)."""
|
|
103
|
+
rows = [a[:, 1:] if skip_first and a.shape[1] > 1 else a for a in acts]
|
|
104
|
+
flat = torch.cat([r.reshape(-1, r.shape[-1]) for r in rows])
|
|
105
|
+
return fit_on(sae, flat)
|
|
106
|
+
|
|
107
|
+
|
|
108
|
+
def _patching(
|
|
109
|
+
spec: Spec,
|
|
110
|
+
backend: ModelBackend,
|
|
111
|
+
prompts: list[PreparedPrompt],
|
|
112
|
+
sae: SAE,
|
|
113
|
+
*,
|
|
114
|
+
on_progress: ProgressFn | None = None,
|
|
115
|
+
on_layer: LayerFn | None = None,
|
|
116
|
+
cancel: threading.Event | None = None,
|
|
117
|
+
on_start: Callable[[list[ResolvedSite], dict[str, Any]], None] | None = None,
|
|
118
|
+
receiver_override: str | None = None,
|
|
119
|
+
source_override: str | None = None,
|
|
120
|
+
) -> EngineResult:
|
|
121
|
+
exp = spec.experiment
|
|
122
|
+
assert isinstance(exp, ActivationPatching | Ablation)
|
|
123
|
+
batch_size = spec.execution.batch_size
|
|
124
|
+
sites, layout = expand_scope(spec, backend.info, prompts)
|
|
125
|
+
_check_sites(sites, sae)
|
|
126
|
+
if on_start is not None:
|
|
127
|
+
on_start(sites, layout)
|
|
128
|
+
groups = group_by_length(prompts)
|
|
129
|
+
baselines = compute_baselines(backend, prompts, groups, batch_size, cancel)
|
|
130
|
+
receiver, source, reference = _directions(exp)
|
|
131
|
+
receiver = receiver_override or receiver
|
|
132
|
+
source = source_override or source
|
|
133
|
+
warnings = check_gap(spec, baselines, prompts, receiver, reference)
|
|
134
|
+
key = (sae.site, sae.layer)
|
|
135
|
+
features = sorted({rs.site.feature for rs in sites if rs.site.feature is not None})
|
|
136
|
+
column = {f: j for j, f in enumerate(features)}
|
|
137
|
+
n = len(prompts)
|
|
138
|
+
patched_ld = np.zeros((len(sites), n))
|
|
139
|
+
patched_prob = np.zeros((len(sites), n))
|
|
140
|
+
result = EngineResult(
|
|
141
|
+
sites=sites,
|
|
142
|
+
layout=layout,
|
|
143
|
+
prompts=prompts,
|
|
144
|
+
baselines=baselines,
|
|
145
|
+
receiver=receiver,
|
|
146
|
+
reference=reference,
|
|
147
|
+
patched_ld=patched_ld,
|
|
148
|
+
patched_prob=patched_prob,
|
|
149
|
+
warnings=warnings,
|
|
150
|
+
)
|
|
151
|
+
receiver_acts: list[torch.Tensor] = []
|
|
152
|
+
total = len(sites) * n
|
|
153
|
+
done = 0
|
|
154
|
+
local_of = {i: li for g in groups for li, i in enumerate(g.members)}
|
|
155
|
+
for group in groups:
|
|
156
|
+
tokens = group.clean if receiver == "clean" else group.corrupt
|
|
157
|
+
# The chosen features' activations in the source prompts, at every position.
|
|
158
|
+
targets: dict[int, torch.Tensor] = {}
|
|
159
|
+
if source:
|
|
160
|
+
source_tokens = group.clean if source == "clean" else group.corrupt
|
|
161
|
+
for sl in _chunks(len(group.members), batch_size):
|
|
162
|
+
x = backend.capture(source_tokens[sl], [key])[key]
|
|
163
|
+
f, _ = sae.encode(x)
|
|
164
|
+
for row, p in enumerate(group.members[sl]):
|
|
165
|
+
targets[p] = f[row][:, features] # [pos, n_features]
|
|
166
|
+
for sl in _chunks(len(group.members), batch_size):
|
|
167
|
+
receiver_acts.append(backend.capture(tokens[sl], [key])[key].float())
|
|
168
|
+
rows = [(rs, p) for rs in sites for p in group.members]
|
|
169
|
+
for sl in _chunks(len(rows), batch_size):
|
|
170
|
+
if cancel is not None and cancel.is_set():
|
|
171
|
+
raise Cancelled()
|
|
172
|
+
chunk = rows[sl]
|
|
173
|
+
local = torch.tensor([local_of[p] for _, p in chunk], dtype=torch.long)
|
|
174
|
+
|
|
175
|
+
def edit(
|
|
176
|
+
x: torch.Tensor,
|
|
177
|
+
chunk: list[tuple[ResolvedSite, int]] = chunk,
|
|
178
|
+
targets: dict[int, torch.Tensor] = targets,
|
|
179
|
+
) -> torch.Tensor:
|
|
180
|
+
f, stats = sae.encode(x)
|
|
181
|
+
out = x.float().clone()
|
|
182
|
+
for b, (rs, p) in enumerate(chunk):
|
|
183
|
+
i = int(rs.site.feature or 0)
|
|
184
|
+
now = f[b, :, i] # [pos]
|
|
185
|
+
if source:
|
|
186
|
+
new = targets[p][:, column[i]].to(now.device)
|
|
187
|
+
else:
|
|
188
|
+
new = torch.zeros_like(now)
|
|
189
|
+
change = new - now
|
|
190
|
+
if not isinstance(rs.site.position, AllPositions):
|
|
191
|
+
at = resolve_position(rs.site.position, prompts[p])
|
|
192
|
+
keep = torch.zeros_like(change)
|
|
193
|
+
keep[at] = change[at]
|
|
194
|
+
change = keep
|
|
195
|
+
row_stats = None if stats is None else (stats[0][b], stats[1][b])
|
|
196
|
+
out[b] = out[b] + change[:, None] * sae.feature_direction(i, row_stats)
|
|
197
|
+
return out
|
|
198
|
+
|
|
199
|
+
logits = backend.edit_logits(tokens[local], sae.site, sae.layer, edit)
|
|
200
|
+
ld, prob = _metric(logits, *_answer_tensors(prompts, [p for _, p in chunk]))
|
|
201
|
+
for (rs, p), value_ld, value_prob in zip(chunk, ld, prob, strict=True):
|
|
202
|
+
patched_ld[rs.index, p] = value_ld
|
|
203
|
+
patched_prob[rs.index, p] = value_prob
|
|
204
|
+
done += len(chunk)
|
|
205
|
+
if on_progress is not None:
|
|
206
|
+
on_progress(done, total, sae.layer)
|
|
207
|
+
result.extra["features"] = {
|
|
208
|
+
"sae": sae.describe() | {"fit": None},
|
|
209
|
+
"fit": _fit(sae, receiver_acts, spec.tokenization.prepend_bos),
|
|
210
|
+
}
|
|
211
|
+
if on_layer is not None:
|
|
212
|
+
on_layer(sae.layer, [rs.index for rs in sites], result)
|
|
213
|
+
return result
|
|
214
|
+
|
|
215
|
+
|
|
216
|
+
def _attribution(
|
|
217
|
+
spec: Spec,
|
|
218
|
+
backend: ModelBackend,
|
|
219
|
+
prompts: list[PreparedPrompt],
|
|
220
|
+
sae: SAE,
|
|
221
|
+
*,
|
|
222
|
+
on_progress: ProgressFn | None = None,
|
|
223
|
+
on_layer: LayerFn | None = None,
|
|
224
|
+
cancel: threading.Event | None = None,
|
|
225
|
+
on_start: Callable[[list[ResolvedSite], dict[str, Any]], None] | None = None,
|
|
226
|
+
receiver_override: str | None = None,
|
|
227
|
+
source_override: str | None = None,
|
|
228
|
+
) -> EngineResult:
|
|
229
|
+
exp = spec.experiment
|
|
230
|
+
assert isinstance(exp, AttributionPatching)
|
|
231
|
+
batch_size = spec.execution.batch_size
|
|
232
|
+
scope = spec.scope
|
|
233
|
+
every = isinstance(scope, FeaturesScope)
|
|
234
|
+
chosen: list[ResolvedSite] = []
|
|
235
|
+
layout: dict[str, Any] = {}
|
|
236
|
+
if not every:
|
|
237
|
+
chosen, layout = expand_scope(spec, backend.info, prompts)
|
|
238
|
+
_check_sites(chosen, sae)
|
|
239
|
+
if on_start is not None:
|
|
240
|
+
on_start(chosen, layout)
|
|
241
|
+
position = scope.position if isinstance(scope, FeaturesScope) else None
|
|
242
|
+
if position is not None and not isinstance(position, AllPositions):
|
|
243
|
+
for prompt in prompts:
|
|
244
|
+
resolve_position(position, prompt)
|
|
245
|
+
groups = group_by_length(prompts)
|
|
246
|
+
baselines = compute_baselines(backend, prompts, groups, batch_size, cancel)
|
|
247
|
+
receiver, source, reference = _directions(exp)
|
|
248
|
+
receiver = receiver_override or receiver
|
|
249
|
+
source = source_override or source
|
|
250
|
+
warnings = check_gap(spec, baselines, prompts, receiver, reference)
|
|
251
|
+
key = (sae.site, sae.layer)
|
|
252
|
+
n, d_sae = len(prompts), sae.d_sae
|
|
253
|
+
W_dec = sae.params.W_dec
|
|
254
|
+
# Every feature's estimate for every prompt (or, for chosen sites, those features' at every
|
|
255
|
+
# position), and the whole site's estimate, to see how much of it the features account for.
|
|
256
|
+
estimates = np.zeros((n, d_sae)) if every else np.zeros((len(chosen), n))
|
|
257
|
+
site_total = np.zeros(n)
|
|
258
|
+
receiver_acts: list[torch.Tensor] = []
|
|
259
|
+
done = 0
|
|
260
|
+
for group in groups:
|
|
261
|
+
receiver_tokens = group.clean if receiver == "clean" else group.corrupt
|
|
262
|
+
source_tokens = group.clean if source == "clean" else group.corrupt
|
|
263
|
+
for sl in _chunks(len(group.members), batch_size):
|
|
264
|
+
if cancel is not None and cancel.is_set():
|
|
265
|
+
raise Cancelled()
|
|
266
|
+
idx = group.members[sl]
|
|
267
|
+
x_src = backend.capture(source_tokens[sl], [key])[key].float()
|
|
268
|
+
acts, grads = backend.gradients(
|
|
269
|
+
receiver_tokens[sl], *_answer_tensors(prompts, idx), [key]
|
|
270
|
+
)
|
|
271
|
+
x_rec, grad = acts[key].float(), grads[key].float()
|
|
272
|
+
receiver_acts.append(x_rec)
|
|
273
|
+
positions = range(x_rec.shape[1])
|
|
274
|
+
per_position = ((x_src - x_rec) * grad).sum(-1) # [B, pos]
|
|
275
|
+
if every:
|
|
276
|
+
rows = torch.zeros(len(idx), d_sae, dtype=torch.float64, device=x_rec.device)
|
|
277
|
+
for pos in positions:
|
|
278
|
+
if position is not None and not isinstance(position, AllPositions):
|
|
279
|
+
at = torch.tensor([resolve_position(position, prompts[p]) for p in idx])
|
|
280
|
+
mask = (at == pos).to(x_rec.device)
|
|
281
|
+
if not bool(mask.any()):
|
|
282
|
+
continue
|
|
283
|
+
else:
|
|
284
|
+
mask = None
|
|
285
|
+
f_src, _ = sae.encode(x_src[:, pos])
|
|
286
|
+
f_rec, stats = sae.encode(x_rec[:, pos])
|
|
287
|
+
g = grad[:, pos]
|
|
288
|
+
if stats is not None:
|
|
289
|
+
g = g * stats[1]
|
|
290
|
+
term = ((f_src - f_rec) * (g @ W_dec.T)).double() # [B, d_sae]
|
|
291
|
+
if mask is not None:
|
|
292
|
+
term = term * mask[:, None]
|
|
293
|
+
rows += term
|
|
294
|
+
estimates[idx] = rows.cpu().numpy()
|
|
295
|
+
if position is None or isinstance(position, AllPositions):
|
|
296
|
+
site_total[idx] = per_position.sum(1).double().cpu().numpy()
|
|
297
|
+
else:
|
|
298
|
+
at = [resolve_position(position, prompts[p]) for p in idx]
|
|
299
|
+
site_total[idx] = (
|
|
300
|
+
per_position[torch.arange(len(idx)), at].double().cpu().numpy()
|
|
301
|
+
)
|
|
302
|
+
else:
|
|
303
|
+
f_src, _ = sae.encode(x_src)
|
|
304
|
+
f_rec, stats = sae.encode(x_rec)
|
|
305
|
+
g = grad if stats is None else grad * stats[1]
|
|
306
|
+
for rs in chosen:
|
|
307
|
+
i = int(rs.site.feature or 0)
|
|
308
|
+
term = (f_src[..., i] - f_rec[..., i]) * (g @ W_dec[i]) # [B, pos]
|
|
309
|
+
if isinstance(rs.site.position, AllPositions):
|
|
310
|
+
values = term.sum(1)
|
|
311
|
+
else:
|
|
312
|
+
at = [resolve_position(rs.site.position, prompts[p]) for p in idx]
|
|
313
|
+
values = term[torch.arange(len(idx)), at]
|
|
314
|
+
estimates[rs.index, idx] = values.double().cpu().numpy()
|
|
315
|
+
site_total[idx] = per_position.sum(1).double().cpu().numpy()
|
|
316
|
+
done += len(idx)
|
|
317
|
+
if on_progress is not None:
|
|
318
|
+
on_progress(done, n, sae.layer)
|
|
319
|
+
|
|
320
|
+
if every:
|
|
321
|
+
assert position is not None
|
|
322
|
+
means = estimates.mean(0)
|
|
323
|
+
order = sorted(range(d_sae), key=lambda i: (-abs(means[i]), i))[: scope.top] # type: ignore[union-attr]
|
|
324
|
+
sites = []
|
|
325
|
+
for j, feature in enumerate(order):
|
|
326
|
+
site = Site(kind="sae_feature", layer=sae.layer, feature=feature, position=position)
|
|
327
|
+
sites.append(ResolvedSite(j, site, j, 0, site_label(site)))
|
|
328
|
+
layout = {
|
|
329
|
+
"kind": "sites",
|
|
330
|
+
"row_title": "Feature",
|
|
331
|
+
"col_title": "",
|
|
332
|
+
"rows": [{"key": str(j), "label": rs.label} for j, rs in enumerate(sites)],
|
|
333
|
+
"cols": [{"key": "effect", "label": "effect"}],
|
|
334
|
+
}
|
|
335
|
+
if on_start is not None:
|
|
336
|
+
on_start(sites, layout)
|
|
337
|
+
delta = estimates[:, order].T.copy()
|
|
338
|
+
features_sum = estimates.sum(1)
|
|
339
|
+
else:
|
|
340
|
+
sites = chosen
|
|
341
|
+
delta = estimates
|
|
342
|
+
features_sum = None
|
|
343
|
+
receiver_ld = baselines.ld(receiver)
|
|
344
|
+
result = EngineResult(
|
|
345
|
+
sites=sites,
|
|
346
|
+
layout=layout,
|
|
347
|
+
prompts=prompts,
|
|
348
|
+
baselines=baselines,
|
|
349
|
+
receiver=receiver,
|
|
350
|
+
reference=reference,
|
|
351
|
+
patched_ld=receiver_ld[None, :] + delta,
|
|
352
|
+
patched_prob=np.full(delta.shape, np.nan),
|
|
353
|
+
warnings=warnings,
|
|
354
|
+
measure="estimate",
|
|
355
|
+
delta=delta,
|
|
356
|
+
extra={
|
|
357
|
+
"features": {
|
|
358
|
+
"sae": sae.describe() | {"fit": None},
|
|
359
|
+
"fit": _fit(sae, receiver_acts, spec.tokenization.prepend_bos),
|
|
360
|
+
"site_estimate": float(site_total.mean()),
|
|
361
|
+
"features_estimate": None if features_sum is None else float(features_sum.mean()),
|
|
362
|
+
"evaluated": d_sae if every else len(chosen),
|
|
363
|
+
}
|
|
364
|
+
},
|
|
365
|
+
)
|
|
366
|
+
if on_layer is not None:
|
|
367
|
+
on_layer(sae.layer, [rs.index for rs in sites], result)
|
|
368
|
+
return result
|
logogram/fileio.py
ADDED
|
@@ -0,0 +1,63 @@
|
|
|
1
|
+
"""Small file helpers shared by everything that writes into project folders."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import contextlib
|
|
6
|
+
import os
|
|
7
|
+
import secrets
|
|
8
|
+
import time
|
|
9
|
+
from collections.abc import Iterator
|
|
10
|
+
from pathlib import Path
|
|
11
|
+
from typing import BinaryIO
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def replace_file(tmp: Path, path: Path) -> None:
|
|
15
|
+
"""Atomically move ``tmp`` over ``path``.
|
|
16
|
+
|
|
17
|
+
On Windows the move fails while another thread has ``path`` open for reading (the history
|
|
18
|
+
panel reads specs while runs write them), so retry briefly.
|
|
19
|
+
"""
|
|
20
|
+
for attempt in range(50):
|
|
21
|
+
try:
|
|
22
|
+
os.replace(tmp, path)
|
|
23
|
+
return
|
|
24
|
+
except PermissionError:
|
|
25
|
+
if attempt == 49:
|
|
26
|
+
raise
|
|
27
|
+
time.sleep(0.02)
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _create_exclusive(path: Path) -> tuple[int, Path]:
|
|
31
|
+
# O_EXCL never opens an existing file or follows a symlink, so a link planted next to the
|
|
32
|
+
# target can't redirect the write. The mode is filtered by the umask, as for any new file.
|
|
33
|
+
flags = os.O_WRONLY | os.O_CREAT | os.O_EXCL | getattr(os, "O_BINARY", 0)
|
|
34
|
+
for _ in range(100):
|
|
35
|
+
tmp = path.with_name(f".{path.name}.{secrets.token_hex(6)}.tmp")
|
|
36
|
+
try:
|
|
37
|
+
return os.open(tmp, flags, 0o666), tmp
|
|
38
|
+
except FileExistsError:
|
|
39
|
+
continue
|
|
40
|
+
raise FileExistsError(f"Can't create a temporary file next to {path.name}.")
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
@contextlib.contextmanager
|
|
44
|
+
def atomic_output(path: Path) -> Iterator[BinaryIO]:
|
|
45
|
+
"""Write ``path`` through a fresh temporary file, then move it into place.
|
|
46
|
+
|
|
47
|
+
Readers never see a half-written file, and the final move replaces ``path`` itself (a symlink
|
|
48
|
+
at ``path`` is replaced, not followed).
|
|
49
|
+
"""
|
|
50
|
+
fd, tmp = _create_exclusive(path)
|
|
51
|
+
try:
|
|
52
|
+
with os.fdopen(fd, "wb") as fh:
|
|
53
|
+
yield fh
|
|
54
|
+
replace_file(tmp, path)
|
|
55
|
+
except BaseException:
|
|
56
|
+
with contextlib.suppress(OSError):
|
|
57
|
+
tmp.unlink()
|
|
58
|
+
raise
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def write_text_atomic(path: Path, text: str) -> None:
|
|
62
|
+
with atomic_output(path) as fh:
|
|
63
|
+
fh.write(text.encode("utf-8"))
|
logogram/ioi.py
ADDED
|
@@ -0,0 +1,220 @@
|
|
|
1
|
+
"""Indirect object identification (IOI) prompt generator.
|
|
2
|
+
|
|
3
|
+
Each clean prompt names two people, repeats one of them (the subject, S), and ends where the
|
|
4
|
+
other (the indirect object, IO) is the natural next word::
|
|
5
|
+
|
|
6
|
+
ABBA: When Mary and John went to the store, John gave a drink to -> " Mary"
|
|
7
|
+
BABA: When John and Mary went to the store, John gave a drink to -> " Mary"
|
|
8
|
+
|
|
9
|
+
The answer is the IO name and the distractor is the S name. Two corruptions are available:
|
|
10
|
+
|
|
11
|
+
* ``flip``: the second mention of the subject (S2) becomes the IO name, so the model should now
|
|
12
|
+
prefer the other name. Corrupt prompts have a negative logit difference.
|
|
13
|
+
* ``abc``: all three names are replaced by unrelated names, so neither answer is supported.
|
|
14
|
+
|
|
15
|
+
Named positions (IO, S1, S2, end) are recorded as character spans in the clean prompt.
|
|
16
|
+
Randomness uses only ``random.Random.random()``, whose sequence Python guarantees across versions.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import random
|
|
22
|
+
from collections.abc import Callable, Sequence
|
|
23
|
+
from dataclasses import dataclass
|
|
24
|
+
from typing import Literal
|
|
25
|
+
|
|
26
|
+
from logogram.datasets import PromptRecord
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
@dataclass(frozen=True)
|
|
30
|
+
class IOITemplate:
|
|
31
|
+
id: str
|
|
32
|
+
text: str
|
|
33
|
+
default: bool = False
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
# The default templates share one token structure, so named positions sit at the same index.
|
|
37
|
+
TEMPLATES: tuple[IOITemplate, ...] = (
|
|
38
|
+
IOITemplate("went", "When {A} and {B} went to the {PLACE}, {C} gave a {OBJECT} to", True),
|
|
39
|
+
IOITemplate(
|
|
40
|
+
"arrived", "After {A} and {B} arrived at the {PLACE}, {C} handed a {OBJECT} to", True
|
|
41
|
+
),
|
|
42
|
+
IOITemplate("got", "When {A} and {B} got to the {PLACE}, {C} brought a {OBJECT} to", True),
|
|
43
|
+
IOITemplate(
|
|
44
|
+
"working", "While {A} and {B} were working at the {PLACE}, {C} passed a {OBJECT} to"
|
|
45
|
+
),
|
|
46
|
+
IOITemplate(
|
|
47
|
+
"walked", "Then {A} and {B} walked into the {PLACE}, and {C} offered a {OBJECT} to"
|
|
48
|
+
),
|
|
49
|
+
IOITemplate("reached", "Once {A} and {B} reached the {PLACE}, {C} threw a {OBJECT} to"),
|
|
50
|
+
IOITemplate("left", "As {A} and {B} left the {PLACE}, {C} gave the {OBJECT} to"),
|
|
51
|
+
IOITemplate("met", "Later, {A} and {B} met at the {PLACE}, where {C} gave a {OBJECT} to"),
|
|
52
|
+
)
|
|
53
|
+
|
|
54
|
+
NAMES: tuple[str, ...] = (
|
|
55
|
+
"Mary", "John", "Alice", "Bob", "Tom", "James", "Sarah", "Emma", "David", "Michael",
|
|
56
|
+
"Daniel", "Laura", "Anna", "Kate", "Paul", "Lisa", "Peter", "Jack", "Rachel", "Eric",
|
|
57
|
+
"Susan", "Sam", "Ben", "Amy", "Henry", "Lucy", "Adam", "Emily", "Kevin", "Linda", "Ryan",
|
|
58
|
+
"Megan", "Jason", "Helen", "Chris", "Nancy", "Brian", "Karen", "George", "Ruth", "Frank",
|
|
59
|
+
"Steve", "Jane", "Matt", "Claire", "Scott", "Diana", "Tim", "Julia", "Andrew", "Victoria",
|
|
60
|
+
"Robert", "Joseph", "Jennifer", "Thomas", "Charles", "Jessica", "Anthony", "Nicole",
|
|
61
|
+
"Joshua", "Amanda", "Justin", "Kelly", "Sean", "Hannah", "Carl", "Fiona", "Leo", "Nora",
|
|
62
|
+
"Ivan", "Olivia", "Jake", "Luke", "Zoe", "Noah", "Sophie", "Oscar", "Isaac", "Vera", "Hugo",
|
|
63
|
+
"Owen", "Alex",
|
|
64
|
+
) # fmt: skip
|
|
65
|
+
|
|
66
|
+
PLACES: tuple[str, ...] = (
|
|
67
|
+
"store", "park", "school", "hospital", "station", "beach", "office", "restaurant", "market",
|
|
68
|
+
"library", "garden", "museum", "airport", "kitchen", "church", "bank", "house", "cafe",
|
|
69
|
+
"theater", "zoo", "mall", "hotel", "lake", "river", "forest", "farm", "gym", "club",
|
|
70
|
+
) # fmt: skip
|
|
71
|
+
|
|
72
|
+
OBJECTS: tuple[str, ...] = (
|
|
73
|
+
"drink", "book", "ring", "bag", "snack", "letter", "basket", "bottle", "gift", "card", "key",
|
|
74
|
+
"ball", "cake", "flower", "pen", "hat", "box", "ticket", "necklace", "computer", "phone",
|
|
75
|
+
"coat", "bowl", "shirt", "toy", "cup", "map", "plate",
|
|
76
|
+
) # fmt: skip
|
|
77
|
+
|
|
78
|
+
Pattern = Literal["ABBA", "BABA"]
|
|
79
|
+
Corruption = Literal["flip", "abc"]
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
class IOIError(ValueError):
|
|
83
|
+
pass
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def default_template_ids() -> list[str]:
|
|
87
|
+
return [t.id for t in TEMPLATES if t.default]
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def _pick(rng: random.Random, items: Sequence[str]) -> str:
|
|
91
|
+
return items[int(rng.random() * len(items))]
|
|
92
|
+
|
|
93
|
+
|
|
94
|
+
def _pick_distinct(rng: random.Random, items: Sequence[str], k: int, avoid: set[str]) -> list[str]:
|
|
95
|
+
pool = [x for x in items if x not in avoid]
|
|
96
|
+
if len(pool) < k:
|
|
97
|
+
raise IOIError(f"Need at least {k} usable names but only {len(pool)} are available.")
|
|
98
|
+
chosen: list[str] = []
|
|
99
|
+
while len(chosen) < k:
|
|
100
|
+
candidate = _pick(rng, pool)
|
|
101
|
+
if candidate not in chosen:
|
|
102
|
+
chosen.append(candidate)
|
|
103
|
+
return chosen
|
|
104
|
+
|
|
105
|
+
|
|
106
|
+
def _fill(template: str, values: dict[str, str]) -> tuple[str, dict[str, tuple[int, int]]]:
|
|
107
|
+
"""Fill ``{A}``-style slots left to right, returning the text and each slot's span."""
|
|
108
|
+
out: list[str] = []
|
|
109
|
+
spans: dict[str, tuple[int, int]] = {}
|
|
110
|
+
cursor = 0
|
|
111
|
+
length = 0
|
|
112
|
+
while cursor < len(template):
|
|
113
|
+
start = template.find("{", cursor)
|
|
114
|
+
if start == -1:
|
|
115
|
+
out.append(template[cursor:])
|
|
116
|
+
length += len(template) - cursor
|
|
117
|
+
break
|
|
118
|
+
out.append(template[cursor:start])
|
|
119
|
+
length += start - cursor
|
|
120
|
+
end = template.index("}", start)
|
|
121
|
+
slot = template[start + 1 : end]
|
|
122
|
+
value = values[slot]
|
|
123
|
+
spans[slot] = (length, length + len(value))
|
|
124
|
+
out.append(value)
|
|
125
|
+
length += len(value)
|
|
126
|
+
cursor = end + 1
|
|
127
|
+
return "".join(out), spans
|
|
128
|
+
|
|
129
|
+
|
|
130
|
+
def generate_ioi(
|
|
131
|
+
n: int,
|
|
132
|
+
seed: int = 0,
|
|
133
|
+
templates: Sequence[str] | None = None,
|
|
134
|
+
patterns: Sequence[Pattern] = ("ABBA", "BABA"),
|
|
135
|
+
corruption: Corruption = "flip",
|
|
136
|
+
single_token: Callable[[str], bool] | None = None,
|
|
137
|
+
) -> list[PromptRecord]:
|
|
138
|
+
"""Generate ``n`` IOI prompt pairs.
|
|
139
|
+
|
|
140
|
+
``single_token``, if given, filters names, places and objects to words that are a single
|
|
141
|
+
token (with a leading space) for the model that will be used.
|
|
142
|
+
"""
|
|
143
|
+
if n < 1:
|
|
144
|
+
raise IOIError("Choose at least one prompt.")
|
|
145
|
+
if n > 100_000:
|
|
146
|
+
raise IOIError("Choose at most 100,000 prompts.")
|
|
147
|
+
chosen_ids = list(templates) if templates else default_template_ids()
|
|
148
|
+
by_id = {t.id: t for t in TEMPLATES}
|
|
149
|
+
unknown = [t for t in chosen_ids if t not in by_id]
|
|
150
|
+
if unknown:
|
|
151
|
+
raise IOIError(f"Unknown template(s): {', '.join(unknown)}.")
|
|
152
|
+
if not patterns:
|
|
153
|
+
raise IOIError("Choose at least one of ABBA and BABA.")
|
|
154
|
+
for p in patterns:
|
|
155
|
+
if p not in ("ABBA", "BABA"):
|
|
156
|
+
raise IOIError(f"Unknown pattern {p!r}; use ABBA or BABA.")
|
|
157
|
+
if corruption not in ("flip", "abc"):
|
|
158
|
+
raise IOIError(f"Unknown corruption {corruption!r}; use flip or abc.")
|
|
159
|
+
|
|
160
|
+
def usable(words: Sequence[str]) -> list[str]:
|
|
161
|
+
if single_token is None:
|
|
162
|
+
return list(words)
|
|
163
|
+
return [w for w in words if single_token(" " + w)]
|
|
164
|
+
|
|
165
|
+
names, places, objects = usable(NAMES), usable(PLACES), usable(OBJECTS)
|
|
166
|
+
if len(names) < 5 or not places or not objects:
|
|
167
|
+
raise IOIError(
|
|
168
|
+
"This model's tokenizer splits too many of the built-in names, places or objects "
|
|
169
|
+
"into several tokens. Import a JSONL dataset written for this model instead."
|
|
170
|
+
)
|
|
171
|
+
|
|
172
|
+
rng = random.Random(seed)
|
|
173
|
+
records: list[PromptRecord] = []
|
|
174
|
+
seen: set[str] = set()
|
|
175
|
+
attempts = 0
|
|
176
|
+
while len(records) < n:
|
|
177
|
+
attempts += 1
|
|
178
|
+
if attempts > n * 50:
|
|
179
|
+
raise IOIError(
|
|
180
|
+
f"Could only make {len(records)} distinct prompts with these settings. "
|
|
181
|
+
"Choose more templates or fewer prompts."
|
|
182
|
+
)
|
|
183
|
+
index = len(records)
|
|
184
|
+
template = by_id[chosen_ids[int(rng.random() * len(chosen_ids))]]
|
|
185
|
+
pattern = patterns[index % len(patterns)]
|
|
186
|
+
io, s = _pick_distinct(rng, names, 2, set())
|
|
187
|
+
place, obj = _pick(rng, places), _pick(rng, objects)
|
|
188
|
+
a, b = (io, s) if pattern == "ABBA" else (s, io)
|
|
189
|
+
clean, spans = _fill(template.text, {"A": a, "B": b, "C": s, "PLACE": place, "OBJECT": obj})
|
|
190
|
+
if clean in seen:
|
|
191
|
+
continue
|
|
192
|
+
if corruption == "flip":
|
|
193
|
+
corrupt, _ = _fill(
|
|
194
|
+
template.text, {"A": a, "B": b, "C": io, "PLACE": place, "OBJECT": obj}
|
|
195
|
+
)
|
|
196
|
+
else:
|
|
197
|
+
x, y, z = _pick_distinct(rng, names, 3, {io, s})
|
|
198
|
+
corrupt, _ = _fill(
|
|
199
|
+
template.text, {"A": x, "B": y, "C": z, "PLACE": place, "OBJECT": obj}
|
|
200
|
+
)
|
|
201
|
+
seen.add(clean)
|
|
202
|
+
io_slot, s1_slot = ("A", "B") if pattern == "ABBA" else ("B", "A")
|
|
203
|
+
end_start = clean.rfind(" ") + 1
|
|
204
|
+
records.append(
|
|
205
|
+
PromptRecord(
|
|
206
|
+
clean=clean,
|
|
207
|
+
corrupt=corrupt,
|
|
208
|
+
answer=" " + io,
|
|
209
|
+
distractor=" " + s,
|
|
210
|
+
positions={
|
|
211
|
+
"IO": spans[io_slot],
|
|
212
|
+
"S1": spans[s1_slot],
|
|
213
|
+
"S2": spans["C"],
|
|
214
|
+
"end": (end_start, len(clean)),
|
|
215
|
+
},
|
|
216
|
+
id=f"ioi-{index:05d}",
|
|
217
|
+
meta={"template": template.id, "pattern": pattern, "corruption": corruption},
|
|
218
|
+
)
|
|
219
|
+
)
|
|
220
|
+
return records
|