cmvr_edge_ai/scripts/export_detection_onnx.py

360 lines
12 KiB
Python

#!/usr/bin/env python3
"""Export one registered YOLOv8 detector into the strict edge ONNX contract."""
from __future__ import annotations
import argparse
import hashlib
import json
import os
import shutil
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Any
# Ultralytics normally tries to install missing export dependencies at runtime.
# That would bypass uv.lock and make deployment artifacts irreproducible.
os.environ["YOLO_AUTOINSTALL"] = "false"
from cmvr_edge_ai.detection.models.mobile_phone import MOBILE_PHONE_LABELS
from cmvr_edge_ai.detection.models.people_talking import PEOPLE_TALKING_LABELS
from cmvr_edge_ai.detection.models.ppe_yolo import (
CONSTRUCTION_PPE_LABELS,
PPE_6CLASS_LABELS,
)
from cmvr_edge_ai.detection.models.yolo_onnx import OnnxYoloModel
@dataclass(frozen=True, slots=True)
class ExportDefinition:
source_model_id: str
target_model_id: str
source: Path
output: Path
labels: tuple[str, ...]
EXPORT_DEFINITIONS = {
"construction-ppe-yolov8@2": ExportDefinition(
source_model_id="construction-ppe-yolov8@1",
target_model_id="construction-ppe-yolov8@2",
source=Path("models/detection/construction-ppe-yolov8/v1/best.pt"),
output=Path("models/detection/construction-ppe-yolov8/v2/model.onnx"),
labels=CONSTRUCTION_PPE_LABELS,
),
"ppe-6classes-yolov8n@2": ExportDefinition(
source_model_id="ppe-6classes-yolov8n@1",
target_model_id="ppe-6classes-yolov8n@2",
source=Path("models/detection/ppe-6classes-yolov8n/v1/best.pt"),
output=Path("models/detection/ppe-6classes-yolov8n/v2/model.onnx"),
labels=PPE_6CLASS_LABELS,
),
"people-talking-yolov8x@2": ExportDefinition(
source_model_id="people-talking-yolov8x@1",
target_model_id="people-talking-yolov8x@2",
source=Path("models/detection/people-talking-yolov8x/v1/best.pt"),
output=Path("models/detection/people-talking-yolov8x/v2/model.onnx"),
labels=PEOPLE_TALKING_LABELS,
),
"yolov8n-mobile-phone@2": ExportDefinition(
source_model_id="yolov8n-mobile-phone@1",
target_model_id="yolov8n-mobile-phone@2",
source=Path(
"models/detection/yolov8n-mobile-phone/yolov8n-mobile-phone.pt"
),
output=Path("models/detection/yolov8n-mobile-phone/v2/model.onnx"),
labels=MOBILE_PHONE_LABELS,
),
}
def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument(
"--model-id",
required=True,
choices=tuple(sorted(EXPORT_DEFINITIONS)),
help="registered @2 ONNX model to build",
)
parser.add_argument(
"--source",
type=Path,
help="override the registered trusted .pt source path",
)
parser.add_argument(
"--output",
type=Path,
help="override the registered v2/model.onnx destination",
)
parser.add_argument(
"--manifest",
type=Path,
help="manifest destination (default: manifest.json beside the ONNX file)",
)
parser.add_argument("--imgsz", type=int, default=640)
parser.add_argument("--opset", type=int, default=17)
parser.add_argument(
"--simplify",
action=argparse.BooleanOptionalAction,
default=True,
)
parser.add_argument(
"--force",
action="store_true",
help="replace an existing validated output and manifest",
)
return parser.parse_args(argv)
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as artifact:
for chunk in iter(lambda: artifact.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _ordered_names(value: Any) -> tuple[str, ...]:
if isinstance(value, dict):
normalized = {int(key): str(label) for key, label in value.items()}
indexes = sorted(normalized)
if indexes != list(range(len(indexes))):
raise ValueError("source checkpoint class IDs must start at zero")
return tuple(normalized[index] for index in indexes)
if isinstance(value, (list, tuple)):
return tuple(str(label) for label in value)
raise ValueError(f"source checkpoint has invalid class names: {value!r}")
def _sanitize_onnx_metadata(
onnx: Any,
path: Path,
*,
definition: ExportDefinition,
source_sha256: str,
) -> Any:
"""Remove build-host paths/timestamps and add stable CMVR identity."""
model = onnx.load(str(path))
metadata = {
str(item.key): str(item.value)
for item in getattr(model, "metadata_props", ())
}
metadata.pop("date", None)
metadata["description"] = (
f"CMVR static ONNX export for {definition.target_model_id}"
)
metadata["cmvr_model_id"] = definition.target_model_id
metadata["cmvr_source_sha256"] = source_sha256
onnx.helper.set_model_props(model, metadata)
onnx.save_model(model, str(path))
return model
def _require_export_dependencies(*, simplify: bool) -> dict[str, Any]:
try:
import onnx
import onnxruntime
import torch
import ultralytics
from ultralytics import YOLO
onnxslim = None
if simplify:
import onnxslim
except ImportError as exc:
raise RuntimeError(
"ONNX export dependencies are incomplete; run uv with the "
"onnx-export-cpu extra"
) from exc
return {
"onnx": onnx,
"onnxruntime": onnxruntime,
"onnxslim": onnxslim,
"torch": torch,
"ultralytics": ultralytics,
"YOLO": YOLO,
}
def _validate_paths(
*,
source: Path,
output: Path,
manifest: Path,
force: bool,
) -> None:
if not source.is_file():
raise FileNotFoundError(f"source checkpoint does not exist: {source}")
if source.suffix.lower() != ".pt":
raise ValueError(f"source checkpoint must end in .pt: {source}")
if output.suffix.lower() != ".onnx":
raise ValueError(f"output artifact must end in .onnx: {output}")
if output.resolve() == source.resolve():
raise ValueError("source and output paths must be different")
if manifest.resolve() in {source.resolve(), output.resolve()}:
raise ValueError("manifest path must differ from source and output")
existing = [path for path in (output, manifest) if path.exists()]
if existing and not force:
raise FileExistsError(
"refusing to replace existing export path(s) without --force: "
+ ", ".join(str(path) for path in existing)
)
with source.open("rb") as artifact:
prefix = artifact.read(128)
if b"version https://git-lfs.github.com/spec/v1" in prefix:
raise ValueError(f"source checkpoint is a Git LFS pointer: {source}")
def _write_manifest(path: Path, payload: dict[str, Any]) -> None:
with path.open("x", encoding="utf-8") as stream:
json.dump(payload, stream, ensure_ascii=False, indent=2, sort_keys=True)
stream.write("\n")
def export_model(args: argparse.Namespace) -> tuple[Path, Path]:
definition = EXPORT_DEFINITIONS[args.model_id]
source = (args.source or definition.source).expanduser()
output = (args.output or definition.output).expanduser()
manifest = (
args.manifest.expanduser()
if args.manifest is not None
else output.with_name("manifest.json")
)
if args.imgsz < 32 or args.imgsz % 32:
raise ValueError("imgsz must be at least 32 and divisible by 32")
if args.opset < 12:
raise ValueError("opset must be at least 12")
_validate_paths(
source=source,
output=output,
manifest=manifest,
force=args.force,
)
dependencies = _require_export_dependencies(simplify=args.simplify)
output.parent.mkdir(parents=True, exist_ok=True)
manifest.parent.mkdir(parents=True, exist_ok=True)
source_sha256 = _sha256(source)
with tempfile.TemporaryDirectory(
dir=output.parent,
prefix=".onnx-export-",
) as temporary:
staging = Path(temporary)
staged_source = staging / "source.pt"
shutil.copy2(source, staged_source)
model = dependencies["YOLO"](str(staged_source), task="detect")
if model.task != "detect":
raise ValueError(f"source checkpoint task must be detect, got {model.task!r}")
actual_labels = _ordered_names(model.names)
if actual_labels != definition.labels:
raise ValueError(
"source checkpoint labels do not match the registered model: "
f"expected {definition.labels!r}, got {actual_labels!r}"
)
exported = model.export(
format="onnx",
imgsz=args.imgsz,
batch=1,
dynamic=False,
simplify=args.simplify,
opset=args.opset,
nms=False,
half=False,
device="cpu",
)
staged_output = Path(str(exported)).resolve()
if not staged_output.is_file() or staged_output.suffix.lower() != ".onnx":
raise RuntimeError(
f"Ultralytics did not produce the expected ONNX file: {exported!r}"
)
try:
staged_output.relative_to(staging.resolve())
except ValueError as exc:
raise RuntimeError(
f"export escaped the atomic staging directory: {staged_output}"
) from exc
onnx = dependencies["onnx"]
exported_model = _sanitize_onnx_metadata(
onnx,
staged_output,
definition=definition,
source_sha256=source_sha256,
)
onnx.checker.check_model(exported_model)
runtime = OnnxYoloModel(
{
"weights": str(staged_output),
"providers": ["CPUExecutionProvider"],
"imgsz": args.imgsz,
},
expected_labels=definition.labels,
)
runtime.load()
runtime.close()
output_sha256 = _sha256(staged_output)
staged_manifest = staging / "manifest.json"
onnxslim = dependencies["onnxslim"]
manifest_payload = {
"schema_version": "cmvr.detection-model-manifest/v1",
"source_model_id": definition.source_model_id,
"model_id": definition.target_model_id,
"backend": "onnxruntime-yolov8",
"source": str(source),
"source_sha256": source_sha256,
"artifact": str(output),
"artifact_sha256": output_sha256,
"task": "detect",
"labels": list(definition.labels),
"input": {
"layout": "NCHW",
"shape": [1, 3, args.imgsz, args.imgsz],
"pixel_format": "RGB",
"normalization": "uint8 / 255",
"letterbox_color": [114, 114, 114],
},
"export": {
"format": "onnx",
"opset": args.opset,
"batch": 1,
"dynamic": False,
"nms": False,
"half": False,
"simplify": bool(args.simplify),
},
"tools": {
"torch": dependencies["torch"].__version__,
"ultralytics": dependencies["ultralytics"].__version__,
"onnx": dependencies["onnx"].__version__,
"onnxruntime": dependencies["onnxruntime"].__version__,
"onnxslim": (
None if onnxslim is None else onnxslim.__version__
),
},
}
_write_manifest(staged_manifest, manifest_payload)
# Each file is published with an atomic same-filesystem replacement.
# The manifest is published last and carries the artifact digest, so a
# process interrupted between replacements fails closed on hash check
# instead of accepting a mismatched artifact pair.
os.replace(staged_output, output)
os.replace(staged_manifest, manifest)
return output, manifest
def main(argv: list[str] | None = None) -> int:
output, manifest = export_model(parse_args(argv))
print(f"exported ONNX artifact: {output}")
print(f"wrote model manifest: {manifest}")
return 0
if __name__ == "__main__":
raise SystemExit(main())