stuntd 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.
- stuntd/__init__.py +5 -0
- stuntd/__main__.py +5 -0
- stuntd/cli.py +665 -0
- stuntd/decisions/__init__.py +0 -0
- stuntd/decisions/schema.py +77 -0
- stuntd/decisions/site.py +54 -0
- stuntd/jev/__init__.py +0 -0
- stuntd/jev/answer.py +58 -0
- stuntd/jev/schema.py +183 -0
- stuntd/jev/state.py +12 -0
- stuntd/modes.py +15 -0
- stuntd/paths.py +67 -0
- stuntd/proxy/__init__.py +0 -0
- stuntd/proxy/app.py +423 -0
- stuntd/proxy/capture.py +61 -0
- stuntd/proxy/headers.py +121 -0
- stuntd/proxy/jev.py +425 -0
- stuntd/serve/__init__.py +0 -0
- stuntd/serve/answer.py +78 -0
- stuntd/serve/decider.py +216 -0
- stuntd/serve/modes.py +87 -0
- stuntd/serve/monitor.py +85 -0
- stuntd/serve/runtime.py +374 -0
- stuntd/settings.py +377 -0
- stuntd/store/__init__.py +0 -0
- stuntd/store/db.py +287 -0
- stuntd/store/redact.py +50 -0
- stuntd/train/__init__.py +0 -0
- stuntd/train/artifacts.py +120 -0
- stuntd/train/dataset.py +147 -0
- stuntd/train/metrics.py +198 -0
- stuntd/train/report.py +79 -0
- stuntd/train/run.py +177 -0
- stuntd/train/trainer.py +309 -0
- stuntd-0.1.0.dist-info/METADATA +495 -0
- stuntd-0.1.0.dist-info/RECORD +40 -0
- stuntd-0.1.0.dist-info/WHEEL +5 -0
- stuntd-0.1.0.dist-info/entry_points.txt +2 -0
- stuntd-0.1.0.dist-info/licenses/LICENSE +201 -0
- stuntd-0.1.0.dist-info/top_level.txt +1 -0
stuntd/__init__.py
ADDED
stuntd/__main__.py
ADDED
stuntd/cli.py
ADDED
|
@@ -0,0 +1,665 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import argparse
|
|
4
|
+
import json
|
|
5
|
+
import math
|
|
6
|
+
import sqlite3
|
|
7
|
+
import sys
|
|
8
|
+
import time
|
|
9
|
+
from collections.abc import Callable
|
|
10
|
+
from dataclasses import asdict, dataclass, fields
|
|
11
|
+
from pathlib import Path
|
|
12
|
+
from typing import TYPE_CHECKING
|
|
13
|
+
|
|
14
|
+
from stuntd.decisions.schema import DecisionSchema, detect_schema
|
|
15
|
+
from stuntd.paths import private_file
|
|
16
|
+
from stuntd.serve.modes import (
|
|
17
|
+
MODE_CHECK,
|
|
18
|
+
MODE_COLLECT,
|
|
19
|
+
MODE_LIVE,
|
|
20
|
+
MODE_SHADOW,
|
|
21
|
+
SiteState,
|
|
22
|
+
site_states,
|
|
23
|
+
write_mode,
|
|
24
|
+
)
|
|
25
|
+
from stuntd.serve.monitor import window_agreement, window_start
|
|
26
|
+
from stuntd.settings import (
|
|
27
|
+
CONFIG_TEMPLATE,
|
|
28
|
+
Settings,
|
|
29
|
+
config_path,
|
|
30
|
+
database_path,
|
|
31
|
+
load_settings,
|
|
32
|
+
models_path,
|
|
33
|
+
)
|
|
34
|
+
from stuntd.train.artifacts import HEAD_FILE, SiteModel, list_models, load_model, site_dir
|
|
35
|
+
from stuntd.train.report import render, render_json
|
|
36
|
+
from stuntd.train.run import NO_CAPTURES, Trainer, TrainResult, train_sites
|
|
37
|
+
|
|
38
|
+
if TYPE_CHECKING:
|
|
39
|
+
from stuntd.serve.runtime import DeciderLike
|
|
40
|
+
from stuntd.store.db import Capture, Store
|
|
41
|
+
|
|
42
|
+
__all__ = ["main"]
|
|
43
|
+
|
|
44
|
+
_COLUMN_WIDTHS = (18, 8, 9, 8, 6)
|
|
45
|
+
_CONFIG_HELP = "settings file to read instead of the one in the data directory"
|
|
46
|
+
_IMPORT_MODEL = "import"
|
|
47
|
+
_COVERAGE_DIGITS = 2
|
|
48
|
+
# Below this share of the holdout a head answers so little that the operator should hear about it.
|
|
49
|
+
_LOW_COVERAGE = 0.10
|
|
50
|
+
_KINDS = ("choice", "boolean", "number")
|
|
51
|
+
_BOOLEAN_ANSWERS = ("true", "false")
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
@dataclass(frozen=True)
|
|
55
|
+
class _SiteStatus:
|
|
56
|
+
"""One row of the status table: how a site serves and what it has answered lately."""
|
|
57
|
+
|
|
58
|
+
site: str
|
|
59
|
+
mode: str
|
|
60
|
+
captures: int
|
|
61
|
+
shadow: int
|
|
62
|
+
live: int
|
|
63
|
+
agreement: float | None
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class _LearningOff(Exception):
|
|
67
|
+
"""A command that would write captures, models or decisions while learning is off."""
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _require_learning(config: str | None, settings: Settings) -> None:
|
|
71
|
+
if not settings.learn:
|
|
72
|
+
raise _LearningOff(f"learning is off in {_config_file(config)}")
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def _make_decider(settings: Settings) -> DeciderLike:
|
|
76
|
+
# The only place serving reaches torch and laya, so every other command runs without the extra.
|
|
77
|
+
from stuntd.serve.decider import Decider
|
|
78
|
+
|
|
79
|
+
return Decider(settings.base_model, settings.device)
|
|
80
|
+
|
|
81
|
+
|
|
82
|
+
def _require_serving() -> None:
|
|
83
|
+
"""Refuses the way serve would when the train extra is missing, without loading a model."""
|
|
84
|
+
try:
|
|
85
|
+
import stuntd.serve.decider # noqa: F401
|
|
86
|
+
except ImportError as exc:
|
|
87
|
+
raise ValueError('serving needs the train extra: pip install "stuntd[train]"') from exc
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def _serve_decider(settings: Settings, serving: int) -> DeciderLike | None:
|
|
91
|
+
# A local Jev answers its questions zero-shot from the base checkpoint, trained head or not,
|
|
92
|
+
# so it needs the decider even where no site serves one of its own.
|
|
93
|
+
local_jev = settings.jev_upstream == ""
|
|
94
|
+
if serving == 0 and not local_jev:
|
|
95
|
+
return None
|
|
96
|
+
try:
|
|
97
|
+
return _make_decider(settings)
|
|
98
|
+
except ImportError:
|
|
99
|
+
# The proxy still records captures, so a missing extra costs the answers, not the run.
|
|
100
|
+
print('stuntd: serving disabled: pip install "stuntd[train]"', file=sys.stderr)
|
|
101
|
+
return None
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
def _serving_sites(models: Path) -> list[tuple[str, SiteModel]]:
|
|
105
|
+
return [
|
|
106
|
+
(state.site, state.model)
|
|
107
|
+
for state in site_states(models)
|
|
108
|
+
if state.mode != MODE_COLLECT and state.model is not None
|
|
109
|
+
]
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
def _serve(args: argparse.Namespace) -> int:
|
|
113
|
+
settings = _load(args.config, args.upstream, args.port)
|
|
114
|
+
import uvicorn
|
|
115
|
+
|
|
116
|
+
from stuntd.proxy.app import build_app
|
|
117
|
+
|
|
118
|
+
models = models_path(settings)
|
|
119
|
+
# With learning off no head answers, so the sites on disk serve nothing.
|
|
120
|
+
serving = _serving_sites(models) if settings.learn else []
|
|
121
|
+
# The base model loads before anything is printed: it takes seconds, and a checkpoint that
|
|
122
|
+
# will not load must fail the command rather than leave a line promising a proxy that is up.
|
|
123
|
+
decider = _serve_decider(settings, len(serving))
|
|
124
|
+
if decider is not None and serving:
|
|
125
|
+
# The first pass through a freshly loaded model is far slower than the rest, so one site
|
|
126
|
+
# pays for it here rather than the first caller of whichever site asks first.
|
|
127
|
+
site, model = serving[0]
|
|
128
|
+
decider.warm(model, site_dir(models, site) / HEAD_FILE)
|
|
129
|
+
elif decider is not None:
|
|
130
|
+
decider.warm_base()
|
|
131
|
+
# uvicorn.run never returns while the daemon is up, so a piped stdout needs the lines now.
|
|
132
|
+
listening = f"stuntd listening on http://{settings.host}:{settings.port}"
|
|
133
|
+
if settings.upstream:
|
|
134
|
+
print(f"{listening} -> {settings.upstream}", flush=True)
|
|
135
|
+
else:
|
|
136
|
+
print(listening, flush=True)
|
|
137
|
+
print("no upstream: only the Jev routes are served", flush=True)
|
|
138
|
+
if decider is not None:
|
|
139
|
+
served = f"{len(serving)} site(s)" if serving else "jev locally"
|
|
140
|
+
print(f"serving {served} with {settings.base_model}", flush=True)
|
|
141
|
+
if not settings.learn:
|
|
142
|
+
print("learning off", flush=True)
|
|
143
|
+
# The relay is byte-exact, so uvicorn must not put its own Date and Server headers next to
|
|
144
|
+
# the ones the provider sent.
|
|
145
|
+
uvicorn.run(
|
|
146
|
+
build_app(settings, decider=decider),
|
|
147
|
+
host=settings.host,
|
|
148
|
+
port=settings.port,
|
|
149
|
+
log_level="warning",
|
|
150
|
+
server_header=False,
|
|
151
|
+
date_header=False,
|
|
152
|
+
)
|
|
153
|
+
return 0
|
|
154
|
+
|
|
155
|
+
|
|
156
|
+
def _enable(args: argparse.Namespace) -> int:
|
|
157
|
+
settings = _load(args.config, None, None)
|
|
158
|
+
_require_learning(args.config, settings)
|
|
159
|
+
models = models_path(settings)
|
|
160
|
+
site = _site_name(args.site)
|
|
161
|
+
model = _one_model(models, site)
|
|
162
|
+
if model.threshold is None:
|
|
163
|
+
print(
|
|
164
|
+
f"stuntd: {site} never reached the target agreement on its holdout;"
|
|
165
|
+
" retrain with more examples",
|
|
166
|
+
file=sys.stderr,
|
|
167
|
+
)
|
|
168
|
+
return 1
|
|
169
|
+
# The curve always answers at least one holdout row, so a useless operating point shows up
|
|
170
|
+
# as the coverage the report rounds to, not as a bare zero.
|
|
171
|
+
coverage = None if model.coverage is None else round(model.coverage, _COVERAGE_DIGITS)
|
|
172
|
+
if coverage == 0:
|
|
173
|
+
print(
|
|
174
|
+
f"stuntd: {site} covers 0% of the holdout at target"
|
|
175
|
+
f" {model.target_agreement:.2f}: nothing will be answered locally;"
|
|
176
|
+
" lower training.target_agreement or retrain",
|
|
177
|
+
file=sys.stderr,
|
|
178
|
+
)
|
|
179
|
+
return 1
|
|
180
|
+
if coverage is not None and coverage < _LOW_COVERAGE:
|
|
181
|
+
print(
|
|
182
|
+
f"stuntd: {site} covers {coverage:.0%} of the holdout at target"
|
|
183
|
+
f" {model.target_agreement:.2f}; most answers will still go to the provider",
|
|
184
|
+
file=sys.stderr,
|
|
185
|
+
)
|
|
186
|
+
_require_serving()
|
|
187
|
+
# The proxy re-reads the file once it changes, so the running daemon needs no restart.
|
|
188
|
+
write_mode(site_dir(models, site), MODE_LIVE, time.time())
|
|
189
|
+
print(f"{site}: {MODE_LIVE}")
|
|
190
|
+
return 0
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _disable(args: argparse.Namespace) -> int:
|
|
194
|
+
settings = _load(args.config, None, None)
|
|
195
|
+
_require_learning(args.config, settings)
|
|
196
|
+
models = models_path(settings)
|
|
197
|
+
site = _site_name(args.site)
|
|
198
|
+
_one_model(models, site)
|
|
199
|
+
write_mode(site_dir(models, site), MODE_SHADOW, time.time())
|
|
200
|
+
print(f"{site}: {MODE_SHADOW}")
|
|
201
|
+
return 0
|
|
202
|
+
|
|
203
|
+
|
|
204
|
+
def _status(args: argparse.Namespace) -> int:
|
|
205
|
+
settings = _load(args.config, None, None)
|
|
206
|
+
rows = _status_rows(settings) if settings.learn else []
|
|
207
|
+
if args.json:
|
|
208
|
+
print(json.dumps({"learning": settings.learn, "sites": [asdict(row) for row in rows]}))
|
|
209
|
+
return 0
|
|
210
|
+
if not settings.learn:
|
|
211
|
+
print("learning off")
|
|
212
|
+
return 0
|
|
213
|
+
if not rows:
|
|
214
|
+
print("no captures yet")
|
|
215
|
+
return 0
|
|
216
|
+
print(_row("site", "mode", "captures", "shadow", "live", "agreement"))
|
|
217
|
+
for row in rows:
|
|
218
|
+
print(_status_line(row))
|
|
219
|
+
return 0
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def _config_init(args: argparse.Namespace) -> int:
|
|
223
|
+
path = Path(args.config) if args.config else config_path()
|
|
224
|
+
if path.exists() and not args.force:
|
|
225
|
+
print(f"{path} exists; use --force to overwrite")
|
|
226
|
+
return 1
|
|
227
|
+
path.parent.mkdir(parents=True, exist_ok=True)
|
|
228
|
+
private_file(path).write_text(CONFIG_TEMPLATE, encoding="utf-8")
|
|
229
|
+
print(f"wrote {path}")
|
|
230
|
+
return 0
|
|
231
|
+
|
|
232
|
+
|
|
233
|
+
def _config_show(args: argparse.Namespace) -> int:
|
|
234
|
+
settings = _load(args.config, args.upstream, args.port)
|
|
235
|
+
flagged = {name for name in ("upstream", "port") if getattr(args, name) is not None}
|
|
236
|
+
defaults = Settings()
|
|
237
|
+
for item in fields(Settings):
|
|
238
|
+
value = getattr(settings, item.name)
|
|
239
|
+
if item.name in flagged:
|
|
240
|
+
source = "flag"
|
|
241
|
+
elif value != getattr(defaults, item.name):
|
|
242
|
+
source = "file"
|
|
243
|
+
else:
|
|
244
|
+
source = "default"
|
|
245
|
+
if item.name == "redaction_patterns":
|
|
246
|
+
shown = ", ".join(settings.redaction_patterns)
|
|
247
|
+
else:
|
|
248
|
+
shown = "" if value is None else str(value)
|
|
249
|
+
print(f"{item.name} = {shown} ({source})")
|
|
250
|
+
return 0
|
|
251
|
+
|
|
252
|
+
|
|
253
|
+
def _make_trainer(settings: Settings) -> Trainer:
|
|
254
|
+
# The only place torch and laya are reached, so every other command runs without the extra.
|
|
255
|
+
from stuntd.train.trainer import LayaTrainer
|
|
256
|
+
|
|
257
|
+
return LayaTrainer(
|
|
258
|
+
settings.base_model,
|
|
259
|
+
settings.device,
|
|
260
|
+
settings.epochs,
|
|
261
|
+
cache_encoder=settings.cache_encoder,
|
|
262
|
+
cache_max_bytes=settings.cache_max_mb * 2**20,
|
|
263
|
+
)
|
|
264
|
+
|
|
265
|
+
|
|
266
|
+
def _load_trainer(settings: Settings) -> Trainer:
|
|
267
|
+
try:
|
|
268
|
+
return _make_trainer(settings)
|
|
269
|
+
except ImportError as exc:
|
|
270
|
+
raise ValueError('training needs the train extra: pip install "stuntd[train]"') from exc
|
|
271
|
+
|
|
272
|
+
|
|
273
|
+
def _base_model_line(settings: Settings) -> str:
|
|
274
|
+
# A checkpoint already on disk is read from there, so only a model name is fetched.
|
|
275
|
+
if Path(settings.base_model).is_dir():
|
|
276
|
+
return f"base model: {settings.base_model}"
|
|
277
|
+
return f"base model: {settings.base_model} (downloaded on first use)"
|
|
278
|
+
|
|
279
|
+
|
|
280
|
+
def _result_line(result: TrainResult) -> str:
|
|
281
|
+
model = result.model
|
|
282
|
+
if model is None:
|
|
283
|
+
return f"{result.site} skipped {result.reason}"
|
|
284
|
+
headline = f"{result.site} trained holdout agreement {model.agreement:.3f}"
|
|
285
|
+
if model.coverage is None or model.covered_agreement is None:
|
|
286
|
+
return f"{headline} coverage - ece {model.ece:.3f}"
|
|
287
|
+
return (
|
|
288
|
+
f"{headline} coverage {model.coverage:.2f}"
|
|
289
|
+
f" at agreement {model.covered_agreement:.3f} ece {model.ece:.3f}"
|
|
290
|
+
)
|
|
291
|
+
|
|
292
|
+
|
|
293
|
+
def _train(args: argparse.Namespace) -> int:
|
|
294
|
+
from stuntd.store.db import Store
|
|
295
|
+
from stuntd.store.redact import Redactor
|
|
296
|
+
|
|
297
|
+
settings = _load(args.config, None, None)
|
|
298
|
+
_require_learning(args.config, settings)
|
|
299
|
+
database = database_path(settings)
|
|
300
|
+
if not database.exists():
|
|
301
|
+
print("no captures yet")
|
|
302
|
+
return 0
|
|
303
|
+
store = Store(database, Redactor())
|
|
304
|
+
try:
|
|
305
|
+
known = {info.site for info in store.sites()}
|
|
306
|
+
sites = [_site_name(site) for site in args.sites]
|
|
307
|
+
if not known and not sites:
|
|
308
|
+
print("no captures yet")
|
|
309
|
+
return 0
|
|
310
|
+
trainable = [site for site in sites if site in known] if sites else list(known)
|
|
311
|
+
if not trainable:
|
|
312
|
+
results = [TrainResult(site, None, NO_CAPTURES) for site in sites]
|
|
313
|
+
else:
|
|
314
|
+
# Building the trainer fetches the checkpoint, which is a long silent wait without
|
|
315
|
+
# the line, so it waits until a site actually needs training.
|
|
316
|
+
print(_base_model_line(settings), flush=True)
|
|
317
|
+
results = train_sites(store, settings, _load_trainer(settings), sites)
|
|
318
|
+
finally:
|
|
319
|
+
store.close()
|
|
320
|
+
for result in results:
|
|
321
|
+
print(_result_line(result))
|
|
322
|
+
return 0
|
|
323
|
+
|
|
324
|
+
|
|
325
|
+
class _InvalidLine(Exception):
|
|
326
|
+
"""A row of an import file that cannot be recorded, already named by its line number."""
|
|
327
|
+
|
|
328
|
+
|
|
329
|
+
def _decision_schema(schema: object) -> DecisionSchema:
|
|
330
|
+
# Read back the way a request is, so an imported site carries the canonical schema and the
|
|
331
|
+
# kind the same decision would have got through the proxy.
|
|
332
|
+
found = detect_schema(
|
|
333
|
+
{"response_format": {"type": "json_schema", "json_schema": {"schema": schema}}}
|
|
334
|
+
)
|
|
335
|
+
if found is None:
|
|
336
|
+
raise ValueError("schema must describe an object with one typed field")
|
|
337
|
+
return found
|
|
338
|
+
|
|
339
|
+
|
|
340
|
+
def _built_schema(site: str, kind: str, labels: str | None) -> dict[str, object]:
|
|
341
|
+
if kind != "choice":
|
|
342
|
+
spec: dict[str, object] = {"type": "boolean" if kind == "boolean" else "integer"}
|
|
343
|
+
else:
|
|
344
|
+
# A label repeated in the list would become a class no example can carry.
|
|
345
|
+
options = dict.fromkeys(
|
|
346
|
+
label.strip() for label in (labels or "").split(",") if label.strip()
|
|
347
|
+
)
|
|
348
|
+
if not options:
|
|
349
|
+
raise ValueError("a choice site needs --labels a,b,c or --schema")
|
|
350
|
+
spec = {"type": "string", "enum": list(options)}
|
|
351
|
+
# The field is what the head is asked for at serve time, so a site built here reads the same
|
|
352
|
+
# way a Jev question of that name does rather than asking the model to "choose answer".
|
|
353
|
+
return {"type": "object", "properties": {site: spec}}
|
|
354
|
+
|
|
355
|
+
|
|
356
|
+
def _import_schema(
|
|
357
|
+
site: str, kind: str | None, raw: str | None, labels: str | None
|
|
358
|
+
) -> DecisionSchema:
|
|
359
|
+
if raw is None:
|
|
360
|
+
return _decision_schema(_built_schema(site, kind or "choice", labels))
|
|
361
|
+
try:
|
|
362
|
+
given = json.loads(raw)
|
|
363
|
+
except json.JSONDecodeError as exc:
|
|
364
|
+
raise ValueError(f"--schema is not valid JSON: {exc.msg}") from exc
|
|
365
|
+
schema = _decision_schema(given)
|
|
366
|
+
if kind is not None and kind != schema.kind:
|
|
367
|
+
raise ValueError(f"--schema describes a {schema.kind} field, not {kind}")
|
|
368
|
+
return schema
|
|
369
|
+
|
|
370
|
+
|
|
371
|
+
def _check_answer(schema: DecisionSchema, answer: str) -> None:
|
|
372
|
+
if schema.kind == "choice":
|
|
373
|
+
if answer not in schema.options:
|
|
374
|
+
raise ValueError(f"answer {answer!r} not in labels")
|
|
375
|
+
elif schema.kind == "boolean":
|
|
376
|
+
if answer not in _BOOLEAN_ANSWERS:
|
|
377
|
+
raise ValueError(f"answer {answer!r} is not true or false")
|
|
378
|
+
else:
|
|
379
|
+
try:
|
|
380
|
+
value = float(answer)
|
|
381
|
+
except ValueError as exc:
|
|
382
|
+
raise ValueError(f"answer {answer!r} is not a number") from exc
|
|
383
|
+
# An infinity or a NaN parses and is then dropped when the dataset is built, which would
|
|
384
|
+
# cost rows the import reported as recorded.
|
|
385
|
+
if not math.isfinite(value):
|
|
386
|
+
raise ValueError(f"answer {answer!r} is not a finite number")
|
|
387
|
+
|
|
388
|
+
|
|
389
|
+
def _labelled_row(line: str, schema: DecisionSchema) -> tuple[str, str]:
|
|
390
|
+
try:
|
|
391
|
+
row = json.loads(line)
|
|
392
|
+
except json.JSONDecodeError as exc:
|
|
393
|
+
raise ValueError(f"invalid JSON: {exc.msg}") from exc
|
|
394
|
+
if not isinstance(row, dict):
|
|
395
|
+
raise ValueError("expected an object with text and answer")
|
|
396
|
+
text, answer = row.get("text"), row.get("answer")
|
|
397
|
+
if not isinstance(text, str) or not text.strip():
|
|
398
|
+
raise ValueError("text must be a non-empty string")
|
|
399
|
+
if not isinstance(answer, str):
|
|
400
|
+
raise ValueError("answer must be a string")
|
|
401
|
+
_check_answer(schema, answer)
|
|
402
|
+
return text, answer
|
|
403
|
+
|
|
404
|
+
|
|
405
|
+
def _import_captures(path: Path, site: str, schema: DecisionSchema) -> list[Capture]:
|
|
406
|
+
from stuntd.store.db import Capture
|
|
407
|
+
|
|
408
|
+
captures = []
|
|
409
|
+
# utf-8-sig so a file written by a Windows editor is not read with its BOM glued to line 1.
|
|
410
|
+
for number, line in enumerate(path.read_text(encoding="utf-8-sig").splitlines(), start=1):
|
|
411
|
+
if not line.strip():
|
|
412
|
+
continue
|
|
413
|
+
try:
|
|
414
|
+
text, answer = _labelled_row(line, schema)
|
|
415
|
+
except ValueError as exc:
|
|
416
|
+
raise _InvalidLine(f"line {number}: {exc}") from exc
|
|
417
|
+
captures.append(
|
|
418
|
+
Capture(site, schema.canonical, schema.kind, text, answer, _IMPORT_MODEL, 0, None, None)
|
|
419
|
+
)
|
|
420
|
+
return captures
|
|
421
|
+
|
|
422
|
+
|
|
423
|
+
def _import(args: argparse.Namespace) -> int:
|
|
424
|
+
from stuntd.store.db import Store
|
|
425
|
+
from stuntd.store.redact import Redactor
|
|
426
|
+
|
|
427
|
+
settings = _load(args.config, None, None)
|
|
428
|
+
_require_learning(args.config, settings)
|
|
429
|
+
site = _site_name(args.site)
|
|
430
|
+
# Nothing is trained here; the name only has to pass the whitelist it would be stored under.
|
|
431
|
+
site_dir(models_path(settings), site)
|
|
432
|
+
schema = _import_schema(site, args.kind, args.schema, args.labels)
|
|
433
|
+
try:
|
|
434
|
+
# The whole file is read before the store is opened, so a rejected row leaves no database.
|
|
435
|
+
captures = _import_captures(Path(args.file), site, schema)
|
|
436
|
+
except _InvalidLine as exc:
|
|
437
|
+
print(exc, file=sys.stderr)
|
|
438
|
+
return 1
|
|
439
|
+
store = Store(
|
|
440
|
+
database_path(settings),
|
|
441
|
+
Redactor(settings.redaction_patterns, builtin=settings.redact),
|
|
442
|
+
max_rows=settings.max_rows,
|
|
443
|
+
max_age_days=settings.max_age_days,
|
|
444
|
+
)
|
|
445
|
+
# Importing the same file twice records its rows twice; training keeps the latest row per text.
|
|
446
|
+
try:
|
|
447
|
+
for capture in captures:
|
|
448
|
+
store.record(capture)
|
|
449
|
+
finally:
|
|
450
|
+
store.close()
|
|
451
|
+
print(f"imported {len(captures)} rows into {site}")
|
|
452
|
+
return 0
|
|
453
|
+
|
|
454
|
+
|
|
455
|
+
def _site_name(site: str) -> str:
|
|
456
|
+
# A Jev question namespaced with a colon lands in a folder named with a dot, so either
|
|
457
|
+
# spelling of the name reaches the same site whichever command is given it.
|
|
458
|
+
return site.replace(":", ".")
|
|
459
|
+
|
|
460
|
+
|
|
461
|
+
def _one_model(models: Path, site: str) -> SiteModel:
|
|
462
|
+
try:
|
|
463
|
+
return load_model(models, site)
|
|
464
|
+
except FileNotFoundError as exc:
|
|
465
|
+
raise ValueError(f"no model for {site}") from exc
|
|
466
|
+
|
|
467
|
+
|
|
468
|
+
def _report(args: argparse.Namespace) -> int:
|
|
469
|
+
settings = _load(args.config, None, None)
|
|
470
|
+
_require_learning(args.config, settings)
|
|
471
|
+
models = models_path(settings)
|
|
472
|
+
site = None if args.site is None else _site_name(args.site)
|
|
473
|
+
found = list_models(models) if site is None else [_one_model(models, site)]
|
|
474
|
+
if args.json:
|
|
475
|
+
print(render_json(found))
|
|
476
|
+
return 0
|
|
477
|
+
if not found:
|
|
478
|
+
print("no trained sites yet")
|
|
479
|
+
return 0
|
|
480
|
+
print("\n\n".join(render(model, args.curve) for model in found))
|
|
481
|
+
return 0
|
|
482
|
+
|
|
483
|
+
|
|
484
|
+
def _row(
|
|
485
|
+
site: object, mode: object, captures: object, shadow: object, live: object, agreement: object
|
|
486
|
+
) -> str:
|
|
487
|
+
# The space past each pad keeps the columns apart when a value fills its width.
|
|
488
|
+
columns = zip((site, mode, captures, shadow, live), _COLUMN_WIDTHS, strict=True)
|
|
489
|
+
return "".join(f"{value:<{width - 1}} " for value, width in columns) + str(agreement)
|
|
490
|
+
|
|
491
|
+
|
|
492
|
+
def _status_line(row: _SiteStatus) -> str:
|
|
493
|
+
# A site with no model has nothing to compare or to serve, which is not the same as none yet.
|
|
494
|
+
served = row.mode != MODE_COLLECT
|
|
495
|
+
return _row(
|
|
496
|
+
row.site,
|
|
497
|
+
row.mode,
|
|
498
|
+
row.captures,
|
|
499
|
+
row.shadow if served else "-",
|
|
500
|
+
row.live if served else "-",
|
|
501
|
+
"-" if row.agreement is None else f"{row.agreement:.3f}",
|
|
502
|
+
)
|
|
503
|
+
|
|
504
|
+
|
|
505
|
+
def _compared(counts: dict[str, int]) -> int:
|
|
506
|
+
# A checked live request is answered by both sides, so it counts as a comparison too.
|
|
507
|
+
return counts.get(MODE_SHADOW, 0) + counts.get(MODE_CHECK, 0)
|
|
508
|
+
|
|
509
|
+
|
|
510
|
+
def _agreement(store: Store, site: str, state: SiteState | None, window: int) -> float | None:
|
|
511
|
+
if state is None or state.model is None or state.model.threshold is None:
|
|
512
|
+
return None
|
|
513
|
+
rows = store.comparisons(site, window, window_start(state))
|
|
514
|
+
return window_agreement(rows, state.model.threshold).agreement
|
|
515
|
+
|
|
516
|
+
|
|
517
|
+
def _row_without_captures(state: SiteState) -> _SiteStatus:
|
|
518
|
+
return _SiteStatus(state.site, state.mode, 0, 0, 0, None)
|
|
519
|
+
|
|
520
|
+
|
|
521
|
+
def _status_rows(settings: Settings) -> list[_SiteStatus]:
|
|
522
|
+
from stuntd.store.db import Store
|
|
523
|
+
from stuntd.store.redact import Redactor
|
|
524
|
+
|
|
525
|
+
states = {state.site: state for state in site_states(models_path(settings))}
|
|
526
|
+
database = database_path(settings)
|
|
527
|
+
if not database.exists():
|
|
528
|
+
# Opening a store would create the database that status is only here to read.
|
|
529
|
+
return [_row_without_captures(state) for state in states.values()]
|
|
530
|
+
store = Store(database, Redactor())
|
|
531
|
+
try:
|
|
532
|
+
captures = {info.site: info.count for info in store.sites()}
|
|
533
|
+
counts = store.decision_counts()
|
|
534
|
+
rows = []
|
|
535
|
+
for site in sorted(set(captures) | set(states)):
|
|
536
|
+
state = states.get(site)
|
|
537
|
+
decisions = counts.get(site, {})
|
|
538
|
+
rows.append(
|
|
539
|
+
_SiteStatus(
|
|
540
|
+
site=site,
|
|
541
|
+
mode=MODE_COLLECT if state is None else state.mode,
|
|
542
|
+
captures=captures.get(site, 0),
|
|
543
|
+
shadow=_compared(decisions),
|
|
544
|
+
live=decisions.get(MODE_LIVE, 0),
|
|
545
|
+
agreement=_agreement(store, site, state, settings.window),
|
|
546
|
+
)
|
|
547
|
+
)
|
|
548
|
+
return rows
|
|
549
|
+
finally:
|
|
550
|
+
store.close()
|
|
551
|
+
|
|
552
|
+
|
|
553
|
+
def _config_file(explicit: str | None) -> Path:
|
|
554
|
+
if explicit is None:
|
|
555
|
+
return config_path()
|
|
556
|
+
path = Path(explicit)
|
|
557
|
+
if path.is_dir():
|
|
558
|
+
raise ValueError(f"{path} is not a file")
|
|
559
|
+
if not path.is_file():
|
|
560
|
+
raise ValueError(f"{path} does not exist")
|
|
561
|
+
return path
|
|
562
|
+
|
|
563
|
+
|
|
564
|
+
def _load(config: str | None, upstream: str | None, port: int | None) -> Settings:
|
|
565
|
+
return load_settings(_config_file(config), {"upstream": upstream, "port": port})
|
|
566
|
+
|
|
567
|
+
|
|
568
|
+
def _build_parser() -> argparse.ArgumentParser:
|
|
569
|
+
parser = argparse.ArgumentParser(
|
|
570
|
+
prog="stuntd", description="Local proxy that records typed LLM decisions."
|
|
571
|
+
)
|
|
572
|
+
commands = parser.add_subparsers(dest="command", required=True)
|
|
573
|
+
|
|
574
|
+
serve = commands.add_parser("serve", help="run the proxy until it is interrupted")
|
|
575
|
+
serve.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
576
|
+
serve.add_argument("--upstream", metavar="URL", help="provider base URL to forward to")
|
|
577
|
+
serve.add_argument("--port", type=int, metavar="N", help="port to listen on")
|
|
578
|
+
serve.set_defaults(handler=_serve)
|
|
579
|
+
|
|
580
|
+
status = commands.add_parser("status", help="summarise the captures recorded so far")
|
|
581
|
+
status.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
582
|
+
status.add_argument("--json", action="store_true", help="print the rows as JSON")
|
|
583
|
+
status.set_defaults(handler=_status)
|
|
584
|
+
|
|
585
|
+
train = commands.add_parser("train", help="train a model for each site with enough captures")
|
|
586
|
+
train.add_argument(
|
|
587
|
+
"sites", nargs="*", metavar="SITE", help="sites to train, every known site by default"
|
|
588
|
+
)
|
|
589
|
+
train.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
590
|
+
train.set_defaults(handler=_train)
|
|
591
|
+
|
|
592
|
+
report = commands.add_parser("report", help="print how the trained models did")
|
|
593
|
+
report.add_argument(
|
|
594
|
+
"site", nargs="?", metavar="SITE", help="site to report on, every trained site by default"
|
|
595
|
+
)
|
|
596
|
+
report.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
597
|
+
report.add_argument("--json", action="store_true", help="print the models as JSON")
|
|
598
|
+
report.add_argument("--curve", action="store_true", help="add the whole threshold curve")
|
|
599
|
+
report.set_defaults(handler=_report)
|
|
600
|
+
|
|
601
|
+
enable = commands.add_parser("enable", help="let a trained site answer from its own model")
|
|
602
|
+
enable.add_argument("site", metavar="SITE", help="site to serve locally")
|
|
603
|
+
enable.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
604
|
+
enable.set_defaults(handler=_enable)
|
|
605
|
+
|
|
606
|
+
disable = commands.add_parser("disable", help="send a site back to the provider")
|
|
607
|
+
disable.add_argument("site", metavar="SITE", help="site to stop serving locally")
|
|
608
|
+
disable.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
609
|
+
disable.set_defaults(handler=_disable)
|
|
610
|
+
|
|
611
|
+
importer = commands.add_parser(
|
|
612
|
+
"import",
|
|
613
|
+
help="record labelled examples for a site",
|
|
614
|
+
description=(
|
|
615
|
+
"Record labelled examples for a site from a JSONL file. Each text is stored and"
|
|
616
|
+
" served exactly as written, so it has to be spelled the way the site sees its"
|
|
617
|
+
" input: for a chat site the 'user: ...' lines the store holds, for a Jev site the"
|
|
618
|
+
" serialised state."
|
|
619
|
+
),
|
|
620
|
+
)
|
|
621
|
+
importer.add_argument("site", metavar="SITE", help="site to record the examples under")
|
|
622
|
+
importer.add_argument(
|
|
623
|
+
"file",
|
|
624
|
+
metavar="FILE",
|
|
625
|
+
help='JSONL file of {"text": ..., "answer": ...} rows, each text written the way the'
|
|
626
|
+
" site sees its input: 'user: ...' lines for a chat site, the serialised state for Jev",
|
|
627
|
+
)
|
|
628
|
+
importer.add_argument("--kind", choices=_KINDS, help="what the answers are, choice by default")
|
|
629
|
+
importer.add_argument("--schema", metavar="JSON", help="JSON schema of the answered field")
|
|
630
|
+
importer.add_argument(
|
|
631
|
+
"--labels", metavar="A,B,C", help="answers a choice site may carry, instead of --schema"
|
|
632
|
+
)
|
|
633
|
+
importer.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
634
|
+
importer.set_defaults(handler=_import)
|
|
635
|
+
|
|
636
|
+
config = commands.add_parser("config", help="create or inspect the settings file")
|
|
637
|
+
config_commands = config.add_subparsers(dest="config_command", required=True)
|
|
638
|
+
|
|
639
|
+
init = config_commands.add_parser("init", help="write a commented settings file")
|
|
640
|
+
init.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
641
|
+
init.add_argument("--force", action="store_true", help="overwrite an existing file")
|
|
642
|
+
init.set_defaults(handler=_config_init)
|
|
643
|
+
|
|
644
|
+
show = config_commands.add_parser("show", help="print every setting with its source")
|
|
645
|
+
show.add_argument("--config", metavar="PATH", help=_CONFIG_HELP)
|
|
646
|
+
show.add_argument("--upstream", metavar="URL", help="provider base URL to forward to")
|
|
647
|
+
show.add_argument("--port", type=int, metavar="N", help="port to listen on")
|
|
648
|
+
show.set_defaults(handler=_config_show)
|
|
649
|
+
return parser
|
|
650
|
+
|
|
651
|
+
|
|
652
|
+
def main(argv: list[str] | None = None) -> int:
|
|
653
|
+
"""Parses one command line and runs the command it names, returning the exit code."""
|
|
654
|
+
args = _build_parser().parse_args(argv)
|
|
655
|
+
handler: Callable[[argparse.Namespace], int] = args.handler
|
|
656
|
+
try:
|
|
657
|
+
return handler(args)
|
|
658
|
+
except _LearningOff as exc:
|
|
659
|
+
print(f"stuntd: {exc}", file=sys.stderr)
|
|
660
|
+
return 1
|
|
661
|
+
# A rejected setting, an unwritable path, a database that will not open and a failed training
|
|
662
|
+
# run are all things the user has to fix, so they are reported rather than raised as a traceback.
|
|
663
|
+
except (ValueError, OSError, RuntimeError, sqlite3.DatabaseError) as exc:
|
|
664
|
+
print(f"stuntd: {exc}", file=sys.stderr)
|
|
665
|
+
return 2
|
|
File without changes
|