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/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)
|