Let upper layer to handle track_id generation
This commit is contained in:
parent
a05882add1
commit
db2aa721f8
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,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
|
||||
Loading…
Reference in New Issue