annotide-training 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.
- annotide_training/__init__.py +10 -0
- annotide_training/cli.py +251 -0
- annotide_training/dataset.py +141 -0
- annotide_training/detector.py +477 -0
- annotide_training/evaluate.py +89 -0
- annotide_training/pipeline.py +274 -0
- annotide_training/tracking.py +110 -0
- annotide_training/trainers.py +213 -0
- annotide_training/webhook.py +94 -0
- annotide_training-0.1.0.dist-info/METADATA +19 -0
- annotide_training-0.1.0.dist-info/RECORD +14 -0
- annotide_training-0.1.0.dist-info/WHEEL +5 -0
- annotide_training-0.1.0.dist-info/entry_points.txt +2 -0
- annotide_training-0.1.0.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,10 @@
|
|
|
1
|
+
"""Reference training pipeline: prepare → split → train → evaluate → register (ML-9, EXP-8).
|
|
2
|
+
|
|
3
|
+
The platform emits `retrain.requested` and trains nothing itself. This
|
|
4
|
+
package is what a customer runs on the other end: it pulls the snapshot's
|
|
5
|
+
export through the REST API, checks it is the data the event named, trains
|
|
6
|
+
with a pluggable `Trainer`, and registers the result as a model version
|
|
7
|
+
carrying its lineage (`snapshot_id`, `snapshot_digest`, `training_run`).
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
__version__ = "0.1.0"
|
annotide_training/cli.py
ADDED
|
@@ -0,0 +1,251 @@
|
|
|
1
|
+
"""`annotide-train run …` for one snapshot, `annotide-train serve` for the webhook.
|
|
2
|
+
|
|
3
|
+
Configuration comes from flags or the environment:
|
|
4
|
+
`ANNOTIDE_API_URL` (or the SDK's `ANNOTIDE_URL`), `ANNOTIDE_API_KEY`
|
|
5
|
+
(a `write`-scoped API key),
|
|
6
|
+
`ANNOTIDE_WEBHOOK_SECRET` (the webhook's signing secret, `serve` only).
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
import argparse
|
|
12
|
+
import json
|
|
13
|
+
import logging
|
|
14
|
+
import os
|
|
15
|
+
import sys
|
|
16
|
+
import threading
|
|
17
|
+
from collections.abc import Callable, Sequence
|
|
18
|
+
from http import HTTPStatus
|
|
19
|
+
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
20
|
+
from pathlib import Path
|
|
21
|
+
from typing import Any
|
|
22
|
+
|
|
23
|
+
from annotide import Client
|
|
24
|
+
|
|
25
|
+
from annotide_training.dataset import SplitConfig
|
|
26
|
+
from annotide_training.pipeline import run_pipeline
|
|
27
|
+
from annotide_training.tracking import MlflowSettings
|
|
28
|
+
from annotide_training.trainers import QUANTIZE_DTYPES, load_trainer
|
|
29
|
+
from annotide_training.webhook import (
|
|
30
|
+
DELIVERY_HEADER,
|
|
31
|
+
SIGNATURE_HEADER,
|
|
32
|
+
RetrainRequest,
|
|
33
|
+
WebhookError,
|
|
34
|
+
parse_retrain,
|
|
35
|
+
verify,
|
|
36
|
+
)
|
|
37
|
+
|
|
38
|
+
log = logging.getLogger("annotide_training")
|
|
39
|
+
|
|
40
|
+
#: Bodies are small JSON events; anything larger is not a delivery.
|
|
41
|
+
MAX_BODY_BYTES = 256 * 1024
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _param(value: str) -> tuple[str, Any]:
|
|
45
|
+
key, sep, raw = value.partition("=")
|
|
46
|
+
if not sep or not key:
|
|
47
|
+
raise argparse.ArgumentTypeError(f"--param {value!r}: use key=value")
|
|
48
|
+
try:
|
|
49
|
+
return key, json.loads(raw)
|
|
50
|
+
except json.JSONDecodeError:
|
|
51
|
+
return key, raw
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _split(value: str) -> SplitConfig:
|
|
55
|
+
try:
|
|
56
|
+
train, val, test = (float(part) for part in value.split(","))
|
|
57
|
+
except ValueError as exc:
|
|
58
|
+
raise argparse.ArgumentTypeError("--split: use train,val,test e.g. 0.8,0.1,0.1") from exc
|
|
59
|
+
return SplitConfig(train=train, val=val, test=test)
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
def _parser() -> argparse.ArgumentParser:
|
|
63
|
+
parser = argparse.ArgumentParser(prog="annotide-train", description=__doc__)
|
|
64
|
+
parser.add_argument(
|
|
65
|
+
"--url", default=os.environ.get("ANNOTIDE_API_URL") or os.environ.get("ANNOTIDE_URL")
|
|
66
|
+
)
|
|
67
|
+
parser.add_argument("--api-key", default=os.environ.get("ANNOTIDE_API_KEY"))
|
|
68
|
+
parser.add_argument(
|
|
69
|
+
"--trainer", default="baseline", help="baseline, fasterrcnn, or module:Class"
|
|
70
|
+
)
|
|
71
|
+
parser.add_argument("--out", type=Path, default=Path("runs"), help="where runs are written")
|
|
72
|
+
parser.add_argument("--param", type=_param, action="append", default=[], metavar="K=V")
|
|
73
|
+
parser.add_argument(
|
|
74
|
+
"--split", type=_split, default=None, help="for unsplit snapshots: train,val,test"
|
|
75
|
+
)
|
|
76
|
+
parser.add_argument("--model", help="model to register into when the event names none")
|
|
77
|
+
parser.add_argument(
|
|
78
|
+
"--quantize",
|
|
79
|
+
choices=QUANTIZE_DTYPES,
|
|
80
|
+
help="also register a copy at this precision, a child of the trained version",
|
|
81
|
+
)
|
|
82
|
+
parser.add_argument(
|
|
83
|
+
"--teacher",
|
|
84
|
+
metavar="VERSION_ID",
|
|
85
|
+
help="register as distilled from this model version (its pre-labels, corrected, "
|
|
86
|
+
"are the snapshot)",
|
|
87
|
+
)
|
|
88
|
+
parser.add_argument(
|
|
89
|
+
"--mlflow-experiment",
|
|
90
|
+
help="also record the run in this MLflow experiment (needs the mlflow extra)",
|
|
91
|
+
)
|
|
92
|
+
parser.add_argument(
|
|
93
|
+
"--mlflow-uri",
|
|
94
|
+
default=os.environ.get("MLFLOW_TRACKING_URI"),
|
|
95
|
+
help="MLflow tracking URI (default: MLFLOW_TRACKING_URI)",
|
|
96
|
+
)
|
|
97
|
+
parser.add_argument(
|
|
98
|
+
"--mlflow-register", metavar="NAME", help="register the run as a version of this model"
|
|
99
|
+
)
|
|
100
|
+
commands = parser.add_subparsers(dest="command", required=True)
|
|
101
|
+
|
|
102
|
+
run = commands.add_parser("run", help="train on one snapshot now")
|
|
103
|
+
run.add_argument("--project", required=True)
|
|
104
|
+
run.add_argument("--snapshot", required=True)
|
|
105
|
+
run.add_argument("--no-register", action="store_true", help="train and evaluate only")
|
|
106
|
+
|
|
107
|
+
serve = commands.add_parser("serve", help="receive retrain.requested webhooks")
|
|
108
|
+
serve.add_argument("--host", default="127.0.0.1")
|
|
109
|
+
serve.add_argument("--port", type=int, default=8088)
|
|
110
|
+
serve.add_argument("--secret", default=os.environ.get("ANNOTIDE_WEBHOOK_SECRET"))
|
|
111
|
+
return parser
|
|
112
|
+
|
|
113
|
+
|
|
114
|
+
def make_handler(
|
|
115
|
+
secret: str, on_request: Callable[[RetrainRequest], None]
|
|
116
|
+
) -> type[BaseHTTPRequestHandler]:
|
|
117
|
+
"""A handler that verifies, de-duplicates and hands off; it never trains inline.
|
|
118
|
+
|
|
119
|
+
The platform times out deliveries (`APP_WEBHOOK_TIMEOUT`), so the answer
|
|
120
|
+
is 202 as soon as the event is accepted and the run happens on a thread.
|
|
121
|
+
`X-Annotation-Delivery` de-duplicates retried deliveries.
|
|
122
|
+
"""
|
|
123
|
+
seen: set[str] = set()
|
|
124
|
+
lock = threading.Lock()
|
|
125
|
+
|
|
126
|
+
class Handler(BaseHTTPRequestHandler):
|
|
127
|
+
def _reply(self, status: HTTPStatus, message: str) -> None:
|
|
128
|
+
body = json.dumps({"detail": message}).encode()
|
|
129
|
+
self.send_response(status)
|
|
130
|
+
self.send_header("Content-Type", "application/json")
|
|
131
|
+
self.send_header("Content-Length", str(len(body)))
|
|
132
|
+
self.end_headers()
|
|
133
|
+
self.wfile.write(body)
|
|
134
|
+
|
|
135
|
+
def do_POST(self) -> None:
|
|
136
|
+
length = int(self.headers.get("Content-Length") or 0)
|
|
137
|
+
if length <= 0 or length > MAX_BODY_BYTES:
|
|
138
|
+
self._reply(HTTPStatus.BAD_REQUEST, "missing or oversized body")
|
|
139
|
+
return
|
|
140
|
+
body = self.rfile.read(length)
|
|
141
|
+
if not verify(secret, self.headers.get(SIGNATURE_HEADER, ""), body):
|
|
142
|
+
self._reply(HTTPStatus.UNAUTHORIZED, "bad or stale signature")
|
|
143
|
+
return
|
|
144
|
+
try:
|
|
145
|
+
request = parse_retrain(body)
|
|
146
|
+
except WebhookError as exc:
|
|
147
|
+
# Signed but not for us (another event, no snapshot): accept so
|
|
148
|
+
# the platform does not retry it for hours.
|
|
149
|
+
self._reply(HTTPStatus.OK, f"ignored: {exc}")
|
|
150
|
+
return
|
|
151
|
+
delivery = self.headers.get(DELIVERY_HEADER) or request.delivery_id
|
|
152
|
+
with lock:
|
|
153
|
+
if delivery and delivery in seen:
|
|
154
|
+
self._reply(HTTPStatus.OK, "duplicate delivery")
|
|
155
|
+
return
|
|
156
|
+
if delivery:
|
|
157
|
+
seen.add(delivery)
|
|
158
|
+
threading.Thread(target=on_request, args=(request,), daemon=True).start()
|
|
159
|
+
self._reply(HTTPStatus.ACCEPTED, "training started")
|
|
160
|
+
|
|
161
|
+
def log_message(self, format: str, *args: Any) -> None: # noqa: A002 - stdlib name
|
|
162
|
+
log.info("webhook %s", format % args)
|
|
163
|
+
|
|
164
|
+
return Handler
|
|
165
|
+
|
|
166
|
+
|
|
167
|
+
def main(argv: Sequence[str] | None = None) -> int:
|
|
168
|
+
logging.basicConfig(level=logging.INFO, format="%(asctime)s %(levelname)s %(message)s")
|
|
169
|
+
args = _parser().parse_args(argv)
|
|
170
|
+
if not args.url or not args.api_key:
|
|
171
|
+
print(
|
|
172
|
+
"--url/ANNOTIDE_API_URL and --api-key/ANNOTIDE_API_KEY are required",
|
|
173
|
+
file=sys.stderr,
|
|
174
|
+
)
|
|
175
|
+
return 2
|
|
176
|
+
trainer = load_trainer(args.trainer)
|
|
177
|
+
params = dict(args.param)
|
|
178
|
+
mlflow_settings = (
|
|
179
|
+
MlflowSettings(
|
|
180
|
+
experiment=args.mlflow_experiment,
|
|
181
|
+
tracking_uri=args.mlflow_uri,
|
|
182
|
+
registered_model=args.mlflow_register,
|
|
183
|
+
)
|
|
184
|
+
if args.mlflow_experiment
|
|
185
|
+
else None
|
|
186
|
+
)
|
|
187
|
+
|
|
188
|
+
def train(request: RetrainRequest, *, register: bool = True) -> None:
|
|
189
|
+
with Client(args.url, args.api_key) as client:
|
|
190
|
+
result = run_pipeline(
|
|
191
|
+
request,
|
|
192
|
+
client,
|
|
193
|
+
trainer,
|
|
194
|
+
out_dir=args.out,
|
|
195
|
+
model_id=args.model,
|
|
196
|
+
params=params,
|
|
197
|
+
split=args.split,
|
|
198
|
+
register=register,
|
|
199
|
+
mlflow=mlflow_settings,
|
|
200
|
+
quantize=args.quantize,
|
|
201
|
+
teacher_version_id=args.teacher,
|
|
202
|
+
)
|
|
203
|
+
summary: dict[str, Any] = {
|
|
204
|
+
"run_id": result.run_id,
|
|
205
|
+
"metrics": result.metrics,
|
|
206
|
+
"version": result.version,
|
|
207
|
+
}
|
|
208
|
+
if result.quantized is not None:
|
|
209
|
+
summary["quantized"] = {
|
|
210
|
+
"metrics": result.quantized.metrics,
|
|
211
|
+
"version": result.quantized.version,
|
|
212
|
+
}
|
|
213
|
+
print(json.dumps(summary, indent=2))
|
|
214
|
+
|
|
215
|
+
if args.command == "run":
|
|
216
|
+
with Client(args.url, args.api_key) as client:
|
|
217
|
+
snapshot = client.get_snapshot(args.project, args.snapshot)
|
|
218
|
+
request = RetrainRequest(
|
|
219
|
+
project_id=args.project,
|
|
220
|
+
snapshot_id=args.snapshot,
|
|
221
|
+
snapshot_digest=str(snapshot["digest"]),
|
|
222
|
+
model_id=None,
|
|
223
|
+
note=None,
|
|
224
|
+
delivery_id=None,
|
|
225
|
+
)
|
|
226
|
+
train(request, register=not args.no_register)
|
|
227
|
+
return 0
|
|
228
|
+
|
|
229
|
+
if not args.secret:
|
|
230
|
+
print("--secret/ANNOTIDE_WEBHOOK_SECRET is required for serve", file=sys.stderr)
|
|
231
|
+
return 2
|
|
232
|
+
|
|
233
|
+
def on_request(request: RetrainRequest) -> None:
|
|
234
|
+
try:
|
|
235
|
+
train(request)
|
|
236
|
+
except Exception: # a failed run must not take the receiver down
|
|
237
|
+
log.exception("training run for snapshot %s failed", request.snapshot_id)
|
|
238
|
+
|
|
239
|
+
server = ThreadingHTTPServer((args.host, args.port), make_handler(args.secret, on_request))
|
|
240
|
+
log.info("listening on http://%s:%d", args.host, args.port)
|
|
241
|
+
try:
|
|
242
|
+
server.serve_forever()
|
|
243
|
+
except KeyboardInterrupt:
|
|
244
|
+
pass
|
|
245
|
+
finally:
|
|
246
|
+
server.server_close()
|
|
247
|
+
return 0
|
|
248
|
+
|
|
249
|
+
|
|
250
|
+
if __name__ == "__main__":
|
|
251
|
+
raise SystemExit(main())
|
|
@@ -0,0 +1,141 @@
|
|
|
1
|
+
"""Prepare and split: turn a `native` export archive into train / val / test records.
|
|
2
|
+
|
|
3
|
+
The archive (CONTRACTS.md, export) holds `manifest.json` — with the
|
|
4
|
+
`snapshot_id` and `digest` it was exported from — and `annotations.jsonl`,
|
|
5
|
+
or one `annotations.jsonl` per `train/`, `val/`, `test/` when the snapshot
|
|
6
|
+
was split on the platform (EXP-3). An unsplit snapshot is split here with
|
|
7
|
+
the platform's own rule, so the same seed gives the same partition.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import hashlib
|
|
13
|
+
import io
|
|
14
|
+
import json
|
|
15
|
+
import zipfile
|
|
16
|
+
from dataclasses import dataclass, field
|
|
17
|
+
from typing import Any, Literal
|
|
18
|
+
|
|
19
|
+
SplitName = Literal["train", "val", "test"]
|
|
20
|
+
SPLITS: tuple[SplitName, ...] = ("train", "val", "test")
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class DatasetError(ValueError):
|
|
24
|
+
"""The export is not the data the retrain request named, or is malformed."""
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
@dataclass(frozen=True, slots=True)
|
|
28
|
+
class SplitConfig:
|
|
29
|
+
train: float = 0.8
|
|
30
|
+
val: float = 0.1
|
|
31
|
+
test: float = 0.1
|
|
32
|
+
seed: int = 0
|
|
33
|
+
|
|
34
|
+
def __post_init__(self) -> None:
|
|
35
|
+
ratios = (self.train, self.val, self.test)
|
|
36
|
+
if any(r < 0 or r > 1 for r in ratios) or abs(sum(ratios) - 1) > 1e-9:
|
|
37
|
+
raise ValueError("split ratios must be in [0, 1] and sum to 1")
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
def assign_split(item_id: str, config: SplitConfig) -> SplitName:
|
|
41
|
+
"""`backend/app/services/datasets.py::assign_split` for an ungrouped item."""
|
|
42
|
+
digest = hashlib.sha256(f"{config.seed}:item:{item_id}".encode()).digest()
|
|
43
|
+
point = int.from_bytes(digest[:8], "big") / 2**64
|
|
44
|
+
if point < config.train:
|
|
45
|
+
return "train"
|
|
46
|
+
if point < config.train + config.val:
|
|
47
|
+
return "val"
|
|
48
|
+
return "test"
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
@dataclass(slots=True)
|
|
52
|
+
class Dataset:
|
|
53
|
+
snapshot_id: str
|
|
54
|
+
digest: str
|
|
55
|
+
manifest: dict[str, Any]
|
|
56
|
+
records: dict[SplitName, list[dict[str, Any]]] = field(
|
|
57
|
+
default_factory=lambda: {name: [] for name in SPLITS}
|
|
58
|
+
)
|
|
59
|
+
#: Whether the partition came from the platform (EXP-3) or `SplitConfig` here.
|
|
60
|
+
split_source: Literal["snapshot", "local"] = "snapshot"
|
|
61
|
+
|
|
62
|
+
@property
|
|
63
|
+
def counts(self) -> dict[str, int]:
|
|
64
|
+
return {name: len(self.records[name]) for name in SPLITS}
|
|
65
|
+
|
|
66
|
+
@property
|
|
67
|
+
def classes(self) -> list[str]:
|
|
68
|
+
seen = {
|
|
69
|
+
str(shape.get("class"))
|
|
70
|
+
for records in self.records.values()
|
|
71
|
+
for record in records
|
|
72
|
+
for shape in record.get("shapes", [])
|
|
73
|
+
}
|
|
74
|
+
return sorted(seen)
|
|
75
|
+
|
|
76
|
+
|
|
77
|
+
def _jsonl(data: bytes, name: str) -> list[dict[str, Any]]:
|
|
78
|
+
records: list[dict[str, Any]] = []
|
|
79
|
+
for number, line in enumerate(data.decode("utf-8").splitlines(), start=1):
|
|
80
|
+
if not line.strip():
|
|
81
|
+
continue
|
|
82
|
+
try:
|
|
83
|
+
record = json.loads(line)
|
|
84
|
+
except json.JSONDecodeError as exc:
|
|
85
|
+
raise DatasetError(f"{name}:{number}: not JSON") from exc
|
|
86
|
+
if not isinstance(record, dict) or "item_id" not in record:
|
|
87
|
+
raise DatasetError(f"{name}:{number}: not a native export record")
|
|
88
|
+
records.append(record)
|
|
89
|
+
return records
|
|
90
|
+
|
|
91
|
+
|
|
92
|
+
def load_export(
|
|
93
|
+
archive: bytes,
|
|
94
|
+
*,
|
|
95
|
+
snapshot_id: str,
|
|
96
|
+
snapshot_digest: str,
|
|
97
|
+
split: SplitConfig | None = None,
|
|
98
|
+
) -> Dataset:
|
|
99
|
+
"""Read a native export archive and check it came from the requested snapshot.
|
|
100
|
+
|
|
101
|
+
The digest the event carried must be the digest the export's manifest
|
|
102
|
+
records: training on anything else would register lineage that is not
|
|
103
|
+
true (the platform would 409 the registration too, but only after the
|
|
104
|
+
training run was spent).
|
|
105
|
+
"""
|
|
106
|
+
try:
|
|
107
|
+
bundle = zipfile.ZipFile(io.BytesIO(archive))
|
|
108
|
+
except zipfile.BadZipFile as exc:
|
|
109
|
+
raise DatasetError("export is not a zip archive") from exc
|
|
110
|
+
with bundle:
|
|
111
|
+
names = set(bundle.namelist())
|
|
112
|
+
if "manifest.json" not in names:
|
|
113
|
+
raise DatasetError("export has no manifest.json")
|
|
114
|
+
manifest: dict[str, Any] = json.loads(bundle.read("manifest.json"))
|
|
115
|
+
if manifest.get("format") != "native":
|
|
116
|
+
raise DatasetError(f"expected a native export, got {manifest.get('format')!r}")
|
|
117
|
+
if str(manifest.get("snapshot_id")) != snapshot_id:
|
|
118
|
+
raise DatasetError(
|
|
119
|
+
f"export is of snapshot {manifest.get('snapshot_id')}, not {snapshot_id}"
|
|
120
|
+
)
|
|
121
|
+
if manifest.get("digest") != snapshot_digest:
|
|
122
|
+
raise DatasetError(
|
|
123
|
+
f"snapshot digest mismatch: export {manifest.get('digest')}, "
|
|
124
|
+
f"request {snapshot_digest}"
|
|
125
|
+
)
|
|
126
|
+
|
|
127
|
+
dataset = Dataset(snapshot_id=snapshot_id, digest=snapshot_digest, manifest=manifest)
|
|
128
|
+
split_files = {name: f"{name}/annotations.jsonl" for name in SPLITS}
|
|
129
|
+
if any(path in names for path in split_files.values()):
|
|
130
|
+
for name, path in split_files.items():
|
|
131
|
+
if path in names:
|
|
132
|
+
dataset.records[name] = _jsonl(bundle.read(path), path)
|
|
133
|
+
return dataset
|
|
134
|
+
|
|
135
|
+
if "annotations.jsonl" not in names:
|
|
136
|
+
raise DatasetError("export has no annotations.jsonl")
|
|
137
|
+
config = split or SplitConfig()
|
|
138
|
+
dataset.split_source = "local"
|
|
139
|
+
for record in _jsonl(bundle.read("annotations.jsonl"), "annotations.jsonl"):
|
|
140
|
+
dataset.records[assign_split(str(record["item_id"]), config)].append(record)
|
|
141
|
+
return dataset
|