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.
Files changed (54) hide show
  1. logogram/__init__.py +6 -0
  2. logogram/__main__.py +5 -0
  3. logogram/analysis.py +419 -0
  4. logogram/atp.py +120 -0
  5. logogram/backends/__init__.py +5 -0
  6. logogram/backends/base.py +202 -0
  7. logogram/backends/hub.py +375 -0
  8. logogram/backends/saes.py +277 -0
  9. logogram/backends/transformer_lens.py +872 -0
  10. logogram/cli.py +496 -0
  11. logogram/compare.py +177 -0
  12. logogram/datasets.py +159 -0
  13. logogram/direct.py +193 -0
  14. logogram/engine.py +550 -0
  15. logogram/examples/ioi-gpt2/.gitignore +3 -0
  16. logogram/examples/ioi-gpt2/datasets/ioi.jsonl +32 -0
  17. logogram/examples/ioi-gpt2/experiments/ioi-head-patching/spec.json +42 -0
  18. logogram/examples/ioi-gpt2/project.json +6 -0
  19. logogram/exports.py +33 -0
  20. logogram/features.py +368 -0
  21. logogram/fileio.py +63 -0
  22. logogram/ioi.py +220 -0
  23. logogram/paths.py +204 -0
  24. logogram/project.py +444 -0
  25. logogram/prompts.py +204 -0
  26. logogram/research.py +84 -0
  27. logogram/results.py +240 -0
  28. logogram/runner.py +396 -0
  29. logogram/runs.py +98 -0
  30. logogram/sae.py +161 -0
  31. logogram/schema.py +302 -0
  32. logogram/server/__init__.py +1 -0
  33. logogram/server/app.py +1083 -0
  34. logogram/server/models.py +426 -0
  35. logogram/server/security.py +212 -0
  36. logogram/server/state.py +585 -0
  37. logogram/sites.py +249 -0
  38. logogram/spec.py +518 -0
  39. logogram/stats.py +171 -0
  40. logogram/steering.py +258 -0
  41. logogram/system.py +379 -0
  42. logogram/updates.py +194 -0
  43. logogram/verify.py +39 -0
  44. logogram/web_dist/assets/index-BvCU-2uy.js +54 -0
  45. logogram/web_dist/assets/index-DTr8_ucV.css +1 -0
  46. logogram/web_dist/assets/instrument-sans-latin-ext-standard-normal-C5E2Gvlv.woff2 +0 -0
  47. logogram/web_dist/assets/instrument-sans-latin-standard-normal-BVScPF0l.woff2 +0 -0
  48. logogram/web_dist/favicon.svg +1 -0
  49. logogram/web_dist/index.html +15 -0
  50. logogram-0.1.0.dist-info/METADATA +550 -0
  51. logogram-0.1.0.dist-info/RECORD +54 -0
  52. logogram-0.1.0.dist-info/WHEEL +4 -0
  53. logogram-0.1.0.dist-info/entry_points.txt +2 -0
  54. 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