Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 23 additions & 1 deletion .github/workflows/detector-export.yml
Original file line number Diff line number Diff line change
Expand Up @@ -4,21 +4,29 @@ on:
pull_request:
paths:
- ".github/workflows/detector-export.yml"
- "complexity/deploy/onnx_detector/**"
- "complexity/generative/detection/**"
- "scripts/check_onnx_parity.py"
- "scripts/onnx_detect.py"
- "scripts/export_onnx.py"
- "scripts/export_tensorrt.py"
- "tests/test_detector_export.py"
- "tests/test_onnx_detect_cli.py"
- "tests/test_onnx_detector_*.py"
- "pyproject.toml"
push:
branches: [main]
paths:
- ".github/workflows/detector-export.yml"
- "complexity/deploy/onnx_detector/**"
- "complexity/generative/detection/**"
- "scripts/check_onnx_parity.py"
- "scripts/onnx_detect.py"
- "scripts/export_onnx.py"
- "scripts/export_tensorrt.py"
- "tests/test_detector_export.py"
- "tests/test_onnx_detect_cli.py"
- "tests/test_onnx_detector_*.py"
- "pyproject.toml"

permissions:
Expand All @@ -42,10 +50,24 @@ jobs:
- name: Lint detector export code
run: >-
ruff check
complexity/deploy/onnx_detector
complexity/generative/detection/exporting.py
scripts/check_onnx_parity.py
scripts/onnx_detect.py
scripts/export_onnx.py
scripts/export_tensorrt.py
tests/test_detector_export.py
tests/test_onnx_detect_cli.py
tests/test_onnx_detector_core.py
tests/test_onnx_detector_metadata.py
tests/test_onnx_detector_pipeline.py
tests/test_onnx_detector_skeleton.py
- name: Test ONNX branches and dynamic batch parity
run: pytest -q tests/test_detector_export.py
run: >-
pytest -q
tests/test_detector_export.py
tests/test_onnx_detect_cli.py
tests/test_onnx_detector_core.py
tests/test_onnx_detector_metadata.py
tests/test_onnx_detector_pipeline.py
tests/test_onnx_detector_skeleton.py
4 changes: 4 additions & 0 deletions complexity/deploy/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,4 @@
"""Deployment helpers for exported Complexity models."""

__all__ = ["onnx_detector"]

31 changes: 31 additions & 0 deletions complexity/deploy/onnx_detector/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,31 @@
"""ONNX Runtime deployment helpers for TR-Hash Vision detectors."""

from .metadata import (
DEFAULT_CONFIDENCE_THRESHOLD,
DEFAULT_IOU_THRESHOLD,
DEFAULT_MAX_DETECTIONS,
BranchType,
OnnxDetectorMetadata,
)
from .pipeline import OnnxDetectorPipeline
from .preprocess import ImageGeometry, PreprocessResult, preprocess_image, restore_boxes
from .session import OnnxDetectorSession, OrtSessionConfig
from .types import Detection, DetectionResult, TimingBreakdown

__all__ = [
"BranchType",
"DEFAULT_CONFIDENCE_THRESHOLD",
"DEFAULT_IOU_THRESHOLD",
"DEFAULT_MAX_DETECTIONS",
"Detection",
"DetectionResult",
"ImageGeometry",
"OnnxDetectorMetadata",
"OnnxDetectorPipeline",
"OnnxDetectorSession",
"OrtSessionConfig",
"PreprocessResult",
"TimingBreakdown",
"preprocess_image",
"restore_boxes",
]
49 changes: 49 additions & 0 deletions complexity/deploy/onnx_detector/dfl.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,49 @@
"""DFL box decode for TR-Hash Vision v8 ONNX detector exports."""

from __future__ import annotations

import numpy as np

from .grid import GridGeometry
from .metadata import OnnxDetectorMetadata


def decode_dfl_boxes(
regression_logits: np.ndarray,
metadata: OnnxDetectorMetadata,
geometry: GridGeometry,
) -> np.ndarray:
"""Decode raw LTRB DFL logits into input-pixel xyxy boxes."""

logits = np.asarray(regression_logits, dtype=np.float32)
if logits.ndim not in {2, 3}:
raise ValueError("regression_logits must have shape [N, C] or [B, N, C]")
if logits.shape[-1] != metadata.regression_width:
raise ValueError(
"regression_logits last dimension does not match metadata.regression_width"
)
if logits.shape[-2] != geometry.centers_xy.shape[0]:
raise ValueError("regression_logits cell count does not match grid geometry")

