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,274 @@
1
+ """prepare → split → train → evaluate → [quantize → evaluate] → register (ML-9, EXP-8).
2
+
3
+ Each registered version names how it came to be (`derivation`): `trained`,
4
+ or `distilled` from a teacher version when the run says the snapshot was
5
+ pre-labelled by one, and `quantized` for the int8 copy, whose parent is the
6
+ version this run just registered. The UI draws that as the model's family.
7
+ """
8
+
9
+ from __future__ import annotations
10
+
11
+ import json
12
+ import logging
13
+ import platform
14
+ import random
15
+ import tempfile
16
+ import uuid
17
+ from collections.abc import Mapping, Sequence
18
+ from dataclasses import dataclass
19
+ from datetime import UTC, datetime
20
+ from pathlib import Path
21
+ from typing import Any
22
+
23
+ from annotide import Client
24
+
25
+ from annotide_training import __version__
26
+ from annotide_training.dataset import DatasetError, SplitConfig, SplitName, load_export
27
+ from annotide_training.evaluate import evaluate
28
+ from annotide_training.tracking import MlflowSettings, log_run
29
+ from annotide_training.trainers import (
30
+ QUANTIZE_DTYPES,
31
+ Quantizer,
32
+ Record,
33
+ TrainedModel,
34
+ Trainer,
35
+ )
36
+ from annotide_training.webhook import RetrainRequest
37
+
38
+ log = logging.getLogger("annotide_training")
39
+
40
+
41
+ @dataclass(frozen=True, slots=True)
42
+ class RunResult:
43
+ run_id: str
44
+ artifact_path: Path
45
+ metrics: dict[str, Any]
46
+ #: The registered version, or `None` for a `register=False` (dry) run.
47
+ version: dict[str, Any] | None
48
+ #: The lower-precision copy when the run was asked to `quantize`.
49
+ quantized: QuantizedResult | None = None
50
+
51
+
52
+ @dataclass(frozen=True, slots=True)
53
+ class QuantizedResult:
54
+ artifact_path: Path
55
+ metrics: dict[str, Any]
56
+ version: dict[str, Any] | None
57
+
58
+
59
+ def run_pipeline(
60
+ request: RetrainRequest,
61
+ client: Client,
62
+ trainer: Trainer,
63
+ *,
64
+ out_dir: Path,
65
+ model_id: str | None = None,
66
+ params: Mapping[str, Any] | None = None,
67
+ split: SplitConfig | None = None,
68
+ register: bool = True,
69
+ job_timeout: float = 600.0,
70
+ mlflow: MlflowSettings | None = None,
71
+ quantize: str | None = None,
72
+ teacher_version_id: str | None = None,
73
+ ) -> RunResult:
74
+ """Train on the snapshot the request names and register the result.
75
+
76
+ `model_id` is the model to add a version to when the event did not name
77
+ one. With `register=False` everything runs except the registration, so a
78
+ trainer can be tried against real data without touching the registry.
79
+ With `mlflow`, the run is also recorded there (API-6) before the version
80
+ is registered, and `training_run["mlflow"]` names it.
81
+
82
+ `quantize="int8"` also stores the trained model at that precision (the
83
+ trainer must implement `Quantizer`), scores it on the same held-out
84
+ split and registers it as a `quantized` child of the version trained
85
+ here. `teacher_version_id` registers the trained version as `distilled`
86
+ from that version: the model whose pre-labels, corrected by people, are
87
+ the snapshot this run trained on.
88
+ """
89
+ target_model = request.model_id or model_id
90
+ if register and not target_model:
91
+ raise ValueError("no model to register a version for: the event names none")
92
+ quantizer: Quantizer | None = None
93
+ if quantize is not None:
94
+ if quantize not in QUANTIZE_DTYPES:
95
+ raise ValueError(f"cannot quantize to {quantize!r}: use one of {QUANTIZE_DTYPES}")
96
+ if not isinstance(trainer, Quantizer):
97
+ raise ValueError(f"trainer {trainer.name} cannot quantize (no quantize method)")
98
+ quantizer = trainer
99
+ run_params = dict(params or {})
100
+ run_id = str(uuid.uuid4())
101
+ started_at = datetime.now(UTC)
102
+
103
+ # prepare: the snapshot row first, so a stale or forged digest fails before
104
+ # an export is queued.
105
+ snapshot = client.get_snapshot(request.project_id, request.snapshot_id)
106
+ if snapshot.get("digest") != request.snapshot_digest:
107
+ raise DatasetError(
108
+ f"snapshot {request.snapshot_id} has digest {snapshot.get('digest')}, "
109
+ f"the request says {request.snapshot_digest}"
110
+ )
111
+ export_job = client.create_export(
112
+ request.project_id, "native", snapshot_id=request.snapshot_id
113
+ )["id"]
114
+ client.wait_for_job(export_job, timeout=job_timeout)
115
+ with tempfile.TemporaryDirectory() as tmp:
116
+ archive = client.download_export(export_job, Path(tmp) / "export.zip").read_bytes()
117
+
118
+ # split: the platform's partition when the snapshot has one (EXP-3).
119
+ dataset = load_export(
120
+ archive,
121
+ snapshot_id=request.snapshot_id,
122
+ snapshot_digest=request.snapshot_digest,
123
+ split=split,
124
+ )
125
+ log.info("prepared run=%s counts=%s source=%s", run_id, dataset.counts, dataset.split_source)
126
+ if not dataset.records["train"]:
127
+ raise DatasetError("the train split is empty; nothing to train on")
128
+
129
+ # train
130
+ model = trainer.train(dataset.records["train"], dataset.records["val"], run_params)
131
+
132
+ # evaluate: on test, or val when the split left test empty.
133
+ held_out: SplitName = "test" if dataset.records["test"] else "val"
134
+ metrics = _metrics(trainer, model, dataset.records[held_out], held_out)
135
+ finished_at = datetime.now(UTC)
136
+
137
+ run_dir = out_dir / run_id
138
+ artifact_path = _write(model, run_dir)
139
+ training_run = {
140
+ "id": run_id,
141
+ "pipeline": f"annotide-training {__version__}",
142
+ "trainer": trainer.name,
143
+ "params": run_params,
144
+ "started_at": started_at.isoformat(),
145
+ "finished_at": finished_at.isoformat(),
146
+ "export_job_id": export_job,
147
+ "split_source": dataset.split_source,
148
+ "split_counts": dataset.counts,
149
+ "classes": model.classes,
150
+ "artifact": str(artifact_path),
151
+ "info": model.info,
152
+ "delivery_id": request.delivery_id,
153
+ "note": request.note,
154
+ }
155
+ if teacher_version_id is not None:
156
+ training_run["teacher_version_id"] = teacher_version_id
157
+ if mlflow is not None:
158
+ training_run["mlflow"] = log_run(
159
+ mlflow,
160
+ project_id=request.project_id,
161
+ snapshot_id=request.snapshot_id,
162
+ snapshot_digest=request.snapshot_digest,
163
+ run_name=f"train {trainer.name} {run_id[:8]}",
164
+ params=run_params,
165
+ metrics=metrics,
166
+ run_dir=run_dir,
167
+ )
168
+ (run_dir / "run.json").write_text(
169
+ json.dumps({"training_run": training_run, "metrics": metrics}, indent=2, sort_keys=True)
170
+ )
171
+
172
+ version: dict[str, Any] | None = None
173
+ if register and target_model:
174
+ # register: the platform re-checks the digest against the row (409).
175
+ registered = client.create_model_version(
176
+ target_model,
177
+ {
178
+ "class_mapping": {},
179
+ "snapshot_id": request.snapshot_id,
180
+ "snapshot_digest": request.snapshot_digest,
181
+ "training_run": training_run,
182
+ "metrics": metrics,
183
+ "parent_version_id": teacher_version_id,
184
+ "derivation": "distilled" if teacher_version_id else "trained",
185
+ },
186
+ )
187
+ version = dict(registered)
188
+ log.info("registered run=%s model=%s version=%s", run_id, target_model, version["version"])
189
+
190
+ quantized: QuantizedResult | None = None
191
+ if quantizer is not None and quantize is not None:
192
+ # Calibrate on a fixed-seed sample of train, not its first rows.
193
+ calibration = list(dataset.records["train"])
194
+ random.Random(0).shuffle(calibration)
195
+ small = quantizer.quantize(model, calibration, quantize)
196
+ small_metrics = _metrics(trainer, small, dataset.records[held_out], held_out)
197
+ small_dir = run_dir / quantize
198
+ small_path = _write(small, small_dir)
199
+ small_run = {
200
+ **training_run,
201
+ "quantized_from": str(artifact_path),
202
+ "artifact": str(small_path),
203
+ "info": small.info,
204
+ "finished_at": datetime.now(UTC).isoformat(),
205
+ }
206
+ small_run.pop("mlflow", None)
207
+ small_run.pop("teacher_version_id", None)
208
+ if mlflow is not None:
209
+ small_run["mlflow"] = log_run(
210
+ mlflow,
211
+ project_id=request.project_id,
212
+ snapshot_id=request.snapshot_id,
213
+ snapshot_digest=request.snapshot_digest,
214
+ run_name=f"quantize {quantize} {trainer.name} {run_id[:8]}",
215
+ params={**run_params, "quantize": quantize},
216
+ metrics=small_metrics,
217
+ run_dir=small_dir,
218
+ )
219
+ (small_dir / "run.json").write_text(
220
+ json.dumps(
221
+ {"training_run": small_run, "metrics": small_metrics}, indent=2, sort_keys=True
222
+ )
223
+ )
224
+ small_version: dict[str, Any] | None = None
225
+ if version is not None and target_model:
226
+ small_version = dict(
227
+ client.create_model_version(
228
+ target_model,
229
+ {
230
+ "class_mapping": {},
231
+ "snapshot_id": request.snapshot_id,
232
+ "snapshot_digest": request.snapshot_digest,
233
+ "training_run": small_run,
234
+ "metrics": small_metrics,
235
+ "parent_version_id": version["id"],
236
+ "derivation": "quantized",
237
+ },
238
+ )
239
+ )
240
+ log.info(
241
+ "registered %s copy run=%s version=%s", quantize, run_id, small_version["version"]
242
+ )
243
+ quantized = QuantizedResult(
244
+ artifact_path=small_path, metrics=small_metrics, version=small_version
245
+ )
246
+ return RunResult(
247
+ run_id=run_id,
248
+ artifact_path=artifact_path,
249
+ metrics=metrics,
250
+ version=version,
251
+ quantized=quantized,
252
+ )
253
+
254
+
255
+ def _metrics(
256
+ trainer: Trainer, model: TrainedModel, records: Sequence[Record], split: SplitName
257
+ ) -> dict[str, Any]:
258
+ """The held-out scores plus what the UI compares versions on (size, dtype)."""
259
+ metrics: dict[str, Any] = {"split": split, **evaluate(trainer, model, records)}
260
+ metrics["size_bytes"] = len(model.artifact)
261
+ metrics["latency_device"] = f"{platform.system()} {platform.machine()}".strip()
262
+ if "dtype" in model.info:
263
+ metrics["dtype"] = model.info["dtype"]
264
+ return metrics
265
+
266
+
267
+ def _write(model: TrainedModel, run_dir: Path) -> Path:
268
+ """The artifact and its sibling files, never outside `run_dir`."""
269
+ run_dir.mkdir(parents=True, exist_ok=True)
270
+ artifact_path = run_dir / model.filename
271
+ artifact_path.write_bytes(model.artifact)
272
+ for name, data in model.extra_files.items():
273
+ (run_dir / Path(name).name).write_bytes(data)
274
+ return artifact_path
@@ -0,0 +1,110 @@
1
+ """Optional MLflow tracking for a training run (API-6, EXP-8).
2
+
3
+ With `--mlflow-experiment`, the pipeline also records its run in the
4
+ customer's MLflow — a local server, Databricks or Azure ML, wherever
5
+ `MLFLOW_TRACKING_URI` (or `--mlflow-uri`) points and however that
6
+ environment authenticates — with the snapshot as the run's dataset input in
7
+ exactly the shape the platform publishes (`source_type`
8
+ `annotation-snapshot`) and the `annotation.snapshot_*` tags. The platform's
9
+ `POST /models/{id}/versions/import` reads the lineage back from either, so a
10
+ version imported from this run, or from the registered model it creates, is
11
+ linked to its snapshot without a separate step.
12
+
13
+ Needs the `mlflow` extra (`mlflow-skinny`); imported only when used.
14
+ """
15
+
16
+ from __future__ import annotations
17
+
18
+ import json
19
+ import os
20
+ from dataclasses import dataclass
21
+ from pathlib import Path
22
+ from typing import Any
23
+
24
+ SNAPSHOT_SOURCE_TYPE = "annotation-snapshot"
25
+
26
+
27
+ @dataclass(frozen=True, slots=True)
28
+ class MlflowSettings:
29
+ experiment: str
30
+ tracking_uri: str | None = None
31
+ #: Register the run's artifacts as a new version of this registered model.
32
+ registered_model: str | None = None
33
+
34
+
35
+ def _numeric(metrics: dict[str, Any], prefix: str = "") -> dict[str, float]:
36
+ """MLflow takes numbers only; nested dicts are flattened with dots."""
37
+ flat: dict[str, float] = {}
38
+ for key, value in metrics.items():
39
+ name = f"{prefix}{key}"
40
+ if isinstance(value, bool):
41
+ continue
42
+ if isinstance(value, int | float):
43
+ flat[name] = float(value)
44
+ elif isinstance(value, dict):
45
+ flat.update(_numeric(value, f"{name}."))
46
+ return flat
47
+
48
+
49
+ def log_run(
50
+ settings: MlflowSettings,
51
+ *,
52
+ project_id: str,
53
+ snapshot_id: str,
54
+ snapshot_digest: str,
55
+ run_name: str,
56
+ params: dict[str, Any],
57
+ metrics: dict[str, Any],
58
+ run_dir: Path,
59
+ ) -> dict[str, Any]:
60
+ """Record the finished run in MLflow; what to keep in `training_run["mlflow"]`."""
61
+ os.environ.setdefault("MLFLOW_DISABLE_AGENT_HINT", "1")
62
+ try:
63
+ import mlflow
64
+ from mlflow.entities import Dataset, DatasetInput, InputTag
65
+ from mlflow.tracking import MlflowClient
66
+ except ImportError as exc: # pragma: no cover - depends on the install
67
+ raise RuntimeError("--mlflow-* needs the `mlflow` extra (mlflow-skinny)") from exc
68
+
69
+ if settings.tracking_uri:
70
+ mlflow.set_tracking_uri(settings.tracking_uri)
71
+ experiment = mlflow.set_experiment(settings.experiment)
72
+ with mlflow.start_run(run_name=run_name) as run:
73
+ run_id = run.info.run_id
74
+ mlflow.set_tags(
75
+ {
76
+ "annotation.project_id": project_id,
77
+ "annotation.snapshot_id": snapshot_id,
78
+ "annotation.snapshot_digest": snapshot_digest,
79
+ }
80
+ )
81
+ if params:
82
+ mlflow.log_params({key: str(value) for key, value in params.items()})
83
+ numbers = _numeric(metrics)
84
+ if numbers:
85
+ mlflow.log_metrics(numbers)
86
+ dataset = Dataset(
87
+ name=f"snapshot-{snapshot_id[:8]}",
88
+ digest=snapshot_digest[:16],
89
+ source_type=SNAPSHOT_SOURCE_TYPE,
90
+ source=json.dumps(
91
+ {"snapshot_id": snapshot_id, "project_id": project_id, "digest": snapshot_digest},
92
+ sort_keys=True,
93
+ ),
94
+ )
95
+ MlflowClient().log_inputs(
96
+ run_id,
97
+ datasets=[DatasetInput(dataset, tags=[InputTag("mlflow.data.context", "training")])],
98
+ )
99
+ mlflow.log_artifacts(str(run_dir), artifact_path="model")
100
+
101
+ record: dict[str, Any] = {
102
+ "run_id": run_id,
103
+ "experiment_id": experiment.experiment_id,
104
+ "tracking_uri": mlflow.get_tracking_uri(),
105
+ }
106
+ if settings.registered_model:
107
+ registered = mlflow.register_model(f"runs:/{run_id}/model", settings.registered_model)
108
+ record["registered_model"] = settings.registered_model
109
+ record["model_version"] = str(registered.version)
110
+ return record
@@ -0,0 +1,213 @@
1
+ """The pluggable part: a `Trainer` turns records into a model and predicts with it.
2
+
3
+ A real trainer (a detector fine-tune exported to ONNX for
4
+ `model-service`'s `onnx` backend, say) implements the same two methods and
5
+ is loaded by `--trainer package.module:ClassName`; its framework is its own
6
+ dependency. `BaselineTrainer` needs nothing: it learns where each class
7
+ usually sits in the frame, which is enough to exercise every stage of the
8
+ pipeline end to end and gives a floor any real model has to beat.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import importlib
14
+ import json
15
+ from collections.abc import Mapping, Sequence
16
+ from dataclasses import dataclass, field
17
+ from typing import Any, Protocol, runtime_checkable
18
+
19
+ Record = Mapping[str, Any]
20
+ Box = tuple[float, float, float, float] # x1, y1, x2, y2 in pixels
21
+
22
+
23
+ @dataclass(frozen=True, slots=True)
24
+ class Prediction:
25
+ class_: str
26
+ bbox: Box
27
+ confidence: float
28
+
29
+
30
+ @dataclass(slots=True)
31
+ class TrainedModel:
32
+ """What a trainer hands back: an artifact to store and what it can emit."""
33
+
34
+ artifact: bytes
35
+ filename: str
36
+ classes: list[str]
37
+ info: dict[str, Any] = field(default_factory=dict)
38
+ #: Files written next to the artifact, by name (an ONNX model's `.names`).
39
+ extra_files: dict[str, bytes] = field(default_factory=dict)
40
+
41
+
42
+ @runtime_checkable
43
+ class Trainer(Protocol):
44
+ name: str
45
+
46
+ def train(
47
+ self, train: Sequence[Record], val: Sequence[Record], params: Mapping[str, Any]
48
+ ) -> TrainedModel: ...
49
+
50
+ def predict(self, model: TrainedModel, record: Record) -> list[Prediction]: ...
51
+
52
+
53
+ #: What `--quantize` accepts. int8 is the one every CPU runtime speeds up.
54
+ QUANTIZE_DTYPES = ("int8",)
55
+
56
+
57
+ @runtime_checkable
58
+ class Quantizer(Protocol):
59
+ """Optional: a trainer that can store a trained model at lower precision.
60
+
61
+ `calibration` is a sample of training records for static quantization
62
+ (activation ranges); the result must load and predict through the same
63
+ trainer's `predict`, so the pipeline scores it exactly like the original.
64
+ """
65
+
66
+ def quantize(
67
+ self, model: TrainedModel, calibration: Sequence[Record], dtype: str
68
+ ) -> TrainedModel: ...
69
+
70
+
71
+ def record_boxes(record: Record) -> list[tuple[str, Box]]:
72
+ """The `bbox` shapes of a native export record, as (class, x1 y1 x2 y2)."""
73
+ boxes: list[tuple[str, Box]] = []
74
+ for shape in record.get("shapes", []):
75
+ bbox = shape.get("bbox")
76
+ if shape.get("type") == "bbox" and isinstance(bbox, list) and len(bbox) == 4:
77
+ x1, y1, x2, y2 = (float(v) for v in bbox)
78
+ boxes.append((str(shape.get("class")), (x1, y1, x2, y2)))
79
+ return boxes
80
+
81
+
82
+ class BaselineTrainer:
83
+ """Per class: how often it appears and its mean normalised box.
84
+
85
+ Predicts, for every class seen in at least `min_frequency` of the
86
+ training images, one box at that mean position with the frequency as
87
+ its confidence. Deterministic, dependency-free, and honest about being
88
+ a baseline.
89
+ """
90
+
91
+ name = "baseline-class-prior"
92
+
93
+ def train(
94
+ self, train: Sequence[Record], val: Sequence[Record], params: Mapping[str, Any]
95
+ ) -> TrainedModel:
96
+ min_frequency = float(params.get("min_frequency", 0.0))
97
+ images = 0
98
+ present: dict[str, int] = {}
99
+ sums: dict[str, list[float]] = {}
100
+ counts: dict[str, int] = {}
101
+ for record in train:
102
+ width = float(record.get("width") or 0)
103
+ height = float(record.get("height") or 0)
104
+ if width <= 0 or height <= 0:
105
+ continue
106
+ images += 1
107
+ seen: set[str] = set()
108
+ for class_name, (x1, y1, x2, y2) in record_boxes(record):
109
+ acc = sums.setdefault(class_name, [0.0, 0.0, 0.0, 0.0])
110
+ for i, value in enumerate((x1 / width, y1 / height, x2 / width, y2 / height)):
111
+ acc[i] += value
112
+ counts[class_name] = counts.get(class_name, 0) + 1
113
+ seen.add(class_name)
114
+ for class_name in seen:
115
+ present[class_name] = present.get(class_name, 0) + 1
116
+
117
+ priors: dict[str, dict[str, Any]] = {
118
+ name: {
119
+ "frequency": round(present.get(name, 0) / images, 6) if images else 0.0,
120
+ "box": [round(v / counts[name], 6) for v in sums[name]],
121
+ }
122
+ for name in sorted(sums)
123
+ }
124
+ emitted = sorted(n for n, p in priors.items() if p["frequency"] >= min_frequency)
125
+ body = {
126
+ "trainer": self.name,
127
+ "images": images,
128
+ "min_frequency": min_frequency,
129
+ "priors": {n: priors[n] for n in emitted},
130
+ }
131
+ return TrainedModel(
132
+ artifact=json.dumps(body, indent=2, sort_keys=True).encode(),
133
+ filename="baseline.json",
134
+ classes=emitted,
135
+ info={"train_images": images, "val_images": len(val), "dtype": "fp64"},
136
+ )
137
+
138
+ def quantize(
139
+ self, model: TrainedModel, calibration: Sequence[Record], dtype: str
140
+ ) -> TrainedModel:
141
+ """Each prior's box and frequency as an 8-bit fraction (k / 255).
142
+
143
+ The baseline's floats are already tiny; this exists so the
144
+ `--quantize` path runs end to end without the torch extra.
145
+ """
146
+ if dtype != "int8":
147
+ raise ValueError(f"{self.name} quantizes to int8 only, not {dtype!r}")
148
+ body = json.loads(model.artifact)
149
+ for prior in body["priors"].values():
150
+ prior["box"] = [_to_uint8(v) for v in prior["box"]]
151
+ prior["frequency"] = _to_uint8(prior["frequency"])
152
+ body["dtype"] = "int8"
153
+ return TrainedModel(
154
+ artifact=json.dumps(body, sort_keys=True, separators=(",", ":")).encode(),
155
+ filename=model.filename,
156
+ classes=list(model.classes),
157
+ info={**model.info, "dtype": "int8"},
158
+ extra_files=dict(model.extra_files),
159
+ )
160
+
161
+ def predict(self, model: TrainedModel, record: Record) -> list[Prediction]:
162
+ body = json.loads(model.artifact)
163
+ scale = 255.0 if body.get("dtype") == "int8" else 1.0
164
+ width = float(record.get("width") or 0)
165
+ height = float(record.get("height") or 0)
166
+ predictions: list[Prediction] = []
167
+ for class_name, prior in body["priors"].items():
168
+ nx1, ny1, nx2, ny2 = (v / scale for v in prior["box"])
169
+ predictions.append(
170
+ Prediction(
171
+ class_=class_name,
172
+ bbox=(nx1 * width, ny1 * height, nx2 * width, ny2 * height),
173
+ confidence=float(prior["frequency"]) / scale,
174
+ )
175
+ )
176
+ return predictions
177
+
178
+
179
+ def _to_uint8(value: float) -> int:
180
+ return max(0, min(255, round(value * 255)))
181
+
182
+
183
+ # Built-ins by name. `fasterrcnn` needs the `torch` extra, so it is imported
184
+ # only when asked for.
185
+ BUILTIN_TRAINERS: dict[str, str] = {
186
+ "baseline": "annotide_training.trainers:BaselineTrainer",
187
+ "fasterrcnn": "annotide_training.detector:FasterRCNNTrainer",
188
+ }
189
+ _TORCH_EXTRA = {"torch", "torchvision", "onnx", "onnxruntime", "numpy", "PIL"}
190
+
191
+
192
+ def load_trainer(spec: str) -> Trainer:
193
+ """`baseline`, `fasterrcnn`, or `package.module:ClassName` for your own."""
194
+ builtin = BUILTIN_TRAINERS.get(spec)
195
+ try:
196
+ return _load(builtin or spec)
197
+ except ModuleNotFoundError as exc:
198
+ if builtin and exc.name in _TORCH_EXTRA:
199
+ raise ValueError(
200
+ f"trainer {spec!r} needs the torch extra: pip install 'annotide-training[torch]'"
201
+ ) from exc
202
+ raise
203
+
204
+
205
+ def _load(spec: str) -> Trainer:
206
+ module_name, sep, attr = spec.partition(":")
207
+ if not sep or not module_name or not attr:
208
+ raise ValueError(f"trainer {spec!r}: use a built-in name or 'module:ClassName'")
209
+ factory = getattr(importlib.import_module(module_name), attr)
210
+ trainer = factory()
211
+ if not isinstance(trainer, Trainer):
212
+ raise TypeError(f"{spec} does not implement Trainer (name, train, predict)")
213
+ return trainer
@@ -0,0 +1,94 @@
1
+ """Verify and parse `retrain.requested` deliveries (CONTRACTS.md, webhooks).
2
+
3
+ Mirrors `backend/app/services/webhooks.py::verify`: the header is
4
+ `t=<unix seconds>,v1=<hex HMAC-SHA256 of "{t}.{body}">`, stale timestamps
5
+ are rejected and the comparison is constant-time.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import hashlib
11
+ import hmac
12
+ import json
13
+ import time
14
+ from dataclasses import dataclass
15
+ from typing import Any
16
+
17
+ SIGNATURE_HEADER = "X-Annotation-Signature"
18
+ EVENT_HEADER = "X-Annotation-Event"
19
+ DELIVERY_HEADER = "X-Annotation-Delivery"
20
+ RETRAIN_EVENT = "retrain.requested"
21
+ TOLERANCE_SECONDS = 300
22
+
23
+
24
+ class WebhookError(ValueError):
25
+ """The delivery is not one this pipeline acts on."""
26
+
27
+
28
+ def sign(secret: str, timestamp: int, body: bytes) -> str:
29
+ mac = hmac.new(secret.encode(), f"{timestamp}.".encode() + body, hashlib.sha256)
30
+ return f"t={timestamp},v1={mac.hexdigest()}"
31
+
32
+
33
+ def verify(
34
+ secret: str,
35
+ header: str,
36
+ body: bytes,
37
+ *,
38
+ now: float | None = None,
39
+ tolerance_seconds: int = TOLERANCE_SECONDS,
40
+ ) -> bool:
41
+ parts = dict(part.split("=", 1) for part in header.split(",") if "=" in part)
42
+ try:
43
+ timestamp = int(parts["t"])
44
+ except (KeyError, ValueError):
45
+ return False
46
+ current = int(time.time() if now is None else now)
47
+ if abs(current - timestamp) > tolerance_seconds:
48
+ return False
49
+ expected = sign(secret, timestamp, body).split("v1=", 1)[1]
50
+ return hmac.compare_digest(expected, parts.get("v1", ""))
51
+
52
+
53
+ @dataclass(frozen=True, slots=True)
54
+ class RetrainRequest:
55
+ """What one `retrain.requested` event asks for."""
56
+
57
+ project_id: str
58
+ snapshot_id: str
59
+ snapshot_digest: str
60
+ model_id: str | None
61
+ note: str | None
62
+ delivery_id: str | None
63
+
64
+
65
+ def parse_retrain(body: bytes) -> RetrainRequest:
66
+ """Parse a verified delivery body; `WebhookError` unless it names a snapshot."""
67
+ try:
68
+ payload: Any = json.loads(body)
69
+ except json.JSONDecodeError as exc:
70
+ raise WebhookError("body is not JSON") from exc
71
+ if not isinstance(payload, dict) or payload.get("event") != RETRAIN_EVENT:
72
+ raise WebhookError(f"not a {RETRAIN_EVENT} event")
73
+ data = payload.get("data")
74
+ if not isinstance(data, dict):
75
+ raise WebhookError("event has no data")
76
+ snapshot = data.get("snapshot")
77
+ if not isinstance(snapshot, dict) or not snapshot.get("id") or not snapshot.get("digest"):
78
+ # The platform allows a retrain request without a snapshot; there is
79
+ # nothing to train on until someone freezes one.
80
+ raise WebhookError("retrain.requested names no snapshot")
81
+ project_id = data.get("project_id") or payload.get("project_id")
82
+ if not project_id:
83
+ raise WebhookError("event has no project_id")
84
+ model_id = data.get("model_id")
85
+ note = data.get("note")
86
+ delivery = payload.get("delivery_id")
87
+ return RetrainRequest(
88
+ project_id=str(project_id),
89
+ snapshot_id=str(snapshot["id"]),
90
+ snapshot_digest=str(snapshot["digest"]),
91
+ model_id=str(model_id) if model_id else None,
92
+ note=str(note) if note else None,
93
+ delivery_id=str(delivery) if delivery else None,
94
+ )