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