bins = metadata.dfl_bins
distances = logits.reshape(*logits.shape[:-1], 4, bins)
if metadata.reg_max:
shifted = distances - distances.max(axis=-1, keepdims=True)
probabilities = np.exp(shifted)
probabilities /= probabilities.sum(axis=-1, keepdims=True)
bucket_indices = np.arange(bins, dtype=np.float32)
distances = (probabilities * bucket_indices).sum(axis=-1)
else:
# DFL disabled mirrors PyTorch's single-bin regression width contract.
distances = np.log1p(np.exp(distances[..., 0]))

distances_px = distances * geometry.strides.reshape(1, -1, 1)
centers = geometry.centers_xy.reshape(1, -1, 2)
top_left = centers - distances_px[..., (0, 1)]
bottom_right = centers + distances_px[..., (2, 3)]
boxes = np.concatenate((top_left, bottom_right), axis=-1)
boxes = np.clip(boxes, 0.0, float(metadata.image_size)).astype(np.float32)

if logits.ndim == 2:
return boxes[0]
return boxes
47 changes: 47 additions & 0 deletions complexity/deploy/onnx_detector/grid.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,47 @@
"""Anchor-center generation for TR-Hash Vision v8 ONNX detector exports."""

from __future__ import annotations

from dataclasses import dataclass
from typing import Sequence

import numpy as np


@dataclass(frozen=True)
class GridGeometry:
"""Concatenated row-major feature-grid geometry."""

centers_xy: np.ndarray
strides: np.ndarray
grid_sizes: tuple[int, ...]


def generate_grid_geometry(image_size: int, grid_sizes: Sequence[int]) -> GridGeometry:
"""Generate concatenated anchor centers and strides for each grid level."""

if image_size <= 0:
raise ValueError("image_size must be positive")

resolved_grid_sizes = tuple(int(grid) for grid in grid_sizes)
if not resolved_grid_sizes:
raise ValueError("grid_sizes must not be empty")
if any(grid <= 0 for grid in resolved_grid_sizes):
raise ValueError("grid_sizes must contain only positive values")

centers: list[np.ndarray] = []
strides: list[np.ndarray] = []
for grid in resolved_grid_sizes:
stride = float(image_size) / float(grid)
coords = (np.arange(grid, dtype=np.float32) + 0.5) * np.float32(stride)
x_centers, y_centers = np.meshgrid(coords, coords, indexing="xy")
centers.append(
np.stack((x_centers.reshape(-1), y_centers.reshape(-1)), axis=1)
)
strides.append(np.full((grid * grid,), stride, dtype=np.float32))

return GridGeometry(
centers_xy=np.concatenate(centers, axis=0).astype(np.float32, copy=False),
strides=np.concatenate(strides, axis=0).astype(np.float32, copy=False),
grid_sizes=resolved_grid_sizes,
)
185 changes: 185 additions & 0 deletions complexity/deploy/onnx_detector/metadata.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,185 @@
"""Metadata sidecar schema for TR-Hash Vision v8 ONNX detector exports."""

from __future__ import annotations

import json
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path
from typing import Literal, Mapping, cast

BranchType = Literal["nms-free", "o2m"]

DEFAULT_CONFIDENCE_THRESHOLD = 0.25
DEFAULT_IOU_THRESHOLD = 0.45
DEFAULT_MAX_DETECTIONS = 300


@dataclass(frozen=True)
class OnnxDetectorMetadata:
"""Validated deployment metadata loaded from an ONNX sidecar JSON file."""

architecture_version: int
image_size: int
num_classes: int
num_cells: int
regression_width: int
reg_max: int
scale_factors: tuple[int, ...]
grid_sizes: tuple[int, ...]
p2_head: bool
branch: BranchType
requires_nms: bool
output_semantics: str
confidence_threshold: float = DEFAULT_CONFIDENCE_THRESHOLD
iou_threshold: float = DEFAULT_IOU_THRESHOLD
max_detections: int = DEFAULT_MAX_DETECTIONS
class_names: tuple[str, ...] | None = None

@property
def dfl_bins(self) -> int:
return self.reg_max + 1 if self.reg_max else 1

@property
def prediction_width(self) -> int:
return self.regression_width + self.num_classes

@property
def strides(self) -> tuple[float, ...]:
return tuple(float(self.image_size) / float(grid) for grid in self.grid_sizes)

def as_dict(self) -> dict[str, object]:
return {
"architecture_version": self.architecture_version,
"image_size": self.image_size,
"num_classes": self.num_classes,
"num_cells": self.num_cells,
"regression_width": self.regression_width,
"reg_max": self.reg_max,
"scale_factors": list(self.scale_factors),
"grid_sizes": list(self.grid_sizes),
"p2_head": self.p2_head,
"branch": self.branch,
"requires_nms": self.requires_nms,
"output_semantics": self.output_semantics,
"confidence_threshold": self.confidence_threshold,
"iou_threshold": self.iou_threshold,
"max_detections": self.max_detections,
"class_names": list(self.class_names) if self.class_names is not None else None,
}


