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/datasets.py
ADDED
|
@@ -0,0 +1,159 @@
|
|
|
1
|
+
"""Prompt datasets: JSONL files of clean/corrupt pairs with an answer and a distractor.
|
|
2
|
+
|
|
3
|
+
Each line is one JSON object::
|
|
4
|
+
|
|
5
|
+
{"clean": "When Mary and John went to the store, John gave a drink to",
|
|
6
|
+
"corrupt": "When Mary and John went to the store, Mary gave a drink to",
|
|
7
|
+
"answer": " Mary", "distractor": " John",
|
|
8
|
+
"positions": {"IO": [5, 9], "S1": [14, 18], "S2": [38, 42]}}
|
|
9
|
+
|
|
10
|
+
``positions`` is optional. It names character spans in the clean prompt; a label points to the
|
|
11
|
+
last token that overlaps its span. ``id`` and ``meta`` are optional and carried through.
|
|
12
|
+
"""
|
|
13
|
+
|
|
14
|
+
from __future__ import annotations
|
|
15
|
+
|
|
16
|
+
import hashlib
|
|
17
|
+
import json
|
|
18
|
+
import re
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
from typing import Any
|
|
21
|
+
|
|
22
|
+
from pydantic import BaseModel, ConfigDict, Field, ValidationError, field_validator
|
|
23
|
+
|
|
24
|
+
from logogram.fileio import write_text_atomic
|
|
25
|
+
|
|
26
|
+
DATASET_NAME_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,99}$")
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
class DatasetError(ValueError):
|
|
30
|
+
"""A dataset file can't be read. The message says where and how to fix it."""
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
class PromptRecord(BaseModel):
|
|
34
|
+
model_config = ConfigDict(extra="forbid")
|
|
35
|
+
|
|
36
|
+
clean: str = Field(min_length=1)
|
|
37
|
+
corrupt: str = Field(min_length=1)
|
|
38
|
+
answer: str = Field(min_length=1)
|
|
39
|
+
distractor: str = Field(min_length=1)
|
|
40
|
+
positions: dict[str, tuple[int, int]] | None = None
|
|
41
|
+
id: str | None = None
|
|
42
|
+
meta: dict[str, Any] | None = None
|
|
43
|
+
|
|
44
|
+
@field_validator("positions")
|
|
45
|
+
@classmethod
|
|
46
|
+
def _spans_ordered(
|
|
47
|
+
cls, value: dict[str, tuple[int, int]] | None
|
|
48
|
+
) -> dict[str, tuple[int, int]] | None:
|
|
49
|
+
if value is None:
|
|
50
|
+
return value
|
|
51
|
+
for label, (start, end) in value.items():
|
|
52
|
+
if not label or not label.strip():
|
|
53
|
+
raise ValueError("position labels must be non-empty")
|
|
54
|
+
if label in ("all", "last") or re.fullmatch(r"-?\d+", label):
|
|
55
|
+
raise ValueError(
|
|
56
|
+
f"{label!r} can't be a position label: 'all', 'last' and numbers already "
|
|
57
|
+
"name positions"
|
|
58
|
+
)
|
|
59
|
+
if start < 0 or end <= start:
|
|
60
|
+
raise ValueError(f"position {label!r} must be a span [start, end) with start < end")
|
|
61
|
+
return value
|
|
62
|
+
|
|
63
|
+
|
|
64
|
+
def parse_jsonl(text: str, *, source: str = "dataset") -> list[PromptRecord]:
|
|
65
|
+
records: list[PromptRecord] = []
|
|
66
|
+
for line_no, line in enumerate(text.splitlines(), start=1):
|
|
67
|
+
if not line.strip():
|
|
68
|
+
continue
|
|
69
|
+
try:
|
|
70
|
+
raw = json.loads(line)
|
|
71
|
+
except json.JSONDecodeError as exc:
|
|
72
|
+
raise DatasetError(
|
|
73
|
+
f"{source}, line {line_no}: not valid JSON ({exc.msg}). "
|
|
74
|
+
"Each line must be one JSON object."
|
|
75
|
+
) from exc
|
|
76
|
+
if not isinstance(raw, dict):
|
|
77
|
+
raise DatasetError(f"{source}, line {line_no}: expected a JSON object.")
|
|
78
|
+
missing = [k for k in ("clean", "corrupt", "answer", "distractor") if k not in raw]
|
|
79
|
+
if missing:
|
|
80
|
+
raise DatasetError(
|
|
81
|
+
f"{source}, line {line_no}: missing {', '.join(repr(m) for m in missing)}. "
|
|
82
|
+
"Each line needs clean, corrupt, answer and distractor."
|
|
83
|
+
)
|
|
84
|
+
if isinstance(raw.get("positions"), dict):
|
|
85
|
+
raw["positions"] = {
|
|
86
|
+
k: tuple(v) if isinstance(v, list) else v for k, v in raw["positions"].items()
|
|
87
|
+
}
|
|
88
|
+
try:
|
|
89
|
+
record = PromptRecord.model_validate(raw)
|
|
90
|
+
except ValidationError as exc:
|
|
91
|
+
first = exc.errors()[0]
|
|
92
|
+
where = ".".join(str(p) for p in first["loc"])
|
|
93
|
+
raise DatasetError(f"{source}, line {line_no}: {where}: {first['msg']}.") from exc
|
|
94
|
+
_check_spans(record, line_no, source)
|
|
95
|
+
records.append(record)
|
|
96
|
+
if not records:
|
|
97
|
+
raise DatasetError(f"{source} has no prompts. Add one JSON object per line.")
|
|
98
|
+
return records
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def _check_spans(record: PromptRecord, line_no: int, source: str) -> None:
|
|
102
|
+
if not record.positions:
|
|
103
|
+
return
|
|
104
|
+
for label, (_start, end) in record.positions.items():
|
|
105
|
+
if end > len(record.clean):
|
|
106
|
+
raise DatasetError(
|
|
107
|
+
f"{source}, line {line_no}: position {label!r} ends at character {end}, but "
|
|
108
|
+
f"the clean prompt has {len(record.clean)} characters."
|
|
109
|
+
)
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def load_dataset(path: Path) -> list[PromptRecord]:
|
|
113
|
+
if not path.exists():
|
|
114
|
+
raise DatasetError(f"Dataset not found: {path.name}. Check the path in the spec.")
|
|
115
|
+
if not path.is_file(): # a folder, device or pipe could fail, block or never end
|
|
116
|
+
raise DatasetError(f"{path.name} isn't a regular file.")
|
|
117
|
+
try:
|
|
118
|
+
text = path.read_text(encoding="utf-8")
|
|
119
|
+
except UnicodeDecodeError as exc:
|
|
120
|
+
raise DatasetError(
|
|
121
|
+
f"{path.name} isn't UTF-8 text (byte {exc.start} can't be decoded). Save it as UTF-8."
|
|
122
|
+
) from exc
|
|
123
|
+
except OSError as exc:
|
|
124
|
+
raise DatasetError(f"{path.name} can't be read: {exc.strerror or exc}.") from exc
|
|
125
|
+
return parse_jsonl(text, source=path.name)
|
|
126
|
+
|
|
127
|
+
|
|
128
|
+
def file_sha256(path: Path) -> str:
|
|
129
|
+
digest = hashlib.sha256()
|
|
130
|
+
with path.open("rb") as fh:
|
|
131
|
+
for chunk in iter(lambda: fh.read(1 << 20), b""):
|
|
132
|
+
digest.update(chunk)
|
|
133
|
+
return digest.hexdigest()
|
|
134
|
+
|
|
135
|
+
|
|
136
|
+
def to_jsonl(records: list[PromptRecord]) -> str:
|
|
137
|
+
lines = []
|
|
138
|
+
for record in records:
|
|
139
|
+
data = record.model_dump(mode="json", exclude_none=True)
|
|
140
|
+
if "positions" in data:
|
|
141
|
+
data["positions"] = {k: list(v) for k, v in data["positions"].items()}
|
|
142
|
+
lines.append(json.dumps(data, ensure_ascii=False))
|
|
143
|
+
return "\n".join(lines) + "\n"
|
|
144
|
+
|
|
145
|
+
|
|
146
|
+
def write_dataset(path: Path, records: list[PromptRecord]) -> None:
|
|
147
|
+
"""Write a dataset file. Callers in a project get ``path`` from ``Project.dataset_file``."""
|
|
148
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
149
|
+
write_text_atomic(path, to_jsonl(records))
|
|
150
|
+
|
|
151
|
+
|
|
152
|
+
def check_dataset_name(name: str) -> str:
|
|
153
|
+
stem = name[:-6] if name.endswith(".jsonl") else name
|
|
154
|
+
if not DATASET_NAME_RE.match(stem):
|
|
155
|
+
raise DatasetError(
|
|
156
|
+
"Dataset names can use letters, digits, '.', '_' and '-', and must start with a "
|
|
157
|
+
"letter or digit."
|
|
158
|
+
)
|
|
159
|
+
return stem + ".jsonl"
|
logogram/direct.py
ADDED
|
@@ -0,0 +1,193 @@
|
|
|
1
|
+
"""Direct logit attribution: what each component writes straight into the logit difference.
|
|
2
|
+
|
|
3
|
+
At the last position the residual stream is the embeddings plus every attention and MLP output.
|
|
4
|
+
With the final normalization's scale held at its value in the run, the logit difference is an
|
|
5
|
+
affine function of that sum, so it splits into one term per component plus a constant from biases.
|
|
6
|
+
A component's term is its *direct* effect: it leaves out everything the component does through
|
|
7
|
+
later components, which patching measures. The backend computes the terms; this module checks that
|
|
8
|
+
the split is valid for the loaded model and the chosen sites, and turns the terms into a result.
|
|
9
|
+
"""
|
|
10
|
+
|
|
11
|
+
from __future__ import annotations
|
|
12
|
+
|
|
13
|
+
import threading
|
|
14
|
+
from collections.abc import Callable
|
|
15
|
+
from typing import Any
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
from logogram.backends.base import Cancelled, ModelBackend, ModelInfo
|
|
20
|
+
from logogram.engine import (
|
|
21
|
+
EngineError,
|
|
22
|
+
EngineResult,
|
|
23
|
+
LayerFn,
|
|
24
|
+
ProgressFn,
|
|
25
|
+
_answer_tensors,
|
|
26
|
+
_chunks,
|
|
27
|
+
compute_baselines,
|
|
28
|
+
)
|
|
29
|
+
from logogram.prompts import PreparedPrompt, group_by_length
|
|
30
|
+
from logogram.sites import ResolvedSite, ScopeError, expand_scope, resolve_position
|
|
31
|
+
from logogram.spec import DirectLogitAttribution, Spec
|
|
32
|
+
|
|
33
|
+
DIRECT_KINDS = ("head", "attn_out", "mlp_out")
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def check_direct_sites(
|
|
37
|
+
sites: list[ResolvedSite], info: ModelInfo, prompts: list[PreparedPrompt]
|
|
38
|
+
) -> None:
|
|
39
|
+
"""Refuse sites, or models, for which a direct effect isn't defined."""
|
|
40
|
+
extra = info.extra
|
|
41
|
+
if not str(extra.get("block_structure", "")).startswith(("sequential", "parallel")):
|
|
42
|
+
raise ScopeError(
|
|
43
|
+
"When this model loaded, its residual stream couldn't be verified to be the sum of its "
|
|
44
|
+
"attention and MLP outputs, so the logit difference can't be split among them."
|
|
45
|
+
)
|
|
46
|
+
if extra.get("logit_soft_cap"):
|
|
47
|
+
raise ScopeError(
|
|
48
|
+
"This model soft-caps its logits (tanh), so the logit difference isn't a sum of "
|
|
49
|
+
"direct effects. Use activation patching for this model."
|
|
50
|
+
)
|
|
51
|
+
for rs in sites:
|
|
52
|
+
if rs.kind not in DIRECT_KINDS:
|
|
53
|
+
raise ScopeError(
|
|
54
|
+
f"{rs.label} is a state of the residual stream, not something written into it. "
|
|
55
|
+
"Direct logit attribution splits the logit difference among heads, attention "
|
|
56
|
+
"outputs and MLP outputs; choose those."
|
|
57
|
+
)
|
|
58
|
+
if any(rs.kind == "head" for rs in sites):
|
|
59
|
+
checks = extra.get("checks") or {}
|
|
60
|
+
heads = checks.get("heads")
|
|
61
|
+
if heads is None or heads > checks.get("tolerance", 0):
|
|
62
|
+
raise ScopeError(
|
|
63
|
+
"In this model the attention output isn't the sum of its heads' outputs (it is "
|
|
64
|
+
"normalized after they are combined), so a single head has no direct effect of its "
|
|
65
|
+
"own. Use attention and MLP outputs per layer."
|
|
66
|
+
)
|
|
67
|
+
checked: set[str] = set()
|
|
68
|
+
for rs in sites:
|
|
69
|
+
key = rs.site.position.model_dump_json()
|
|
70
|
+
if key in checked:
|
|
71
|
+
continue
|
|
72
|
+
checked.add(key)
|
|
73
|
+
for prompt in prompts:
|
|
74
|
+
if resolve_position(rs.site.position, prompt) != prompt.length - 1:
|
|
75
|
+
raise ScopeError(
|
|
76
|
+
"Direct effects are read at the last token, where the logit difference is "
|
|
77
|
+
f"measured, but {rs.label} is at another position in prompt {prompt.index}. "
|
|
78
|
+
"Set the position to the last token."
|
|
79
|
+
)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def run_direct_effects(
|
|
83
|
+
spec: Spec,
|
|
84
|
+
backend: ModelBackend,
|
|
85
|
+
prompts: list[PreparedPrompt],
|
|
86
|
+
*,
|
|
87
|
+
on_progress: ProgressFn | None = None,
|
|
88
|
+
on_layer: LayerFn | None = None,
|
|
89
|
+
cancel: threading.Event | None = None,
|
|
90
|
+
on_start: Callable[[list[ResolvedSite], dict[str, Any]], None] | None = None,
|
|
91
|
+
) -> EngineResult:
|
|
92
|
+
exp = spec.experiment
|
|
93
|
+
assert isinstance(exp, DirectLogitAttribution)
|
|
94
|
+
info = backend.info
|
|
95
|
+
batch_size = spec.execution.batch_size
|
|
96
|
+
sites, layout = expand_scope(spec, info, prompts)
|
|
97
|
+
check_direct_sites(sites, info, prompts)
|
|
98
|
+
if on_start is not None:
|
|
99
|
+
on_start(sites, layout)
|
|
100
|
+
groups = group_by_length(prompts)
|
|
101
|
+
baselines = compute_baselines(backend, prompts, groups, batch_size, cancel)
|
|
102
|
+
which = exp.prompts
|
|
103
|
+
n, n_layers = len(prompts), info.n_layers
|
|
104
|
+
heads = any(rs.kind == "head" for rs in sites)
|
|
105
|
+
|
|
106
|
+
terms: dict[str, np.ndarray] = {
|
|
107
|
+
"embed": np.zeros(n),
|
|
108
|
+
"attn_out": np.zeros((n, n_layers)),
|
|
109
|
+
"mlp_out": np.zeros((n, n_layers)),
|
|
110
|
+
"logit_diff": np.zeros(n),
|
|
111
|
+
"remainder": np.zeros(n),
|
|
112
|
+
}
|
|
113
|
+
if heads:
|
|
114
|
+
terms["head"] = np.zeros((n, n_layers, info.n_heads))
|
|
115
|
+
done = 0
|
|
116
|
+
for group in groups:
|
|
117
|
+
tokens = group.clean if which == "clean" else group.corrupt
|
|
118
|
+
for sl in _chunks(len(group.members), batch_size):
|
|
119
|
+
if cancel is not None and cancel.is_set():
|
|
120
|
+
raise Cancelled()
|
|
121
|
+
idx = group.members[sl]
|
|
122
|
+
out = backend.direct_effects(tokens[sl], *_answer_tensors(prompts, idx), heads)
|
|
123
|
+
for key, values in out.items():
|
|
124
|
+
terms[key][idx] = values.numpy()
|
|
125
|
+
done += len(idx)
|
|
126
|
+
if on_progress is not None:
|
|
127
|
+
on_progress(done, n, n_layers - 1)
|
|
128
|
+
|
|
129
|
+
gap = baselines.ld(which)
|
|
130
|
+
warnings: list[str] = []
|
|
131
|
+
mean_gap = float(gap.mean())
|
|
132
|
+
if spec.metric.normalization == "dataset_gap" and abs(mean_gap) < 1e-3:
|
|
133
|
+
raise EngineError(
|
|
134
|
+
f"The {which} prompts' mean logit difference is almost zero ({mean_gap:.4f}), so a "
|
|
135
|
+
"share of it is undefined. Normalize by each prompt's own logit difference, or check "
|
|
136
|
+
"the baseline."
|
|
137
|
+
)
|
|
138
|
+
if spec.metric.normalization == "prompt_gap":
|
|
139
|
+
zero = [p.index for p, g in zip(prompts, gap, strict=True) if abs(g) < 1e-6]
|
|
140
|
+
if zero:
|
|
141
|
+
raise EngineError(
|
|
142
|
+
f"Prompt(s) {', '.join(map(str, zero[:5]))} have a logit difference of zero, so a "
|
|
143
|
+
"share of it is undefined. Normalize by the dataset's mean, or fix those prompts."
|
|
144
|
+
)
|
|
145
|
+
small = int((np.abs(gap) < 0.1).sum())
|
|
146
|
+
if small:
|
|
147
|
+
warnings.append(
|
|
148
|
+
f"{small} prompt(s) have a logit difference below 0.1, so their per-prompt shares "
|
|
149
|
+
"are unstable."
|
|
150
|
+
)
|
|
151
|
+
# The split is of the model's own logit difference; the two ways of reading it out must agree.
|
|
152
|
+
drift = float(np.abs(terms["logit_diff"] - gap).max()) if n else 0.0
|
|
153
|
+
if drift > 1e-3 * max(1.0, float(np.abs(gap).max())):
|
|
154
|
+
warnings.append(
|
|
155
|
+
f"The decomposed logit difference differs from the measured one by up to {drift:.2g}, "
|
|
156
|
+
"more than rounding explains."
|
|
157
|
+
)
|
|
158
|
+
|
|
159
|
+
delta = np.zeros((len(sites), n))
|
|
160
|
+
for rs in sites:
|
|
161
|
+
if rs.kind == "head":
|
|
162
|
+
delta[rs.index] = terms["head"][:, rs.layer, rs.head]
|
|
163
|
+
else:
|
|
164
|
+
delta[rs.index] = terms[rs.kind][:, rs.layer]
|
|
165
|
+
nan = np.full((len(sites), n), np.nan)
|
|
166
|
+
result = EngineResult(
|
|
167
|
+
sites=sites,
|
|
168
|
+
layout=layout,
|
|
169
|
+
prompts=prompts,
|
|
170
|
+
baselines=baselines,
|
|
171
|
+
receiver=which,
|
|
172
|
+
reference=which,
|
|
173
|
+
patched_ld=nan,
|
|
174
|
+
patched_prob=nan.copy(),
|
|
175
|
+
warnings=warnings,
|
|
176
|
+
measure="attribution",
|
|
177
|
+
delta=delta,
|
|
178
|
+
gap=gap,
|
|
179
|
+
extra={
|
|
180
|
+
"direct": {
|
|
181
|
+
"prompts": which,
|
|
182
|
+
"logit_diff": float(terms["logit_diff"].mean()),
|
|
183
|
+
"embeddings": float(terms["embed"].mean()),
|
|
184
|
+
"attention": float(terms["attn_out"].sum(1).mean()),
|
|
185
|
+
"mlp": float(terms["mlp_out"].sum(1).mean()),
|
|
186
|
+
"biases": float(terms["remainder"].mean()),
|
|
187
|
+
}
|
|
188
|
+
},
|
|
189
|
+
)
|
|
190
|
+
if on_layer is not None:
|
|
191
|
+
for layer in sorted({rs.layer for rs in sites}):
|
|
192
|
+
on_layer(layer, [rs.index for rs in sites if rs.layer == layer], result)
|
|
193
|
+
return result
|