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,477 @@
1
+ """A real trainer: Faster R-CNN (MobileNetV3-FPN) fine-tuned and exported to ONNX.
2
+
3
+ `--trainer fasterrcnn`, with the `torch` extra installed. The artifact is a
4
+ `.onnx` file plus a sibling `.names` file, in exactly the shape
5
+ `model-service`'s `onnx` backend loads (`MODEL_BACKEND=onnx`,
6
+ `MODEL_PATH=…/fasterrcnn.onnx`): one `1x3x640x640` float input in `[0, 1]`,
7
+ letterboxed with grey padding, and three outputs — boxes (`x1 y1 x2 y2` in
8
+ letterboxed pixels), scores, and 0-based class indices into the `.names`
9
+ lines. Normalisation and NMS are inside the graph.
10
+
11
+ Why torchvision: its code is BSD-3-Clause, so a customer can ship the
12
+ trainer and the model in a commercial product. Ultralytics YOLO, the obvious
13
+ alternative, is AGPL-3.0. The COCO weights used for `weights=coco` are
14
+ torchvision's own; check their terms against your use before shipping a
15
+ model fine-tuned from them, or train from `imagenet` / `none`.
16
+
17
+ Media never passes through the platform: records carry only the item's
18
+ `path` in the customer's storage, so the trainer reads images from
19
+ `image_root` — a directory where that storage is mounted or synced
20
+ (blobfuse2, gcsfuse, `aws s3 sync`, …).
21
+
22
+ `predict` runs the exported ONNX file through onnxruntime rather than the
23
+ torch model, so the metrics the pipeline registers are for the artifact that
24
+ ships, including the export and the letterbox.
25
+ """
26
+
27
+ from __future__ import annotations
28
+
29
+ import random
30
+ import tempfile
31
+ from collections.abc import Mapping, Sequence
32
+ from dataclasses import dataclass
33
+ from pathlib import Path
34
+ from typing import Any
35
+
36
+ import numpy as np
37
+ import onnxruntime
38
+ import torch
39
+ from onnxruntime.quantization import (
40
+ CalibrationDataReader,
41
+ QuantFormat,
42
+ QuantType,
43
+ quantize_static,
44
+ )
45
+ from onnxruntime.quantization.shape_inference import quant_pre_process
46
+ from PIL import Image
47
+ from torchvision.models import MobileNet_V3_Large_Weights
48
+ from torchvision.models.detection import (
49
+ FasterRCNN_MobileNet_V3_Large_FPN_Weights,
50
+ fasterrcnn_mobilenet_v3_large_fpn,
51
+ )
52
+ from torchvision.models.detection.faster_rcnn import FastRCNNPredictor
53
+
54
+ from annotide_training.trainers import Box, Prediction, Record, TrainedModel, record_boxes
55
+
56
+ # Both must match model-service/app/backends/onnx.py (`letterbox`).
57
+ INPUT_SIZE = 640
58
+ PAD_GREY = 114
59
+ WEIGHTS = ("coco", "imagenet", "none")
60
+
61
+
62
+ @dataclass(frozen=True, slots=True)
63
+ class Letterbox:
64
+ """Original-image ↔ letterboxed-input coordinates (scale, then centre-pad)."""
65
+
66
+ scale: float
67
+ pad_x: float
68
+ pad_y: float
69
+
70
+ @classmethod
71
+ def fit(cls, width: int, height: int, size: int = INPUT_SIZE) -> Letterbox:
72
+ scale = min(size / max(width, 1), size / max(height, 1))
73
+ return cls(scale, (size - width * scale) / 2, (size - height * scale) / 2)
74
+
75
+ def box_to_input(self, box: Box) -> Box:
76
+ x1, y1, x2, y2 = box
77
+ s, px, py = self.scale, self.pad_x, self.pad_y
78
+ return (x1 * s + px, y1 * s + py, x2 * s + px, y2 * s + py)
79
+
80
+ def box_to_original(self, box: Box, width: int, height: int) -> Box:
81
+ x1, y1, x2, y2 = box
82
+ s, px, py = self.scale, self.pad_x, self.pad_y
83
+ xs = sorted(((x1 - px) / s, (x2 - px) / s))
84
+ ys = sorted(((y1 - py) / s, (y2 - py) / s))
85
+ return (
86
+ max(0.0, min(xs[0], float(width))),
87
+ max(0.0, min(ys[0], float(height))),
88
+ max(0.0, min(xs[1], float(width))),
89
+ max(0.0, min(ys[1], float(height))),
90
+ )
91
+
92
+
93
+ def letterbox(image: Image.Image, size: int = INPUT_SIZE) -> tuple[np.ndarray, Letterbox]:
94
+ """A `3xsizexsize` float32 array in `[0, 1]` and the transform that made it."""
95
+ transform = Letterbox.fit(image.width, image.height, size)
96
+ new_w = max(1, round(image.width * transform.scale))
97
+ new_h = max(1, round(image.height * transform.scale))
98
+ resized = image.resize((new_w, new_h), Image.Resampling.BILINEAR)
99
+ canvas = Image.new("RGB", (size, size), (PAD_GREY, PAD_GREY, PAD_GREY))
100
+ canvas.paste(resized, (int(transform.pad_x), int(transform.pad_y)))
101
+ array = np.asarray(canvas, dtype=np.float32) / 255.0
102
+ return np.ascontiguousarray(np.transpose(array, (2, 0, 1))), transform
103
+
104
+
105
+ class ImageRoot:
106
+ """Reads a record's image from a local directory, never outside it."""
107
+
108
+ def __init__(self, root: str | Path) -> None:
109
+ self.root = Path(root).resolve()
110
+ if not self.root.is_dir():
111
+ raise ValueError(f"image_root {self.root} is not a directory")
112
+
113
+ def path(self, record: Record) -> Path:
114
+ relative = str(record.get("path") or "").lstrip("/")
115
+ path = (self.root / relative).resolve()
116
+ # An export record is data from elsewhere: a `../` must not walk out.
117
+ if not relative or not path.is_relative_to(self.root):
118
+ raise ValueError(f"item {record.get('item_id')}: path {relative!r} escapes image_root")
119
+ return path
120
+
121
+ def open(self, record: Record) -> Image.Image:
122
+ with Image.open(self.path(record)) as image:
123
+ image.load()
124
+ return image.convert("RGB")
125
+
126
+ def check(self, records: Sequence[Record]) -> None:
127
+ """Fail before training, not an hour into it, when images are missing."""
128
+ missing = [str(r.get("path")) for r in records if not self.path(r).is_file()]
129
+ if missing:
130
+ shown = ", ".join(missing[:5]) + (
131
+ f" and {len(missing) - 5} more" if len(missing) > 5 else ""
132
+ )
133
+ raise FileNotFoundError(f"{len(missing)} image(s) not under {self.root}: {shown}")
134
+
135
+
136
+ def build_model(num_classes: int, weights: str) -> torch.nn.Module:
137
+ """Faster R-CNN with `num_classes` foreground classes (+ background)."""
138
+ if weights not in WEIGHTS:
139
+ raise ValueError(f"weights {weights!r}: use one of {', '.join(WEIGHTS)}")
140
+ # No RPN score floor (torchvision's MobileNet default is 0.05): an image
141
+ # with no proposals left reaches a `Reshape(0, -1)` the ONNX graph cannot run.
142
+ size: dict[str, Any] = {"min_size": INPUT_SIZE, "max_size": INPUT_SIZE, "rpn_score_thresh": 0.0}
143
+ model: torch.nn.Module
144
+ if weights == "coco":
145
+ coco: Any = fasterrcnn_mobilenet_v3_large_fpn(
146
+ weights=FasterRCNN_MobileNet_V3_Large_FPN_Weights.COCO_V1, **size
147
+ )
148
+ in_features = coco.roi_heads.box_predictor.cls_score.in_features
149
+ coco.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes + 1)
150
+ model = coco
151
+ return model
152
+ backbone = MobileNet_V3_Large_Weights.IMAGENET1K_V1 if weights == "imagenet" else None
153
+ model = fasterrcnn_mobilenet_v3_large_fpn(
154
+ weights=None, weights_backbone=backbone, num_classes=num_classes + 1, **size
155
+ )
156
+ return model
157
+
158
+
159
+ class _OnnxExport(torch.nn.Module):
160
+ """Batch-of-one tensor in, (boxes, scores, 0-based labels) out."""
161
+
162
+ def __init__(self, model: torch.nn.Module) -> None:
163
+ super().__init__()
164
+ self.model = model
165
+
166
+ def forward(self, images: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
167
+ detections = self.model([images[0]])[0]
168
+ # torchvision reserves label 0 for background; the .names file does not.
169
+ return detections["boxes"], detections["scores"], detections["labels"] - 1
170
+
171
+
172
+ def export_onnx(model: torch.nn.Module, samples: Sequence[np.ndarray]) -> bytes:
173
+ """Trace the model to ONNX on a real image and check it reproduces torch.
174
+
175
+ The trace input matters: traced on an image with no detections (noise, a
176
+ blank frame), the ROI head's box reshape is folded into a constant and the
177
+ graph fails on any image that has some. So each candidate is traced on a
178
+ training image and accepted only when onnxruntime matches torch on all the
179
+ samples and on a blank frame (nothing to detect).
180
+ """
181
+ if not samples:
182
+ raise ValueError("export needs at least one sample image")
183
+ model.eval().cpu()
184
+ blank = np.full((3, INPUT_SIZE, INPUT_SIZE), PAD_GREY / 255.0, dtype=np.float32)
185
+ checks = [*samples, blank]
186
+ problem = ""
187
+ for trace in samples[:3]:
188
+ with tempfile.TemporaryDirectory() as tmp:
189
+ path = Path(tmp) / "model.onnx"
190
+ # The TorchScript exporter: torchvision's detection models are
191
+ # written for it. `.eval()` on the wrapper too, or export restores
192
+ # it — and the detector inside — to training mode afterwards.
193
+ torch.onnx.export(
194
+ _OnnxExport(model).eval(),
195
+ (torch.from_numpy(trace)[None],),
196
+ path,
197
+ input_names=["images"],
198
+ output_names=["boxes", "scores", "labels"],
199
+ dynamic_axes={"boxes": {0: "n"}, "scores": {0: "n"}, "labels": {0: "n"}},
200
+ opset_version=17,
201
+ dynamo=False,
202
+ )
203
+ artifact = path.read_bytes()
204
+ problem = _mismatch(model, artifact, checks)
205
+ if not problem:
206
+ return artifact
207
+ raise RuntimeError(f"the ONNX export does not reproduce the torch model: {problem}")
208
+
209
+
210
+ def _mismatch(model: torch.nn.Module, artifact: bytes, samples: Sequence[np.ndarray]) -> str:
211
+ """Empty when onnxruntime gives the torch model's detections on every sample."""
212
+ session = _session(artifact)
213
+ for index, sample in enumerate(samples):
214
+ with torch.no_grad():
215
+ expected = model([torch.from_numpy(sample)])[0]
216
+ try:
217
+ boxes, scores, _ = session.run(None, {"images": sample[None, ...]})
218
+ except Exception as exc: # onnxruntime raises its own Fail type
219
+ return f"sample {index}: {type(exc).__name__}: {exc}"
220
+ want_scores = expected["scores"].numpy()
221
+ want_boxes = expected["boxes"].numpy()
222
+ if len(scores) != len(want_scores):
223
+ return f"sample {index}: {len(scores)} detections, torch has {len(want_scores)}"
224
+ if not len(scores):
225
+ continue
226
+ # Among near-equal scores, float noise decides which proposals survive
227
+ # top-k and NMS, differently in each runtime (a barely trained head
228
+ # scores everything alike). So only boxes whose score stands clear of
229
+ # every other are compared, each to its nearest counterpart.
230
+ distance = np.abs(boxes[:, None, :] - want_boxes[None, :, :]).max(axis=-1)
231
+ own, want_own = _untied(scores), _untied(want_scores)
232
+ score_gap = float(np.abs(np.sort(scores) - np.sort(want_scores)).max())
233
+ box_gap = max(
234
+ float(distance[own].min(axis=1).max()) if own.any() else 0.0,
235
+ float(distance[:, want_own].min(axis=0).max()) if want_own.any() else 0.0,
236
+ )
237
+ if score_gap > 1e-3 or box_gap > 0.5:
238
+ return f"sample {index}: score differs by {score_gap:.2g}, a box by {box_gap:.2g} px"
239
+ return ""
240
+
241
+
242
+ def _untied(scores: np.ndarray) -> np.ndarray:
243
+ """Mask of the scores no other score is within 1e-3 of."""
244
+ gaps = np.abs(scores[:, None] - scores[None, :]) < 1e-3
245
+ return np.asarray(gaps.sum(axis=1) == 1)
246
+
247
+
248
+ class _Calibration(CalibrationDataReader): # type: ignore[misc] # onnxruntime is untyped
249
+ """Letterboxed training images, one batch of one at a time."""
250
+
251
+ def __init__(self, samples: Sequence[np.ndarray]) -> None:
252
+ self._samples = iter(samples)
253
+
254
+ def get_next(self) -> dict[str, np.ndarray] | None:
255
+ sample = next(self._samples, None)
256
+ return None if sample is None else {"images": sample[None, ...]}
257
+
258
+
259
+ def quantize_onnx(artifact: bytes, samples: Sequence[np.ndarray]) -> bytes:
260
+ """Static int8 (QDQ) quantization, calibrated on real letterboxed images.
261
+
262
+ Conv covers the backbone and FPN; Gemm / MatMul the box head's two fully
263
+ connected layers, which hold most of the weights. NMS, RoiAlign and the
264
+ box decoding stay float. Weights per channel, activations from the
265
+ calibration images' ranges, so they must look like what the model will see.
266
+ """
267
+ if not samples:
268
+ raise ValueError("int8 quantization needs at least one calibration image")
269
+ with tempfile.TemporaryDirectory() as tmp:
270
+ source, prepared, output = (Path(tmp) / n for n in ("in.onnx", "pre.onnx", "q.onnx"))
271
+ source.write_bytes(artifact)
272
+ # Shape inference and constant folding so the quantizer sees every Conv;
273
+ # the symbolic pass fails on the detection head's dynamic shapes.
274
+ quant_pre_process(str(source), str(prepared), skip_symbolic_shape=True)
275
+ quantize_static(
276
+ str(prepared),
277
+ str(output),
278
+ _Calibration(samples),
279
+ quant_format=QuantFormat.QDQ,
280
+ per_channel=True,
281
+ activation_type=QuantType.QUInt8,
282
+ weight_type=QuantType.QInt8,
283
+ op_types_to_quantize=["Conv", "Gemm", "MatMul"],
284
+ )
285
+ return output.read_bytes()
286
+
287
+
288
+ def _session(artifact: bytes) -> Any:
289
+ options = onnxruntime.SessionOptions()
290
+ options.log_severity_level = 3 # the exported graph's shape warnings are noise
291
+ return onnxruntime.InferenceSession(artifact, options, providers=["CPUExecutionProvider"])
292
+
293
+
294
+ class FasterRCNNTrainer:
295
+ """Fine-tunes Faster R-CNN on the snapshot's boxes and exports it to ONNX.
296
+
297
+ `--quantize int8` stores the exported graph as static int8 (QDQ),
298
+ calibrated on up to `calibration_images` (32) training images.
299
+
300
+ Params (`--param key=value`): `image_root` (required unless given to the
301
+ constructor), `epochs` (10), `batch_size` (4), `lr` (0.01, SGD with
302
+ cosine decay), `weights` (`coco` | `imagenet` | `none`), `hflip` (true),
303
+ `seed` (0), `device` (`cuda` when available, else `cpu`),
304
+ `score_threshold` (0.5, applied in `predict` and so in the metrics).
305
+ """
306
+
307
+ name = "fasterrcnn-mobilenet-v3-onnx"
308
+
309
+ def __init__(self, image_root: str | Path | None = None) -> None:
310
+ self._images = ImageRoot(image_root) if image_root is not None else None
311
+ self._score_threshold = 0.5
312
+ self._calibration_images = 32
313
+ self._session: tuple[bytes, Any] | None = None
314
+
315
+ def train(
316
+ self, train: Sequence[Record], val: Sequence[Record], params: Mapping[str, Any]
317
+ ) -> TrainedModel:
318
+ if "image_root" in params:
319
+ self._images = ImageRoot(str(params["image_root"]))
320
+ if self._images is None:
321
+ raise ValueError("fasterrcnn needs image_root: where the item paths are mounted")
322
+ images = self._images
323
+ epochs = int(params.get("epochs", 10))
324
+ batch_size = int(params.get("batch_size", 4))
325
+ lr = float(params.get("lr", 0.01))
326
+ weights = str(params.get("weights", "coco"))
327
+ hflip = bool(params.get("hflip", True))
328
+ seed = int(params.get("seed", 0))
329
+ device = torch.device(
330
+ str(params.get("device") or ("cuda" if torch.cuda.is_available() else "cpu"))
331
+ )
332
+ self._score_threshold = float(params.get("score_threshold", 0.5))
333
+ self._calibration_images = int(params.get("calibration_images", 32))
334
+ if epochs < 1 or batch_size < 1:
335
+ raise ValueError("epochs and batch_size must be at least 1")
336
+
337
+ classes = sorted({name for record in train for name, _ in record_boxes(record)})
338
+ if not classes:
339
+ raise ValueError("the train split has no bbox shapes to learn from")
340
+ images.check([*train, *val])
341
+
342
+ torch.manual_seed(seed)
343
+ rng = random.Random(seed)
344
+ model = build_model(len(classes), weights).to(device)
345
+ model.train()
346
+ trainable = [p for p in model.parameters() if p.requires_grad]
347
+ optimizer = torch.optim.SGD(trainable, lr=lr, momentum=0.9, weight_decay=5e-4)
348
+ steps = epochs * -(-len(train) // batch_size)
349
+ scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=steps)
350
+
351
+ epoch_losses: list[float] = []
352
+ order = list(range(len(train)))
353
+ for _ in range(epochs):
354
+ rng.shuffle(order)
355
+ total, batches = 0.0, 0
356
+ for start in range(0, len(order), batch_size):
357
+ inputs, targets = [], []
358
+ for index in order[start : start + batch_size]:
359
+ flip = hflip and rng.random() < 0.5
360
+ array, target = self._sample(images, train[index], classes, flip)
361
+ inputs.append(torch.from_numpy(array).to(device))
362
+ targets.append({k: v.to(device) for k, v in target.items()})
363
+ losses = model(inputs, targets)
364
+ loss = torch.stack(list(losses.values())).sum()
365
+ optimizer.zero_grad()
366
+ torch.autograd.backward(loss)
367
+ optimizer.step()
368
+ scheduler.step()
369
+ total += float(loss.detach())
370
+ batches += 1
371
+ epoch_losses.append(round(total / batches, 4))
372
+
373
+ # Trace on the images with the most boxes: they are sure to detect something.
374
+ busiest = sorted(train, key=lambda r: -len(record_boxes(r)))[:4]
375
+ artifact = export_onnx(model, [letterbox(images.open(r))[0] for r in busiest])
376
+ return TrainedModel(
377
+ artifact=artifact,
378
+ filename="fasterrcnn.onnx",
379
+ classes=classes,
380
+ info={
381
+ "architecture": "fasterrcnn_mobilenet_v3_large_fpn",
382
+ "weights": weights,
383
+ "input_size": INPUT_SIZE,
384
+ "epochs": epochs,
385
+ "epoch_losses": epoch_losses,
386
+ "train_images": len(train),
387
+ "val_images": len(val),
388
+ "score_threshold": self._score_threshold,
389
+ "onnx_bytes": len(artifact),
390
+ "dtype": "fp32",
391
+ },
392
+ extra_files={"fasterrcnn.names": ("\n".join(classes) + "\n").encode()},
393
+ )
394
+
395
+ def quantize(
396
+ self, model: TrainedModel, calibration: Sequence[Record], dtype: str
397
+ ) -> TrainedModel:
398
+ if dtype != "int8":
399
+ raise ValueError(f"{self.name} quantizes to int8 only, not {dtype!r}")
400
+ if self._images is None:
401
+ raise ValueError("fasterrcnn needs image_root to calibrate")
402
+ images = self._images
403
+ samples = [letterbox(images.open(r))[0] for r in calibration[: self._calibration_images]]
404
+ artifact = quantize_onnx(model.artifact, samples)
405
+ return TrainedModel(
406
+ artifact=artifact,
407
+ filename=model.filename,
408
+ classes=list(model.classes),
409
+ info={
410
+ **model.info,
411
+ "dtype": "int8",
412
+ "quantization": "static QDQ, per-channel weights; Conv, Gemm, MatMul",
413
+ "calibration_images": len(samples),
414
+ "onnx_bytes": len(artifact),
415
+ },
416
+ extra_files=dict(model.extra_files),
417
+ )
418
+
419
+ def _sample(
420
+ self, images: ImageRoot, record: Record, classes: list[str], flip: bool
421
+ ) -> tuple[np.ndarray, dict[str, torch.Tensor]]:
422
+ image = images.open(record)
423
+ # Boxes are in the record's pixel frame; rescale if the file differs.
424
+ sx = image.width / float(record.get("width") or image.width)
425
+ sy = image.height / float(record.get("height") or image.height)
426
+ array, transform = letterbox(image)
427
+ boxes: list[Box] = []
428
+ labels: list[int] = []
429
+ for name, (x1, y1, x2, y2) in record_boxes(record):
430
+ if name not in classes:
431
+ continue
432
+ bx1, by1, bx2, by2 = transform.box_to_input((x1 * sx, y1 * sy, x2 * sx, y2 * sy))
433
+ if flip:
434
+ bx1, bx2 = INPUT_SIZE - bx2, INPUT_SIZE - bx1
435
+ # Faster R-CNN rejects zero-area targets.
436
+ if bx2 - bx1 >= 1 and by2 - by1 >= 1:
437
+ boxes.append((bx1, by1, bx2, by2))
438
+ labels.append(classes.index(name) + 1)
439
+ if flip:
440
+ array = np.ascontiguousarray(array[:, :, ::-1])
441
+ return array, {
442
+ "boxes": torch.tensor(boxes, dtype=torch.float32).reshape(-1, 4),
443
+ "labels": torch.tensor(labels, dtype=torch.int64),
444
+ }
445
+
446
+ def predict(self, model: TrainedModel, record: Record) -> list[Prediction]:
447
+ if self._images is None:
448
+ raise ValueError("fasterrcnn needs image_root to predict")
449
+ session = self._onnx_session(model.artifact)
450
+ image = self._images.open(record)
451
+ array, transform = letterbox(image)
452
+ boxes, scores, labels = session.run(None, {"images": array[None, ...]})
453
+ width = int(record.get("width") or image.width)
454
+ height = int(record.get("height") or image.height)
455
+ sx, sy = width / image.width, height / image.height
456
+ predictions: list[Prediction] = []
457
+ for box, score, label in zip(boxes, scores, labels, strict=True):
458
+ if float(score) < self._score_threshold or not 0 <= int(label) < len(model.classes):
459
+ continue
460
+ x1, y1, x2, y2 = transform.box_to_original(
461
+ (float(box[0]), float(box[1]), float(box[2]), float(box[3])),
462
+ image.width,
463
+ image.height,
464
+ )
465
+ predictions.append(
466
+ Prediction(
467
+ class_=model.classes[int(label)],
468
+ bbox=(x1 * sx, y1 * sy, x2 * sx, y2 * sy),
469
+ confidence=float(score),
470
+ )
471
+ )
472
+ return predictions
473
+
474
+ def _onnx_session(self, artifact: bytes) -> Any:
475
+ if self._session is None or self._session[0] is not artifact:
476
+ self._session = (artifact, _session(artifact))
477
+ return self._session[1]
@@ -0,0 +1,89 @@
1
+ """Evaluate: box precision / recall / F1 at an IoU threshold, per class and overall.
2
+
3
+ Also the median time `predict` takes per item, so a quantized or distilled
4
+ version can be compared with its parent on speed as well as accuracy.
5
+
6
+ Framework-independent, so every trainer is scored the same way on the same
7
+ held-out split. Greedy matching by confidence, one prediction per ground
8
+ truth box, classes never cross.
9
+ """
10
+
11
+ from __future__ import annotations
12
+
13
+ import statistics
14
+ import time
15
+ from collections.abc import Sequence
16
+ from typing import Any
17
+
18
+ from annotide_training.trainers import Box, Record, TrainedModel, Trainer, record_boxes
19
+
20
+
21
+ def iou(a: Box, b: Box) -> float:
22
+ ix1, iy1 = max(a[0], b[0]), max(a[1], b[1])
23
+ ix2, iy2 = min(a[2], b[2]), min(a[3], b[3])
24
+ inter = max(0.0, ix2 - ix1) * max(0.0, iy2 - iy1)
25
+ union = (a[2] - a[0]) * (a[3] - a[1]) + (b[2] - b[0]) * (b[3] - b[1]) - inter
26
+ return inter / union if union > 0 else 0.0
27
+
28
+
29
+ def _prf(tp: int, fp: int, fn: int) -> dict[str, float | int]:
30
+ precision = tp / (tp + fp) if tp + fp else 0.0
31
+ recall = tp / (tp + fn) if tp + fn else 0.0
32
+ f1 = 2 * precision * recall / (precision + recall) if precision + recall else 0.0
33
+ return {
34
+ "tp": tp,
35
+ "fp": fp,
36
+ "fn": fn,
37
+ "precision": round(precision, 4),
38
+ "recall": round(recall, 4),
39
+ "f1": round(f1, 4),
40
+ }
41
+
42
+
43
+ def evaluate(
44
+ trainer: Trainer,
45
+ model: TrainedModel,
46
+ records: Sequence[Record],
47
+ *,
48
+ iou_threshold: float = 0.5,
49
+ min_confidence: float = 0.0,
50
+ ) -> dict[str, Any]:
51
+ tallies: dict[str, list[int]] = {} # class -> [tp, fp, fn]
52
+ latencies: list[float] = []
53
+ for record in records:
54
+ truths = record_boxes(record)
55
+ started = time.perf_counter()
56
+ raw = trainer.predict(model, record)
57
+ latencies.append(time.perf_counter() - started)
58
+ predictions = sorted(
59
+ (p for p in raw if p.confidence >= min_confidence),
60
+ key=lambda p: -p.confidence,
61
+ )
62
+ matched: set[int] = set()
63
+ for prediction in predictions:
64
+ tally = tallies.setdefault(prediction.class_, [0, 0, 0])
65
+ best, best_iou = -1, iou_threshold
66
+ for index, (class_name, box) in enumerate(truths):
67
+ if index in matched or class_name != prediction.class_:
68
+ continue
69
+ overlap = iou(prediction.bbox, box)
70
+ if overlap >= best_iou:
71
+ best, best_iou = index, overlap
72
+ if best >= 0:
73
+ matched.add(best)
74
+ tally[0] += 1
75
+ else:
76
+ tally[1] += 1
77
+ for index, (class_name, _) in enumerate(truths):
78
+ if index not in matched:
79
+ tallies.setdefault(class_name, [0, 0, 0])[2] += 1
80
+
81
+ totals = [sum(t[i] for t in tallies.values()) for i in range(3)]
82
+ return {
83
+ "iou_threshold": iou_threshold,
84
+ "images": len(records),
85
+ # Per item, reading the image included, on the machine that ran this.
86
+ "latency_ms_p50": round(statistics.median(latencies) * 1000) if latencies else None,
87
+ **_prf(*totals),
88
+ "per_class": {name: _prf(*tallies[name]) for name in sorted(tallies)},
89
+ }