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/server/app.py ADDED
@@ -0,0 +1,1083 @@
1
+ """The local web server: JSON API, WebSocket event stream and the bundled web app."""
2
+
3
+ from __future__ import annotations
4
+
5
+ import asyncio
6
+ import contextlib
7
+ import hashlib
8
+ import logging
9
+ import os
10
+ import threading
11
+ from importlib import resources
12
+ from pathlib import Path
13
+ from typing import Any, Literal
14
+
15
+ import anyio
16
+ from fastapi import FastAPI, HTTPException, Request, WebSocket, WebSocketDisconnect
17
+ from fastapi.responses import FileResponse, JSONResponse, Response, StreamingResponse
18
+ from pydantic import BaseModel, Field, ValidationError
19
+
20
+ from logogram import __version__
21
+ from logogram.backends.base import BackendError
22
+ from logogram.datasets import (
23
+ DatasetError,
24
+ PromptRecord,
25
+ check_dataset_name,
26
+ file_sha256,
27
+ load_dataset,
28
+ parse_jsonl,
29
+ write_dataset,
30
+ )
31
+ from logogram.project import (
32
+ Project,
33
+ ProjectError,
34
+ default_projects_parent,
35
+ load_recent,
36
+ open_example,
37
+ )
38
+ from logogram.research import (
39
+ Notebook,
40
+ NoteConflict,
41
+ NoteInput,
42
+ ResearchNote,
43
+ delete_note,
44
+ read_notebook,
45
+ save_note,
46
+ )
47
+ from logogram.schema import Manifest, RunListing, Summary
48
+ from logogram.server import models as M
49
+ from logogram.server.security import SecurityConfig, SecurityMiddleware
50
+ from logogram.server.state import AppState, Conflict, EventHub, Missing, request_project
51
+ from logogram.spec import NAME_MAX, ModelRef, PredictionSettings, Spec, describe_intervention
52
+ from logogram.system import SystemReport
53
+
54
+ log = logging.getLogger(__name__)
55
+
56
+ # Starting points that load through TransformerLens and fit an ordinary machine. Every model is
57
+ # checked again when it loads (backends/transformer_lens.check_model); only GPT-2 small has been
58
+ # run end to end with the bundled example.
59
+ MODEL_PRESETS = [
60
+ {
61
+ "id": "openai-community/gpt2",
62
+ "label": "GPT-2 small",
63
+ "detail": "124M · the model the IOI example was made for",
64
+ "tested": True,
65
+ "gated": False,
66
+ },
67
+ {
68
+ "id": "openai-community/gpt2-medium",
69
+ "label": "GPT-2 medium",
70
+ "detail": "355M · 24 layers",
71
+ "tested": False,
72
+ "gated": False,
73
+ },
74
+ {
75
+ "id": "EleutherAI/pythia-160m",
76
+ "label": "Pythia 160M",
77
+ "detail": "Attention and MLP in parallel · training checkpoints as revisions, like step3000",
78
+ "tested": False,
79
+ "gated": False,
80
+ },
81
+ {
82
+ "id": "EleutherAI/pythia-410m",
83
+ "label": "Pythia 410M",
84
+ "detail": "Attention and MLP in parallel · training checkpoints as revisions",
85
+ "tested": False,
86
+ "gated": False,
87
+ },
88
+ {
89
+ "id": "HuggingFaceTB/SmolLM2-135M",
90
+ "label": "SmolLM2 135M",
91
+ "detail": "Llama architecture, small enough for a CPU",
92
+ "tested": False,
93
+ "gated": False,
94
+ },
95
+ {
96
+ "id": "Qwen/Qwen2.5-0.5B",
97
+ "label": "Qwen2.5 0.5B",
98
+ "detail": "No beginning-of-sequence token",
99
+ "tested": False,
100
+ "gated": False,
101
+ },
102
+ {
103
+ "id": "allenai/OLMo-2-0425-1B",
104
+ "label": "OLMo 2 1B",
105
+ "detail": "Normalization after each sublayer",
106
+ "tested": False,
107
+ "gated": False,
108
+ },
109
+ {
110
+ "id": "meta-llama/Llama-3.2-1B",
111
+ "label": "Llama 3.2 1B",
112
+ "detail": "Gated: request access on Hugging Face first",
113
+ "tested": False,
114
+ "gated": True,
115
+ },
116
+ {
117
+ "id": "google/gemma-2-2b",
118
+ "label": "Gemma 2 2B",
119
+ "detail": "Gated: request access on Hugging Face first · soft-capped logits",
120
+ "tested": False,
121
+ "gated": True,
122
+ },
123
+ ]
124
+ MODEL_SUGGESTIONS = [
125
+ *(p["id"] for p in MODEL_PRESETS),
126
+ "EleutherAI/pythia-70m",
127
+ "EleutherAI/pythia-1b",
128
+ "HuggingFaceTB/SmolLM2-360M",
129
+ "Qwen/Qwen3-0.6B-Base",
130
+ "google/gemma-3-270m",
131
+ "microsoft/phi-1_5",
132
+ ]
133
+
134
+
135
+ # Published SAEs for the suggested models, by model id. Each is checked against the loaded model when
136
+ # it loads and measured on the project's prompts before its features are trusted.
137
+ SAE_SUGGESTIONS: dict[str, list[dict[str, str]]] = {
138
+ "openai-community/gpt2": [
139
+ {
140
+ "repo": "jbloom/GPT2-Small-SAEs-Reformatted",
141
+ "detail": "Residual stream before each layer · 24,576 features",
142
+ },
143
+ {
144
+ "repo": "jbloom/GPT2-Small-OAI-v5-32k-resid-post-SAEs",
145
+ "detail": "OpenAI's TopK SAEs · residual stream after each layer · 32,768 features",
146
+ },
147
+ ],
148
+ "EleutherAI/pythia-70m": [
149
+ {
150
+ "repo": "EleutherAI/sae-pythia-70m-32k",
151
+ "detail": "TopK SAEs · each layer's output, attention and MLP · 32,768 features",
152
+ }
153
+ ],
154
+ "EleutherAI/pythia-70m-deduped": [
155
+ {
156
+ "repo": "EleutherAI/sae-pythia-70m-deduped-32k",
157
+ "detail": "TopK SAEs · each layer's output, attention and MLP · 32,768 features",
158
+ }
159
+ ],
160
+ "HuggingFaceTB/SmolLM2-135M": [
161
+ {"repo": "EleutherAI/sae-smollm2-135m-64x", "detail": "TopK SAEs · MLP outputs"}
162
+ ],
163
+ "meta-llama/Llama-3.2-1B": [
164
+ {
165
+ "repo": "EleutherAI/sae-llama-3.2-1b-131k",
166
+ "detail": "TopK SAEs · MLP outputs · 131,072 features",
167
+ }
168
+ ],
169
+ }
170
+
171
+
172
+ def web_dist() -> Path:
173
+ return Path(str(resources.files("logogram") / "web_dist"))
174
+
175
+
176
+ # -- request bodies ----------------------------------------------------------------------------
177
+
178
+
179
+ class CreateProject(BaseModel):
180
+ name: str
181
+ parent: str | None = None
182
+
183
+
184
+ class OpenProject(BaseModel):
185
+ path: str
186
+
187
+
188
+ class IOIRequest(BaseModel):
189
+ name: str = "ioi"
190
+ n: int = Field(default=32, ge=1, le=100_000)
191
+ seed: int = 0
192
+ templates: list[str] | None = None
193
+ patterns: list[Literal["ABBA", "BABA"]] = ["ABBA", "BABA"]
194
+ corruption: Literal["flip", "abc"] = "flip"
195
+ overwrite: bool = False
196
+
197
+
198
+ class ImportRequest(BaseModel):
199
+ name: str
200
+ text: str
201
+ overwrite: bool = False
202
+
203
+
204
+ class PairRequest(BaseModel):
205
+ name: str = "pair"
206
+ clean: str
207
+ corrupt: str
208
+ answer: str
209
+ distractor: str
210
+ overwrite: bool = False
211
+
212
+
213
+ class AnalysisRequest(BaseModel):
214
+ model: ModelRef | None = None
215
+ dataset_sha256: str | None = None
216
+ prepend_bos: bool = True
217
+ limit: int | None = Field(default=None, ge=1)
218
+ batch_size: int = Field(default=64, ge=1)
219
+
220
+
221
+ class TokenizeRequest(AnalysisRequest):
222
+ dataset: str | None = None
223
+ index: int = 0
224
+ record: PromptRecord | None = None
225
+
226
+
227
+ class EstimateRequest(BaseModel):
228
+ id: str
229
+ revision: str | None = None
230
+ dtype: Literal["float32", "float16", "bfloat16"] = "float32"
231
+ device: Literal["auto", "cpu", "cuda", "mps"] = "auto"
232
+
233
+
234
+ class BaselineRequest(AnalysisRequest):
235
+ dataset: str
236
+
237
+
238
+ class AttentionRequest(AnalysisRequest):
239
+ dataset: str
240
+ index: int = 0
241
+ layer: int
242
+ head: int
243
+ which: Literal["clean", "corrupt"] = "clean"
244
+
245
+
246
+ class PredictionRequest(AnalysisRequest):
247
+ dataset: str
248
+ index: int = Field(ge=0)
249
+ settings: PredictionSettings
250
+
251
+
252
+ class SAELoadRequest(BaseModel):
253
+ repo: str = Field(min_length=1)
254
+ path: str = ""
255
+ revision: str | None = None
256
+
257
+
258
+ class SAEAnalysisRequest(AnalysisRequest):
259
+ dataset: str
260
+
261
+
262
+ class TokenFeaturesRequest(SAEAnalysisRequest):
263
+ index: int = Field(default=0, ge=0)
264
+ which: Literal["clean", "corrupt"] = "clean"
265
+ top_k: int = Field(default=8, ge=1, le=32)
266
+
267
+
268
+ class FeatureRequest(SAEAnalysisRequest):
269
+ index: int = Field(default=0, ge=0)
270
+ which: Literal["clean", "corrupt"] = "clean"
271
+ feature: int = Field(ge=0)
272
+
273
+
274
+ class NoteUpdate(NoteInput):
275
+ revision: int = Field(ge=1)
276
+
277
+
278
+ class RunRequest(BaseModel):
279
+ spec: dict[str, Any]
280
+ draft_id: str | None = None # run a saved, not-yet-run experiment in its own folder
281
+
282
+
283
+ class RobustnessRequest(BaseModel):
284
+ experiment: dict[str, Any]
285
+
286
+
287
+ class VerifyRequest(BaseModel):
288
+ top: int = Field(default=10, ge=1, le=200)
289
+
290
+
291
+ class SettingsRequest(BaseModel):
292
+ system_check_seen: bool | None = None
293
+ theme: Literal["light", "dark", "system"] | None = None
294
+ # Whether Logogram may ask PyPI about new versions once a day. Off until the user says yes.
295
+ update_check: bool | None = None
296
+
297
+
298
+ # -- app factory -------------------------------------------------------------------------------
299
+
300
+
301
+ def create_app(
302
+ security: SecurityConfig,
303
+ *,
304
+ state: AppState | None = None,
305
+ initial_project: Path | None = None,
306
+ serve_web: bool = True,
307
+ check_updates: bool = False,
308
+ ) -> FastAPI:
309
+ hub = state.hub if state else EventHub()
310
+ state = state or AppState(hub)
311
+
312
+ @contextlib.asynccontextmanager
313
+ async def lifespan(app: FastAPI): # type: ignore[no-untyped-def]
314
+ hub.bind(asyncio.get_running_loop())
315
+ if check_updates:
316
+ state.start_update_checks()
317
+ yield
318
+ state.stop()
319
+
320
+ app = FastAPI(
321
+ title="Logogram",
322
+ version=__version__,
323
+ docs_url=None,
324
+ redoc_url=None,
325
+ openapi_url=None,
326
+ lifespan=lifespan,
327
+ )
328
+ app.state.logogram = state
329
+ if initial_project is not None:
330
+ state.open_project(Project.open(initial_project))
331
+
332
+ @app.exception_handler(Missing)
333
+ async def _missing(_: Request, exc: Missing) -> JSONResponse:
334
+ return JSONResponse({"error": str(exc)}, status_code=409)
335
+
336
+ @app.exception_handler(Conflict)
337
+ @app.exception_handler(NoteConflict)
338
+ async def _conflict(_: Request, exc: Conflict) -> JSONResponse:
339
+ return JSONResponse({"error": str(exc)}, status_code=409)
340
+
341
+ for exc_type in (ProjectError, DatasetError, BackendError, ValueError):
342
+
343
+ @app.exception_handler(exc_type)
344
+ async def _bad(_: Request, exc: Exception) -> JSONResponse:
345
+ return JSONResponse({"error": str(exc)}, status_code=400)
346
+
347
+ @app.exception_handler(OSError)
348
+ async def _os_error(_: Request, exc: OSError) -> JSONResponse:
349
+ name = Path(exc.filename).name if isinstance(exc.filename, str) else None
350
+ what = exc.strerror or type(exc).__name__
351
+ return JSONResponse({"error": f"{what}: {name}" if name else what}, status_code=400)
352
+
353
+ @app.exception_handler(ValidationError)
354
+ async def _invalid(_: Request, exc: ValidationError) -> JSONResponse:
355
+ first = exc.errors()[0]
356
+ where = ".".join(str(p) for p in first["loc"])
357
+ return JSONResponse({"error": f"{where}: {first['msg']}"}, status_code=400)
358
+
359
+ # -- state -------------------------------------------------------------------------------
360
+
361
+ @app.get("/api/state", response_model=M.ServerState)
362
+ def get_state() -> dict[str, Any]:
363
+ settings = state.settings()
364
+ return {
365
+ "version": __version__,
366
+ "project": state.project.to_dict() if state.project else None,
367
+ "model": state.model_payload(),
368
+ "job": state.job.to_dict() if state.job else None,
369
+ "first_run": not settings.get("system_check_seen", False),
370
+ "theme": settings.get("theme", "light"),
371
+ "projects_parent": str(default_projects_parent()),
372
+ "update": state.update_status(),
373
+ "sae": state.sae_payload(),
374
+ }
375
+
376
+ @app.post("/api/settings", response_model=M.Settings)
377
+ def post_settings(body: SettingsRequest) -> dict[str, Any]:
378
+ values = {k: v for k, v in body.model_dump().items() if v is not None}
379
+ state.update_settings(**values)
380
+ if values.get("update_check") is True:
381
+ # Allowed just now: check right away rather than at the next hourly look.
382
+ threading.Thread(
383
+ target=state.check_updates, name="logogram-update", daemon=True
384
+ ).start()
385
+ elif "update_check" in values:
386
+ state.hub.publish("update", state.update_status())
387
+ return state.settings()
388
+
389
+ @app.get("/api/update", response_model=M.UpdateStatus)
390
+ def get_update() -> dict[str, Any]:
391
+ """What is known about newer versions. Never contacts the network."""
392
+ return state.update_status()
393
+
394
+ @app.post("/api/update/check", response_model=M.UpdateStatus)
395
+ def check_update() -> dict[str, Any]:
396
+ """Ask PyPI now: the user pressed Check now."""
397
+ return state.check_updates()
398
+
399
+ @app.get("/api/system", response_model=SystemReport)
400
+ def get_system() -> dict[str, Any]:
401
+ from logogram.system import system_report
402
+
403
+ return system_report().to_dict()
404
+
405
+ # -- files and projects -----------------------------------------------------------------
406
+
407
+ def _is_project(folder: Path) -> bool:
408
+ try:
409
+ return (folder / "project.json").is_file()
410
+ except OSError: # e.g. a folder we may list but not enter
411
+ return False
412
+
413
+ @app.get("/api/fs", response_model=M.FolderListing)
414
+ def list_dirs(path: str | None = None) -> dict[str, Any]:
415
+ base = Path(path).expanduser() if path else Path.home()
416
+ base = base.resolve()
417
+ try:
418
+ if not base.is_dir():
419
+ raise ProjectError(f"{base} isn't a folder.")
420
+ children = sorted(base.iterdir(), key=lambda p: p.name.lower())
421
+ except OSError as exc:
422
+ raise ProjectError(f"Logogram isn't allowed to read {base}.") from exc
423
+ entries = []
424
+ for child in children:
425
+ try:
426
+ if child.name.startswith(".") or not child.is_dir():
427
+ continue
428
+ except OSError:
429
+ continue
430
+ entries.append(
431
+ {"name": child.name, "path": str(child), "is_project": _is_project(child)}
432
+ )
433
+ return {
434
+ "path": str(base),
435
+ "parent": str(base.parent) if base.parent != base else None,
436
+ "is_project": _is_project(base),
437
+ "entries": entries,
438
+ }
439
+
440
+ @app.get("/api/projects/recent", response_model=list[M.RecentProject])
441
+ def recent() -> list[dict[str, Any]]:
442
+ return load_recent()
443
+
444
+ def _enter(project: Project) -> dict[str, Any]:
445
+ info = project.to_dict() # read it first: a project that can't be listed isn't opened
446
+ state.open_project(project)
447
+ return info
448
+
449
+ @app.post("/api/projects", response_model=M.ProjectInfo)
450
+ def create_project(body: CreateProject) -> dict[str, Any]:
451
+ parent = Path(body.parent).expanduser() if body.parent else default_projects_parent()
452
+ return _enter(Project.create(parent, body.name))
453
+
454
+ @app.post("/api/projects/open", response_model=M.ProjectInfo)
455
+ def open_project(body: OpenProject) -> dict[str, Any]:
456
+ return _enter(Project.open(body.path))
457
+
458
+ @app.post("/api/projects/example", response_model=M.ProjectInfo)
459
+ def example() -> dict[str, Any]:
460
+ return _enter(open_example())
461
+
462
+ @app.post("/api/projects/close", response_model=M.OkOut)
463
+ def close_project() -> dict[str, Any]:
464
+ state.close_project()
465
+ return {"ok": True}
466
+
467
+ @app.get("/api/project", response_model=M.ProjectInfo)
468
+ def get_project() -> dict[str, Any]:
469
+ return state.require_project().to_dict()
470
+
471
+ @app.get("/api/research", response_model=Notebook)
472
+ def research_notes() -> Notebook:
473
+ return read_notebook(state.require_project())
474
+
475
+ @app.post("/api/research", response_model=ResearchNote)
476
+ def new_research_note(body: NoteInput) -> ResearchNote:
477
+ with state._job_lock:
478
+ project = state.require_project()
479
+ note = save_note(project, body)
480
+ state.hub.publish("research.updated", {"project_session": project.session_id})
481
+ return note
482
+
483
+ @app.put("/api/research/{note_id}", response_model=ResearchNote)
484
+ def update_research_note(note_id: str, body: NoteUpdate) -> ResearchNote:
485
+ with state._job_lock:
486
+ project = state.require_project()
487
+ value = NoteInput.model_validate(body.model_dump(exclude={"revision"}))
488
+ note = save_note(project, value, note_id=note_id, revision=body.revision)
489
+ state.hub.publish("research.updated", {"project_session": project.session_id})
490
+ return note
491
+
492
+ @app.delete("/api/research/{note_id}", response_model=M.OkOut)
493
+ def remove_research_note(note_id: str, revision: int) -> dict[str, bool]:
494
+ with state._job_lock:
495
+ project = state.require_project()
496
+ delete_note(project, note_id, revision)
497
+ state.hub.publish("research.updated", {"project_session": project.session_id})
498
+ return {"ok": True}
499
+
500
+ # -- datasets ----------------------------------------------------------------------------
501
+
502
+ def _dataset_path(project: Project, name: str) -> Path:
503
+ """A dataset to read: it must really be inside the project (not linked from elsewhere)."""
504
+ return project.readable(project.datasets_dir / check_dataset_name(Path(name).name))
505
+
506
+ def _new_dataset_path(project: Project, name: str) -> Path:
507
+ return project.dataset_file(check_dataset_name(Path(name).name))
508
+
509
+ def _write_new(path: Path, records: list[PromptRecord], overwrite: bool) -> dict[str, Any]:
510
+ if os.path.lexists(path) and not overwrite:
511
+ raise DatasetError(
512
+ f"{path.name} already exists. Choose another name, or replace it explicitly."
513
+ )
514
+ write_dataset(path, records)
515
+ return {"name": path.name, "path": f"datasets/{path.name}", "n": len(records)}
516
+
517
+ def _records(
518
+ dataset: str, limit: int | None = None, sha256: str | None = None
519
+ ) -> list[PromptRecord]:
520
+ project = state.require_project()
521
+ path = project.resolve_dataset(dataset)
522
+ content = path.read_bytes()
523
+ if sha256 is not None and hashlib.sha256(content).hexdigest() != sha256:
524
+ raise Conflict(
525
+ "The dataset changed since these settings were saved. Restore the original dataset or choose the updated dataset for a new experiment."
526
+ )
527
+ records = parse_jsonl(content.decode("utf-8"), source=dataset)
528
+ return records[:limit] if limit else records
529
+
530
+ @app.get("/api/datasets/{name}", response_model=M.DatasetDetail)
531
+ @app.get("/api/dataset", response_model=M.DatasetDetail)
532
+ def get_dataset(
533
+ name: str = "", path: str | None = None, prepend_bos: bool = True, limit: int | None = None
534
+ ) -> dict[str, Any]:
535
+ project = state.require_project()
536
+ file = project.resolve_dataset(path) if path else _dataset_path(project, name)
537
+ records = load_dataset(file)
538
+ if limit is not None:
539
+ if limit < 1:
540
+ raise ValueError("Prompt limit must be positive.")
541
+ records = records[:limit]
542
+ out: dict[str, Any] = {
543
+ "name": file.name,
544
+ "path": file.relative_to(project.root).as_posix(),
545
+ "sha256": file_sha256(file),
546
+ "n": len(records),
547
+ "records": [r.model_dump(mode="json", exclude_none=True) for r in records],
548
+ }
549
+ if state.backend is not None:
550
+ from logogram.analysis import prepare_with_issues
551
+
552
+ prepared, issues = prepare_with_issues(state.backend, records, prepend_bos)
553
+ out["issues"] = [i.to_dict() for i in issues]
554
+ out["lengths"] = sorted({p.length for p in prepared})
555
+ return out
556
+
557
+ @app.get("/api/ioi/templates", response_model=list[M.IOITemplateOut])
558
+ def ioi_templates() -> list[dict[str, Any]]:
559
+ from logogram.ioi import TEMPLATES
560
+
561
+ return [{"id": t.id, "text": t.text, "default": t.default} for t in TEMPLATES]
562
+
563
+ @app.post("/api/datasets/ioi", response_model=M.DatasetCreated)
564
+ def make_ioi(body: IOIRequest) -> dict[str, Any]:
565
+ from logogram.ioi import generate_ioi
566
+
567
+ project = state.require_project()
568
+ backend = state.backend
569
+ single = (lambda w: backend.single_token_id(w) is not None) if backend else None
570
+ records = generate_ioi(
571
+ body.n,
572
+ seed=body.seed,
573
+ templates=body.templates,
574
+ patterns=body.patterns,
575
+ corruption=body.corruption,
576
+ single_token=single,
577
+ )
578
+ return _write_new(_new_dataset_path(project, body.name), records, body.overwrite)
579
+
580
+ @app.post("/api/datasets/import", response_model=M.DatasetCreated)
581
+ def import_dataset(body: ImportRequest) -> dict[str, Any]:
582
+ project = state.require_project()
583
+ path = _new_dataset_path(project, body.name)
584
+ records = parse_jsonl(body.text, source=path.name)
585
+ return _write_new(path, records, body.overwrite)
586
+
587
+ @app.post("/api/datasets/pair", response_model=M.DatasetCreated)
588
+ def make_pair(body: PairRequest) -> dict[str, Any]:
589
+ project = state.require_project()
590
+ record = PromptRecord(
591
+ clean=body.clean, corrupt=body.corrupt, answer=body.answer, distractor=body.distractor
592
+ )
593
+ return _write_new(_new_dataset_path(project, body.name), [record], body.overwrite)
594
+
595
+ @app.post("/api/tokenize", response_model=M.TokenStrip)
596
+ def tokenize(body: TokenizeRequest) -> dict[str, Any]:
597
+ from logogram.analysis import tokenize_pair
598
+
599
+ backend = _analysis_backend(body)
600
+ if body.record is not None:
601
+ record = body.record
602
+ elif body.dataset:
603
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
604
+ if not 0 <= body.index < len(records):
605
+ raise ValueError(f"There is no prompt {body.index}.")
606
+ record = records[body.index]
607
+ else:
608
+ raise ValueError("Send a dataset and index, or a prompt pair.")
609
+ return tokenize_pair(backend, record, body.prepend_bos)
610
+
611
+ # -- models ------------------------------------------------------------------------------
612
+
613
+ @app.get("/api/models/presets", response_model=M.Presets)
614
+ def presets() -> dict[str, Any]:
615
+ return {"presets": MODEL_PRESETS, "suggestions": MODEL_SUGGESTIONS}
616
+
617
+ @app.post("/api/models/estimate", response_model=M.EstimateOut)
618
+ def estimate(body: EstimateRequest) -> dict[str, Any]:
619
+ from logogram.backends import hub
620
+ from logogram.backends.transformer_lens import architecture_support, resolve_device
621
+ from logogram.system import estimate_memory
622
+
623
+ device = resolve_device(body.device)
624
+ repo = hub.resolve(body.id, body.revision)
625
+ config = hub.fetch_config(body.id, repo.revision)
626
+ arch = hub.read_architecture(config)
627
+ support_note = architecture_support(arch.architecture)
628
+ n_params = repo.n_params
629
+ if n_params is None:
630
+ weights = sum(s for f, s in repo.files if f.endswith(".safetensors"))
631
+ n_params = weights // 2 # stored size is a fallback; most checkpoints are 16-bit
632
+ download = 0
633
+ from huggingface_hub import try_to_load_from_cache
634
+
635
+ for filename, size in repo.files:
636
+ if not isinstance(
637
+ try_to_load_from_cache(body.id, filename, revision=repo.revision), str
638
+ ):
639
+ download += size
640
+ loaded = 0
641
+ if state.backend is not None and state.backend.info.device == device:
642
+ loaded = state.backend.memory_in_use() or 0
643
+ est = estimate_memory(
644
+ n_params=n_params,
645
+ n_layers=arch.n_layers,
646
+ n_heads=arch.n_heads,
647
+ d_model=arch.d_model,
648
+ d_mlp=arch.d_mlp,
649
+ d_vocab=arch.d_vocab,
650
+ dtype=body.dtype,
651
+ device=device,
652
+ loaded_bytes=loaded,
653
+ )
654
+ return {
655
+ "id": body.id,
656
+ "revision": repo.revision,
657
+ "architecture": arch.__dict__,
658
+ "download_bytes": download,
659
+ "total_bytes": sum(s for _, s in repo.files),
660
+ "gated": repo.gated,
661
+ "supported": support_note is None,
662
+ "support_note": support_note,
663
+ "estimate": est.to_dict(),
664
+ }
665
+
666
+ @app.post("/api/models/load", response_model=M.JobInfo)
667
+ def load(body: ModelRef) -> dict[str, Any]:
668
+ return state.load_model_job(body).to_dict()
669
+
670
+ @app.post("/api/models/unload", response_model=M.ModelStatus)
671
+ def unload() -> dict[str, Any]:
672
+ state.unload_model()
673
+ return state.model_payload()
674
+
675
+ # -- analyses ----------------------------------------------------------------------------
676
+
677
+ def _analysis_backend(body: AnalysisRequest): # type: ignore[no-untyped-def]
678
+ from logogram.runner import model_matches
679
+
680
+ backend = state.require_backend()
681
+ if body.model is not None and not model_matches(backend, body.model):
682
+ raise Conflict(
683
+ "The loaded model doesn't match these experiment settings. Load the experiment's model, revision, dtype and weight processing before inspecting it."
684
+ )
685
+ return backend
686
+
687
+ @app.post("/api/baseline", response_model=M.BaselineReport)
688
+ def baseline(body: BaselineRequest) -> dict[str, Any]:
689
+ from logogram.analysis import baseline_report
690
+
691
+ backend = _analysis_backend(body)
692
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
693
+ with backend.lock: # unloading waits until the analysis is done
694
+ return baseline_report(
695
+ backend, records, prepend_bos=body.prepend_bos, batch_size=body.batch_size
696
+ )
697
+
698
+ @app.post("/api/attention", response_model=M.AttentionData)
699
+ def attention(body: AttentionRequest) -> dict[str, Any]:
700
+ from logogram.analysis import attention_report
701
+
702
+ backend = _analysis_backend(body)
703
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
704
+ with backend.lock:
705
+ return attention_report(
706
+ backend,
707
+ records,
708
+ index=body.index,
709
+ layer=body.layer,
710
+ head=body.head,
711
+ which=body.which,
712
+ prepend_bos=body.prepend_bos,
713
+ batch_size=body.batch_size,
714
+ )
715
+
716
+ @app.post("/api/predictions", response_model=M.PredictionReport)
717
+ def predictions(body: PredictionRequest) -> dict[str, Any]:
718
+ from logogram.analysis import prediction_report
719
+
720
+ backend = _analysis_backend(body)
721
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
722
+ with backend.lock:
723
+ return prediction_report(
724
+ backend,
725
+ records,
726
+ index=body.index,
727
+ settings=body.settings,
728
+ prepend_bos=body.prepend_bos,
729
+ batch_size=body.batch_size,
730
+ )
731
+
732
+ # -- sparse autoencoders -----------------------------------------------------------------
733
+
734
+ @app.get("/api/sae/suggestions")
735
+ def sae_suggestions(model: str) -> list[dict[str, str]]:
736
+ return SAE_SUGGESTIONS.get(model, [])
737
+
738
+ @app.get("/api/sae/folders")
739
+ def sae_folders(repo: str, revision: str | None = None) -> dict[str, Any]:
740
+ """The SAEs in a Hugging Face repository (one per folder), at its exact revision."""
741
+ from logogram.backends.saes import list_saes
742
+
743
+ sha, folders = list_saes(repo, revision)
744
+ if not folders:
745
+ raise ValueError(
746
+ f"{repo} has no SAE that Logogram can read: it looks for cfg.json with "
747
+ "sae_weights.safetensors or sae.safetensors."
748
+ )
749
+ return {"repo": repo, "revision": sha, "folders": folders}
750
+
751
+ @app.post("/api/sae/load", response_model=M.JobInfo)
752
+ def sae_load(body: SAELoadRequest) -> dict[str, Any]:
753
+ from logogram.spec import SAERef
754
+
755
+ state.require_backend()
756
+ ref = SAERef(repo=body.repo, path=body.path, revision=body.revision)
757
+ return state.load_sae_job(ref).to_dict()
758
+
759
+ @app.post("/api/sae/unload")
760
+ def sae_unload() -> dict[str, Any]:
761
+ state.unload_sae()
762
+ return state.sae_payload()
763
+
764
+ def _sae() -> Any:
765
+ if state.sae is None:
766
+ raise Missing("Load an SAE first.")
767
+ return state.sae
768
+
769
+ @app.post("/api/sae/fit")
770
+ def sae_fit(body: SAEAnalysisRequest) -> dict[str, Any]:
771
+ from logogram.analysis import sae_fit_report
772
+
773
+ backend, sae = _analysis_backend(body), _sae()
774
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
775
+ with backend.lock:
776
+ fit = sae_fit_report(
777
+ backend, sae, records, prepend_bos=body.prepend_bos, batch_size=body.batch_size
778
+ )
779
+ state.hub.publish("sae", state.sae_payload())
780
+ return fit
781
+
782
+ @app.post("/api/sae/tokens")
783
+ def sae_tokens(body: TokenFeaturesRequest) -> dict[str, Any]:
784
+ from logogram.analysis import token_features_report
785
+
786
+ backend, sae = _analysis_backend(body), _sae()
787
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
788
+ with backend.lock:
789
+ return token_features_report(
790
+ backend,
791
+ sae,
792
+ records,
793
+ index=body.index,
794
+ which=body.which,
795
+ prepend_bos=body.prepend_bos,
796
+ top_k=body.top_k,
797
+ )
798
+
799
+ @app.post("/api/sae/feature")
800
+ def sae_feature(body: FeatureRequest) -> dict[str, Any]:
801
+ from logogram.analysis import feature_report
802
+
803
+ backend, sae = _analysis_backend(body), _sae()
804
+ records = _records(body.dataset, body.limit, body.dataset_sha256)
805
+ with backend.lock:
806
+ return feature_report(
807
+ backend,
808
+ sae,
809
+ records,
810
+ feature=body.feature,
811
+ index=body.index,
812
+ which=body.which,
813
+ prepend_bos=body.prepend_bos,
814
+ batch_size=body.batch_size,
815
+ )
816
+
817
+ # -- runs --------------------------------------------------------------------------------
818
+
819
+ def _read_run(project: Project, run_id: str) -> dict[str, Any]:
820
+ folder = project.run_dir(run_id)
821
+ if not (folder / "spec.json").is_file():
822
+ raise Missing(f"There is no run {run_id}.")
823
+ out: dict[str, Any] = {"id": run_id}
824
+ for name in ("spec", "summary", "manifest", "predictions"):
825
+ path = folder / f"{name}.json"
826
+ try:
827
+ data = project.read_json(path, optional=name != "spec")
828
+ schema = {
829
+ "spec": Spec,
830
+ "summary": Summary,
831
+ "manifest": Manifest,
832
+ "predictions": M.PredictionReport,
833
+ }[name]
834
+ out[name] = (
835
+ schema.model_validate(data).model_dump(mode="json")
836
+ if data is not None
837
+ else None
838
+ )
839
+ except ProjectError:
840
+ raise
841
+ except ValueError:
842
+ if name == "spec":
843
+ raise
844
+ out[name] = None
845
+ listing = project.run_listing(run_id)
846
+ out["listing"] = listing.to_dict() if listing else None
847
+ out["folder"] = f"experiments/{run_id}"
848
+ return out
849
+
850
+ @app.get("/api/runs", response_model=list[RunListing])
851
+ def runs() -> list[dict[str, Any]]:
852
+ project = state.require_project()
853
+ out = [r.to_dict() for r in project.list_runs()]
854
+ job = state.job
855
+ if job is not None and job.status == "running" and job.run_id:
856
+ for r in out:
857
+ # The job outlives its run by a moment (it writes the manifest, then reports);
858
+ # a run whose manifest is written has ended, whatever the job says.
859
+ if r["id"] == job.run_id and r["status"] == "draft":
860
+ r["status"] = "running"
861
+ return out
862
+
863
+ @app.get("/api/runs/{run_id}", response_model=M.RunDetail)
864
+ def get_run(run_id: str) -> dict[str, Any]:
865
+ return _read_run(state.require_project(), run_id)
866
+
867
+ @app.get("/api/runs/{run_id}/export.csv")
868
+ def export_run(run_id: str) -> StreamingResponse:
869
+ from logogram.exports import results_csv
870
+
871
+ project = state.require_project()
872
+ listing = project.run_listing(run_id)
873
+ if listing is None or listing.status != "finished":
874
+ raise Conflict("Only finished runs can be exported. Select a finished run first.")
875
+ path = project.readable(project.run_dir(run_id) / "results.parquet")
876
+ return StreamingResponse(
877
+ results_csv(path),
878
+ media_type="text/csv",
879
+ headers={
880
+ "content-disposition": f'attachment; filename="{run_id}.csv"',
881
+ },
882
+ )
883
+
884
+ @app.post("/api/runs", response_model=M.StartedRun)
885
+ def start_run(body: RunRequest) -> dict[str, Any]:
886
+ spec = Spec.model_validate(body.spec)
887
+ state.require_project()
888
+ job = state.run_spec_job(spec, draft_id=body.draft_id)
889
+ return {"run_id": job.run_id, "job": job.to_dict()}
890
+
891
+ @app.post("/api/drafts", response_model=M.DraftSaved)
892
+ def save_draft(body: RunRequest) -> dict[str, Any]:
893
+ """Save a spec without running it, for `logogram run` or later. With ``draft_id``, update
894
+ that saved experiment, as long as it hasn't run."""
895
+ spec = Spec.model_validate(body.spec)
896
+ run_id = state.save_draft(spec, body.draft_id)
897
+ return {"run_id": run_id, "path": f"experiments/{run_id}/spec.json"}
898
+
899
+ @app.post("/api/runs/{run_id}/rerun", response_model=M.StartedRun)
900
+ def rerun(run_id: str) -> dict[str, Any]:
901
+ project = state.require_project()
902
+ spec = Spec.from_path(project.readable(project.run_dir(run_id) / "spec.json"))
903
+ job = state.run_spec_job(spec, derived_from={"run": run_id, "kind": "rerun"})
904
+ return {"run_id": job.run_id, "job": job.to_dict()}
905
+
906
+ @app.post("/api/runs/{run_id}/robustness", response_model=M.StartedRun)
907
+ def robustness(run_id: str, body: RobustnessRequest) -> dict[str, Any]:
908
+ project = state.require_project()
909
+ original = Spec.from_path(project.readable(project.run_dir(run_id) / "spec.json"))
910
+ data = original.model_dump(mode="json")
911
+ data["experiment"] = body.experiment
912
+ suffix = " · robustness"
913
+ data["name"] = original.name[: NAME_MAX - len(suffix)].rstrip() + suffix
914
+ variant = Spec.model_validate(data)
915
+ change = describe_intervention(variant.experiment)
916
+ job = state.run_spec_job(
917
+ variant, derived_from={"run": run_id, "kind": "robustness", "change": change}
918
+ )
919
+ return {"run_id": job.run_id, "job": job.to_dict()}
920
+
921
+ @app.post("/api/runs/{run_id}/verify", response_model=M.StartedRun)
922
+ def verify(run_id: str, body: VerifyRequest) -> dict[str, Any]:
923
+ """Patch, for real, the sites an attribution patching run estimated to matter most."""
924
+ from logogram.verify import verification_spec
925
+
926
+ project = state.require_project()
927
+ run = _read_run(project, run_id)
928
+ if not run["summary"]:
929
+ raise ValueError("Only a finished run can be verified.")
930
+ spec = verification_spec(Spec.model_validate(run["spec"]), run["summary"], body.top)
931
+ n = len(spec.scope.sites) # type: ignore[union-attr]
932
+ job = state.run_spec_job(
933
+ spec,
934
+ derived_from={
935
+ "run": run_id,
936
+ "kind": "verification",
937
+ "change": f"the {n} strongest estimated site{'s' if n != 1 else ''}, patched",
938
+ },
939
+ )
940
+ return {"run_id": job.run_id, "job": job.to_dict()}
941
+
942
+ @app.get("/api/runs/{run_id}/derived", response_model=list[RunListing])
943
+ def derived(run_id: str) -> list[dict[str, Any]]:
944
+ project = state.require_project()
945
+ return [
946
+ r.to_dict()
947
+ for r in project.list_runs()
948
+ if r.derived_from and r.derived_from.run == run_id
949
+ ]
950
+
951
+ @app.get("/api/runs/{run_id}/sites/{site}", response_model=M.SiteDetail)
952
+ def site_detail(run_id: str, site: int) -> dict[str, Any]:
953
+ from logogram.runs import site_detail as detail
954
+
955
+ return detail(state.require_project(), run_id, site)
956
+
957
+ @app.get("/api/compare", response_model=M.Comparison)
958
+ def compare(a: str, b: str) -> dict[str, Any]:
959
+ from logogram.compare import compare_summaries
960
+
961
+ project = state.require_project()
962
+ ra, rb = _read_run(project, a), _read_run(project, b)
963
+ if not ra["summary"] or not rb["summary"]:
964
+ raise ValueError("Both runs need to have finished.")
965
+ return compare_summaries(
966
+ ra["summary"],
967
+ rb["summary"],
968
+ Spec.model_validate(ra["spec"]),
969
+ Spec.model_validate(rb["spec"]),
970
+ )
971
+
972
+ @app.post("/api/jobs/cancel", response_model=M.CancelOut)
973
+ def cancel() -> dict[str, Any]:
974
+ job = state.cancel_job()
975
+ return {"job": job.to_dict() if job else None}
976
+
977
+ # -- events ------------------------------------------------------------------------------
978
+
979
+ @app.websocket("/ws")
980
+ async def events(ws: WebSocket) -> None:
981
+ await ws.accept()
982
+ # Subscribe before saying hello: the tab refetches state when the stream opens, and every
983
+ # event after this point (and the current run so far) is queued for it.
984
+ queue = hub.subscribe()
985
+
986
+ async def send_events() -> None:
987
+ while True:
988
+ message = await queue.get()
989
+ if message is None: # fell too far behind: close, so the tab reconnects
990
+ await ws.close(code=1013)
991
+ return
992
+ await ws.send_text(message)
993
+
994
+ async def until_closed() -> None:
995
+ # Returns when the tab closes or the server shuts down, so nothing lingers.
996
+ while (await ws.receive())["type"] != "websocket.disconnect":
997
+ pass
998
+
999
+ async def first_to_finish(job: Any, scope: anyio.CancelScope) -> None:
1000
+ await job()
1001
+ scope.cancel()
1002
+
1003
+ try:
1004
+ await ws.send_json({"type": "hello", "version": __version__})
1005
+ # A task group, not bare asyncio tasks, so a server-side cancel (shutdown, or the
1006
+ # test client closing) stays inside this handler's scope.
1007
+ async with anyio.create_task_group() as group:
1008
+ group.start_soon(first_to_finish, send_events, group.cancel_scope)
1009
+ group.start_soon(first_to_finish, until_closed, group.cancel_scope)
1010
+ except* (WebSocketDisconnect, RuntimeError):
1011
+ pass
1012
+ finally:
1013
+ hub.unsubscribe(queue)
1014
+
1015
+ # -- web app -----------------------------------------------------------------------------
1016
+
1017
+ if serve_web:
1018
+ dist = web_dist()
1019
+
1020
+ @app.get("/{path:path}", include_in_schema=False)
1021
+ def spa(path: str) -> Response:
1022
+ if path.startswith("api/"):
1023
+ raise HTTPException(status_code=404)
1024
+ target = (dist / path).resolve()
1025
+ if path and target.is_file() and target.is_relative_to(dist.resolve()):
1026
+ headers = (
1027
+ {"cache-control": "public, max-age=31536000, immutable"}
1028
+ if path.startswith("assets/")
1029
+ else {}
1030
+ )
1031
+ return FileResponse(target, headers=headers)
1032
+ index = dist / "index.html"
1033
+ if not index.is_file():
1034
+ return Response(
1035
+ "The web app hasn't been built. Run `npm run build` in web/.",
1036
+ media_type="text/plain",
1037
+ status_code=503,
1038
+ )
1039
+ return FileResponse(index, headers={"cache-control": "no-store"})
1040
+
1041
+ @app.middleware("http")
1042
+ async def headers(request: Request, call_next): # type: ignore[no-untyped-def]
1043
+ project = state.project
1044
+ expected = request.headers.get("x-logogram-project")
1045
+ # Browser requests carry an ephemeral project identity. Capture the object as well:
1046
+ # an open/close concurrent with dispatch must never retarget a pending operation.
1047
+ if expected is not None and expected != (project.session_id if project else "none"):
1048
+ return JSONResponse(
1049
+ {
1050
+ "error": "The project changed in another tab. Wait for this tab to refresh, then try again."
1051
+ },
1052
+ status_code=409,
1053
+ )
1054
+ token = request_project.set((project,))
1055
+ try:
1056
+ response = await call_next(request)
1057
+ finally:
1058
+ request_project.reset(token)
1059
+ response.headers.setdefault("x-content-type-options", "nosniff")
1060
+ response.headers.setdefault("referrer-policy", "no-referrer")
1061
+ response.headers.setdefault("x-frame-options", "DENY")
1062
+ host = request.headers.get("host", "")
1063
+ response.headers.setdefault(
1064
+ "content-security-policy",
1065
+ f"default-src 'self'; connect-src 'self' ws://{host}; img-src 'self' data: blob:; "
1066
+ "style-src 'self' 'unsafe-inline'; font-src 'self'; object-src 'none'; "
1067
+ "base-uri 'none'; frame-ancestors 'none'",
1068
+ )
1069
+ return response
1070
+
1071
+ return SecuredApp(app, security) # type: ignore[return-value]
1072
+
1073
+
1074
+ class SecuredApp:
1075
+ """The FastAPI app wrapped in the security middleware (outermost, so it sees everything)."""
1076
+
1077
+ def __init__(self, app: FastAPI, security: SecurityConfig) -> None:
1078
+ self.inner = app
1079
+ self.state = app.state
1080
+ self._app = SecurityMiddleware(app, security)
1081
+
1082
+ async def __call__(self, scope, receive, send): # type: ignore[no-untyped-def]
1083
+ await self._app(scope, receive, send)