Correct but useless: Moved shared kalman to ByteTrack

* It's not stateful! All it keeps is the initial params - motion and update matrices
This commit is contained in:
LinasKo 2024-10-18 02:23:08 +03:00
parent 89be6f92d2
commit b380bd2df7
1 changed files with 8 additions and 5 deletions

View File

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