annotide-training 0.1.0__tar.gz
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-0.1.0/PKG-INFO +19 -0
- annotide_training-0.1.0/README.md +177 -0
- annotide_training-0.1.0/annotide_training/__init__.py +10 -0
- annotide_training-0.1.0/annotide_training/cli.py +251 -0
- annotide_training-0.1.0/annotide_training/dataset.py +141 -0
- annotide_training-0.1.0/annotide_training/detector.py +477 -0
- annotide_training-0.1.0/annotide_training/evaluate.py +89 -0
- annotide_training-0.1.0/annotide_training/pipeline.py +274 -0
- annotide_training-0.1.0/annotide_training/tracking.py +110 -0
- annotide_training-0.1.0/annotide_training/trainers.py +213 -0
- annotide_training-0.1.0/annotide_training/webhook.py +94 -0
- annotide_training-0.1.0/annotide_training.egg-info/PKG-INFO +19 -0
- annotide_training-0.1.0/annotide_training.egg-info/SOURCES.txt +22 -0
- annotide_training-0.1.0/annotide_training.egg-info/dependency_links.txt +1 -0
- annotide_training-0.1.0/annotide_training.egg-info/entry_points.txt +2 -0
- annotide_training-0.1.0/annotide_training.egg-info/requires.txt +17 -0
- annotide_training-0.1.0/annotide_training.egg-info/top_level.txt +1 -0
- annotide_training-0.1.0/pyproject.toml +75 -0
- annotide_training-0.1.0/setup.cfg +4 -0
- annotide_training-0.1.0/tests/test_dataset.py +62 -0
- annotide_training-0.1.0/tests/test_detector.py +174 -0
- annotide_training-0.1.0/tests/test_pipeline.py +271 -0
- annotide_training-0.1.0/tests/test_trainers.py +89 -0
- annotide_training-0.1.0/tests/test_webhook.py +96 -0
|
@@ -0,0 +1,19 @@
|
|
|
1
|
+
Metadata-Version: 2.4
|
|
2
|
+
Name: annotide-training
|
|
3
|
+
Version: 0.1.0
|
|
4
|
+
Summary: Cloud-agnostic annotation platform — reference customer-side training pipeline (ML-9, EXP-8)
|
|
5
|
+
Requires-Python: >=3.12
|
|
6
|
+
Requires-Dist: annotide>=0.1
|
|
7
|
+
Provides-Extra: torch
|
|
8
|
+
Requires-Dist: torch>=2.5; extra == "torch"
|
|
9
|
+
Requires-Dist: torchvision>=0.20; extra == "torch"
|
|
10
|
+
Requires-Dist: onnx>=1.16; extra == "torch"
|
|
11
|
+
Requires-Dist: onnxruntime>=1.20; extra == "torch"
|
|
12
|
+
Requires-Dist: numpy>=2.1; extra == "torch"
|
|
13
|
+
Requires-Dist: pillow>=11.0; extra == "torch"
|
|
14
|
+
Provides-Extra: mlflow
|
|
15
|
+
Requires-Dist: mlflow-skinny>=3.1; extra == "mlflow"
|
|
16
|
+
Provides-Extra: dev
|
|
17
|
+
Requires-Dist: pytest>=8.3; extra == "dev"
|
|
18
|
+
Requires-Dist: ruff>=0.8; extra == "dev"
|
|
19
|
+
Requires-Dist: mypy>=1.13; extra == "dev"
|
|
@@ -0,0 +1,177 @@
|
|
|
1
|
+
# Reference training pipeline
|
|
2
|
+
|
|
3
|
+
The platform never trains a model itself (ML-9). It freezes a snapshot,
|
|
4
|
+
emits `retrain.requested` to the project's webhooks, and records whatever
|
|
5
|
+
version comes back with its lineage (EXP-8). This package is the other end:
|
|
6
|
+
a pipeline a customer runs next to their own compute.
|
|
7
|
+
|
|
8
|
+
```
|
|
9
|
+
retrain.requested ─▶ prepare ─▶ split ─▶ train ─▶ evaluate ─▶ register
|
|
10
|
+
(webhook) export EXP-3 Trainer box P/R/F1 POST /models/{id}/versions
|
|
11
|
+
+ digest snapshot_id + digest + training_run
|
|
12
|
+
```
|
|
13
|
+
|
|
14
|
+
1. **prepare** — reads the snapshot row, refuses a digest that is not the
|
|
15
|
+
snapshot's, queues a `native` export of it, downloads the archive through
|
|
16
|
+
its signed URL (the API key never goes to the storage host) and checks
|
|
17
|
+
the archive's manifest names the same snapshot and digest.
|
|
18
|
+
2. **split** — uses the snapshot's own train / val / test partition when it
|
|
19
|
+
has one; otherwise splits with the platform's rule
|
|
20
|
+
(`sha256("{seed}:item:{id}")`), so the same seed gives the same sets.
|
|
21
|
+
3. **train** — a pluggable `Trainer` (`annotide_training/trainers.py`).
|
|
22
|
+
4. **evaluate** — box precision / recall / F1 at IoU 0.5, per class and
|
|
23
|
+
overall, on `test` (or `val` when `test` is empty). The same scorer for
|
|
24
|
+
every trainer.
|
|
25
|
+
5. **register** — adds a model version with `snapshot_id`,
|
|
26
|
+
`snapshot_digest`, the metrics and a `training_run` record (run id,
|
|
27
|
+
trainer, params, timings, export job, split counts, artifact location,
|
|
28
|
+
webhook delivery id). The platform re-checks the digest (409 otherwise).
|
|
29
|
+
|
|
30
|
+
Artifacts and a `run.json` land in `runs/{run_id}/`. Where the artifact goes
|
|
31
|
+
next — an ONNX file for `model-service`'s `onnx` backend, a model registry —
|
|
32
|
+
is the trainer's business; the platform only stores the lineage.
|
|
33
|
+
|
|
34
|
+
## Run it
|
|
35
|
+
|
|
36
|
+
```sh
|
|
37
|
+
cd training
|
|
38
|
+
python3.12 -m venv .venv
|
|
39
|
+
.venv/bin/pip install -e ../sdk # first: the SDK is not on a package index
|
|
40
|
+
.venv/bin/pip install -e '.[dev]'
|
|
41
|
+
|
|
42
|
+
export ANNOTIDE_API_URL=http://localhost:8000 # site root, not /api/v1
|
|
43
|
+
export ANNOTIDE_API_KEY=... # an API key with the `write` scope (AUTH-4)
|
|
44
|
+
|
|
45
|
+
# One snapshot, now. --no-register trains and evaluates without adding a version.
|
|
46
|
+
.venv/bin/annotide-train --model <model-id> run --project <project-id> --snapshot <snapshot-id>
|
|
47
|
+
|
|
48
|
+
# Or receive webhooks: subscribe http://<host>:8088/ to retrain.requested
|
|
49
|
+
export ANNOTIDE_WEBHOOK_SECRET=... # the webhook's signing secret
|
|
50
|
+
.venv/bin/annotide-train --model <fallback-model-id> serve --host 0.0.0.0 --port 8088
|
|
51
|
+
```
|
|
52
|
+
|
|
53
|
+
`serve` verifies `X-Annotation-Signature` (5 min tolerance, constant-time),
|
|
54
|
+
answers `202` at once and trains on a thread (deliveries time out), skips a
|
|
55
|
+
repeated `X-Annotation-Delivery`, and acknowledges signed events it does not
|
|
56
|
+
act on (another event, no snapshot) with `200` so they are not retried.
|
|
57
|
+
The event's `model_id` wins over `--model`.
|
|
58
|
+
|
|
59
|
+
Options: `--trainer` (`baseline`, `fasterrcnn` or `package.module:ClassName`),
|
|
60
|
+
`--param key=value` (JSON values; repeatable, passed to the trainer),
|
|
61
|
+
`--split 0.8,0.1,0.1` (unsplit snapshots only), `--out runs`,
|
|
62
|
+
`--quantize int8` and `--teacher VERSION_ID` (below).
|
|
63
|
+
|
|
64
|
+
Every registered version carries the metrics the Models page compares
|
|
65
|
+
versions on: `precision` / `recall` / `f1` and `per_class` on the held-out
|
|
66
|
+
split, `size_bytes`, `latency_ms_p50` (median `predict` time per item on this
|
|
67
|
+
machine, `latency_device`) and `dtype`.
|
|
68
|
+
|
|
69
|
+
### Distillation and quantization
|
|
70
|
+
|
|
71
|
+
Both are recorded as the version's `derivation` and `parent_version_id`
|
|
72
|
+
(`docs/CONTRACTS.md`), which the Models page draws as the model's family.
|
|
73
|
+
|
|
74
|
+
**Distillation** here is the annotation loop with a big model as the
|
|
75
|
+
teacher: pre-label with the teacher's version (`POST /projects/{id}/prelabel`),
|
|
76
|
+
let people correct the drafts, take a snapshot, and train a smaller model on
|
|
77
|
+
it with `--teacher <the teacher's version id>`. The new version is registered
|
|
78
|
+
as `distilled` from the teacher, so its F1, size and latency sit next to the
|
|
79
|
+
teacher's. The teacher can be any model of the organisation, including one
|
|
80
|
+
the platform only calls through an endpoint.
|
|
81
|
+
|
|
82
|
+
**Quantization**: `--quantize int8` stores the trained model again at 8 bits,
|
|
83
|
+
scores it on the same held-out split and registers it as a `quantized` child
|
|
84
|
+
of the version this run registered, in `runs/{run_id}/int8/`. The trainer
|
|
85
|
+
implements `quantize(model, calibration, dtype)`; the calibration records are
|
|
86
|
+
a fixed-seed sample of the train split. `fasterrcnn` uses onnxruntime's static
|
|
87
|
+
QDQ quantization (Conv, Gemm, MatMul, per-channel weights); in a smoke test
|
|
88
|
+
the 76 MB model became 20 MB. How much faster it runs depends on the CPU:
|
|
89
|
+
x86 with VNNI / AMX gains most, Apple silicon little. Read the int8
|
|
90
|
+
version's F1 before switching prelabelling to it. `baseline` stores its
|
|
91
|
+
priors as 8-bit fractions, so the path runs without torch.
|
|
92
|
+
|
|
93
|
+
### Recording runs in MLflow (API-6)
|
|
94
|
+
|
|
95
|
+
With the `mlflow` extra (`pip install -e '.[dev,mlflow]'`), `--mlflow-experiment
|
|
96
|
+
NAME` also records each run in MLflow — a local server (`make mlflow` at the
|
|
97
|
+
repo root, `--mlflow-uri http://localhost:5001`), Databricks or Azure ML,
|
|
98
|
+
wherever `MLFLOW_TRACKING_URI` points and however that environment signs in.
|
|
99
|
+
The run gets the params, the numeric metrics, the artifacts under `model/`,
|
|
100
|
+
the tags `annotation.snapshot_id` / `annotation.snapshot_digest` and the
|
|
101
|
+
snapshot as its dataset input (`source_type` `annotation-snapshot`).
|
|
102
|
+
`--mlflow-register NAME` also registers it as a version of that registered
|
|
103
|
+
model. The platform's *Import from MLflow* (`POST /models/{id}/versions/import`)
|
|
104
|
+
reads the lineage back from the run or the registered version, so a version
|
|
105
|
+
imported that way is linked to its snapshot as one registered by this
|
|
106
|
+
pipeline is. `training_run.mlflow` names the run either way.
|
|
107
|
+
|
|
108
|
+
## The Faster R-CNN trainer
|
|
109
|
+
|
|
110
|
+
`--trainer fasterrcnn` fine-tunes torchvision's Faster R-CNN
|
|
111
|
+
(MobileNetV3-Large FPN) on the snapshot's boxes and exports it to ONNX in the
|
|
112
|
+
format `model-service`'s `onnx` backend loads:
|
|
113
|
+
|
|
114
|
+
```sh
|
|
115
|
+
pip install -e '.[torch]' # torch, torchvision, onnx, onnxruntime (CPU wheels:
|
|
116
|
+
# --index-url https://download.pytorch.org/whl/cpu)
|
|
117
|
+
annotide-train run … --trainer fasterrcnn \
|
|
118
|
+
--param image_root=/mnt/images --param epochs=20
|
|
119
|
+
|
|
120
|
+
# runs/{run_id}/fasterrcnn.onnx + fasterrcnn.names → the model service:
|
|
121
|
+
MODEL_BACKEND=onnx MODEL_PATH=runs/{run_id}/fasterrcnn.onnx uvicorn app.main:app
|
|
122
|
+
```
|
|
123
|
+
|
|
124
|
+
- **Images** come from `image_root`, a directory where the customer's own
|
|
125
|
+
storage is mounted or synced (blobfuse2, gcsfuse, `aws s3 sync`): export
|
|
126
|
+
records carry only each item's `path`, and the platform never serves media
|
|
127
|
+
to the trainer. Missing files fail the run before training starts; a path
|
|
128
|
+
that resolves outside `image_root` is refused.
|
|
129
|
+
- **Params**: `epochs` (10), `batch_size` (4), `lr` (0.01, SGD + cosine),
|
|
130
|
+
`weights` (`coco` | `imagenet` | `none`), `hflip` (true), `seed` (0),
|
|
131
|
+
`device` (`cuda` if available, else `cpu`), `score_threshold` (0.5).
|
|
132
|
+
- **Input** is letterboxed to 640 × 640 with grey padding, the same transform
|
|
133
|
+
the model service applies, for training, export and evaluation alike.
|
|
134
|
+
- **Evaluation runs the exported ONNX file** through onnxruntime, not the
|
|
135
|
+
torch model, so the registered metrics are for the artifact that ships.
|
|
136
|
+
The export is also checked against torch on a few training images and a
|
|
137
|
+
blank frame before it is accepted.
|
|
138
|
+
- **Licences**: torchvision, onnx and onnxruntime are BSD / Apache / MIT.
|
|
139
|
+
`weights=coco` starts from torchvision's COCO checkpoint; check its terms
|
|
140
|
+
for your use, or start from `imagenet` / `none`.
|
|
141
|
+
- **Cost**: the artifact is ~76 MB. On an Apple M-series CPU, 48 images ×
|
|
142
|
+
6 epochs take about 90 s; use `device=cuda` for real datasets.
|
|
143
|
+
|
|
144
|
+
## Bring your own trainer
|
|
145
|
+
|
|
146
|
+
```python
|
|
147
|
+
from annotide_training.trainers import Prediction, TrainedModel
|
|
148
|
+
|
|
149
|
+
|
|
150
|
+
class MyDetector:
|
|
151
|
+
name = "my-detector"
|
|
152
|
+
|
|
153
|
+
def train(self, train, val, params) -> TrainedModel:
|
|
154
|
+
# train: native export records — item_id, path, width, height, shapes[…]
|
|
155
|
+
# Media is the customer's own storage: read it from there.
|
|
156
|
+
...
|
|
157
|
+
return TrainedModel(artifact=onnx_bytes, filename="model.onnx", classes=[...])
|
|
158
|
+
|
|
159
|
+
def predict(self, model, record) -> list[Prediction]: ...
|
|
160
|
+
|
|
161
|
+
# Optional, for --quantize: the same model at lower precision, loadable
|
|
162
|
+
# by predict above.
|
|
163
|
+
def quantize(self, model, calibration, dtype) -> TrainedModel: ...
|
|
164
|
+
```
|
|
165
|
+
|
|
166
|
+
`annotide-train --trainer my_package.detector:MyDetector …`. The framework
|
|
167
|
+
it needs is its own dependency, never the platform's.
|
|
168
|
+
|
|
169
|
+
`BaselineTrainer` needs nothing: it learns how often each class appears and
|
|
170
|
+
where it usually sits, and predicts that. It exists to exercise every stage
|
|
171
|
+
end to end and to give a floor a real model has to beat.
|
|
172
|
+
|
|
173
|
+
## Develop
|
|
174
|
+
|
|
175
|
+
```sh
|
|
176
|
+
.venv/bin/ruff check . && .venv/bin/ruff format --check . && .venv/bin/mypy annotide_training tests && .venv/bin/pytest
|
|
177
|
+
```
|
|
@@ -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
|