jevstiller 0.2.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.
- jevstiller/__init__.py +19 -0
- jevstiller/admin.py +232 -0
- jevstiller/backup.py +112 -0
- jevstiller/calibrate.py +133 -0
- jevstiller/cli.py +307 -0
- jevstiller/core.py +948 -0
- jevstiller/encoders/__init__.py +57 -0
- jevstiller/encoders/batching.py +121 -0
- jevstiller/encoders/hashing.py +43 -0
- jevstiller/encoders/hf.py +49 -0
- jevstiller/encoders/onnx.py +72 -0
- jevstiller/manager.py +666 -0
- jevstiller/metrics.py +151 -0
- jevstiller/ood.py +73 -0
- jevstiller/py.typed +0 -0
- jevstiller/registry.py +196 -0
- jevstiller/scheduler.py +208 -0
- jevstiller/server.py +845 -0
- jevstiller/settings.py +315 -0
- jevstiller/store.py +448 -0
- jevstiller/student.py +104 -0
- jevstiller/task.py +154 -0
- jevstiller/teachers/__init__.py +49 -0
- jevstiller/teachers/jev.py +92 -0
- jevstiller/teachers/replay.py +78 -0
- jevstiller/teachers/synthetic.py +94 -0
- jevstiller/training.py +203 -0
- jevstiller-0.2.0.dist-info/METADATA +270 -0
- jevstiller-0.2.0.dist-info/RECORD +32 -0
- jevstiller-0.2.0.dist-info/WHEEL +4 -0
- jevstiller-0.2.0.dist-info/entry_points.txt +2 -0
- jevstiller-0.2.0.dist-info/licenses/LICENSE +21 -0
jevstiller/__init__.py
ADDED
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
from importlib.metadata import PackageNotFoundError
|
|
2
|
+
from importlib.metadata import version as _version
|
|
3
|
+
|
|
4
|
+
from .core import Jevstiller, Result, Status, TeacherError, TrainReport
|
|
5
|
+
from .encoders import BatchingEncoder, Encoder, HashEncoder, load_encoder
|
|
6
|
+
from .manager import Admission, TaskManager
|
|
7
|
+
from .scheduler import TrainScheduler
|
|
8
|
+
from .task import Config, Task
|
|
9
|
+
from .teachers import CachedTeacher, ReplayTeacher, SyntheticTeacher, SyntheticWorld, Teacher, TeacherOutput
|
|
10
|
+
|
|
11
|
+
__all__ = ["Task", "Config", "Jevstiller", "Result", "Status", "TrainReport", "TeacherError", "Teacher",
|
|
12
|
+
"TeacherOutput", "SyntheticTeacher", "SyntheticWorld", "ReplayTeacher", "CachedTeacher", "Encoder",
|
|
13
|
+
"HashEncoder", "BatchingEncoder", "load_encoder", "TaskManager", "Admission",
|
|
14
|
+
"TrainScheduler"]
|
|
15
|
+
|
|
16
|
+
try:
|
|
17
|
+
__version__ = _version("jevstiller")
|
|
18
|
+
except PackageNotFoundError: # pragma: no cover - running from a source tree that was never installed
|
|
19
|
+
__version__ = "0.0.0+unknown"
|
jevstiller/admin.py
ADDED
|
@@ -0,0 +1,232 @@
|
|
|
1
|
+
"""Admin API, under `/jevstiller/v1/` (never forwarded to Jev). Off unless an admin token is configured; every
|
|
2
|
+
call needs `Authorization: Bearer <admin token>`.
|
|
3
|
+
|
|
4
|
+
GET /jevstiller/v1/tasks[?tenant=] list tasks
|
|
5
|
+
GET /jevstiller/v1/tasks/{key} status (+ text report) of one task
|
|
6
|
+
GET /jevstiller/v1/tasks/{key}/versions its student versions
|
|
7
|
+
POST /jevstiller/v1/tasks/{key}/mode {"mode": "auto" | "teacher_only" | "cascade"} (persisted)
|
|
8
|
+
POST /jevstiller/v1/tasks/{key}/target {"target_agreement": 0.99} (persisted)
|
|
9
|
+
POST /jevstiller/v1/tasks/{key}/train train a candidate now (waits for it)
|
|
10
|
+
POST /jevstiller/v1/tasks/{key}/promote {"version": "student:v3"}
|
|
11
|
+
POST /jevstiller/v1/tasks/{key}/rollback
|
|
12
|
+
DELETE /jevstiller/v1/tasks/{key} delete the task and all its data
|
|
13
|
+
DELETE /jevstiller/v1/tenants/{tenant} delete every task of a tenant
|
|
14
|
+
GET /jevstiller/v1/stats proxy, manager and scheduler counters
|
|
15
|
+
"""
|
|
16
|
+
from __future__ import annotations
|
|
17
|
+
|
|
18
|
+
import dataclasses
|
|
19
|
+
import hmac
|
|
20
|
+
import math
|
|
21
|
+
import re
|
|
22
|
+
from collections.abc import Callable
|
|
23
|
+
from typing import Any
|
|
24
|
+
|
|
25
|
+
import anyio
|
|
26
|
+
from starlette.requests import Request
|
|
27
|
+
from starlette.responses import JSONResponse, Response
|
|
28
|
+
from starlette.routing import Route
|
|
29
|
+
|
|
30
|
+
from .manager import TaskManager
|
|
31
|
+
|
|
32
|
+
PREFIX = "/jevstiller/v1"
|
|
33
|
+
_KEY = re.compile(r"^[0-9a-f]{20}$")
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
def _json_safe(x: Any) -> Any:
|
|
37
|
+
if isinstance(x, float) and (math.isnan(x) or math.isinf(x)):
|
|
38
|
+
return None
|
|
39
|
+
if isinstance(x, dict):
|
|
40
|
+
return {str(k): _json_safe(v) for k, v in x.items()}
|
|
41
|
+
if isinstance(x, (list, tuple)):
|
|
42
|
+
return [_json_safe(v) for v in x]
|
|
43
|
+
if dataclasses.is_dataclass(x) and not isinstance(x, type):
|
|
44
|
+
return _json_safe(dataclasses.asdict(x))
|
|
45
|
+
return x
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _ok(data: Any, status: int = 200) -> JSONResponse:
|
|
49
|
+
return JSONResponse(_json_safe(data), status_code=status)
|
|
50
|
+
|
|
51
|
+
|
|
52
|
+
def _err(status: int, detail: str) -> JSONResponse:
|
|
53
|
+
return JSONResponse({"detail": detail}, status_code=status)
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
class Admin:
|
|
57
|
+
def __init__(self, manager: TaskManager, token: str | None, stats: Callable[[], dict] | None = None):
|
|
58
|
+
self.manager, self.token, self.stats_fn = manager, token, stats
|
|
59
|
+
|
|
60
|
+
def authorized(self, request: Request) -> Response | None:
|
|
61
|
+
"""None when the call may proceed, else the response to send."""
|
|
62
|
+
if not self.token:
|
|
63
|
+
return _err(404, "not found") # the admin API is off
|
|
64
|
+
auth = request.headers.get("authorization", "")
|
|
65
|
+
given = auth[7:].strip() if auth.lower().startswith("bearer ") else ""
|
|
66
|
+
if not hmac.compare_digest(given.encode(), self.token.encode()):
|
|
67
|
+
return _err(401, "admin token required")
|
|
68
|
+
return None
|
|
69
|
+
|
|
70
|
+
def _key(self, request: Request) -> str | None:
|
|
71
|
+
key = request.path_params.get("key", "")
|
|
72
|
+
return key if _KEY.match(key) and key in {i.key for i in self.manager.tasks()} else None
|
|
73
|
+
|
|
74
|
+
async def _with_engine(self, key: str, fn: Callable[[Any], Any]) -> Any:
|
|
75
|
+
"""`fn(engine)` in a worker thread with the task's engine held open throughout (an unload or delete
|
|
76
|
+
closing it under a running call crashed the process: security audit run 2). None if the task is gone."""
|
|
77
|
+
def run():
|
|
78
|
+
try:
|
|
79
|
+
with self.manager.hold(key) as engine:
|
|
80
|
+
return fn(engine)
|
|
81
|
+
except KeyError: # deleted meanwhile
|
|
82
|
+
return None
|
|
83
|
+
return await anyio.to_thread.run_sync(run)
|
|
84
|
+
|
|
85
|
+
async def _body(self, request: Request) -> dict:
|
|
86
|
+
try:
|
|
87
|
+
body = await request.json()
|
|
88
|
+
except ValueError:
|
|
89
|
+
return {}
|
|
90
|
+
return body if isinstance(body, dict) else {}
|
|
91
|
+
|
|
92
|
+
# ---- handlers -----------------------------------------------------------------
|
|
93
|
+
async def tasks(self, request: Request) -> Response:
|
|
94
|
+
if (denied := self.authorized(request)) is not None:
|
|
95
|
+
return denied
|
|
96
|
+
loaded = set(self.manager.loaded())
|
|
97
|
+
tenant = request.query_params.get("tenant")
|
|
98
|
+
return _ok([{"key": i.key, "tenant": i.tenant, "model": i.model, "classes": len(i.classes),
|
|
99
|
+
"target_agreement": i.target_agreement, "mode": i.mode, "created": i.created,
|
|
100
|
+
"last_seen": i.last_seen, "loaded": i.key in loaded} for i in self.manager.tasks(tenant)])
|
|
101
|
+
|
|
102
|
+
async def task(self, request: Request) -> Response:
|
|
103
|
+
if (denied := self.authorized(request)) is not None:
|
|
104
|
+
return denied
|
|
105
|
+
if (key := self._key(request)) is None:
|
|
106
|
+
return _err(404, "unknown task")
|
|
107
|
+
st = await self._with_engine(key, lambda e: e.status())
|
|
108
|
+
info = next((i for i in self.manager.tasks() if i.key == key), None)
|
|
109
|
+
if st is None or info is None:
|
|
110
|
+
return _err(404, "unknown task")
|
|
111
|
+
return _ok({"info": info, "status": st, "report": st.report()})
|
|
112
|
+
|
|
113
|
+
async def versions(self, request: Request) -> Response:
|
|
114
|
+
if (denied := self.authorized(request)) is not None:
|
|
115
|
+
return denied
|
|
116
|
+
if (key := self._key(request)) is None:
|
|
117
|
+
return _err(404, "unknown task")
|
|
118
|
+
versions = await self._with_engine(key, lambda e: e.versions())
|
|
119
|
+
return _ok(versions) if versions is not None else _err(404, "unknown task")
|
|
120
|
+
|
|
121
|
+
async def mode(self, request: Request) -> Response:
|
|
122
|
+
if (denied := self.authorized(request)) is not None:
|
|
123
|
+
return denied
|
|
124
|
+
if (key := self._key(request)) is None:
|
|
125
|
+
return _err(404, "unknown task")
|
|
126
|
+
mode = (await self._body(request)).get("mode")
|
|
127
|
+
try:
|
|
128
|
+
await anyio.to_thread.run_sync(self.manager.set_mode, key, mode)
|
|
129
|
+
except ValueError as e:
|
|
130
|
+
return _err(422, str(e))
|
|
131
|
+
return _ok({"key": key, "mode": mode})
|
|
132
|
+
|
|
133
|
+
async def target(self, request: Request) -> Response:
|
|
134
|
+
if (denied := self.authorized(request)) is not None:
|
|
135
|
+
return denied
|
|
136
|
+
if (key := self._key(request)) is None:
|
|
137
|
+
return _err(404, "unknown task")
|
|
138
|
+
value = (await self._body(request)).get("target_agreement")
|
|
139
|
+
if not isinstance(value, (int, float)) or isinstance(value, bool):
|
|
140
|
+
return _err(422, "target_agreement must be a number in [0.5, 1)")
|
|
141
|
+
try:
|
|
142
|
+
await anyio.to_thread.run_sync(self.manager.set_target, key, float(value))
|
|
143
|
+
except ValueError as e:
|
|
144
|
+
return _err(422, str(e))
|
|
145
|
+
return _ok({"key": key, "target_agreement": float(value)})
|
|
146
|
+
|
|
147
|
+
async def train(self, request: Request) -> Response:
|
|
148
|
+
if (denied := self.authorized(request)) is not None:
|
|
149
|
+
return denied
|
|
150
|
+
if (key := self._key(request)) is None:
|
|
151
|
+
return _err(404, "unknown task")
|
|
152
|
+
try:
|
|
153
|
+
report = await self._with_engine(key, lambda e: e.train_now())
|
|
154
|
+
except Exception as e:
|
|
155
|
+
return _err(500, f"training failed: {type(e).__name__}")
|
|
156
|
+
return _ok(report) if report is not None else _err(404, "unknown task")
|
|
157
|
+
|
|
158
|
+
async def promote(self, request: Request) -> Response:
|
|
159
|
+
if (denied := self.authorized(request)) is not None:
|
|
160
|
+
return denied
|
|
161
|
+
if (key := self._key(request)) is None:
|
|
162
|
+
return _err(404, "unknown task")
|
|
163
|
+
version = (await self._body(request)).get("version")
|
|
164
|
+
try:
|
|
165
|
+
if await self._with_engine(key, lambda e: e.promote(str(version)) or True) is None:
|
|
166
|
+
return _err(404, "unknown task")
|
|
167
|
+
except ValueError as e:
|
|
168
|
+
return _err(422, str(e))
|
|
169
|
+
return _ok({"key": key, "production": version})
|
|
170
|
+
|
|
171
|
+
async def rollback(self, request: Request) -> Response:
|
|
172
|
+
if (denied := self.authorized(request)) is not None:
|
|
173
|
+
return denied
|
|
174
|
+
if (key := self._key(request)) is None:
|
|
175
|
+
return _err(404, "unknown task")
|
|
176
|
+
target = await self._with_engine(key, lambda e: (e.rollback(),))
|
|
177
|
+
if target is None:
|
|
178
|
+
return _err(404, "unknown task")
|
|
179
|
+
return _ok({"key": key, "production": target[0]})
|
|
180
|
+
|
|
181
|
+
async def delete_task(self, request: Request) -> Response:
|
|
182
|
+
if (denied := self.authorized(request)) is not None:
|
|
183
|
+
return denied
|
|
184
|
+
if (key := self._key(request)) is None:
|
|
185
|
+
return _err(404, "unknown task")
|
|
186
|
+
if not await anyio.to_thread.run_sync(self.manager.delete, key, "admin API"):
|
|
187
|
+
return _err(409, "a request is in flight for this task; retry")
|
|
188
|
+
return _ok({"deleted": key})
|
|
189
|
+
|
|
190
|
+
async def delete_tenant(self, request: Request) -> Response:
|
|
191
|
+
if (denied := self.authorized(request)) is not None:
|
|
192
|
+
return denied
|
|
193
|
+
tenant = request.path_params["tenant"]
|
|
194
|
+
deleted = await anyio.to_thread.run_sync(self.manager.delete_tenant, tenant)
|
|
195
|
+
left = [i.key for i in self.manager.tasks(tenant)]
|
|
196
|
+
return _ok({"deleted": deleted, "remaining": left}, 200 if not left else 409)
|
|
197
|
+
|
|
198
|
+
async def stats(self, request: Request) -> Response:
|
|
199
|
+
if (denied := self.authorized(request)) is not None:
|
|
200
|
+
return denied
|
|
201
|
+
m = self.manager
|
|
202
|
+
data = {"tasks": len(m.tasks()), "loaded": len(m.loaded()), "loads": m.loads, "unloads": m.unloads,
|
|
203
|
+
"student_memory_mb": m.memory_mb()}
|
|
204
|
+
sched = getattr(m.train_executor, "stats", None)
|
|
205
|
+
if callable(sched):
|
|
206
|
+
data["training"] = sched()
|
|
207
|
+
if self.stats_fn:
|
|
208
|
+
data.update(self.stats_fn())
|
|
209
|
+
return _ok(data)
|
|
210
|
+
|
|
211
|
+
def routes(self) -> list[Route]:
|
|
212
|
+
p = PREFIX
|
|
213
|
+
return [
|
|
214
|
+
Route(f"{p}/tasks", self.tasks, methods=["GET"]),
|
|
215
|
+
Route(f"{p}/tasks/{{key}}", self.task, methods=["GET"]),
|
|
216
|
+
Route(f"{p}/tasks/{{key}}", self.delete_task, methods=["DELETE"]),
|
|
217
|
+
Route(f"{p}/tasks/{{key}}/versions", self.versions, methods=["GET"]),
|
|
218
|
+
Route(f"{p}/tasks/{{key}}/mode", self.mode, methods=["POST"]),
|
|
219
|
+
Route(f"{p}/tasks/{{key}}/target", self.target, methods=["POST"]),
|
|
220
|
+
Route(f"{p}/tasks/{{key}}/train", self.train, methods=["POST"]),
|
|
221
|
+
Route(f"{p}/tasks/{{key}}/promote", self.promote, methods=["POST"]),
|
|
222
|
+
Route(f"{p}/tasks/{{key}}/rollback", self.rollback, methods=["POST"]),
|
|
223
|
+
Route(f"{p}/tenants/{{tenant}}", self.delete_tenant, methods=["DELETE"]),
|
|
224
|
+
Route(f"{p}/stats", self.stats, methods=["GET"]),
|
|
225
|
+
Route(f"{p}/{{rest:path}}", self._not_found,
|
|
226
|
+
methods=["GET", "POST", "PUT", "PATCH", "DELETE", "HEAD", "OPTIONS"]),
|
|
227
|
+
]
|
|
228
|
+
|
|
229
|
+
async def _not_found(self, request: Request) -> Response:
|
|
230
|
+
if (denied := self.authorized(request)) is not None:
|
|
231
|
+
return denied
|
|
232
|
+
return _err(404, "not found")
|
jevstiller/backup.py
ADDED
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
"""`jevstiller backup` / `jevstiller restore`: a consistent copy of a data directory, safe while serving.
|
|
2
|
+
|
|
3
|
+
Per task: SQLite's online backup API for `samples.sqlite` (a consistent snapshot even with the writer running),
|
|
4
|
+
`task.json`, and `versions/` (the registry index is copied first, so every version it lists is present in the
|
|
5
|
+
copy). The deployment's `key-salt` is copied too: without it, key hashes in a tenants file would no longer match.
|
|
6
|
+
|
|
7
|
+
Backups hold request text like the data dir, so they get the same modes (files 0600, directories 0700), and a
|
|
8
|
+
restored data dir is set to them too, whatever the source's modes (security audit run 2).
|
|
9
|
+
"""
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import contextlib
|
|
13
|
+
import json
|
|
14
|
+
import os
|
|
15
|
+
import shutil
|
|
16
|
+
import sqlite3
|
|
17
|
+
import time
|
|
18
|
+
from collections.abc import Iterator
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@contextlib.contextmanager
|
|
23
|
+
def _private() -> Iterator[None]:
|
|
24
|
+
"""Everything created inside is owner-only (umask 077), as `jevstiller serve` does."""
|
|
25
|
+
old = os.umask(0o077)
|
|
26
|
+
try:
|
|
27
|
+
yield
|
|
28
|
+
finally:
|
|
29
|
+
os.umask(old)
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
def _lock_down(root: Path) -> None:
|
|
33
|
+
"""Owner-only modes for `root` and everything below it (copies keep their source's modes)."""
|
|
34
|
+
if not root.exists():
|
|
35
|
+
return
|
|
36
|
+
os.chmod(root, 0o700 if root.is_dir() else 0o600)
|
|
37
|
+
for dirpath, dirnames, filenames in os.walk(root):
|
|
38
|
+
for d in dirnames:
|
|
39
|
+
os.chmod(os.path.join(dirpath, d), 0o700)
|
|
40
|
+
for f in filenames:
|
|
41
|
+
os.chmod(os.path.join(dirpath, f), 0o600)
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _copy_versions(src: Path, dst: Path) -> None:
|
|
45
|
+
dst.mkdir(parents=True, exist_ok=True)
|
|
46
|
+
if (src / "registry.json").exists():
|
|
47
|
+
shutil.copy2(src / "registry.json", dst / "registry.json")
|
|
48
|
+
for d in src.iterdir():
|
|
49
|
+
if d.is_dir() and not d.name.startswith("."): # skip staging dirs of running jobs
|
|
50
|
+
shutil.copytree(d, dst / d.name, dirs_exist_ok=True)
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
def backup(data_dir: str | Path, out: str | Path) -> dict:
|
|
54
|
+
with _private():
|
|
55
|
+
manifest = _backup(Path(data_dir), Path(out))
|
|
56
|
+
_lock_down(Path(out))
|
|
57
|
+
return manifest
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def _backup(src: Path, dst: Path) -> dict:
|
|
61
|
+
if dst.exists() and any(dst.iterdir()):
|
|
62
|
+
raise FileExistsError(f"{dst} is not empty")
|
|
63
|
+
dst.mkdir(parents=True, exist_ok=True)
|
|
64
|
+
tasks = 0
|
|
65
|
+
for tdir in sorted((src / "tasks").glob("*")) if (src / "tasks").exists() else []:
|
|
66
|
+
if not (tdir / "task.json").exists():
|
|
67
|
+
continue
|
|
68
|
+
out_t = dst / "tasks" / tdir.name
|
|
69
|
+
out_t.mkdir(parents=True)
|
|
70
|
+
shutil.copy2(tdir / "task.json", out_t / "task.json")
|
|
71
|
+
if (tdir / "samples.sqlite").exists():
|
|
72
|
+
s = sqlite3.connect(f"{(tdir / 'samples.sqlite').resolve().as_uri()}?mode=ro", uri=True)
|
|
73
|
+
d = sqlite3.connect(out_t / "samples.sqlite")
|
|
74
|
+
try:
|
|
75
|
+
s.backup(d)
|
|
76
|
+
finally:
|
|
77
|
+
d.close()
|
|
78
|
+
s.close()
|
|
79
|
+
if (tdir / "versions").exists():
|
|
80
|
+
_copy_versions(tdir / "versions", out_t / "versions")
|
|
81
|
+
tasks += 1
|
|
82
|
+
if (src / "key-salt").exists():
|
|
83
|
+
shutil.copy2(src / "key-salt", dst / "key-salt")
|
|
84
|
+
manifest = {"created": time.time(), "source": str(src.resolve()), "tasks": tasks}
|
|
85
|
+
(dst / "backup.json").write_text(json.dumps(manifest, indent=2))
|
|
86
|
+
return manifest
|
|
87
|
+
|
|
88
|
+
|
|
89
|
+
def restore(backup_dir: str | Path, data_dir: str | Path, force: bool = False) -> dict:
|
|
90
|
+
"""Copy a backup into an empty data directory (or over an existing one with `force`). Stop the server first."""
|
|
91
|
+
with _private():
|
|
92
|
+
manifest = _restore(Path(backup_dir), Path(data_dir), force)
|
|
93
|
+
dst = Path(data_dir)
|
|
94
|
+
os.chmod(dst, 0o700)
|
|
95
|
+
_lock_down(dst / "tasks")
|
|
96
|
+
_lock_down(dst / "key-salt")
|
|
97
|
+
return manifest
|
|
98
|
+
|
|
99
|
+
|
|
100
|
+
def _restore(src: Path, dst: Path, force: bool) -> dict:
|
|
101
|
+
if not (src / "backup.json").exists():
|
|
102
|
+
raise FileNotFoundError(f"{src} is not a jevstiller backup (no backup.json)")
|
|
103
|
+
if dst.exists() and any(dst.iterdir()) and not force:
|
|
104
|
+
raise FileExistsError(f"{dst} is not empty (use force to overwrite)")
|
|
105
|
+
dst.mkdir(parents=True, exist_ok=True)
|
|
106
|
+
if (dst / "tasks").exists() and force:
|
|
107
|
+
shutil.rmtree(dst / "tasks")
|
|
108
|
+
if (src / "tasks").exists():
|
|
109
|
+
shutil.copytree(src / "tasks", dst / "tasks")
|
|
110
|
+
if (src / "key-salt").exists():
|
|
111
|
+
shutil.copy2(src / "key-salt", dst / "key-salt")
|
|
112
|
+
return json.loads((src / "backup.json").read_text())
|
jevstiller/calibrate.py
ADDED
|
@@ -0,0 +1,133 @@
|
|
|
1
|
+
"""Threshold selection under a disagreement budget, with a finite-sample bound."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import json
|
|
5
|
+
import math
|
|
6
|
+
from collections.abc import Sequence
|
|
7
|
+
from dataclasses import asdict, dataclass, field
|
|
8
|
+
from pathlib import Path
|
|
9
|
+
|
|
10
|
+
import numpy as np
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _log_binom_cdf(k: int, n: int, p: float) -> float:
|
|
14
|
+
"""log P(X <= k) for X ~ Binomial(n, p)."""
|
|
15
|
+
if p <= 0:
|
|
16
|
+
return 0.0
|
|
17
|
+
if p >= 1:
|
|
18
|
+
return 0.0 if k >= n else -math.inf
|
|
19
|
+
lp, lq = math.log(p), math.log1p(-p)
|
|
20
|
+
terms = [math.lgamma(n + 1) - math.lgamma(i + 1) - math.lgamma(n - i + 1) + i * lp + (n - i) * lq
|
|
21
|
+
for i in range(k + 1)]
|
|
22
|
+
m = max(terms)
|
|
23
|
+
return m + math.log(sum(math.exp(t - m) for t in terms))
|
|
24
|
+
|
|
25
|
+
|
|
26
|
+
def clopper_pearson_upper(k: int, n: int, delta: float = 0.05) -> float:
|
|
27
|
+
"""One-sided upper confidence bound on a binomial proportion: smallest p with P(X<=k|n,p) <= delta."""
|
|
28
|
+
if n <= 0:
|
|
29
|
+
return 1.0
|
|
30
|
+
if k >= n:
|
|
31
|
+
return 1.0
|
|
32
|
+
lo, hi = k / n, 1.0
|
|
33
|
+
log_delta = math.log(delta)
|
|
34
|
+
for _ in range(60):
|
|
35
|
+
mid = (lo + hi) / 2
|
|
36
|
+
if _log_binom_cdf(k, n, mid) > log_delta:
|
|
37
|
+
lo = mid
|
|
38
|
+
else:
|
|
39
|
+
hi = mid
|
|
40
|
+
return hi
|
|
41
|
+
|
|
42
|
+
|
|
43
|
+
def clopper_pearson_lower(k: int, n: int, delta: float = 0.05) -> float:
|
|
44
|
+
"""One-sided lower confidence bound on a binomial proportion."""
|
|
45
|
+
if n <= 0 or k <= 0:
|
|
46
|
+
return 0.0
|
|
47
|
+
return 1.0 - clopper_pearson_upper(n - k, n, delta)
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
@dataclass
|
|
51
|
+
class RoutingPolicy:
|
|
52
|
+
conf_threshold: float | None # None -> student never answers
|
|
53
|
+
ood_threshold: float
|
|
54
|
+
expected_coverage: float # share of calibration rows the student answers
|
|
55
|
+
disagreement_ub: float # upper bound (at 1 - delta) on P(student answers and disagrees), per request
|
|
56
|
+
expected_system_disagreement: float # that rate on the calibration rows (the bound is what meets the budget)
|
|
57
|
+
budget: float
|
|
58
|
+
delta: float
|
|
59
|
+
n_calib: int
|
|
60
|
+
deferred_labels: list[str] = field(default_factory=list) # the student never answers these (rare classes)
|
|
61
|
+
|
|
62
|
+
@property
|
|
63
|
+
def usable(self) -> bool:
|
|
64
|
+
return self.conf_threshold is not None
|
|
65
|
+
|
|
66
|
+
def accepts(self, conf: np.ndarray, ood: np.ndarray, pred: np.ndarray | None = None) -> np.ndarray:
|
|
67
|
+
"""Which rows the student answers. `pred` (the student's predicted labels) is required when the
|
|
68
|
+
policy defers labels."""
|
|
69
|
+
if self.conf_threshold is None:
|
|
70
|
+
return np.zeros(len(conf), dtype=bool)
|
|
71
|
+
ok = (conf >= self.conf_threshold) & (ood <= self.ood_threshold)
|
|
72
|
+
if self.deferred_labels:
|
|
73
|
+
if pred is None:
|
|
74
|
+
raise ValueError("this policy defers labels; pass the predicted labels")
|
|
75
|
+
ok &= ~np.isin(np.asarray(pred, dtype=object), self.deferred_labels)
|
|
76
|
+
return ok
|
|
77
|
+
|
|
78
|
+
def save(self, path: Path) -> None:
|
|
79
|
+
path.write_text(json.dumps(asdict(self), indent=2))
|
|
80
|
+
|
|
81
|
+
@classmethod
|
|
82
|
+
def load(cls, path: Path) -> RoutingPolicy:
|
|
83
|
+
return cls(**json.loads(path.read_text()))
|
|
84
|
+
|
|
85
|
+
|
|
86
|
+
def threshold_grid(n_labels: int, size: int = 400) -> np.ndarray:
|
|
87
|
+
"""Candidate confidence thresholds, strictest first. Fixed before any calibration row is seen, so testing
|
|
88
|
+
them in sequence costs no confidence. Dense near 1, where a student's usable answers are."""
|
|
89
|
+
floor = 1.0 / max(n_labels, 2)
|
|
90
|
+
return 1.0 - np.geomspace(1e-4, 1.0 - floor, size)
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
def fit_policy(conf: np.ndarray, agree: np.ndarray, ood: np.ndarray, budget: float, delta: float = 0.05,
|
|
94
|
+
ood_threshold: float = math.inf, candidates: np.ndarray | None = None,
|
|
95
|
+
eligible: np.ndarray | None = None, deferred_labels: Sequence[str] = ()) -> RoutingPolicy:
|
|
96
|
+
"""The loosest confidence threshold whose rate of *answered and disagreeing* requests is at most `budget`,
|
|
97
|
+
with probability at least 1 - `delta`.
|
|
98
|
+
|
|
99
|
+
Rows come from an IID calibration set: the student's max-prob, whether it agrees with the teacher, and its
|
|
100
|
+
OOD score. The loss per row is 1[the student answers and disagrees], so its rate over all rows is exactly
|
|
101
|
+
what the budget limits (disagreement over all requests).
|
|
102
|
+
|
|
103
|
+
Candidates are tested strictest first with an exact Clopper-Pearson bound at `delta`, and the scan stops at
|
|
104
|
+
the first failure (fixed-sequence testing, as in "Learn Then Test"). The loss can only grow as the
|
|
105
|
+
threshold loosens, so the chance that the chosen threshold breaks the budget is at most `delta`, with no
|
|
106
|
+
multiple-testing penalty. That needs everything else fixed without these rows: `candidates` (default: a
|
|
107
|
+
fixed grid) and `ood_threshold` (the caller picks it on training data). `eligible` marks rows the student
|
|
108
|
+
may answer at all (False where it predicts one of `deferred_labels`); the rate still counts every row.
|
|
109
|
+
"""
|
|
110
|
+
conf = np.asarray(conf, dtype=float)
|
|
111
|
+
agree = np.asarray(agree, dtype=bool)
|
|
112
|
+
ood = np.asarray(ood, dtype=float)
|
|
113
|
+
N = len(conf)
|
|
114
|
+
in_dist = ood <= ood_threshold
|
|
115
|
+
if eligible is not None:
|
|
116
|
+
in_dist &= np.asarray(eligible, dtype=bool)
|
|
117
|
+
grid = threshold_grid(2) if candidates is None else np.asarray(candidates, dtype=float)
|
|
118
|
+
deferred = list(deferred_labels)
|
|
119
|
+
best = None
|
|
120
|
+
if N:
|
|
121
|
+
for t in sorted(grid, reverse=True):
|
|
122
|
+
sel = in_dist & (conf >= t)
|
|
123
|
+
k = int((sel & ~agree).sum())
|
|
124
|
+
ub = clopper_pearson_upper(k, N, delta)
|
|
125
|
+
if ub > budget:
|
|
126
|
+
break # fixed sequence: every looser threshold is untested
|
|
127
|
+
if sel.any(): # a threshold that answers nothing is no policy
|
|
128
|
+
best = (float(t), int(sel.sum()) / N, ub, k / N)
|
|
129
|
+
ood_thr = float(ood_threshold) if math.isfinite(ood_threshold) else float(np.max(ood, initial=0.0))
|
|
130
|
+
if best is None:
|
|
131
|
+
return RoutingPolicy(None, ood_thr, 0.0, 1.0, 0.0, budget, delta, N, deferred)
|
|
132
|
+
t, cov, ub, rate = best
|
|
133
|
+
return RoutingPolicy(t, ood_thr, cov, ub, rate, budget, delta, N, deferred)
|