diff --git a/supervision/tracker/byte_tracker/basetrack.py b/supervision/tracker/byte_tracker/basetrack.py index d280274c..b78bc596 100644 --- a/supervision/tracker/byte_tracker/basetrack.py +++ b/supervision/tracker/byte_tracker/basetrack.py @@ -13,7 +13,6 @@ class TrackState(Enum): class BaseTrack: def __init__(self): - self._count = 0 self.track_id = 0 self.is_activated = False self.state = TrackState.New @@ -33,12 +32,7 @@ class BaseTrack: def end_frame(self) -> int: return self.frame_id - def next_id(self) -> int: - self._count += 1 - return self._count - def reset_counter(self): - self._count = 0 self.track_id = 0 self.start_frame = 0 self.frame_id = 0 diff --git a/supervision/tracker/byte_tracker/core.py b/supervision/tracker/byte_tracker/core.py index 89e1e2f2..cd8ea3e5 100644 --- a/supervision/tracker/byte_tracker/core.py +++ b/supervision/tracker/byte_tracker/core.py @@ -1,4 +1,4 @@ -from typing import List, Tuple +from typing import List, Optional, Tuple import numpy as np @@ -11,10 +11,11 @@ from supervision.tracker.byte_tracker.kalman_filter import KalmanFilter class STrack(BaseTrack): shared_kalman = KalmanFilter() - _external_count = 0 def __init__(self, tlwh, score, class_ids, minimum_consecutive_frames): + super().__init__() # wait activate + self._external_count = 0 self._tlwh = np.asarray(tlwh, dtype=np.float32) self.kalman_filter = None self.mean, self.covariance = None, None @@ -54,10 +55,10 @@ class STrack(BaseTrack): stracks[i].mean = mean stracks[i].covariance = cov - def activate(self, kalman_filter, frame_id): + def activate(self, kalman_filter, frame_id, track_id): """Start a new tracklet""" self.kalman_filter = kalman_filter - self.internal_track_id = self.next_id() + self.internal_track_id = track_id self.mean, self.covariance = self.kalman_filter.initiate( self.tlwh_to_xyah(self._tlwh) ) @@ -68,12 +69,12 @@ class STrack(BaseTrack): self.is_activated = True if self.minimum_consecutive_frames == 1: - self.external_track_id = self.next_external_id() + self.external_track_id = track_id self.frame_id = frame_id self.start_frame = frame_id - def re_activate(self, new_track, frame_id, new_id=False): + def re_activate(self, new_track, frame_id, new_id: Optional[int] = None): self.mean, self.covariance = self.kalman_filter.update( self.mean, self.covariance, self.tlwh_to_xyah(new_track.tlwh) ) @@ -82,10 +83,10 @@ class STrack(BaseTrack): self.frame_id = frame_id if new_id: - self.internal_track_id = self.next_id() + self.internal_track_id = new_id self.score = new_track.score - def update(self, new_track, frame_id): + def update(self, new_track, frame_id, track_id): """ Update a matched track :type new_track: STrack @@ -104,7 +105,7 @@ class STrack(BaseTrack): if self.tracklet_len == self.minimum_consecutive_frames: self.is_activated = True if self.external_track_id == -1: - self.external_track_id = self.next_external_id() + self.external_track_id = track_id self.score = new_track.score @@ -142,15 +143,6 @@ class STrack(BaseTrack): def to_xyah(self): return self.tlwh_to_xyah(self.tlwh) - @staticmethod - def next_external_id(): - STrack._external_count += 1 - return STrack._external_count - - @staticmethod - def reset_external_counter(): - STrack._external_count = 0 - @staticmethod def tlbr_to_tlwh(tlbr): ret = np.asarray(tlbr).copy() @@ -225,6 +217,7 @@ class ByteTrack: self.track_activation_threshold = track_activation_threshold self.minimum_matching_threshold = minimum_matching_threshold + self._count = 0 self.frame_id = 0 self.det_thresh = self.track_activation_threshold + 0.1 self.max_time_lost = int(frame_rate / 30.0 * lost_track_buffer) @@ -235,6 +228,10 @@ class ByteTrack: self.lost_tracks: List[STrack] = [] self.removed_tracks: List[STrack] = [] + def _next_id(self) -> int: + self._count += 1 + return self._count + def update_with_detections(self, detections: Detections) -> Detections: """ Updates the tracker with the provided detections and returns the updated @@ -314,8 +311,6 @@ class ByteTrack: self.tracked_tracks: List[STrack] = [] self.lost_tracks: List[STrack] = [] self.removed_tracks: List[STrack] = [] - BaseTrack.reset_counter() - STrack.reset_external_counter() def update_with_tensors(self, tensors: np.ndarray) -> List[STrack]: """ @@ -384,10 +379,10 @@ class ByteTrack: track = strack_pool[itracked] det = detections[idet] if track.state == TrackState.Tracked: - track.update(detections[idet], self.frame_id) + track.update(detections[idet], self.frame_id, self._next_id()) activated_starcks.append(track) else: - track.re_activate(det, self.frame_id, new_id=False) + track.re_activate(det, self.frame_id) refind_stracks.append(track) """ Step 3: Second association, with low score detection boxes""" @@ -413,10 +408,10 @@ class ByteTrack: track = r_tracked_stracks[itracked] det = detections_second[idet] if track.state == TrackState.Tracked: - track.update(det, self.frame_id) + track.update(det, self.frame_id, self._next_id()) activated_starcks.append(track) else: - track.re_activate(det, self.frame_id, new_id=False) + track.re_activate(det, self.frame_id) refind_stracks.append(track) for it in u_track: @@ -434,7 +429,7 @@ class ByteTrack: dists, thresh=0.7 ) for itracked, idet in matches: - unconfirmed[itracked].update(detections[idet], self.frame_id) + unconfirmed[itracked].update(detections[idet], self.frame_id, self._next_id()) activated_starcks.append(unconfirmed[itracked]) for it in u_unconfirmed: track = unconfirmed[it] @@ -446,7 +441,7 @@ class ByteTrack: track = detections[inew] if track.score < self.det_thresh: continue - track.activate(self.kalman_filter, self.frame_id) + track.activate(self.kalman_filter, self.frame_id, self._next_id()) activated_starcks.append(track) """ Step 5: Update state""" for track in self.lost_tracks: diff --git a/test/tracker/__init__.py b/test/tracker/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/test/tracker/test_byte_tracker.py b/test/tracker/test_byte_tracker.py new file mode 100644 index 00000000..673cd90b --- /dev/null +++ b/test/tracker/test_byte_tracker.py @@ -0,0 +1,37 @@ +import numpy as np +import pytest +import supervision as sv + + +@pytest.mark.parametrize( + "detections, expected_results", + [ + ( + [ + sv.Detections( + xyxy=np.array([[10, 10, 20, 20], [30, 30, 40, 40]]), + class_id=np.array([1, 1]), + confidence=np.array([1, 1]), + ), + sv.Detections( + xyxy=np.array([[10, 10, 20, 20], [30, 30, 40, 40]]), + class_id=np.array([1, 1]), + confidence=np.array([1, 1]), + ), + ], + sv.Detections( + xyxy=np.array([[10, 10, 20, 20], [30, 30, 40, 40]]), + class_id=np.array([1, 1]), + confidence=np.array([1, 1]), + tracker_id=np.array([1, 2]), + ) + ), + ], +) +def test_byte_tracker( + detections: list[sv.Detections], + expected_results: sv.Detections, +) -> None: + byte_tracker = sv.ByteTrack() + tracked_detections = [byte_tracker.update_with_detections(d) for d in detections] + assert tracked_detections[-1] == expected_results