Files
feedmill-recounter/src/tracking.py
T

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)