def load_metadata(path: str | Path) -> OnnxDetectorMetadata:
"""Load and validate an ONNX detector metadata sidecar."""

metadata_path = Path(path)
return metadata_from_mapping(json.loads(metadata_path.read_text()))


def metadata_from_mapping(values: Mapping[str, object]) -> OnnxDetectorMetadata:
"""Create validated metadata from a mapping."""

branch = _branch(values.get("branch"))
requires_nms = bool(values.get("requires_nms", branch == "o2m"))
if requires_nms is not (branch == "o2m"):
raise ValueError("requires_nms must be true only for the o2m branch")

class_names_value = values.get("class_names")
class_names = None
if class_names_value is not None:
if not isinstance(class_names_value, Sequence) or isinstance(
class_names_value, (str, bytes)
):
raise ValueError("class_names must be a sequence of strings")
class_names = tuple(str(name) for name in class_names_value)

metadata = OnnxDetectorMetadata(
architecture_version=_positive_int(values, "architecture_version"),
image_size=_positive_int(values, "image_size"),
num_classes=_positive_int(values, "num_classes"),
num_cells=_positive_int(values, "num_cells"),
regression_width=_positive_int(values, "regression_width"),
reg_max=_non_negative_int(values, "reg_max"),
scale_factors=_positive_int_tuple(values, "scale_factors"),
grid_sizes=_positive_int_tuple(values, "grid_sizes"),
p2_head=bool(values.get("p2_head", False)),
branch=branch,
requires_nms=requires_nms,
output_semantics=str(values.get("output_semantics", "")),
confidence_threshold=float(
values.get("confidence_threshold", DEFAULT_CONFIDENCE_THRESHOLD)
),
iou_threshold=float(values.get("iou_threshold", DEFAULT_IOU_THRESHOLD)),
max_detections=int(values.get("max_detections", DEFAULT_MAX_DETECTIONS)),
class_names=class_names,
)
_validate_metadata(metadata)
return metadata


def validate_output_shape(
metadata: OnnxDetectorMetadata,
output_shape: Sequence[int | str | None],
) -> None:
"""Validate an ONNX output shape against the sidecar metadata."""

if len(output_shape) != 3:
raise ValueError(f"ONNX output must be rank 3, got shape {tuple(output_shape)}")
_validate_axis(output_shape[1], metadata.num_cells, "num_cells")
_validate_axis(output_shape[2], metadata.prediction_width, "prediction_width")


def _validate_metadata(metadata: OnnxDetectorMetadata) -> None:
if metadata.architecture_version != 8:
raise ValueError("only TR-Hash detector architecture v8 metadata is supported")
if metadata.num_cells != sum(grid * grid for grid in metadata.grid_sizes):
raise ValueError("num_cells must equal sum(grid ** 2 for grid_sizes)")
if metadata.regression_width != 4 * metadata.dfl_bins:
raise ValueError("regression_width must equal 4 * DFL bin count")
if metadata.max_detections <= 0:
raise ValueError("max_detections must be positive")
if not 0.0 <= metadata.confidence_threshold <= 1.0:
raise ValueError("confidence_threshold must be in [0, 1]")
if not 0.0 <= metadata.iou_threshold <= 1.0:
raise ValueError("iou_threshold must be in [0, 1]")
if (
metadata.class_names is not None
and len(metadata.class_names) != metadata.num_classes
):
raise ValueError("class_names length must match num_classes")


def _branch(value: object) -> BranchType:
normalized = "nms-free" if value == "nms_free" else value
if normalized not in {"nms-free", "o2m"}:
raise ValueError(f"unsupported ONNX detector branch: {value!r}")
return cast(BranchType, normalized)


def _positive_int(values: Mapping[str, object], key: str) -> int:
value = int(values[key])
if value <= 0:
raise ValueError(f"{key} must be positive")
return value


def _non_negative_int(values: Mapping[str, object], key: str) -> int:
value = int(values[key])
if value < 0:
raise ValueError(f"{key} must be non-negative")
return value


def _positive_int_tuple(values: Mapping[str, object], key: str) -> tuple[int, ...]:
raw = values[key]
if not isinstance(raw, Sequence) or isinstance(raw, (str, bytes)):
raise ValueError(f"{key} must be a sequence of positive integers")
resolved = tuple(int(item) for item in raw)
if not resolved or any(item <= 0 for item in resolved):
raise ValueError(f"{key} must contain positive integers")
return resolved


def _validate_axis(axis: int | str | None, expected: int, name: str) -> None:
if isinstance(axis, int) and axis != expected:
raise ValueError(f"ONNX output {name} mismatch: expected {expected}, got {axis}")
Loading
Loading