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