File size: 11,006 Bytes
37e3e2c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
"""Standalone inference returning OPDB identifiers and the ``__unknown__`` sentinel."""

from __future__ import annotations

import argparse
import hashlib
import json
import os
import re
from pathlib import Path
from typing import Any, BinaryIO

import numpy as np
import onnxruntime as ort
from PIL import Image

HEADS = ("group", "machine", "exact")
OUTPUT_NAMES = tuple(f"{head}_logits" for head in HEADS)
MEAN = np.asarray([0.485, 0.456, 0.406], dtype=np.float32)[:, None, None]
STD = np.asarray([0.229, 0.224, 0.225], dtype=np.float32)[:, None, None]
_TOKEN = r"[A-Za-z0-9]+"
_GROUP_PATTERN = re.compile(rf"^G({_TOKEN})$")
_MACHINE_PATTERN = re.compile(rf"^G({_TOKEN})-M({_TOKEN})$")
_EXACT_PATTERN = re.compile(rf"^G({_TOKEN})-M({_TOKEN})(?:-A({_TOKEN}))?$")


def _validate_canonical_vocabularies(vocabularies: dict[str, Any]) -> None:
    """Validate cumulative OPDB IDs and their hierarchy without reordering them."""
    if not isinstance(vocabularies, dict) or set(vocabularies) != set(HEADS):
        raise ValueError(
            "classifier vocabularies must contain group, machine, and exact"
        )
    if any(not isinstance(vocabularies[head], list) for head in HEADS):
        raise ValueError("classifier vocabularies must be arrays")

    groups = set(vocabularies["group"])
    machines = set(vocabularies["machine"])
    for group_id in vocabularies["group"]:
        if not isinstance(group_id, str) or _GROUP_PATTERN.fullmatch(group_id) is None:
            raise ValueError(f"noncanonical group vocabulary ID: {group_id!r}")
    for machine_id in vocabularies["machine"]:
        if not isinstance(machine_id, str):
            raise ValueError(f"noncanonical machine vocabulary ID: {machine_id!r}")
        match = _MACHINE_PATTERN.fullmatch(machine_id)
        if match is None:
            raise ValueError(f"noncanonical machine vocabulary ID: {machine_id!r}")
        if f"G{match.group(1)}" not in groups:
            raise ValueError(
                f"machine vocabulary is missing parent group: {machine_id!r}"
            )
    for exact_id in vocabularies["exact"]:
        if exact_id == "__unknown__":
            continue
        if not isinstance(exact_id, str):
            raise ValueError(f"invalid canonical OPDB exact ID: {exact_id!r}")
        match = _EXACT_PATTERN.fullmatch(exact_id)
        if match is None:
            raise ValueError(f"invalid canonical OPDB exact ID: {exact_id!r}")
        group_id = f"G{match.group(1)}"
        machine_id = f"{group_id}-M{match.group(2)}"
        if group_id not in groups:
            raise ValueError(f"exact vocabulary is missing parent group: {exact_id!r}")
        if machine_id not in machines:
            raise ValueError(
                f"exact vocabulary is missing parent machine: {exact_id!r}"
            )


def _stable_softmax(logits: np.ndarray) -> np.ndarray:
    shifted = logits - logits.max(axis=1, keepdims=True)
    exponentials = np.exp(shifted)
    return exponentials / exponentials.sum(axis=1, keepdims=True)


def _verify_sha256(path: Path, expected: str) -> None:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(chunk)
    if digest.hexdigest() != expected:
        raise ValueError(f"classifier model checksum mismatch: {path}")


def preprocess_image(
    source: str | Path | BinaryIO, image_size: int = 256
) -> np.ndarray:
    """Decode, resize, center-crop, and normalize an image for the classifier."""
    with Image.open(source) as opened:
        image = opened.convert("RGB")
        width, height = image.size
        if width <= height:
            resized_size = (image_size, int(image_size * height / width))
        else:
            resized_size = (int(image_size * width / height), image_size)
        resized = image.resize(resized_size, Image.Resampling.BICUBIC)
        left = round((resized.width - image_size) / 2.0)
        top = round((resized.height - image_size) / 2.0)
        cropped = resized.crop((left, top, left + image_size, top + image_size))
        tensor = np.asarray(cropped, dtype=np.float32).transpose(2, 0, 1) / 255.0
    return np.ascontiguousarray((tensor - MEAN) / STD, dtype=np.float32)


def _provider_name(provider: Any) -> str:
    return provider[0] if isinstance(provider, tuple) else provider


def _resolve_providers(device: str | None) -> list[Any]:
    available = set(ort.get_available_providers())
    cpu = "CPUExecutionProvider"
    cuda = "CUDAExecutionProvider"
    if device in (None, "cpu"):
        if device is None and cuda in available:
            return [(cuda, {"device_id": 0}), cpu]
        if cpu not in available:
            raise RuntimeError("ONNX Runtime CPUExecutionProvider is unavailable")
        return [cpu]
    if device == "cuda":
        device_id = 0
    elif device.startswith("cuda:") and device[5:].isdecimal():
        device_id = int(device[5:])
    else:
        raise ValueError("device must be 'cpu', 'cuda', or 'cuda:N'")
    if cuda not in available:
        raise RuntimeError("CUDA was requested but CUDAExecutionProvider is unavailable")
    return [(cuda, {"device_id": device_id}), cpu]


