Let upper layer to handle track_id generation

This commit is contained in:
Grzegorz Klimaszewski 2024-09-19 13:38:16 +02:00
parent a05882add1
commit db2aa721f8
No known key found for this signature in database
4 changed files with 58 additions and 32 deletions

View File

@ -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

View File

@ -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:

0
test/tracker/__init__.py Normal file
View File

View File

@ -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