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,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
|
+
)
|