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/prompts.py
ADDED
|
@@ -0,0 +1,204 @@
|
|
|
1
|
+
"""Tokenize prompt pairs for a model and check that they can be used in an experiment.
|
|
2
|
+
|
|
3
|
+
A usable pair has clean and corrupt prompts of the same token length (so positions line up), and
|
|
4
|
+
an answer and distractor that are each a single token.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from collections import defaultdict
|
|
10
|
+
from dataclasses import dataclass, field
|
|
11
|
+
|
|
12
|
+
import torch
|
|
13
|
+
|
|
14
|
+
from logogram.backends.base import ModelBackend, Tokenized
|
|
15
|
+
from logogram.datasets import PromptRecord
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@dataclass
|
|
19
|
+
class PromptIssue:
|
|
20
|
+
index: int
|
|
21
|
+
kind: str
|
|
22
|
+
message: str
|
|
23
|
+
|
|
24
|
+
def to_dict(self) -> dict[str, object]:
|
|
25
|
+
return {"index": self.index, "kind": self.kind, "message": self.message}
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@dataclass
|
|
29
|
+
class PreparedPrompt:
|
|
30
|
+
index: int
|
|
31
|
+
record: PromptRecord
|
|
32
|
+
clean: Tokenized
|
|
33
|
+
corrupt: Tokenized
|
|
34
|
+
answer_id: int
|
|
35
|
+
distractor_id: int
|
|
36
|
+
labels: dict[str, int] = field(default_factory=dict)
|
|
37
|
+
|
|
38
|
+
@property
|
|
39
|
+
def length(self) -> int:
|
|
40
|
+
return len(self.clean.ids)
|
|
41
|
+
|
|
42
|
+
def differing_positions(self) -> list[int]:
|
|
43
|
+
return [
|
|
44
|
+
i
|
|
45
|
+
for i, (a, b) in enumerate(zip(self.clean.ids, self.corrupt.ids, strict=True))
|
|
46
|
+
if a != b
|
|
47
|
+
]
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class PromptError(ValueError):
|
|
51
|
+
def __init__(self, issues: list[PromptIssue]):
|
|
52
|
+
self.issues = issues
|
|
53
|
+
shown = "; ".join(f"prompt {i.index}: {i.message}" for i in issues[:3])
|
|
54
|
+
more = f" (and {len(issues) - 3} more)" if len(issues) > 3 else ""
|
|
55
|
+
super().__init__(f"{len(issues)} prompt(s) can't be used. {shown}{more}")
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def _show_tokens(backend: ModelBackend, text: str) -> str:
|
|
59
|
+
pieces = backend.tokenize(text, prepend_bos=False).tokens
|
|
60
|
+
return " + ".join(repr(p) for p in pieces)
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def label_position(tokenized: Tokenized, span: tuple[int, int]) -> int | None:
|
|
64
|
+
"""The last token whose character span overlaps ``span``."""
|
|
65
|
+
start, end = span
|
|
66
|
+
found = None
|
|
67
|
+
for i, (a, b) in enumerate(tokenized.offsets):
|
|
68
|
+
if b > a and a < end and b > start:
|
|
69
|
+
found = i
|
|
70
|
+
return found
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def prepare_prompt(
|
|
74
|
+
backend: ModelBackend, record: PromptRecord, index: int, prepend_bos: bool
|
|
75
|
+
) -> tuple[PreparedPrompt | None, list[PromptIssue], Tokenized, Tokenized]:
|
|
76
|
+
issues: list[PromptIssue] = []
|
|
77
|
+
clean = backend.tokenize(record.clean, prepend_bos)
|
|
78
|
+
corrupt = backend.tokenize(record.corrupt, prepend_bos)
|
|
79
|
+
if len(clean.ids) != len(corrupt.ids):
|
|
80
|
+
issues.append(
|
|
81
|
+
PromptIssue(
|
|
82
|
+
index,
|
|
83
|
+
"length_mismatch",
|
|
84
|
+
f"clean has {len(clean.ids)} tokens and corrupt has {len(corrupt.ids)}, so "
|
|
85
|
+
"positions can't be aligned. Make both prompts tokenize to the same length.",
|
|
86
|
+
)
|
|
87
|
+
)
|
|
88
|
+
answer_id = backend.single_token_id(record.answer)
|
|
89
|
+
if answer_id is None:
|
|
90
|
+
issues.append(
|
|
91
|
+
PromptIssue(
|
|
92
|
+
index,
|
|
93
|
+
"answer_tokens",
|
|
94
|
+
f"the answer {record.answer!r} is several tokens "
|
|
95
|
+
f"({_show_tokens(backend, record.answer)}). Use a single-token answer, often "
|
|
96
|
+
"with a leading space.",
|
|
97
|
+
)
|
|
98
|
+
)
|
|
99
|
+
distractor_id = backend.single_token_id(record.distractor)
|
|
100
|
+
if distractor_id is None:
|
|
101
|
+
issues.append(
|
|
102
|
+
PromptIssue(
|
|
103
|
+
index,
|
|
104
|
+
"distractor_tokens",
|
|
105
|
+
f"the distractor {record.distractor!r} is several tokens "
|
|
106
|
+
f"({_show_tokens(backend, record.distractor)}). Use a single-token distractor.",
|
|
107
|
+
)
|
|
108
|
+
)
|
|
109
|
+
if clean.ids == corrupt.ids:
|
|
110
|
+
issues.append(
|
|
111
|
+
PromptIssue(
|
|
112
|
+
index,
|
|
113
|
+
"identical",
|
|
114
|
+
"the clean and corrupt prompts are the same, so there is nothing to patch.",
|
|
115
|
+
)
|
|
116
|
+
)
|
|
117
|
+
if answer_id is not None and answer_id == distractor_id:
|
|
118
|
+
issues.append(
|
|
119
|
+
PromptIssue(index, "same_answer", "the answer and distractor are the same token.")
|
|
120
|
+
)
|
|
121
|
+
n_ctx = backend.info.n_ctx
|
|
122
|
+
if max(len(clean.ids), len(corrupt.ids)) > n_ctx:
|
|
123
|
+
issues.append(
|
|
124
|
+
PromptIssue(index, "too_long", f"the prompt is longer than the model's {n_ctx} tokens.")
|
|
125
|
+
)
|
|
126
|
+
labels: dict[str, int] = {}
|
|
127
|
+
for label, span in (record.positions or {}).items():
|
|
128
|
+
pos = label_position(clean, span)
|
|
129
|
+
if pos is None:
|
|
130
|
+
issues.append(
|
|
131
|
+
PromptIssue(
|
|
132
|
+
index,
|
|
133
|
+
"label",
|
|
134
|
+
f"position {label!r} (characters {span[0]}–{span[1]}) doesn't cover any token.",
|
|
135
|
+
)
|
|
136
|
+
)
|
|
137
|
+
else:
|
|
138
|
+
labels[label] = pos
|
|
139
|
+
if issues:
|
|
140
|
+
return None, issues, clean, corrupt
|
|
141
|
+
assert answer_id is not None and distractor_id is not None
|
|
142
|
+
prepared = PreparedPrompt(
|
|
143
|
+
index=index,
|
|
144
|
+
record=record,
|
|
145
|
+
clean=clean,
|
|
146
|
+
corrupt=corrupt,
|
|
147
|
+
answer_id=answer_id,
|
|
148
|
+
distractor_id=distractor_id,
|
|
149
|
+
labels=labels,
|
|
150
|
+
)
|
|
151
|
+
return prepared, [], clean, corrupt
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def prepare_prompts(
|
|
155
|
+
backend: ModelBackend, records: list[PromptRecord], prepend_bos: bool
|
|
156
|
+
) -> list[PreparedPrompt]:
|
|
157
|
+
prepared: list[PreparedPrompt] = []
|
|
158
|
+
issues: list[PromptIssue] = []
|
|
159
|
+
for i, record in enumerate(records):
|
|
160
|
+
p, prompt_issues, _, _ = prepare_prompt(backend, record, i, prepend_bos)
|
|
161
|
+
issues.extend(prompt_issues)
|
|
162
|
+
if p is not None:
|
|
163
|
+
prepared.append(p)
|
|
164
|
+
if issues:
|
|
165
|
+
raise PromptError(issues)
|
|
166
|
+
return prepared
|
|
167
|
+
|
|
168
|
+
|
|
169
|
+
@dataclass
|
|
170
|
+
class LengthGroup:
|
|
171
|
+
"""Prompts of one token length. Batches never mix lengths, so no padding is needed."""
|
|
172
|
+
|
|
173
|
+
length: int
|
|
174
|
+
members: list[int] # indices into the prepared prompt list
|
|
175
|
+
clean: torch.Tensor # [n_g, length]
|
|
176
|
+
corrupt: torch.Tensor
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
def group_by_length(prompts: list[PreparedPrompt]) -> list[LengthGroup]:
|
|
180
|
+
by_len: dict[int, list[int]] = defaultdict(list)
|
|
181
|
+
for i, p in enumerate(prompts):
|
|
182
|
+
by_len[p.length].append(i)
|
|
183
|
+
groups = []
|
|
184
|
+
for length in sorted(by_len):
|
|
185
|
+
members = by_len[length]
|
|
186
|
+
groups.append(
|
|
187
|
+
LengthGroup(
|
|
188
|
+
length=length,
|
|
189
|
+
members=members,
|
|
190
|
+
clean=torch.tensor([prompts[i].clean.ids for i in members], dtype=torch.long),
|
|
191
|
+
corrupt=torch.tensor([prompts[i].corrupt.ids for i in members], dtype=torch.long),
|
|
192
|
+
)
|
|
193
|
+
)
|
|
194
|
+
return groups
|
|
195
|
+
|
|
196
|
+
|
|
197
|
+
def common_labels(prompts: list[PreparedPrompt]) -> list[str]:
|
|
198
|
+
"""Labels present in every prompt, ordered by their mean token position."""
|
|
199
|
+
if not prompts:
|
|
200
|
+
return []
|
|
201
|
+
shared = set(prompts[0].labels)
|
|
202
|
+
for p in prompts[1:]:
|
|
203
|
+
shared &= set(p.labels)
|
|
204
|
+
return sorted(shared, key=lambda lab: (sum(p.labels[lab] for p in prompts) / len(prompts), lab))
|
logogram/research.py
ADDED
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""Portable research notes and named selections, with optimistic edit revisions."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import uuid
|
|
6
|
+
from typing import Literal
|
|
7
|
+
|
|
8
|
+
from pydantic import BaseModel, ConfigDict, Field
|
|
9
|
+
|
|
10
|
+
from logogram.fileio import write_text_atomic
|
|
11
|
+
from logogram.project import Project, ProjectError, now_iso
|
|
12
|
+
from logogram.spec import ModelRef, Site
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
class NoteInput(BaseModel):
|
|
16
|
+
model_config = ConfigDict(extra="forbid")
|
|
17
|
+
|
|
18
|
+
title: str = Field(min_length=1, max_length=200)
|
|
19
|
+
body: str = Field(default="", max_length=50_000)
|
|
20
|
+
model: ModelRef
|
|
21
|
+
sites: list[Site] = Field(min_length=1, max_length=512)
|
|
22
|
+
run_id: str | None = Field(default=None, pattern=r"^[A-Za-z0-9][A-Za-z0-9._-]*$")
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class ResearchNote(NoteInput):
|
|
26
|
+
id: str
|
|
27
|
+
revision: int = Field(ge=1)
|
|
28
|
+
created: str
|
|
29
|
+
updated: str
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class Notebook(BaseModel):
|
|
33
|
+
model_config = ConfigDict(extra="forbid")
|
|
34
|
+
|
|
35
|
+
logogram_research: Literal[1] = 1
|
|
36
|
+
notes: list[ResearchNote] = Field(default_factory=list)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
class NoteConflict(ValueError):
|
|
40
|
+
pass
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def read_notebook(project: Project) -> Notebook:
|
|
44
|
+
data = project.read_json(project.root / "research.json", optional=True)
|
|
45
|
+
return Notebook.model_validate(data) if data is not None else Notebook()
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def save_note(
|
|
49
|
+
project: Project, value: NoteInput, *, note_id: str | None = None, revision: int | None = None
|
|
50
|
+
) -> ResearchNote:
|
|
51
|
+
book = read_notebook(project)
|
|
52
|
+
old = next((n for n in book.notes if n.id == note_id), None)
|
|
53
|
+
if note_id is not None and (old is None or old.revision != revision):
|
|
54
|
+
raise NoteConflict("This note changed in another tab. Reload the note before saving again.")
|
|
55
|
+
if value.run_id is not None:
|
|
56
|
+
project.readable(project.run_dir(value.run_id) / "spec.json")
|
|
57
|
+
now = now_iso()
|
|
58
|
+
note = ResearchNote(
|
|
59
|
+
**value.model_dump(),
|
|
60
|
+
id=old.id if old else uuid.uuid4().hex,
|
|
61
|
+
revision=old.revision + 1 if old else 1,
|
|
62
|
+
created=old.created if old else now,
|
|
63
|
+
updated=now,
|
|
64
|
+
)
|
|
65
|
+
book.notes = [note if n.id == note.id else n for n in book.notes]
|
|
66
|
+
if old is None:
|
|
67
|
+
book.notes.append(note)
|
|
68
|
+
write_text_atomic(
|
|
69
|
+
project.writable(project.root / "research.json"), book.model_dump_json(indent=2) + "\n"
|
|
70
|
+
)
|
|
71
|
+
return note
|
|
72
|
+
|
|
73
|
+
|
|
74
|
+
def delete_note(project: Project, note_id: str, revision: int) -> None:
|
|
75
|
+
book = read_notebook(project)
|
|
76
|
+
note = next((n for n in book.notes if n.id == note_id), None)
|
|
77
|
+
if note is None:
|
|
78
|
+
raise ProjectError("This note no longer exists. Refresh the notebook.")
|
|
79
|
+
if note.revision != revision:
|
|
80
|
+
raise NoteConflict("This note changed in another tab. Reload it before removing it.")
|
|
81
|
+
book.notes = [n for n in book.notes if n.id != note_id]
|
|
82
|
+
write_text_atomic(
|
|
83
|
+
project.writable(project.root / "research.json"), book.model_dump_json(indent=2) + "\n"
|
|
84
|
+
)
|
logogram/results.py
ADDED
|
@@ -0,0 +1,240 @@
|
|
|
1
|
+
"""Turn engine output into the files a run leaves behind: summary.json and results.parquet."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import math
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Any
|
|
8
|
+
|
|
9
|
+
import numpy as np
|
|
10
|
+
import pyarrow as pa
|
|
11
|
+
import pyarrow.parquet as pq
|
|
12
|
+
|
|
13
|
+
from logogram.engine import EngineResult
|
|
14
|
+
from logogram.fileio import atomic_output
|
|
15
|
+
from logogram.schema import Summary
|
|
16
|
+
from logogram.spec import Spec, describe_experiment
|
|
17
|
+
from logogram.stats import SiteStats, compute_site_stats, resample_counts
|
|
18
|
+
|
|
19
|
+
SUMMARY_VERSION = 1
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _f(x: Any) -> float | None:
|
|
23
|
+
"""A JSON-safe float (None for NaN or infinity)."""
|
|
24
|
+
value = float(x)
|
|
25
|
+
return value if math.isfinite(value) else None
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
def compute_stats(
|
|
29
|
+
spec: Spec,
|
|
30
|
+
result: EngineResult,
|
|
31
|
+
counts: np.ndarray | None = None,
|
|
32
|
+
indices: list[int] | None = None,
|
|
33
|
+
) -> SiteStats:
|
|
34
|
+
"""Statistics for all sites, or for ``indices`` only (rows then follow ``indices``)."""
|
|
35
|
+
n = result.patched_ld.shape[1]
|
|
36
|
+
if counts is None:
|
|
37
|
+
counts = resample_counts(n, spec.statistics.bootstrap, spec.statistics.seed)
|
|
38
|
+
rows = slice(None) if indices is None else np.asarray(indices, dtype=np.int64)
|
|
39
|
+
return compute_site_stats(
|
|
40
|
+
patched_ld=result.patched_ld[rows],
|
|
41
|
+
patched_prob=result.patched_prob[rows],
|
|
42
|
+
receiver_ld=result.receiver_ld,
|
|
43
|
+
source_ld=result.reference_ld,
|
|
44
|
+
receiver_prob=result.receiver_prob,
|
|
45
|
+
normalization=spec.metric.normalization,
|
|
46
|
+
counts=counts,
|
|
47
|
+
ci=spec.statistics.ci,
|
|
48
|
+
delta=None if result.delta is None else result.delta[rows],
|
|
49
|
+
gap=result.gap,
|
|
50
|
+
)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def site_payload(
|
|
54
|
+
result: EngineResult, stats: SiteStats, indices: list[int]
|
|
55
|
+
) -> list[dict[str, Any]]:
|
|
56
|
+
"""Payload for sites ``indices``; row ``i`` of ``stats`` belongs to ``indices[i]``."""
|
|
57
|
+
out = []
|
|
58
|
+
for i, site_index in enumerate(indices):
|
|
59
|
+
site = result.sites[site_index]
|
|
60
|
+
out.append(
|
|
61
|
+
{
|
|
62
|
+
**site.to_dict(),
|
|
63
|
+
"n": stats.n,
|
|
64
|
+
"effect": {
|
|
65
|
+
"mean": _f(stats.effect_mean[i]),
|
|
66
|
+
"sd": _f(stats.effect_sd[i]),
|
|
67
|
+
"lo": _f(stats.effect_lo[i]),
|
|
68
|
+
"hi": _f(stats.effect_hi[i]),
|
|
69
|
+
},
|
|
70
|
+
"delta": {
|
|
71
|
+
"mean": _f(stats.delta_mean[i]),
|
|
72
|
+
"sd": _f(stats.delta_sd[i]),
|
|
73
|
+
"lo": _f(stats.delta_lo[i]),
|
|
74
|
+
"hi": _f(stats.delta_hi[i]),
|
|
75
|
+
},
|
|
76
|
+
"patched_logit_diff": _f(stats.patched_mean[i]),
|
|
77
|
+
"answer_prob": _f(stats.prob_mean[i]),
|
|
78
|
+
"answer_prob_delta": _f(stats.prob_delta_mean[i]),
|
|
79
|
+
"sign_flips": int(stats.sign_flips[i]),
|
|
80
|
+
"opposite_sign": int(stats.opposite_sign[i]),
|
|
81
|
+
}
|
|
82
|
+
)
|
|
83
|
+
return out
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def _group_stats(values: np.ndarray) -> dict[str, float | None]:
|
|
87
|
+
sd = float(values.std(ddof=1)) if len(values) > 1 else 0.0
|
|
88
|
+
return {"mean": _f(values.mean()), "sd": _f(sd)}
|
|
89
|
+
|
|
90
|
+
|
|
91
|
+
def baseline_payload(result: EngineResult) -> dict[str, Any]:
|
|
92
|
+
b = result.baselines
|
|
93
|
+
gap = result.reference_ld - result.receiver_ld
|
|
94
|
+
return {
|
|
95
|
+
"clean": {
|
|
96
|
+
"logit_diff": _group_stats(b.clean_ld),
|
|
97
|
+
"answer_prob": _group_stats(b.clean_prob),
|
|
98
|
+
"prefers_answer": int((b.clean_ld > 0).sum()),
|
|
99
|
+
},
|
|
100
|
+
"corrupt": {
|
|
101
|
+
"logit_diff": _group_stats(b.corrupt_ld),
|
|
102
|
+
"answer_prob": _group_stats(b.corrupt_prob),
|
|
103
|
+
"prefers_answer": int((b.corrupt_ld > 0).sum()),
|
|
104
|
+
},
|
|
105
|
+
"gap": _group_stats(gap),
|
|
106
|
+
}
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
def build_summary(
|
|
110
|
+
spec: Spec,
|
|
111
|
+
result: EngineResult,
|
|
112
|
+
stats: SiteStats,
|
|
113
|
+
run_id: str,
|
|
114
|
+
model: dict[str, Any] | None = None,
|
|
115
|
+
) -> dict[str, Any]:
|
|
116
|
+
b = result.baselines
|
|
117
|
+
normalization = spec.metric.normalization
|
|
118
|
+
description = "logit(answer) − logit(distractor) at the last position"
|
|
119
|
+
if result.measure == "attribution":
|
|
120
|
+
description = (
|
|
121
|
+
"direct contribution to logit(answer) − logit(distractor) at the last position, with "
|
|
122
|
+
"the final normalization's scale held at its value in the run"
|
|
123
|
+
)
|
|
124
|
+
if normalization == "dataset_gap":
|
|
125
|
+
norm_text = (
|
|
126
|
+
f"contribution ÷ mean {result.receiver} logit difference; "
|
|
127
|
+
f"mean = {stats.denominator:.4f}"
|
|
128
|
+
)
|
|
129
|
+
else:
|
|
130
|
+
norm_text = f"contribution ÷ each prompt's own {result.receiver} logit difference"
|
|
131
|
+
elif normalization == "dataset_gap":
|
|
132
|
+
norm_text = (
|
|
133
|
+
f"(patched − {result.receiver}) ÷ mean({result.reference} − {result.receiver}) "
|
|
134
|
+
f"logit difference; mean gap = {stats.denominator:.4f}"
|
|
135
|
+
)
|
|
136
|
+
else:
|
|
137
|
+
norm_text = (
|
|
138
|
+
f"(patched − {result.receiver}) ÷ ({result.reference} − {result.receiver}) "
|
|
139
|
+
"logit difference, per prompt"
|
|
140
|
+
)
|
|
141
|
+
if "steering" in result.extra:
|
|
142
|
+
norm_text = norm_text.replace("patched", "steered")
|
|
143
|
+
if result.measure == "estimate":
|
|
144
|
+
description = (
|
|
145
|
+
"first-order estimate of the change patching would cause in logit(answer) − "
|
|
146
|
+
"logit(distractor) at the last position: (source − receiver activation) · its "
|
|
147
|
+
"gradient at the receiver run"
|
|
148
|
+
)
|
|
149
|
+
norm_text = f"estimated {norm_text}"
|
|
150
|
+
summary = {
|
|
151
|
+
"logogram_summary": SUMMARY_VERSION,
|
|
152
|
+
"run_id": run_id,
|
|
153
|
+
"name": spec.name,
|
|
154
|
+
"description": describe_experiment(spec),
|
|
155
|
+
"model": model,
|
|
156
|
+
"n_prompts": len(result.prompts),
|
|
157
|
+
"n_sites": len(result.sites),
|
|
158
|
+
"layout": result.layout,
|
|
159
|
+
"receiver": result.receiver,
|
|
160
|
+
"reference": result.reference,
|
|
161
|
+
"measure": result.measure,
|
|
162
|
+
"metric": {
|
|
163
|
+
"kind": spec.metric.kind,
|
|
164
|
+
"normalization": normalization,
|
|
165
|
+
"denominator": _f(stats.denominator) if normalization == "dataset_gap" else None,
|
|
166
|
+
"description": description,
|
|
167
|
+
"normalized_effect": norm_text,
|
|
168
|
+
},
|
|
169
|
+
"statistics": {
|
|
170
|
+
"bootstrap": spec.statistics.bootstrap,
|
|
171
|
+
"ci": spec.statistics.ci,
|
|
172
|
+
"seed": spec.statistics.seed,
|
|
173
|
+
"method": "percentile bootstrap over prompts",
|
|
174
|
+
},
|
|
175
|
+
"baseline": baseline_payload(result),
|
|
176
|
+
"sites": site_payload(result, stats, list(range(len(result.sites)))),
|
|
177
|
+
"per_prompt": {
|
|
178
|
+
"clean_logit_diff": [_f(x) for x in b.clean_ld],
|
|
179
|
+
"corrupt_logit_diff": [_f(x) for x in b.corrupt_ld],
|
|
180
|
+
"clean_answer_prob": [_f(x) for x in b.clean_prob],
|
|
181
|
+
"corrupt_answer_prob": [_f(x) for x in b.corrupt_prob],
|
|
182
|
+
},
|
|
183
|
+
"donors": result.donors,
|
|
184
|
+
"warnings": result.warnings,
|
|
185
|
+
**({"direct": result.extra["direct"]} if "direct" in result.extra else {}),
|
|
186
|
+
**({"steering": result.extra["steering"]} if "steering" in result.extra else {}),
|
|
187
|
+
**({"features": result.extra["features"]} if "features" in result.extra else {}),
|
|
188
|
+
}
|
|
189
|
+
return Summary.model_validate(summary).model_dump(mode="json")
|
|
190
|
+
|
|
191
|
+
|
|
192
|
+
def results_table(result: EngineResult, stats: SiteStats) -> pa.Table:
|
|
193
|
+
"""Long format: one row per (site, prompt), with everything needed to recompute effects."""
|
|
194
|
+
n_sites, n = result.patched_ld.shape
|
|
195
|
+
site_idx = np.repeat(np.arange(n_sites, dtype=np.int32), n)
|
|
196
|
+
# The prompt's index in the dataset (a method may measure only some prompts, as steering
|
|
197
|
+
# measures the held-out ones).
|
|
198
|
+
prompt_idx = np.tile(np.array([p.index for p in result.prompts], dtype=np.int32), n_sites)
|
|
199
|
+
sites = result.sites
|
|
200
|
+
columns = {
|
|
201
|
+
"site": site_idx,
|
|
202
|
+
"kind": pa.array(
|
|
203
|
+
[sites[i].kind for i in site_idx], type=pa.dictionary(pa.int8(), pa.string())
|
|
204
|
+
),
|
|
205
|
+
"layer": np.array([sites[i].layer for i in site_idx], dtype=np.int16),
|
|
206
|
+
"head": np.array(
|
|
207
|
+
[-1 if sites[i].head is None else sites[i].head for i in site_idx], dtype=np.int16
|
|
208
|
+
),
|
|
209
|
+
"feature": np.array(
|
|
210
|
+
[-1 if sites[i].site.feature is None else sites[i].site.feature for i in site_idx],
|
|
211
|
+
dtype=np.int32,
|
|
212
|
+
),
|
|
213
|
+
"position": pa.array(
|
|
214
|
+
[sites[i].position_key() for i in site_idx], type=pa.dictionary(pa.int16(), pa.string())
|
|
215
|
+
),
|
|
216
|
+
"variant": pa.array(
|
|
217
|
+
[sites[i].variant_key or "" for i in site_idx],
|
|
218
|
+
type=pa.dictionary(pa.int16(), pa.string()),
|
|
219
|
+
),
|
|
220
|
+
"prompt": prompt_idx,
|
|
221
|
+
"patched_logit_diff": result.patched_ld.ravel(),
|
|
222
|
+
"patched_answer_prob": result.patched_prob.ravel(),
|
|
223
|
+
"receiver_logit_diff": np.tile(result.receiver_ld, n_sites),
|
|
224
|
+
"reference_logit_diff": np.tile(result.reference_ld, n_sites),
|
|
225
|
+
"receiver_answer_prob": np.tile(result.receiver_prob, n_sites),
|
|
226
|
+
"delta": stats.delta.ravel(),
|
|
227
|
+
"effect": stats.effect.ravel(),
|
|
228
|
+
}
|
|
229
|
+
return pa.table(columns)
|
|
230
|
+
|
|
231
|
+
|
|
232
|
+
def write_results(path: Path, table: pa.Table) -> None:
|
|
233
|
+
with atomic_output(path) as fh:
|
|
234
|
+
pq.write_table(table, fh, compression="zstd", write_statistics=False)
|
|
235
|
+
|
|
236
|
+
|
|
237
|
+
def read_site_rows(path: Path, site: int) -> dict[str, list[Any]]:
|
|
238
|
+
table = pq.read_table(path, filters=[("site", "=", site)])
|
|
239
|
+
table = table.sort_by("prompt")
|
|
240
|
+
return {name: table.column(name).to_pylist() for name in table.column_names}
|