From b380bd2df79a167ffd2b768896eda49eb048c7b5 Mon Sep 17 00:00:00 2001 From: LinasKo Date: Fri, 18 Oct 2024 02:23:08 +0300 Subject: [PATCH] Correct but useless: Moved shared kalman to ByteTrack * It's not stateful! All it keeps is the initial params - motion and update matrices --- supervision/tracker/byte_tracker/core.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/supervision/tracker/byte_tracker/core.py b/supervision/tracker/byte_tracker/core.py index a1ae63f4..019ce890 100644 --- a/supervision/tracker/byte_tracker/core.py +++ b/supervision/tracker/byte_tracker/core.py @@ -37,13 +37,12 @@ class IdCounter: class STrack: - shared_kalman = KalmanFilter() - def __init__( self, tlwh, score, minimum_consecutive_frames, + shared_kalman: KalmanFilter, internal_id_counter: IdCounter, external_id_counter: IdCounter, ): @@ -54,6 +53,7 @@ class STrack: self._tlwh = np.asarray(tlwh, dtype=np.float32) self.kalman_filter = None + self.shared_kalman = shared_kalman self.mean, self.covariance = None, None self.is_activated = False @@ -76,7 +76,7 @@ class STrack: ) @staticmethod - def multi_predict(stracks): + def multi_predict(stracks, shared_kalman: KalmanFilter): if len(stracks) > 0: multi_mean = [] multi_covariance = [] @@ -86,7 +86,7 @@ class STrack: if st.state != TrackState.Tracked: multi_mean[i][7] = 0 - multi_mean, multi_covariance = STrack.shared_kalman.multi_predict( + multi_mean, multi_covariance = shared_kalman.multi_predict( np.asarray(multi_mean), np.asarray(multi_covariance) ) for i, (mean, cov) in enumerate(zip(multi_mean, multi_covariance)): @@ -241,6 +241,7 @@ class ByteTrack: self.max_time_lost = int(frame_rate / 30.0 * lost_track_buffer) self.minimum_consecutive_frames = minimum_consecutive_frames self.kalman_filter = KalmanFilter() + self.shared_kalman = KalmanFilter() self.tracked_tracks: List[STrack] = [] self.lost_tracks: List[STrack] = [] @@ -374,6 +375,7 @@ class ByteTrack: STrack.tlbr_to_tlwh(tlbr), s, self.minimum_consecutive_frames, + self.shared_kalman, self.internal_id_counter, self.external_id_counter, ) @@ -395,7 +397,7 @@ class ByteTrack: """ Step 2: First association, with high score detection boxes""" strack_pool = joint_tracks(tracked_stracks, self.lost_tracks) # Predict the current location with KF - STrack.multi_predict(strack_pool) + STrack.multi_predict(strack_pool, self.shared_kalman) dists = matching.iou_distance(strack_pool, detections) dists = matching.fuse_score(dists, detections) @@ -422,6 +424,7 @@ class ByteTrack: STrack.tlbr_to_tlwh(tlbr), s, self.minimum_consecutive_frames, + self.shared_kalman, self.internal_id_counter, self.external_id_counter, )