TRACK_MATCH_THRESH

This commit is contained in:
proitlab committed 2026-07-15 10:46:06 +07:00
1 parent 8f7ecb4321
commit 657f4f80d9
2 files changed
+26

No files matched your search

+25
View File
@@ -441,6 +441,7 @@ class ByteTracker:
cost_mat = 1.0 - iou_mat
matches = _greedy_match(cost_mat, threshold=1.0 - self.match_thresh)
matched_det_idx = set()
for di, ti in matches:
det_global = int(high_idx[di])
orig_idx = int(remain_orig_idx[det_global])
@@ -450,6 +451,30 @@ class ByteTracker:
det_to_track[orig_idx] = track_pool[ti].track_id
tracked_map[track_pool[ti].track_id] = (track_pool[ti].get_cx(), track_pool[ti].get_cy())
match_pairs_high.append((det_global, ti))
matched_det_idx.add(di)
# --- second pass: any unmatched high-score detection tries to
# pair with ANY track that overlaps above threshold, even one
# that already has a match. This prevents a close/overlapping
# object from spawning a spurious new ID when the greedy 1:1
# match paired its detection elsewhere. ---
for di in range(len(high_dets)):
if di in matched_det_idx:
continue
best_ti = None
best_iou = 0.0
for ti in range(num_tracks):
if iou_mat[di, ti] > best_iou and iou_mat[di, ti] >= self.match_thresh:
best_iou = iou_mat[di, ti]
best_ti = ti
if best_ti is not None:
det_global = int(high_idx[di])
orig_idx = int(remain_orig_idx[det_global])
track_pool[best_ti].update(dets[det_global])
track_pool[best_ti].hit_streak = max(1, track_pool[best_ti].hit_streak)
matched_track_idx.add(best_ti)
det_to_track[orig_idx] = track_pool[best_ti].track_id
tracked_map[track_pool[best_ti].track_id] = (track_pool[best_ti].get_cx(), track_pool[best_ti].get_cy())
unmatched_tracks = [
t for t in range(num_tracks) if t not in matched_track_idx