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/cli.py
ADDED
|
@@ -0,0 +1,496 @@
|
|
|
1
|
+
"""Command line: ``logogram`` (start the app), ``open``, ``run``, ``doctor``, ``serve``."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import contextlib
|
|
6
|
+
import os
|
|
7
|
+
import socket
|
|
8
|
+
import sys
|
|
9
|
+
import threading
|
|
10
|
+
import time
|
|
11
|
+
import webbrowser
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
from typing import Annotated, Any
|
|
14
|
+
|
|
15
|
+
import typer
|
|
16
|
+
|
|
17
|
+
from logogram import __version__
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
def _quiet_environment() -> None:
|
|
21
|
+
"""No telemetry from libraries we depend on, and deterministic cuBLAS."""
|
|
22
|
+
os.environ.setdefault("HF_HUB_DISABLE_TELEMETRY", "1")
|
|
23
|
+
os.environ.setdefault("DO_NOT_TRACK", "1")
|
|
24
|
+
os.environ.setdefault("WANDB_MODE", "disabled")
|
|
25
|
+
os.environ.setdefault("WANDB_SILENT", "true")
|
|
26
|
+
os.environ.setdefault("TOKENIZERS_PARALLELISM", "false")
|
|
27
|
+
os.environ.setdefault("TRANSFORMERS_NO_ADVISORY_WARNINGS", "1")
|
|
28
|
+
os.environ.setdefault("CUBLAS_WORKSPACE_CONFIG", ":4096:8")
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
app = typer.Typer(
|
|
32
|
+
add_completion=False,
|
|
33
|
+
invoke_without_command=True,
|
|
34
|
+
help="Logogram: a local workbench for causal experiments inside language models.",
|
|
35
|
+
rich_markup_mode=None,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _version(value: bool) -> None:
|
|
40
|
+
if value:
|
|
41
|
+
typer.echo(f"logogram {__version__}")
|
|
42
|
+
raise typer.Exit()
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _bind(preferred: int) -> socket.socket:
|
|
46
|
+
"""Bind 127.0.0.1 on ``preferred`` or the next free port (any free port for 0).
|
|
47
|
+
|
|
48
|
+
The socket is handed to the server as it is, so no other process can take the port between
|
|
49
|
+
choosing it and serving on it.
|
|
50
|
+
"""
|
|
51
|
+
candidates = [preferred, *range(preferred + 1, min(preferred + 20, 65536))] if preferred else []
|
|
52
|
+
for port in [*candidates, 0]:
|
|
53
|
+
sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
|
54
|
+
if sys.platform != "win32":
|
|
55
|
+
# Like uvicorn: reuse a port a just-stopped server left in TIME_WAIT.
|
|
56
|
+
sock.setsockopt(socket.SOL_SOCKET, socket.SO_REUSEADDR, 1)
|
|
57
|
+
try:
|
|
58
|
+
sock.bind(("127.0.0.1", port))
|
|
59
|
+
return sock
|
|
60
|
+
except OSError:
|
|
61
|
+
sock.close()
|
|
62
|
+
raise RuntimeError("No free port on 127.0.0.1.")
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
def _open_when_started(server: Any, security: Any) -> None:
|
|
66
|
+
"""Open the browser once the server is serving, with a single-use launch code."""
|
|
67
|
+
|
|
68
|
+
def wait_and_open() -> None:
|
|
69
|
+
while not server.started:
|
|
70
|
+
if server.should_exit:
|
|
71
|
+
return
|
|
72
|
+
time.sleep(0.1)
|
|
73
|
+
code = security.new_launch_code()
|
|
74
|
+
webbrowser.open(f"http://127.0.0.1:{security.port}/?token={code}")
|
|
75
|
+
|
|
76
|
+
threading.Thread(target=wait_and_open, daemon=True).start()
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
def _serve(
|
|
80
|
+
*,
|
|
81
|
+
port: int,
|
|
82
|
+
open_browser: bool,
|
|
83
|
+
project: Path | None = None,
|
|
84
|
+
dev: bool = False,
|
|
85
|
+
) -> None:
|
|
86
|
+
from logogram.server.security import SecurityConfig, new_token
|
|
87
|
+
|
|
88
|
+
sock = _bind(port)
|
|
89
|
+
port = int(sock.getsockname()[1])
|
|
90
|
+
token = os.environ.get("LOGOGRAM_TOKEN") if dev else None
|
|
91
|
+
dev_list = (
|
|
92
|
+
[o.strip() for o in os.environ.get("LOGOGRAM_DEV_ORIGINS", "").split(",") if o.strip()]
|
|
93
|
+
or ["http://localhost:5173", "http://127.0.0.1:5173"]
|
|
94
|
+
if dev
|
|
95
|
+
else []
|
|
96
|
+
)
|
|
97
|
+
security = SecurityConfig(token=token or new_token(), port=port, dev_origins=dev_list)
|
|
98
|
+
url = f"http://127.0.0.1:{port}/?token={security.token}"
|
|
99
|
+
# Print the address first: loading PyTorch takes a few seconds.
|
|
100
|
+
typer.echo(f"Logogram {__version__} is running at:\n\n {url}\n")
|
|
101
|
+
if dev:
|
|
102
|
+
typer.echo(
|
|
103
|
+
f"Development UI (npm run dev):\n\n {dev_list[0]}/api/session?token={security.token}\n"
|
|
104
|
+
)
|
|
105
|
+
typer.echo("Everything runs on this machine. Press Ctrl+C to stop.")
|
|
106
|
+
_update_notice()
|
|
107
|
+
|
|
108
|
+
import uvicorn
|
|
109
|
+
|
|
110
|
+
from logogram.runner import configure_determinism
|
|
111
|
+
from logogram.server.app import create_app
|
|
112
|
+
|
|
113
|
+
configure_determinism()
|
|
114
|
+
application = create_app(
|
|
115
|
+
security, initial_project=project, serve_web=not dev, check_updates=True
|
|
116
|
+
)
|
|
117
|
+
config = uvicorn.Config(
|
|
118
|
+
application,
|
|
119
|
+
host="127.0.0.1",
|
|
120
|
+
port=port,
|
|
121
|
+
log_level="warning",
|
|
122
|
+
ws="auto",
|
|
123
|
+
# Open event streams (browser tabs) must not keep Ctrl+C from stopping the server.
|
|
124
|
+
timeout_graceful_shutdown=2,
|
|
125
|
+
)
|
|
126
|
+
server = uvicorn.Server(config)
|
|
127
|
+
if open_browser:
|
|
128
|
+
_open_when_started(server, security)
|
|
129
|
+
with contextlib.suppress(KeyboardInterrupt):
|
|
130
|
+
server.run(sockets=[sock])
|
|
131
|
+
|
|
132
|
+
|
|
133
|
+
PortOption = typer.Option(
|
|
134
|
+
min=0, max=65535, help="Port to listen on (127.0.0.1 only; 0 picks any free port)."
|
|
135
|
+
)
|
|
136
|
+
|
|
137
|
+
|
|
138
|
+
def _update_notice() -> None:
|
|
139
|
+
"""One line about newer versions, from what is already known (no network here)."""
|
|
140
|
+
from logogram import updates
|
|
141
|
+
from logogram.project import config_dir
|
|
142
|
+
|
|
143
|
+
choice = _settings().get("update_check")
|
|
144
|
+
info = updates.status(
|
|
145
|
+
choice if isinstance(choice, bool) else None,
|
|
146
|
+
updates.read_cache(config_dir() / "updates.json"),
|
|
147
|
+
)
|
|
148
|
+
if info.available:
|
|
149
|
+
typer.echo(
|
|
150
|
+
f"\nLogogram {info.latest} is out (you have {info.current}). Update: {info.command}"
|
|
151
|
+
)
|
|
152
|
+
elif info.old:
|
|
153
|
+
typer.echo(
|
|
154
|
+
f"\nThis version of Logogram is from {info.released}. "
|
|
155
|
+
"Newer ones may be out: logogram check-updates"
|
|
156
|
+
)
|
|
157
|
+
|
|
158
|
+
|
|
159
|
+
def _settings() -> dict: # type: ignore[type-arg]
|
|
160
|
+
import json
|
|
161
|
+
|
|
162
|
+
from logogram.project import config_dir
|
|
163
|
+
|
|
164
|
+
try:
|
|
165
|
+
data = json.loads((config_dir() / "settings.json").read_text(encoding="utf-8"))
|
|
166
|
+
except (OSError, ValueError):
|
|
167
|
+
return {}
|
|
168
|
+
return data if isinstance(data, dict) else {}
|
|
169
|
+
|
|
170
|
+
|
|
171
|
+
@app.callback()
|
|
172
|
+
def main_callback(
|
|
173
|
+
ctx: typer.Context,
|
|
174
|
+
version: Annotated[
|
|
175
|
+
bool,
|
|
176
|
+
typer.Option("--version", callback=_version, is_eager=True, help="Show the version."),
|
|
177
|
+
] = False,
|
|
178
|
+
port: Annotated[int, PortOption] = 8765,
|
|
179
|
+
no_browser: Annotated[bool, typer.Option("--no-browser", help="Don't open a browser.")] = False,
|
|
180
|
+
) -> None:
|
|
181
|
+
"""Start Logogram and open it in your browser."""
|
|
182
|
+
_quiet_environment()
|
|
183
|
+
if ctx.invoked_subcommand is None:
|
|
184
|
+
_serve(port=port, open_browser=not no_browser)
|
|
185
|
+
|
|
186
|
+
|
|
187
|
+
@app.command("open")
|
|
188
|
+
def open_cmd(
|
|
189
|
+
path: Annotated[Path, typer.Argument(help="A project folder.")],
|
|
190
|
+
port: Annotated[int, PortOption] = 8765,
|
|
191
|
+
no_browser: Annotated[bool, typer.Option("--no-browser")] = False,
|
|
192
|
+
) -> None:
|
|
193
|
+
"""Start Logogram with a project open."""
|
|
194
|
+
from logogram.project import Project, ProjectError
|
|
195
|
+
|
|
196
|
+
try:
|
|
197
|
+
project = Project.open(path)
|
|
198
|
+
except ProjectError as exc:
|
|
199
|
+
typer.echo(f"Error: {exc}", err=True)
|
|
200
|
+
raise typer.Exit(1) from exc
|
|
201
|
+
_serve(port=port, open_browser=not no_browser, project=project.root)
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
@app.command()
|
|
205
|
+
def serve(
|
|
206
|
+
port: Annotated[int, PortOption] = 8765,
|
|
207
|
+
dev: Annotated[
|
|
208
|
+
bool, typer.Option("--dev", help="API only, for the Vite dev server in web/.")
|
|
209
|
+
] = False,
|
|
210
|
+
no_browser: Annotated[bool, typer.Option("--no-browser")] = False,
|
|
211
|
+
project: Annotated[Path | None, typer.Option(help="Open this project.")] = None,
|
|
212
|
+
) -> None:
|
|
213
|
+
"""Run the server (use --dev with `npm run dev` when working on the web app)."""
|
|
214
|
+
_serve(port=port, open_browser=not (no_browser or dev), project=project, dev=dev)
|
|
215
|
+
|
|
216
|
+
|
|
217
|
+
@app.command("check-updates")
|
|
218
|
+
def check_updates_cmd() -> None:
|
|
219
|
+
"""Ask PyPI whether a newer Logogram is out (sends nothing about you or your work)."""
|
|
220
|
+
from logogram import updates
|
|
221
|
+
from logogram.project import config_dir
|
|
222
|
+
|
|
223
|
+
path = config_dir() / "updates.json"
|
|
224
|
+
cache = updates.check(path)
|
|
225
|
+
choice = _settings().get("update_check")
|
|
226
|
+
info = updates.status(choice if isinstance(choice, bool) else None, cache)
|
|
227
|
+
if info.error:
|
|
228
|
+
typer.echo(info.error)
|
|
229
|
+
if info.available:
|
|
230
|
+
typer.echo(f"Logogram {info.latest} is out (you have {info.current}).")
|
|
231
|
+
typer.echo(f"Update: {info.command}")
|
|
232
|
+
if info.notes_url:
|
|
233
|
+
typer.echo(f"What's new: {info.notes_url}")
|
|
234
|
+
elif info.latest:
|
|
235
|
+
typer.echo(f"Logogram {info.current} is the newest version.")
|
|
236
|
+
|
|
237
|
+
|
|
238
|
+
@app.command()
|
|
239
|
+
def doctor() -> None:
|
|
240
|
+
"""Report the environment and hardware, with fixes for common problems."""
|
|
241
|
+
from logogram.system import format_bytes, system_report
|
|
242
|
+
|
|
243
|
+
report = system_report()
|
|
244
|
+
rows = [
|
|
245
|
+
("Logogram", __version__),
|
|
246
|
+
("Operating system", f"{report.os} ({report.machine})"),
|
|
247
|
+
("Python", report.python),
|
|
248
|
+
(
|
|
249
|
+
"Processor",
|
|
250
|
+
f"{report.cpu} · {report.cpu_cores or '?'} cores, {report.cpu_threads} threads",
|
|
251
|
+
),
|
|
252
|
+
(
|
|
253
|
+
"Memory",
|
|
254
|
+
f"{format_bytes(report.memory_total)} total, {format_bytes(report.memory_available)} available",
|
|
255
|
+
),
|
|
256
|
+
]
|
|
257
|
+
for gpu in report.gpus:
|
|
258
|
+
free = f", {format_bytes(gpu.memory_free)} free" if gpu.memory_free is not None else ""
|
|
259
|
+
rows.append(("GPU", f"{gpu.name} · {format_bytes(gpu.memory_total)}{free}"))
|
|
260
|
+
if not report.gpus:
|
|
261
|
+
rows.append(("GPU", "none detected"))
|
|
262
|
+
backend = {"cuda": "CUDA", "mps": "Apple Metal (MPS)", "cpu": "CPU"}[report.backend]
|
|
263
|
+
rows += [
|
|
264
|
+
("Compute backend", backend),
|
|
265
|
+
("Recommended precision", report.recommended_dtype),
|
|
266
|
+
("PyTorch", report.torch + (f" (CUDA {report.torch_cuda})" if report.torch_cuda else "")),
|
|
267
|
+
("TransformerLens", report.transformer_lens or "not installed"),
|
|
268
|
+
("Transformers", report.transformers or "not installed"),
|
|
269
|
+
("Updates", _updates_row()),
|
|
270
|
+
]
|
|
271
|
+
width = max(len(k) for k, _ in rows)
|
|
272
|
+
for key, value in rows:
|
|
273
|
+
typer.echo(f"{key:<{width}} {value}")
|
|
274
|
+
typer.echo("")
|
|
275
|
+
typer.echo(report.precision_note)
|
|
276
|
+
typer.echo("")
|
|
277
|
+
if not report.issues:
|
|
278
|
+
typer.echo("No problems found.")
|
|
279
|
+
return
|
|
280
|
+
for issue in report.issues:
|
|
281
|
+
typer.echo(f"[{issue.severity}] {issue.title}")
|
|
282
|
+
typer.echo(f" {issue.detail}")
|
|
283
|
+
if issue.fix:
|
|
284
|
+
typer.echo(f" Fix: {issue.fix}")
|
|
285
|
+
typer.echo("")
|
|
286
|
+
|
|
287
|
+
|
|
288
|
+
def _updates_row() -> str:
|
|
289
|
+
from logogram import updates
|
|
290
|
+
from logogram.project import config_dir
|
|
291
|
+
|
|
292
|
+
choice = _settings().get("update_check")
|
|
293
|
+
info = updates.status(
|
|
294
|
+
choice if isinstance(choice, bool) else None,
|
|
295
|
+
updates.read_cache(config_dir() / "updates.json"),
|
|
296
|
+
)
|
|
297
|
+
automatic = {True: "checked daily", False: "automatic checks off", None: "not set up"}[
|
|
298
|
+
info.automatic
|
|
299
|
+
]
|
|
300
|
+
known = (
|
|
301
|
+
f"{info.latest} available" if info.available else "up to date" if info.latest else "unknown"
|
|
302
|
+
)
|
|
303
|
+
return f"{known} · {automatic} · check now with logogram check-updates"
|
|
304
|
+
|
|
305
|
+
|
|
306
|
+
@app.command()
|
|
307
|
+
def run(
|
|
308
|
+
spec_path: Annotated[Path, typer.Argument(metavar="SPEC.json", help="An experiment spec.")],
|
|
309
|
+
project_dir: Annotated[
|
|
310
|
+
Path | None,
|
|
311
|
+
typer.Option("--project", help="Project folder (default: the one containing the spec)."),
|
|
312
|
+
] = None,
|
|
313
|
+
top: Annotated[int, typer.Option(help="How many sites to list.")] = 10,
|
|
314
|
+
) -> None:
|
|
315
|
+
"""Run an experiment headlessly and print a summary."""
|
|
316
|
+
from logogram.project import Project, ProjectError
|
|
317
|
+
from logogram.runner import run_spec
|
|
318
|
+
from logogram.spec import Spec, describe_experiment
|
|
319
|
+
|
|
320
|
+
if not spec_path.exists():
|
|
321
|
+
typer.echo(f"Error: {spec_path} doesn't exist.", err=True)
|
|
322
|
+
raise typer.Exit(1)
|
|
323
|
+
root = project_dir or Project.find_root(spec_path.parent)
|
|
324
|
+
if root is None:
|
|
325
|
+
typer.echo("Error: no project.json found above the spec. Pass --project PATH.", err=True)
|
|
326
|
+
raise typer.Exit(1)
|
|
327
|
+
try:
|
|
328
|
+
project = Project.open(root)
|
|
329
|
+
spec = Spec.from_path(project.readable(spec_path.absolute()))
|
|
330
|
+
except (ProjectError, ValueError, OSError) as exc:
|
|
331
|
+
typer.echo(f"Error: {exc}", err=True)
|
|
332
|
+
raise typer.Exit(1) from exc
|
|
333
|
+
|
|
334
|
+
original = spec_path.resolve().parent
|
|
335
|
+
original_results = original / "results.parquet"
|
|
336
|
+
compare_to = (
|
|
337
|
+
project.readable(original_results)
|
|
338
|
+
if project.inside(original_results) and original_results.is_file()
|
|
339
|
+
else None
|
|
340
|
+
)
|
|
341
|
+
|
|
342
|
+
typer.echo(f"{spec.name}")
|
|
343
|
+
typer.echo(f" {describe_experiment(spec)}")
|
|
344
|
+
typer.echo(f" model {spec.model.id} · dataset {spec.dataset.path}")
|
|
345
|
+
printer = _ProgressPrinter()
|
|
346
|
+
outcome = run_spec(spec, project, on_event=printer)
|
|
347
|
+
printer.close()
|
|
348
|
+
if outcome.status != "finished":
|
|
349
|
+
typer.echo(f"Run {outcome.status}: {outcome.manifest.get('error')}", err=True)
|
|
350
|
+
raise typer.Exit(1)
|
|
351
|
+
_print_summary(outcome.summary or {}, top)
|
|
352
|
+
rel = outcome.folder.relative_to(project.root)
|
|
353
|
+
typer.echo(f"\nSaved to {rel} ({outcome.manifest.get('wall_time_s')} s)")
|
|
354
|
+
if compare_to is not None:
|
|
355
|
+
_report_reproduction(compare_to, outcome.folder / "results.parquet", original.name)
|
|
356
|
+
|
|
357
|
+
|
|
358
|
+
class _ProgressPrinter:
|
|
359
|
+
def __init__(self) -> None:
|
|
360
|
+
self.last = 0.0
|
|
361
|
+
self.active = False
|
|
362
|
+
|
|
363
|
+
def __call__(self, kind: str, data: dict) -> None: # type: ignore[type-arg]
|
|
364
|
+
now = time.monotonic()
|
|
365
|
+
if kind == "model":
|
|
366
|
+
stage = data.get("stage")
|
|
367
|
+
if stage == "downloading" and data.get("total"):
|
|
368
|
+
if now - self.last > 0.25 or data.get("done") == data.get("total"):
|
|
369
|
+
self.last = now
|
|
370
|
+
pct = 100 * data["done"] / data["total"]
|
|
371
|
+
self._line(f" downloading {pct:5.1f}% of {data['total'] / 1e6:,.0f} MB")
|
|
372
|
+
elif stage in ("resolving", "loading", "processing"):
|
|
373
|
+
self._line(f" model: {stage}")
|
|
374
|
+
elif kind == "progress":
|
|
375
|
+
if now - self.last > 0.2 or data["done"] == data["total"]:
|
|
376
|
+
self.last = now
|
|
377
|
+
pct = 100 * data["done"] / data["total"]
|
|
378
|
+
self._line(f" running layer {data['layer']} · {pct:5.1f}%")
|
|
379
|
+
elif kind == "status":
|
|
380
|
+
self._line(f" {data.get('message')}")
|
|
381
|
+
|
|
382
|
+
def _line(self, text: str) -> None:
|
|
383
|
+
if sys.stderr.isatty():
|
|
384
|
+
sys.stderr.write("\r" + text.ljust(60))
|
|
385
|
+
sys.stderr.flush()
|
|
386
|
+
self.active = True
|
|
387
|
+
else:
|
|
388
|
+
sys.stderr.write(text + "\n")
|
|
389
|
+
|
|
390
|
+
def close(self) -> None:
|
|
391
|
+
if self.active:
|
|
392
|
+
sys.stderr.write("\r" + " " * 60 + "\r")
|
|
393
|
+
sys.stderr.flush()
|
|
394
|
+
|
|
395
|
+
|
|
396
|
+
def _fmt(x: float | None, signed: bool = True) -> str:
|
|
397
|
+
if x is None:
|
|
398
|
+
return "—"
|
|
399
|
+
return f"{x:+.3f}" if signed else f"{x:.3f}"
|
|
400
|
+
|
|
401
|
+
|
|
402
|
+
def _print_summary(summary: dict, top: int) -> None: # type: ignore[type-arg]
|
|
403
|
+
base = summary["baseline"]
|
|
404
|
+
clean = base["clean"]["logit_diff"]["mean"]
|
|
405
|
+
corrupt = base["corrupt"]["logit_diff"]["mean"]
|
|
406
|
+
typer.echo("")
|
|
407
|
+
typer.echo(
|
|
408
|
+
f"Baseline over {summary['n_prompts']} prompts: logit difference clean {_fmt(clean)}, "
|
|
409
|
+
f"corrupt {_fmt(corrupt)}"
|
|
410
|
+
)
|
|
411
|
+
typer.echo(f"Normalized effect: {summary['metric']['normalized_effect']}")
|
|
412
|
+
stats = summary["statistics"]
|
|
413
|
+
ci = round(stats["ci"] * 100)
|
|
414
|
+
typer.echo(
|
|
415
|
+
f"Confidence intervals: {ci}% {stats['method']}, {stats['bootstrap']} resamples, seed {stats['seed']}"
|
|
416
|
+
)
|
|
417
|
+
# What the values are depends on the method: patched runs, estimates of them, or terms of a
|
|
418
|
+
# split; only patched runs can flip a prompt.
|
|
419
|
+
measure = summary.get("measure", "intervention")
|
|
420
|
+
value, change = {
|
|
421
|
+
"intervention": ("effect", "Δ logit diff"),
|
|
422
|
+
"estimate": ("estimated", "estimated Δ"),
|
|
423
|
+
"attribution": ("share", "direct"),
|
|
424
|
+
}.get(measure, ("effect", "Δ logit diff"))
|
|
425
|
+
flips = measure == "intervention"
|
|
426
|
+
sites = sorted(summary["sites"], key=lambda s: -abs(s["effect"]["mean"] or 0.0))[:top]
|
|
427
|
+
typer.echo("")
|
|
428
|
+
header = f"{'site':<22} {value:>9} {f'{ci}% CI':<19} {change:>12}"
|
|
429
|
+
typer.echo(header + (f" {'flipped':>8}" if flips else ""))
|
|
430
|
+
for s in sites:
|
|
431
|
+
e = s["effect"]
|
|
432
|
+
interval = f"[{_fmt(e['lo'])}, {_fmt(e['hi'])}]"
|
|
433
|
+
line = (
|
|
434
|
+
f"{s['label']:<22} {_fmt(e['mean']):>9} {interval:<19} {_fmt(s['delta']['mean']):>12}"
|
|
435
|
+
)
|
|
436
|
+
typer.echo(line + (f" {s['sign_flips']:>3} / {s['n']:<3}" if flips else ""))
|
|
437
|
+
if summary.get("direct"):
|
|
438
|
+
d = summary["direct"]
|
|
439
|
+
typer.echo(
|
|
440
|
+
f"\nThe mean {d['prompts']} logit difference {_fmt(d['logit_diff'])} splits into attention "
|
|
441
|
+
f"{_fmt(d['attention'])}, MLPs {_fmt(d['mlp'])}, embeddings {_fmt(d['embeddings'])} "
|
|
442
|
+
f"and biases {_fmt(d['biases'])}."
|
|
443
|
+
)
|
|
444
|
+
if summary.get("steering"):
|
|
445
|
+
st = summary["steering"]
|
|
446
|
+
typer.echo(
|
|
447
|
+
f"\nDirections from {len(st['train'])} training pairs; measured on the "
|
|
448
|
+
f"{len(st['test'])} held-out pairs."
|
|
449
|
+
)
|
|
450
|
+
if summary.get("features"):
|
|
451
|
+
f = summary["features"]
|
|
452
|
+
fit = f.get("fit") or {}
|
|
453
|
+
typer.echo(
|
|
454
|
+
f"\nSAE {f['sae']['repo']} {f['sae']['path']}: explains "
|
|
455
|
+
f"{fit.get('variance_explained', float('nan')):.1%} of the variance on these prompts."
|
|
456
|
+
)
|
|
457
|
+
if f.get("features_estimate") is not None:
|
|
458
|
+
typer.echo(
|
|
459
|
+
f"Estimated effect of the whole site {_fmt(f['site_estimate'])}, of which the "
|
|
460
|
+
f"features carry {_fmt(f['features_estimate'])}."
|
|
461
|
+
)
|
|
462
|
+
for warning in summary.get("warnings") or []:
|
|
463
|
+
typer.echo(f"\nNote: {warning}")
|
|
464
|
+
|
|
465
|
+
|
|
466
|
+
def _report_reproduction(original: Path, new: Path, original_id: str) -> None:
|
|
467
|
+
import pyarrow.parquet as pq
|
|
468
|
+
|
|
469
|
+
a = pq.read_table(original)
|
|
470
|
+
b = pq.read_table(new)
|
|
471
|
+
if a.equals(b):
|
|
472
|
+
typer.echo(
|
|
473
|
+
f"Identical to {original_id}: all {a.num_rows:,} per-prompt values match exactly."
|
|
474
|
+
)
|
|
475
|
+
return
|
|
476
|
+
import numpy as np
|
|
477
|
+
|
|
478
|
+
if a.num_rows != b.num_rows:
|
|
479
|
+
typer.echo(f"Differs from {original_id}: {a.num_rows} rows before, {b.num_rows} now.")
|
|
480
|
+
return
|
|
481
|
+
diff = np.abs(
|
|
482
|
+
np.asarray(a.column("patched_logit_diff")) - np.asarray(b.column("patched_logit_diff"))
|
|
483
|
+
)
|
|
484
|
+
typer.echo(
|
|
485
|
+
f"Differs from {original_id}: largest change in a patched logit difference is "
|
|
486
|
+
f"{diff.max():.3g}. Check the device, dtype and library versions in the two manifests."
|
|
487
|
+
)
|
|
488
|
+
|
|
489
|
+
|
|
490
|
+
def main() -> None:
|
|
491
|
+
_quiet_environment()
|
|
492
|
+
app()
|
|
493
|
+
|
|
494
|
+
|
|
495
|
+
if __name__ == "__main__":
|
|
496
|
+
main()
|
logogram/compare.py
ADDED
|
@@ -0,0 +1,177 @@
|
|
|
1
|
+
"""Compare two runs: the robustness check and side-by-side comparisons share this.
|
|
2
|
+
|
|
3
|
+
Sites are matched by (kind, layer, head, position). The comparison reports Spearman's rank
|
|
4
|
+
correlation of mean effects, the overlap of the top-k components (ranked by |effect|), and the
|
|
5
|
+
components whose conclusion changed: their sign reversed (both confidence intervals exclude zero,
|
|
6
|
+
on opposite sides) or they entered or left the top k.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import math
|
|
12
|
+
from typing import Any
|
|
13
|
+
|
|
14
|
+
import numpy as np
|
|
15
|
+
|
|
16
|
+
from logogram.spec import Spec
|
|
17
|
+
from logogram.stats import spearman
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class CompareError(ValueError):
|
|
21
|
+
pass
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def site_key(site: dict[str, Any]) -> tuple[Any, ...]:
|
|
25
|
+
return (
|
|
26
|
+
site["kind"],
|
|
27
|
+
site["layer"],
|
|
28
|
+
site.get("head"),
|
|
29
|
+
site.get("feature"),
|
|
30
|
+
site["position_key"],
|
|
31
|
+
site.get("variant_key"),
|
|
32
|
+
)
|
|
33
|
+
|
|
34
|
+
|
|
35
|
+
def _excludes_zero(effect: dict[str, Any]) -> bool:
|
|
36
|
+
lo, hi = effect.get("lo"), effect.get("hi")
|
|
37
|
+
if lo is None or hi is None or lo >= hi:
|
|
38
|
+
return False # no interval (n < 2), or a degenerate one, is never confident
|
|
39
|
+
return lo > 0 or hi < 0
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def _ranks_by_magnitude(values: np.ndarray) -> np.ndarray:
|
|
43
|
+
order = np.argsort(-np.abs(values), kind="mergesort")
|
|
44
|
+
ranks = np.empty(len(values), dtype=np.int64)
|
|
45
|
+
ranks[order] = np.arange(1, len(values) + 1)
|
|
46
|
+
return ranks
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
def default_top_k(n_sites: int) -> int:
|
|
50
|
+
"""About the top fifth of the sites, at most 10, so membership can actually change."""
|
|
51
|
+
return max(1, min(10, math.ceil(n_sites / 5)))
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def spec_differences(
|
|
55
|
+
a: dict[str, Any], b: dict[str, Any], prefix: str = ""
|
|
56
|
+
) -> list[dict[str, Any]]:
|
|
57
|
+
"""Paths where two spec dicts differ, ignoring the name and notes."""
|
|
58
|
+
out: list[dict[str, Any]] = []
|
|
59
|
+
for key in sorted(set(a) | set(b)):
|
|
60
|
+
if not prefix and key in ("name", "notes"):
|
|
61
|
+
continue
|
|
62
|
+
path = f"{prefix}.{key}" if prefix else key
|
|
63
|
+
va, vb = a.get(key), b.get(key)
|
|
64
|
+
if isinstance(va, dict) and isinstance(vb, dict):
|
|
65
|
+
out.extend(spec_differences(va, vb, path))
|
|
66
|
+
elif va != vb:
|
|
67
|
+
out.append({"path": path, "a": va, "b": vb})
|
|
68
|
+
return out
|
|
69
|
+
|
|
70
|
+
|
|
71
|
+
def compare_summaries(
|
|
72
|
+
summary_a: dict[str, Any],
|
|
73
|
+
summary_b: dict[str, Any],
|
|
74
|
+
spec_a: Spec | None = None,
|
|
75
|
+
spec_b: Spec | None = None,
|
|
76
|
+
top_k: int | None = None,
|
|
77
|
+
) -> dict[str, Any]:
|
|
78
|
+
sites_a = {site_key(s): s for s in summary_a["sites"]}
|
|
79
|
+
sites_b = {site_key(s): s for s in summary_b["sites"]}
|
|
80
|
+
# Sites without a value in either run (an undefined statistic) can't be ranked.
|
|
81
|
+
common = [
|
|
82
|
+
k
|
|
83
|
+
for k in sites_a
|
|
84
|
+
if k in sites_b
|
|
85
|
+
and sites_a[k]["effect"]["mean"] is not None
|
|
86
|
+
and sites_b[k]["effect"]["mean"] is not None
|
|
87
|
+
]
|
|
88
|
+
if not common:
|
|
89
|
+
raise CompareError(
|
|
90
|
+
"These runs share no sites (for example a head sweep and a layer × position sweep), "
|
|
91
|
+
"so they can't be compared."
|
|
92
|
+
)
|
|
93
|
+
ea = np.array([sites_a[k]["effect"]["mean"] for k in common], dtype=np.float64)
|
|
94
|
+
eb = np.array([sites_b[k]["effect"]["mean"] for k in common], dtype=np.float64)
|
|
95
|
+
k = top_k or default_top_k(len(common))
|
|
96
|
+
rank_a, rank_b = _ranks_by_magnitude(ea), _ranks_by_magnitude(eb)
|
|
97
|
+
top_a = {i for i in range(len(common)) if rank_a[i] <= k}
|
|
98
|
+
top_b = {i for i in range(len(common)) if rank_b[i] <= k}
|
|
99
|
+
|
|
100
|
+
changes = []
|
|
101
|
+
flagged = []
|
|
102
|
+
for i, key in enumerate(common):
|
|
103
|
+
sa, sb = sites_a[key], sites_b[key]
|
|
104
|
+
flags = []
|
|
105
|
+
# A reversal only counts when both runs are confident about the sign.
|
|
106
|
+
if (
|
|
107
|
+
np.sign(ea[i]) != np.sign(eb[i])
|
|
108
|
+
and _excludes_zero(sa["effect"])
|
|
109
|
+
and _excludes_zero(sb["effect"])
|
|
110
|
+
):
|
|
111
|
+
flags.append("sign")
|
|
112
|
+
if i in top_a and i not in top_b:
|
|
113
|
+
flags.append("left_top")
|
|
114
|
+
if i in top_b and i not in top_a:
|
|
115
|
+
flags.append("entered_top")
|
|
116
|
+
entry = {
|
|
117
|
+
"label": sa["label"],
|
|
118
|
+
"kind": sa["kind"],
|
|
119
|
+
"layer": sa["layer"],
|
|
120
|
+
"head": sa.get("head"),
|
|
121
|
+
"feature": sa.get("feature"),
|
|
122
|
+
"position_key": sa["position_key"],
|
|
123
|
+
"variant_key": sa.get("variant_key"),
|
|
124
|
+
"index_a": sa["index"],
|
|
125
|
+
"index_b": sb["index"],
|
|
126
|
+
"row": sa["row"],
|
|
127
|
+
"col": sa["col"],
|
|
128
|
+
"effect_a": sa["effect"],
|
|
129
|
+
"effect_b": sb["effect"],
|
|
130
|
+
"rank_a": int(rank_a[i]),
|
|
131
|
+
"rank_b": int(rank_b[i]),
|
|
132
|
+
"flags": flags,
|
|
133
|
+
}
|
|
134
|
+
if flags or i in top_a or i in top_b:
|
|
135
|
+
changes.append(entry)
|
|
136
|
+
if flags:
|
|
137
|
+
flagged.append(sa["index"])
|
|
138
|
+
changes.sort(key=lambda c: (min(c["rank_a"], c["rank_b"]), c["rank_a"]))
|
|
139
|
+
|
|
140
|
+
same_layout = (
|
|
141
|
+
summary_a["layout"]["kind"] == summary_b["layout"]["kind"]
|
|
142
|
+
and len(summary_a["layout"]["rows"]) == len(summary_b["layout"]["rows"])
|
|
143
|
+
and len(summary_a["layout"]["cols"]) == len(summary_b["layout"]["cols"])
|
|
144
|
+
)
|
|
145
|
+
diff = [
|
|
146
|
+
{
|
|
147
|
+
"index_a": sites_a[key]["index"],
|
|
148
|
+
"row": sites_a[key]["row"],
|
|
149
|
+
"col": sites_a[key]["col"],
|
|
150
|
+
"value": float(eb[i] - ea[i]),
|
|
151
|
+
}
|
|
152
|
+
for i, key in enumerate(common)
|
|
153
|
+
]
|
|
154
|
+
differences = []
|
|
155
|
+
if spec_a is not None and spec_b is not None:
|
|
156
|
+
differences = spec_differences(
|
|
157
|
+
spec_a.model_dump(mode="json"), spec_b.model_dump(mode="json")
|
|
158
|
+
)
|
|
159
|
+
rho = spearman(ea, eb)
|
|
160
|
+
return {
|
|
161
|
+
"run_a": summary_a["run_id"],
|
|
162
|
+
"run_b": summary_b["run_id"],
|
|
163
|
+
"n_common": len(common),
|
|
164
|
+
"n_a": len(sites_a),
|
|
165
|
+
"n_b": len(sites_b),
|
|
166
|
+
"spearman": rho if np.isfinite(rho) else None,
|
|
167
|
+
"top_k": k,
|
|
168
|
+
"top_overlap": len(top_a & top_b),
|
|
169
|
+
"top_a": [sites_a[common[i]]["label"] for i in sorted(top_a, key=lambda i: rank_a[i])],
|
|
170
|
+
"top_b": [sites_b[common[i]]["label"] for i in sorted(top_b, key=lambda i: rank_b[i])],
|
|
171
|
+
"changes": changes,
|
|
172
|
+
"flagged": flagged,
|
|
173
|
+
"n_sign_changes": sum(1 for c in changes if "sign" in c["flags"]),
|
|
174
|
+
"diff": diff,
|
|
175
|
+
"same_layout": same_layout,
|
|
176
|
+
"spec_differences": differences,
|
|
177
|
+
}
|