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.
@@ -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"
@@ -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