MCPcopy Create free account
hub / github.com/FoundationVision/ByteTrack / STrack

Class STrack

tutorials/fairmot/tracker.py:23–175  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

21from models.utils import _tranpose_and_gather_feat
22
23class STrack(BaseTrack):
24 shared_kalman = KalmanFilter()
25 def __init__(self, tlwh, score, temp_feat, buffer_size=30):
26
27 # wait activate
28 self._tlwh = np.asarray(tlwh, dtype=np.float)
29 self.kalman_filter = None
30 self.mean, self.covariance = None, None
31 self.is_activated = False
32
33 self.score = score
34 self.score_list = []
35 self.tracklet_len = 0
36
37 self.smooth_feat = None
38 self.update_features(temp_feat)
39 self.features = deque([], maxlen=buffer_size)
40 self.alpha = 0.9
41
42 def update_features(self, feat):
43 feat /= np.linalg.norm(feat)
44 self.curr_feat = feat
45 if self.smooth_feat is None:
46 self.smooth_feat = feat
47 else:
48 self.smooth_feat = self.alpha * self.smooth_feat + (1 - self.alpha) * feat
49 self.features.append(feat)
50 self.smooth_feat /= np.linalg.norm(self.smooth_feat)
51
52 def predict(self):
53 mean_state = self.mean.copy()
54 if self.state != TrackState.Tracked:
55 mean_state[7] = 0
56 self.mean, self.covariance = self.kalman_filter.predict(mean_state, self.covariance)
57
58 @staticmethod
59 def multi_predict(stracks):
60 if len(stracks) > 0:
61 multi_mean = np.asarray([st.mean.copy() for st in stracks])
62 multi_covariance = np.asarray([st.covariance for st in stracks])
63 for i, st in enumerate(stracks):
64 if st.state != TrackState.Tracked:
65 multi_mean[i][7] = 0
66 multi_mean, multi_covariance = STrack.shared_kalman.multi_predict(multi_mean, multi_covariance)
67 for i, (mean, cov) in enumerate(zip(multi_mean, multi_covariance)):
68 stracks[i].mean = mean
69 stracks[i].covariance = cov
70
71 def activate(self, kalman_filter, frame_id):
72 """Start a new tracklet"""
73 self.kalman_filter = kalman_filter
74 self.track_id = self.next_id()
75 self.mean, self.covariance = self.kalman_filter.initiate(self.tlwh_to_xyah(self._tlwh))
76
77 self.tracklet_len = 0
78 self.state = TrackState.Tracked
79 if frame_id == 1:
80 self.is_activated = True

Callers 1

updateMethod · 0.70

Calls 1

KalmanFilterClass · 0.90

Tested by

no test coverage detected