cmvr_edge_ai/tests/unit/test_yolo_onnx_model.py

519 lines
15 KiB
Python
Raw Normal View History

from __future__ import annotations
import hashlib
import json
from pathlib import Path
from types import SimpleNamespace
from typing import Any
import pytest
np = pytest.importorskip("numpy")
Image = pytest.importorskip("PIL.Image")
from cmvr_edge_ai.application import (
create_default_capability_registry,
create_default_model_registry,
)
from cmvr_edge_ai.contracts import BoundingBox, Detection, ImageFrame
from cmvr_edge_ai.detection.models import ONNX_MODEL_ARTIFACTS
from cmvr_edge_ai.detection.models.yolo_onnx import (
OnnxYoloDependencyError,
OnnxYoloModel,
)
class _NodeArg:
def __init__(self, name: str, shape: list[Any], type_: str = "tensor(float)") -> None:
self.name = name
self.shape = shape
self.type = type_
class _Session:
def __init__(
self,
raw: Any,
*,
labels: tuple[str, ...],
imgsz: int = 32,
input_shape: list[Any] | None = None,
output_shape: list[Any] | None = None,
metadata: dict[str, str] | None = None,
) -> None:
anchors = sum((imgsz // stride) ** 2 for stride in (8, 16, 32))
self._inputs = [
_NodeArg("images", input_shape or [1, 3, imgsz, imgsz])
]
self._outputs = [
_NodeArg(
"output0",
output_shape or [1, 4 + len(labels), anchors],
)
]
self._metadata = metadata or {
"task": "detect",
"names": repr(dict(enumerate(labels))),
"imgsz": repr([imgsz, imgsz]),
"args": repr({"batch": 1, "dynamic": False, "nms": False}),
}
self.raw = raw
self.run_calls: list[tuple[list[str], dict[str, Any]]] = []
def get_inputs(self) -> list[_NodeArg]:
return self._inputs
def get_outputs(self) -> list[_NodeArg]:
return self._outputs
def get_modelmeta(self) -> Any:
return SimpleNamespace(custom_metadata_map=self._metadata)
def run(self, outputs: list[str], inputs: dict[str, Any]) -> list[Any]:
self.run_calls.append((outputs, inputs))
return [self.raw]
class _SessionOptions:
def __init__(self) -> None:
self.intra_op_num_threads: int | None = None
self.inter_op_num_threads: int | None = None
self.graph_optimization_level: object | None = None
self.execution_mode: object | None = None
class _FakeOrt:
GraphOptimizationLevel = SimpleNamespace(ORT_ENABLE_ALL="all")
ExecutionMode = SimpleNamespace(ORT_SEQUENTIAL="sequential")
SessionOptions = _SessionOptions
def __init__(
self,
session: _Session,
*,
available: tuple[str, ...] = ("CPUExecutionProvider",),
) -> None:
self.session = session
self.available = available
self.created: list[tuple[str, _SessionOptions, list[str]]] = []
def get_available_providers(self) -> list[str]:
return list(self.available)
def InferenceSession(
self,
path: str,
*,
sess_options: _SessionOptions,
providers: list[str],
) -> _Session:
self.created.append((path, sess_options, providers))
return self.session
def _weights(tmp_path: Path) -> Path:
tmp_path.mkdir(parents=True, exist_ok=True)
weights = tmp_path / "model.onnx"
weights.write_bytes(b"test-only-onnx-placeholder")
return weights
def _write_manifest(
weights: Path,
*,
model_id: str,
labels: tuple[str, ...],
source_sha256: str = "1" * 64,
artifact_sha256: str | None = None,
) -> Path:
manifest = weights.with_name("manifest.json")
manifest.write_text(
json.dumps(
{
"schema_version": "cmvr.detection-model-manifest/v1",
"model_id": model_id,
"source_model_id": model_id.replace("@2", "@1"),
"backend": "onnxruntime-yolov8",
"task": "detect",
"labels": list(labels),
"input": {"shape": [1, 3, 32, 32]},
"artifact": str(weights),
"artifact_sha256": artifact_sha256
or hashlib.sha256(weights.read_bytes()).hexdigest(),
"source_sha256": source_sha256,
}
),
encoding="utf-8",
)
return manifest
def _loader(ort: _FakeOrt):
def load(name: str) -> Any:
if name == "numpy":
return np
if name == "PIL.Image":
return Image
if name == "onnxruntime":
return ort
raise AssertionError(f"unexpected module request: {name}")
return load
def _model(
tmp_path: Path,
session: _Session,
*,
labels: tuple[str, ...],
**options: Any,
) -> tuple[OnnxYoloModel, _FakeOrt]:
ort = _FakeOrt(session)
model = OnnxYoloModel(
{
"weights": str(_weights(tmp_path)),
"imgsz": 32,
**options,
},
expected_labels=labels,
module_loader=_loader(ort),
)
return model, ort
def _frame(*, pixel_format: str = "BGR8") -> ImageFrame:
# BGR (10, 20, 30) becomes RGB (30, 20, 10) during preprocessing.
return ImageFrame(
data=bytes((10, 20, 30)) * 8,
width=4,
height=2,
pixel_format=pixel_format,
)
def _raw(labels: tuple[str, ...], imgsz: int = 32) -> Any:
anchors = sum((imgsz // stride) ** 2 for stride in (8, 16, 32))
return np.zeros((1, 4 + len(labels), anchors), dtype=np.float32)
def test_load_is_lazy_and_configures_a_bounded_cpu_session(tmp_path: Path) -> None:
labels = ("mobile_phone",)
session = _Session(_raw(labels), labels=labels)
model, ort = _model(
tmp_path,
session,
labels=labels,
intra_op_threads=2,
inter_op_threads=3,
)
assert ort.created == []
model.load()
assert len(ort.created) == 1
path, options, providers = ort.created[0]
assert path.endswith("model.onnx")
assert providers == ["CPUExecutionProvider"]
assert options.intra_op_num_threads == 2
assert options.inter_op_num_threads == 3
assert options.graph_optimization_level == "all"
assert options.execution_mode == "sequential"
def test_bgr_letterbox_and_box_projection_match_original_frame(tmp_path: Path) -> None:
labels = ("mobile_phone",)
raw = _raw(labels)
# Original box [1, 0, 3, 2] maps to [8, 8, 24, 24] after scale=8,
# top padding=8, represented by YOLO xywh [16, 16, 16, 16].
raw[0, :4, 0] = [16, 16, 16, 16]
raw[0, 4, 0] = 0.9
session = _Session(raw, labels=labels)
model, _ = _model(tmp_path, session, labels=labels, iou=0.5)
model.load()
detections = model.predict(_frame(), labels, 0.5)
assert detections == (
Detection(
label="mobile_phone",
confidence=pytest.approx(0.9),
box=BoundingBox(1.0, 0.0, 3.0, 2.0),
),
)
_, inputs = session.run_calls[0]
tensor = inputs["images"]
assert tensor.shape == (1, 3, 32, 32)
assert tensor.dtype == np.float32
assert tensor[0, :, 0, 0] == pytest.approx(np.array([114, 114, 114]) / 255)
assert tensor[0, :, 8, 0] == pytest.approx(np.array([30, 20, 10]) / 255)
def test_rgb_input_is_not_channel_reversed(tmp_path: Path) -> None:
labels = ("mobile_phone",)
session = _Session(_raw(labels), labels=labels)
model, _ = _model(tmp_path, session, labels=labels)
model.load()
assert model.predict(_frame(pixel_format="RGB8"), labels, 0.5) == ()
tensor = session.run_calls[0][1]["images"]
assert tensor[0, :, 8, 0] == pytest.approx(np.array([10, 20, 30]) / 255)
def test_class_aware_and_agnostic_nms_have_distinct_semantics(tmp_path: Path) -> None:
labels = ("phone", "person")
raw = _raw(labels)
raw[0, :4, 0] = [16, 16, 12, 12]
raw[0, 4, 0] = 0.90
raw[0, :4, 1] = [16, 16, 12, 12]
raw[0, 5, 1] = 0.80
raw[0, :4, 2] = [16, 16, 10, 10]
raw[0, 4, 2] = 0.70
aware_session = _Session(raw, labels=labels)
aware, _ = _model(
tmp_path / "aware",
aware_session,
labels=labels,
iou=0.5,
agnostic_nms=False,
)
aware.load()
aware_results = aware.predict(
ImageFrame(bytes((0, 0, 0)) * 16, 4, 4, "BGR8"),
labels,
0.5,
)
agnostic_session = _Session(raw, labels=labels)
agnostic, _ = _model(
tmp_path / "agnostic",
agnostic_session,
labels=labels,
iou=0.5,
agnostic_nms=True,
)
agnostic.load()
agnostic_results = agnostic.predict(
ImageFrame(bytes((0, 0, 0)) * 16, 4, 4, "BGR8"),
labels,
0.5,
)
assert [item.label for item in aware_results] == ["phone", "person"]
assert [item.label for item in agnostic_results] == ["phone"]
def test_selected_labels_do_not_relabel_an_anchor_from_an_excluded_class(
tmp_path: Path,
) -> None:
labels = ("phone", "person")
raw = _raw(labels)
raw[0, :4, 0] = [16, 16, 12, 12]
raw[0, 4, 0] = 0.70
raw[0, 5, 0] = 0.95
session = _Session(raw, labels=labels)
model, _ = _model(tmp_path, session, labels=labels)
model.load()
detections = model.predict(
ImageFrame(bytes((0, 0, 0)) * 16, 4, 4, "BGR8"),
("phone",),
0.5,
)
assert detections == ()
@pytest.mark.parametrize(
("session_kwargs", "message"),
[
({"input_shape": ["batch", 3, 32, 32]}, "fully static"),
({"input_shape": [2, 3, 32, 32]}, "batch=1"),
({"output_shape": [1, 21, 6]}, "raw YOLOv8"),
(
{
"metadata": {
"task": "detect",
"imgsz": "[32, 32]",
}
},
"metadata.names is required",
),
(
{
"metadata": {
"task": "detect",
"names": "{0: 'other'}",
"imgsz": "[32, 32]",
}
},
"labels do not match",
),
(
{
"metadata": {
"task": "detect",
"names": "{0: 'phone'}",
"imgsz": "[32, 32]",
"args": "{'batch': 1, 'dynamic': False, 'nms': True}",
}
},
"nms=false",
),
],
)
def test_load_rejects_incompatible_graphs_and_metadata(
tmp_path: Path,
session_kwargs: dict[str, Any],
message: str,
) -> None:
labels = ("phone",)
session = _Session(_raw(labels), labels=labels, **session_kwargs)
model, _ = _model(tmp_path, session, labels=labels)
with pytest.raises(ValueError, match=message):
model.load()
def test_missing_dependency_and_provider_errors_are_actionable(tmp_path: Path) -> None:
weights = _weights(tmp_path)
def missing_runtime(name: str) -> Any:
if name == "numpy":
return np
if name == "PIL.Image":
return Image
raise ImportError("onnxruntime unavailable")
model = OnnxYoloModel(
{"weights": str(weights), "imgsz": 32},
expected_labels=("phone",),
module_loader=missing_runtime,
)
with pytest.raises(OnnxYoloDependencyError, match=r"\[onnx-cpu\]"):
model.load()
session = _Session(_raw(("phone",)), labels=("phone",))
ort = _FakeOrt(session, available=("CUDAExecutionProvider",))
unavailable = OnnxYoloModel(
{"weights": str(weights), "imgsz": 32},
expected_labels=("phone",),
module_loader=_loader(ort),
)
with pytest.raises(ValueError, match="provider.*unavailable"):
unavailable.load()
def test_close_is_idempotent_and_requires_reload_before_predict(tmp_path: Path) -> None:
labels = ("phone",)
session = _Session(_raw(labels), labels=labels)
model, _ = _model(tmp_path, session, labels=labels)
model.load()
model.close()
model.close()
with pytest.raises(RuntimeError, match="has not been loaded"):
model.predict(_frame(), labels, 0.5)
def test_builtin_manifest_is_required_and_hash_checked_before_session_creation(
tmp_path: Path,
) -> None:
labels = ("phone",)
model_id = "phone@2"
weights = _weights(tmp_path)
session = _Session(_raw(labels), labels=labels)
ort = _FakeOrt(session)
model = OnnxYoloModel(
{"weights": str(weights), "imgsz": 32},
expected_labels=labels,
expected_model_id=model_id,
module_loader=_loader(ort),
)
with pytest.raises(FileNotFoundError, match="manifest does not exist"):
model.load()
assert ort.created == []
_write_manifest(
weights,
model_id=model_id,
labels=labels,
artifact_sha256="0" * 64,
)
with pytest.raises(ValueError, match="artifact SHA256 mismatch"):
model.load()
assert ort.created == []
def test_valid_builtin_manifest_and_embedded_identity_allow_session_load(
tmp_path: Path,
) -> None:
labels = ("phone",)
model_id = "phone@2"
source_sha256 = "2" * 64
weights = _weights(tmp_path)
_write_manifest(
weights,
model_id=model_id,
labels=labels,
source_sha256=source_sha256,
)
session = _Session(
_raw(labels),
labels=labels,
metadata={
"task": "detect",
"names": repr(dict(enumerate(labels))),
"imgsz": repr([32, 32]),
"args": repr({"batch": 1, "dynamic": False, "nms": False}),
"cmvr_model_id": model_id,
"cmvr_source_sha256": source_sha256,
},
)
ort = _FakeOrt(session)
model = OnnxYoloModel(
{"weights": str(weights), "imgsz": 32},
expected_labels=labels,
expected_model_id=model_id,
module_loader=_loader(ort),
)
model.load()
assert len(ort.created) == 1
def test_default_registry_contains_four_onnx_successors() -> None:
registry = create_default_model_registry(discover_entry_points=False)
expected = {
"construction-ppe-yolov8@2",
"ppe-6classes-yolov8n@2",
"people-talking-yolov8x@2",
"yolov8n-mobile-phone@2",
}
assert expected <= {spec.model_id for spec in registry.specs()}
for model_id in expected:
spec = registry.resolve(model_id)
assert spec.backend == "onnxruntime-yolov8"
model = spec.factory({})
assert isinstance(model, OnnxYoloModel)
assert model.weights_path == Path(ONNX_MODEL_ARTIFACTS[model_id])
def test_default_capability_registry_exposes_onnx_successors_by_same_categories() -> None:
registry = create_default_capability_registry(discover_entry_points=False)
expected = {
"construction-ppe-yolov8@2": "detect.ppe",
"ppe-6classes-yolov8n@2": "detect.ppe_6classes",
"people-talking-yolov8x@2": "detect.phone_use",
"yolov8n-mobile-phone@2": "detect.mobile_phone",
}
for model_id, category in expected.items():
spec = registry.resolve(model_id)
assert spec.category == category
assert spec.backend == "onnxruntime-yolov8"