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/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,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
|
+
}
|