156 lines
5.2 KiB
Python
156 lines
5.2 KiB
Python
"""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",
|
|
]
|