You can not select more than 25 topics
Topics must start with a letter or number, can include dashes ('-') and can be up to 35 characters long.
58 lines
2.0 KiB
58 lines
2.0 KiB
import sys
|
|
from argparse import Namespace
|
|
from types import SimpleNamespace
|
|
|
|
_inserted_numpy_stub = "numpy" not in sys.modules
|
|
if _inserted_numpy_stub:
|
|
sys.modules["numpy"] = SimpleNamespace(ndarray=object)
|
|
|
|
from src.config_loader import TrackerSettings
|
|
from src.tracker import PersonTracker
|
|
|
|
if _inserted_numpy_stub:
|
|
sys.modules.pop("numpy", None)
|
|
|
|
|
|
def settings() -> TrackerSettings:
|
|
return TrackerSettings(False, "bytetrack", .5, .1, .6, .8, 30, 2, 30, True)
|
|
|
|
|
|
def test_external_rows_are_normalized_and_ids_are_stable() -> None:
|
|
tracker = PersonTracker(settings())
|
|
rows = [[-2, 1, 20, 30, 7, .9, 0, 0]]
|
|
first = tracker._normalize(rows, 100, 100)
|
|
second = tracker._normalize(rows, 100, 100)
|
|
assert first[0].track_id == second[0].track_id
|
|
assert first[0].bbox.x1 == 0
|
|
assert not first[0].is_confirmed
|
|
assert second[0].is_confirmed
|
|
|
|
|
|
def test_empty_non_person_and_invalid_boxes_are_filtered() -> None:
|
|
tracker = PersonTracker(settings())
|
|
assert tracker._normalize([], 100, 100) == []
|
|
rows = [[1, 1, 1, 20, 1, .9, 0, 0], [1, 1, 20, 20, 2, .9, 3, 0]]
|
|
assert tracker._normalize(rows, 100, 100) == []
|
|
|
|
|
|
def test_reset_does_not_reuse_public_ids() -> None:
|
|
tracker = PersonTracker(settings())
|
|
first = tracker._normalize([[1, 1, 20, 20, 1, .9, 0, 0]], 100, 100)[0].track_id
|
|
tracker.reset()
|
|
second = tracker._normalize([[1, 1, 20, 20, 1, .9, 0, 0]], 100, 100)[0].track_id
|
|
assert second > first
|
|
|
|
|
|
def test_backend_constructor_supports_new_and_old_ultralytics_apis() -> None:
|
|
class NewBackend:
|
|
def __init__(self, args: Namespace) -> None:
|
|
self.args = args
|
|
|
|
class OldBackend:
|
|
def __init__(self, args: Namespace, frame_rate: int) -> None:
|
|
self.args, self.frame_rate = args, frame_rate
|
|
|
|
args = Namespace(track_buffer=30)
|
|
assert PersonTracker._create_backend(NewBackend, args, 25).args is args
|
|
old = PersonTracker._create_backend(OldBackend, args, 25)
|
|
assert old.args is args and old.frame_rate == 25
|
|
|