85 lines
2.9 KiB
Python
85 lines
2.9 KiB
Python
"""FastTrack wrapper — occlusion-aware tracker with custom tuning.
|
|
|
|
Uses Ultralytics FastTrack which handles:
|
|
- Kalman rollback on occlusion onset (restores pre-occlusion velocity)
|
|
- Enlarged search region during occlusion
|
|
- Re-identification of occluded tracks after reappearance
|
|
|
|
Our custom cfg/tracker.yaml tunes:
|
|
- track_buffer=60 (hold lost tracks ~2.4s to survive worker occlusion)
|
|
- new_track_thresh=0.3 (prevent duplicate IDs from spawning)
|
|
- active_occ_to_lost_thresh=15 (tolerate 15 occluded frames)
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import os
|
|
|
|
import numpy as np
|
|
from ultralytics import YOLO
|
|
|
|
from src.interfaces import Detection
|
|
|
|
_TRACKER_CFG = os.path.join(
|
|
os.path.dirname(os.path.dirname(__file__)), "cfg", "tracker.yaml"
|
|
)
|
|
|
|
|
|
class ByteTrackTracker:
|
|
"""Tracks sacks across frames using FastTrack (occlusion-aware)."""
|
|
|
|
def __init__(self, model_path: str | YOLO, conf: float = 0.35) -> None:
|
|
if isinstance(model_path, YOLO):
|
|
self._model = model_path
|
|
self._model_path = getattr(model_path, "ckpt_path", str(model_path))
|
|
else:
|
|
self._model = YOLO(model_path)
|
|
self._model_path = model_path
|
|
self._conf = conf
|
|
self._tracker_cfg = _TRACKER_CFG if os.path.exists(_TRACKER_CFG) else "bytetrack.yaml"
|
|
|
|
def update(
|
|
self, frame: np.ndarray, detections: list[Detection]
|
|
) -> list[Detection]:
|
|
"""Run tracking on the frame, return detections with track IDs."""
|
|
results = self._model.track(
|
|
frame,
|
|
conf=self._conf,
|
|
persist=True,
|
|
tracker=self._tracker_cfg,
|
|
verbose=False,
|
|
)
|
|
return self._parse(results[0])
|
|
|
|
def _parse(self, result) -> list[Detection]:
|
|
tracked: list[Detection] = []
|
|
if result.boxes is None or len(result.boxes) == 0:
|
|
return tracked
|
|
ids = result.boxes.id
|
|
for i, box in enumerate(result.boxes):
|
|
cls_id = int(box.cls[0])
|
|
name = self._model.names.get(cls_id, str(cls_id)) if isinstance(self._model.names, dict) else self._model.names[cls_id]
|
|
if name not in ("sack", "truck", "box"):
|
|
continue
|
|
track_id = int(ids[i]) if ids is not None else None
|
|
x1, y1, x2, y2 = box.xyxy[0].tolist()
|
|
|
|
mask = None
|
|
|
|
tracked.append(
|
|
Detection(
|
|
bbox=(x1, y1, x2, y2),
|
|
confidence=float(box.conf[0]),
|
|
class_id=cls_id,
|
|
class_name=name,
|
|
track_id=track_id,
|
|
mask=mask,
|
|
)
|
|
)
|
|
return tracked
|
|
|
|
def reset(self) -> None:
|
|
"""Reset tracker state (new batch / new truck)."""
|
|
if isinstance(self._model_path, str) and os.path.exists(self._model_path):
|
|
self._model = YOLO(self._model_path)
|