You cannot select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
38 lines
1.0 KiB
Python
38 lines
1.0 KiB
Python
import sys
|
|
import unittest
|
|
from pathlib import Path
|
|
|
|
import numpy as np
|
|
|
|
PROJECT_ROOT = Path(__file__).resolve().parents[1]
|
|
sys.path.insert(0, str(PROJECT_ROOT))
|
|
|
|
from bytetrack_min_aggressive import BYTETracker, hungarian
|
|
|
|
|
|
class ByteTrackAssignmentTests(unittest.TestCase):
|
|
def test_hungarian_matches_single_pair(self):
|
|
self.assertEqual(hungarian(np.array([[0.0]], dtype=np.float32)), [(0, 0)])
|
|
|
|
def test_bytetrack_keeps_id_for_same_box(self):
|
|
tracker = BYTETracker(
|
|
track_high_thresh=0.02,
|
|
track_low_thresh=0.01,
|
|
new_track_thresh=0.02,
|
|
match_thresh=0.10,
|
|
min_hits=1,
|
|
)
|
|
det = np.array([[10.0, 10.0, 40.0, 40.0, 0.05]], dtype=np.float32)
|
|
|
|
first = tracker.update(det, dt=0.04)
|
|
second = tracker.update(det, dt=0.04)
|
|
|
|
self.assertEqual(len(first), 1)
|
|
self.assertEqual(len(second), 1)
|
|
self.assertEqual(second[0].track_id, first[0].track_id)
|
|
self.assertEqual(second[0].hits, 2)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|