"""Frozen raw-KB4 YOLOX provider for class-agnostic object proposals.""" from __future__ import annotations import time from collections import Counter from collections.abc import Callable from dataclasses import dataclass from threading import Lock from typing import Final import numpy as np from numpy.typing import NDArray from .contracts import BoundingRegion2D, ObjectProposal2D from .providers import SourcePacket from .yolox_object_detector import ( FROZEN_YOLOX_CONFIG, YOLOX_MODEL_ID, YOLOX_MODEL_VERSION, FrozenYoloxConfig, ImageResizer, InferenceBackend, YoloxDetection, postprocess_yolox, preprocess_raw_kb4, ) FROZEN_YOLOX_PROVIDER_ID: Final = "triton-yolox-s-raw-kb4/v1" FROZEN_YOLOX_MODEL_ID: Final = f"{YOLOX_MODEL_ID}:{YOLOX_MODEL_VERSION}" FROZEN_YOLOX_PREPROCESS_ID: Final = "raw-kb4-valid-fov-letterbox/v1" class DetectorProviderError(RuntimeError): """The detector input, frozen inference or proposal output is incompatible.""" @dataclass(frozen=True, slots=True) class DetectorProviderSnapshot: input_frames: int completed_frames: int failed_frames: int zero_proposal_frames: int proposal_count: int rejected: tuple[tuple[str, int], ...] core_duration_ns: int class FrozenYoloxDetectorProvider: """One image payload produces one frozen inference request and proposal tuple.""" provider_id: str = FROZEN_YOLOX_PROVIDER_ID def __init__( self, *, mask: NDArray[np.bool_], backend: InferenceBackend, resizer: ImageResizer | None = None, config: FrozenYoloxConfig = FROZEN_YOLOX_CONFIG, clock_ns: Callable[[], int] = time.perf_counter_ns, ) -> None: if mask.shape != (600, 800) or mask.dtype != np.bool_ or not np.any(mask): raise DetectorProviderError("frozen valid-FOV mask is incompatible") self.mask = np.asarray(mask, dtype=np.bool_) self.backend = backend self.resizer = resizer self.config = config self._clock_ns = clock_ns self._lock = Lock() self._input_frames = 0 self._completed_frames = 0 self._failed_frames = 0 self._zero_proposal_frames = 0 self._proposal_count = 0 self._rejected: Counter[str] = Counter() self._core_duration_ns = 0 def detect(self, packet: SourcePacket) -> tuple[ObjectProposal2D, ...]: payload = packet.image_payload with self._lock: self._input_frames += 1 started_ns = int(self._clock_ns()) try: if not isinstance(payload, np.ndarray): raise DetectorProviderError("detector requires a decoded BGR image payload") image = np.asarray(payload) if image.dtype != np.uint8: raise DetectorProviderError("decoded BGR image must be uint8") tensor = preprocess_raw_kb4( image, self.mask, config=self.config, resizer=self.resizer, ) output = self.backend.infer(tensor) postprocessed = postprocess_yolox(output, self.mask, config=self.config) proposals = proposals_from_detections(packet, postprocessed.detections) except Exception: with self._lock: self._failed_frames += 1 self._core_duration_ns += max(0, int(self._clock_ns()) - started_ns) raise with self._lock: self._completed_frames += 1 self._proposal_count += len(proposals) self._zero_proposal_frames += not proposals self._rejected.update(dict(postprocessed.rejected)) self._core_duration_ns += max(0, int(self._clock_ns()) - started_ns) return proposals def snapshot(self) -> DetectorProviderSnapshot: with self._lock: return DetectorProviderSnapshot( input_frames=self._input_frames, completed_frames=self._completed_frames, failed_frames=self._failed_frames, zero_proposal_frames=self._zero_proposal_frames, proposal_count=self._proposal_count, rejected=tuple(sorted(self._rejected.items())), core_duration_ns=self._core_duration_ns, ) def proposals_from_detections( packet: SourcePacket, detections: tuple[YoloxDetection, ...], ) -> tuple[ObjectProposal2D, ...]: envelope = packet.envelope return tuple( ObjectProposal2D( proposal_id=f"proposal-{envelope.sequence}-{index}", source_id=envelope.source_id, frame_id=envelope.frame_id, region=BoundingRegion2D(*detection.bbox_xyxy), objectness=detection.score, provider_id=FROZEN_YOLOX_PROVIDER_ID, model_id=FROZEN_YOLOX_MODEL_ID, preprocess_id=FROZEN_YOLOX_PREPROCESS_ID, semantic_hint=detection.label, provider_tracklet=None, ) for index, detection in enumerate(detections) ) __all__ = [ "FROZEN_YOLOX_MODEL_ID", "FROZEN_YOLOX_PREPROCESS_ID", "FROZEN_YOLOX_PROVIDER_ID", "DetectorProviderError", "DetectorProviderSnapshot", "FrozenYoloxDetectorProvider", "proposals_from_detections", ]