class PinballClassifier:
    """A validated ONNX session for hierarchical pinball classification."""

    def __init__(
        self,
        models_dir: str | Path = Path(__file__).parent,
        device: str | None = None,
        threads: int | None = None,
    ) -> None:
        if threads is not None and threads <= 0:
            raise ValueError("threads must be positive")

        models_path = Path(models_dir)
        metadata = json.loads((models_path / "onnx-metadata.json").read_text())
        self.model_version = metadata["model_version"]
        self.encoder_model = metadata["encoder_model"]
        self.label_schema_version = metadata["label_schema_version"]
        self.vocabularies = metadata["vocabularies"]
        if self.label_schema_version != 2:
            raise ValueError(
                "classifier metadata must use canonical label schema version 2"
            )
        _validate_canonical_vocabularies(self.vocabularies)
        if tuple(metadata["outputs"]) != OUTPUT_NAMES:
            raise ValueError("classifier metadata output order is invalid")

        model_filename = metadata["onnx"]["file"]
        if not isinstance(model_filename, str) or Path(model_filename).name != model_filename:
            raise ValueError("classifier metadata ONNX file must be a basename")
        model_path = models_path / model_filename
        _verify_sha256(model_path, metadata["onnx"]["sha256"])

        providers = _resolve_providers(device)
        options = ort.SessionOptions()
        if _provider_name(providers[0]) == "CPUExecutionProvider":
            configured_threads = threads
            if configured_threads is None:
                configured_threads = int(
                    os.environ.get(
                        "PINBALL_CLASSIFIER_THREADS", min(4, os.cpu_count() or 1)
                    )
                )
            if configured_threads <= 0:
                raise ValueError("threads must be positive")
            options.intra_op_num_threads = configured_threads
            options.inter_op_num_threads = 1
            options.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
        options.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
        self.output_names = list(OUTPUT_NAMES)
        self.session = ort.InferenceSession(
            str(model_path), sess_options=options, providers=providers
        )
        self._validate_model_contract()

    def _validate_model_contract(self) -> None:
        inputs = self.session.get_inputs()
        outputs = self.session.get_outputs()
        if len(inputs) != 1 or inputs[0].name != "images":
            raise ValueError("classifier must expose one input named 'images'")
        input_shape = list(inputs[0].shape)
        if (
            inputs[0].type != "tensor(float)"
            or len(input_shape) != 4
            or input_shape[1:] != [3, 256, 256]
        ):
            raise ValueError("classifier input must be float32 [B,3,256,256]")
        if not (input_shape[0] is None or isinstance(input_shape[0], str)):
            raise ValueError("classifier batch dimension must be dynamic")
        if [output.name for output in outputs] != self.output_names:
            raise ValueError("classifier ONNX output names do not match metadata")
        for head, output in zip(HEADS, outputs, strict=True):
            shape = list(output.shape)
            expected_classes = len(self.vocabularies[head])
            if (
                output.type != "tensor(float)"
                or len(shape) != 2
                or shape[1] != expected_classes
            ):
                raise ValueError(
                    f"classifier {head} output does not match its vocabulary"
                )

    def predict(
        self,
        image_source: str | Path | BinaryIO,
        top_count: int = 5,
    ) -> dict[str, Any]:
        """Rank OPDB identifiers, plus the exact head's ``__unknown__`` sentinel."""
        if not 1 <= top_count <= 20:
            raise ValueError("top_count must be between 1 and 20")
        batch = preprocess_image(image_source)[None]
        logits = self.session.run(self.output_names, {"images": batch})
        result: dict[str, Any] = {
            "model_version": self.model_version,
            "encoder_model": self.encoder_model,
            "label_schema_version": self.label_schema_version,
        }
        for head, values in zip(HEADS, logits, strict=True):
            probabilities = _stable_softmax(values)[0]
            indices = np.argsort(-probabilities, kind="stable")[:top_count]
            vocabulary = self.vocabularies[head]
            result[head] = [
                {
                    "id": vocabulary[int(index)],
                    "confidence": round(float(probabilities[index]), 6),
                }
                for index in indices
            ]
        return result


def _positive_int(value: str) -> int:
    parsed = int(value)
    if parsed <= 0:
        raise argparse.ArgumentTypeError("must be positive")
    return parsed


def _main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("image", type=Path)
    parser.add_argument(
        "--model-dir", type=Path, default=Path(__file__).parent
    )
    parser.add_argument("--top-k", type=int, default=5)
    parser.add_argument("--device", default="cpu")
    parser.add_argument("--threads", type=_positive_int)
    args = parser.parse_args()
    classifier = PinballClassifier(
        models_dir=args.model_dir, device=args.device, threads=args.threads
    )
    prediction = classifier.predict(args.image, top_count=args.top_k)
    print(json.dumps(prediction, indent=2, ensure_ascii=False))


if __name__ == "__main__":
    _main()