- detect_interval=2 on both pipelines; stabilizer 10-frame hold bridges skipped frames (GPU inference cut ~half during active batch) - sup_det_interval 5 -> 10 (still <= stabilizer hold / synthetic max_age) - preview_every_n default 2 -> 5 - AnnotatedVideoWriter: codec auto-chain gstreamer_nvenc -> avc1 -> mp4v, max_height downscale (even dims, INTER_AREA); backend logged - output_max_height plumbed Job -> web checkbox (720p) -> CLI --output-height - README: 720p option, CLI flag - 153 tests pass (+13)
348 lines
12 KiB
Python
348 lines
12 KiB
Python
"""Tests for pipeline runner (src/pipeline.py)."""
|
|
|
|
import os
|
|
|
|
import cv2
|
|
import numpy as np
|
|
import pytest
|
|
from src.counting import MultiClassLineCounter
|
|
from src.pipeline import apply_class_filter, run_pipeline, PipelineResult
|
|
from src.interfaces import Detection
|
|
from src.model_registry import ModelConfig
|
|
|
|
|
|
def test_pipeline_result_dataclass():
|
|
"""PipelineResult has correct fields."""
|
|
r = PipelineResult(
|
|
output_path="/tmp/out.mp4",
|
|
frame_count=100,
|
|
loading_count=5,
|
|
unloading_count=2,
|
|
batch_count=1,
|
|
duration_seconds=10.0,
|
|
model_name="v4-best.pt",
|
|
class_filter=None,
|
|
)
|
|
assert r.loading_count == 5
|
|
assert r.unloading_count == 2
|
|
assert r.net_count == 3
|
|
|
|
|
|
def test_run_pipeline_processes_video(tmp_path):
|
|
"""run_pipeline processes a 3-frame video and writes output."""
|
|
# Create a test video
|
|
video_path = str(tmp_path / "test.mp4")
|
|
writer = cv2.VideoWriter(video_path, cv2.VideoWriter_fourcc(*"mp4v"), 25.0, (320, 240))
|
|
for _ in range(3):
|
|
writer.write(np.zeros((240, 320, 3), dtype=np.uint8))
|
|
writer.release()
|
|
|
|
# Strengthened (controller ruling): verify the fixture video is valid.
|
|
assert os.path.exists(video_path)
|
|
cap = cv2.VideoCapture(video_path)
|
|
assert cap.isOpened()
|
|
assert int(cap.get(cv2.CAP_PROP_FRAME_COUNT)) == 3
|
|
cap.release()
|
|
|
|
# Create a minimal .pt file placeholder (YOLO will fail to load, but we test the pipeline structure)
|
|
# For unit testing without real models, we test PipelineResult directly
|
|
pass # See integration test below for end-to-end with real models
|
|
|
|
|
|
def test_apply_class_filter_keeps_only_selected():
|
|
"""apply_class_filter keeps only classes in the filter; None/[] keep all."""
|
|
dets = [
|
|
Detection(bbox=(0, 0, 1, 1), confidence=0.9, class_id=0, class_name="sack"),
|
|
Detection(bbox=(0, 0, 1, 1), confidence=0.9, class_id=1, class_name="box"),
|
|
Detection(bbox=(0, 0, 1, 1), confidence=0.9, class_id=0, class_name="sack"),
|
|
]
|
|
filtered = apply_class_filter(dets, ["sack"])
|
|
assert [d.class_name for d in filtered] == ["sack", "sack"]
|
|
assert apply_class_filter(dets, None) == dets
|
|
assert apply_class_filter(dets, []) == dets
|
|
|
|
|
|
def test_run_pipeline_no_model_raises(tmp_path):
|
|
"""run_pipeline raises RuntimeError if video can't be opened."""
|
|
with pytest.raises(RuntimeError, match="Cannot open video"):
|
|
run_pipeline(
|
|
video_path=str(tmp_path / "nonexistent.mp4"),
|
|
model_config=ModelConfig(filename="test.pt", path="/nonexistent.pt", stem="test", known_classes=["sack"]),
|
|
output_path=str(tmp_path / "out.mp4"),
|
|
)
|
|
|
|
|
|
def test_run_pipeline_accepts_frame_callback(tmp_path):
|
|
"""run_pipeline signature accepts frame_callback parameter."""
|
|
import inspect
|
|
sig = inspect.signature(run_pipeline)
|
|
assert "frame_callback" in sig.parameters
|
|
param = sig.parameters["frame_callback"]
|
|
assert param.default is None
|
|
|
|
|
|
def test_run_pipeline_accepts_cancel_check(tmp_path):
|
|
"""run_pipeline signature accepts cancel_check parameter."""
|
|
import inspect
|
|
sig = inspect.signature(run_pipeline)
|
|
assert "cancel_check" in sig.parameters
|
|
param = sig.parameters["cancel_check"]
|
|
assert param.default is None
|
|
|
|
|
|
def test_cancel_check_breaks_pipeline(tmp_path):
|
|
"""run_pipeline stops early when cancel_check returns True."""
|
|
# Create a test video with many frames
|
|
video_path = str(tmp_path / "test_cancel.mp4")
|
|
writer = cv2.VideoWriter(video_path, cv2.VideoWriter_fourcc(*"mp4v"), 25.0, (320, 240))
|
|
for _ in range(120):
|
|
writer.write(np.zeros((240, 320, 3), dtype=np.uint8))
|
|
writer.release()
|
|
|
|
cancel_called = [0]
|
|
|
|
def fake_cancel_check():
|
|
cancel_called[0] += 1
|
|
return True # cancel immediately
|
|
|
|
# Without a real model this will fail at YOLO load, but we verify cancel_check is called
|
|
# by checking the function accepts it and the loop would break
|
|
import inspect
|
|
sig = inspect.signature(run_pipeline)
|
|
assert "cancel_check" in sig.parameters
|
|
|
|
|
|
# ── Truck class must never be counted ────────────────────────────────────
|
|
|
|
|
|
def _cross_det(class_name: str, track_id: int, y1: float, x: float) -> Detection:
|
|
"""Detection whose top edge (y1) sits above/below a line at y=100."""
|
|
return Detection(
|
|
bbox=(x, y1, x + 40, y1 + 40),
|
|
confidence=0.9,
|
|
class_id=0,
|
|
class_name=class_name,
|
|
track_id=track_id,
|
|
)
|
|
|
|
|
|
def test_countable_detections_excludes_truck():
|
|
"""Pipeline's count filter keeps sack/box, drops truck."""
|
|
from src.pipeline import countable_detections
|
|
|
|
dets = [
|
|
_cross_det("sack", 1, 50.0, 100.0),
|
|
_cross_det("truck", 2, 50.0, 200.0),
|
|
_cross_det("box", 3, 50.0, 150.0),
|
|
]
|
|
assert [d.class_name for d in countable_detections(dets)] == ["sack", "box"]
|
|
|
|
|
|
def test_truck_crossing_not_counted_after_filter():
|
|
"""Sack + truck both cross the line; only the sack increments the counter."""
|
|
from src.pipeline import countable_detections
|
|
|
|
counter = MultiClassLineCounter(line_y=100, line_x_start=0, line_x_end=320, margin=20)
|
|
# frame 1: both above the line
|
|
counter.update(countable_detections([
|
|
_cross_det("sack", 1, 50.0, 100.0),
|
|
_cross_det("truck", 2, 50.0, 200.0),
|
|
]))
|
|
# frame 2: both below the line (crossing)
|
|
events = counter.update(countable_detections([
|
|
_cross_det("sack", 1, 150.0, 100.0),
|
|
_cross_det("truck", 2, 150.0, 200.0),
|
|
]))
|
|
|
|
assert counter.loading_count == 1
|
|
assert counter.unloading_count == 0
|
|
assert [ev["track_id"] for ev in events] == [1]
|
|
assert all(ev["track_id"] != 2 for ev in events)
|
|
|
|
|
|
def test_multiclass_line_counter_ignores_truck_class():
|
|
"""MultiClassLineCounter drops truck detections entirely (defense in depth)."""
|
|
counter = MultiClassLineCounter(line_y=100, line_x_start=0, line_x_end=320, margin=20)
|
|
counter.update([_cross_det("truck", 9, 50.0, 100.0)])
|
|
events = counter.update([_cross_det("truck", 9, 150.0, 100.0)])
|
|
|
|
assert events == []
|
|
assert counter.loading_count == 0
|
|
assert counter.unloading_count == 0
|
|
assert counter.box_loading_count == 0
|
|
assert counter.box_unloading_count == 0
|
|
assert counter.net_count == 0
|
|
|
|
|
|
def test_dashboard_label_uses_class_name(monkeypatch):
|
|
"""_draw_detections renders det.class_name, not hardcoded 'sack'."""
|
|
from src import dashboard as dash_mod
|
|
|
|
drawn: list[str] = []
|
|
monkeypatch.setattr(
|
|
dash_mod.cv2, "putText",
|
|
lambda frame, text, *args, **kwargs: drawn.append(text),
|
|
)
|
|
frame = np.zeros((240, 320, 3), dtype=np.uint8)
|
|
dash_mod.DashboardOverlay()._draw_detections(
|
|
frame, [_cross_det("truck", 7, 50.0, 100.0)]
|
|
)
|
|
assert any(t.startswith("truck") for t in drawn), drawn
|
|
|
|
|
|
# ── Box count propagation to results (Task 2) ───────────────────────────
|
|
|
|
|
|
def _box_result_kwargs() -> dict:
|
|
return dict(
|
|
output_path="/tmp/out.mp4",
|
|
frame_count=100,
|
|
loading_count=5,
|
|
unloading_count=2,
|
|
batch_count=1,
|
|
duration_seconds=10.0,
|
|
class_filter=None,
|
|
)
|
|
|
|
|
|
def test_pipeline_result_box_fields_default_zero():
|
|
"""PipelineResult box fields default to 0, box_net_count derives."""
|
|
r = PipelineResult(model_name="v4-best.pt", **_box_result_kwargs())
|
|
assert r.box_loading_count == 0
|
|
assert r.box_unloading_count == 0
|
|
assert r.box_net_count == 0
|
|
|
|
|
|
def test_pipeline_result_box_fields_accept_values():
|
|
"""PipelineResult accepts box counts, box_net_count = in - out."""
|
|
r = PipelineResult(
|
|
model_name="v4-best.pt",
|
|
**_box_result_kwargs(),
|
|
box_loading_count=4,
|
|
box_unloading_count=1,
|
|
)
|
|
assert (r.box_loading_count, r.box_unloading_count, r.box_net_count) == (4, 1, 3)
|
|
|
|
|
|
def test_merged_pipeline_result_box_fields():
|
|
"""MergedPipelineResult has same box fields with defaults."""
|
|
from src.pipeline import MergedPipelineResult
|
|
|
|
kwargs = _box_result_kwargs()
|
|
r = MergedPipelineResult(
|
|
model_names=["a.pt", "b.pt"],
|
|
**kwargs,
|
|
box_loading_count=2,
|
|
box_unloading_count=5,
|
|
)
|
|
assert (r.box_loading_count, r.box_unloading_count, r.box_net_count) == (2, 5, -3)
|
|
r0 = MergedPipelineResult(model_names=["a.pt"], **kwargs)
|
|
assert (r0.box_loading_count, r0.box_unloading_count, r0.box_net_count) == (0, 0, 0)
|
|
|
|
|
|
class _FakeCounter:
|
|
"""Stands in for MultiClassLineCounter with fixed counts."""
|
|
|
|
def __init__(self, line_y=0, line_x_start=0, line_x_end=0, margin=0):
|
|
self.line_y = line_y
|
|
self.line_x_start = line_x_start
|
|
self.line_x_end = line_x_end
|
|
self.margin = margin
|
|
self.loading_count = 7
|
|
self.unloading_count = 3
|
|
self.box_loading_count = 4
|
|
self.box_unloading_count = 1
|
|
|
|
def update(self, detections):
|
|
return []
|
|
|
|
|
|
class _FakeYOLO:
|
|
def __init__(self, path="fake.pt"):
|
|
self.names = {0: "sack", 1: "box", 2: "truck"}
|
|
self.ckpt_path = path
|
|
|
|
|
|
def _patch_pipeline_fakes(monkeypatch):
|
|
from src import detection as det_mod
|
|
from src import pipeline as pipe_mod
|
|
from src import tracking as trk_mod
|
|
|
|
monkeypatch.setattr(pipe_mod, "YOLO", _FakeYOLO)
|
|
monkeypatch.setattr(det_mod, "YOLO", _FakeYOLO)
|
|
monkeypatch.setattr(trk_mod, "YOLO", _FakeYOLO)
|
|
monkeypatch.setattr(pipe_mod, "MultiClassLineCounter", _FakeCounter)
|
|
|
|
|
|
def _write_video(path: str, frames: int = 3) -> None:
|
|
writer = cv2.VideoWriter(path, cv2.VideoWriter_fourcc(*"mp4v"), 25.0, (320, 240))
|
|
for _ in range(frames):
|
|
writer.write(np.zeros((240, 320, 3), dtype=np.uint8))
|
|
writer.release()
|
|
|
|
|
|
def test_run_pipeline_propagates_box_counts(tmp_path, monkeypatch):
|
|
"""run_pipeline passes counter box counts into PipelineResult."""
|
|
_patch_pipeline_fakes(monkeypatch)
|
|
video_path = str(tmp_path / "t.mp4")
|
|
_write_video(video_path)
|
|
|
|
result = run_pipeline(
|
|
video_path=video_path,
|
|
model_config=ModelConfig(
|
|
filename="test.pt",
|
|
path=str(tmp_path / "test.pt"),
|
|
stem="test",
|
|
known_classes=["sack", "box"],
|
|
),
|
|
output_path=str(tmp_path / "out.mp4"),
|
|
preview_enabled=False,
|
|
)
|
|
assert result.loading_count == 7
|
|
assert result.unloading_count == 3
|
|
assert result.box_loading_count == 4
|
|
assert result.box_unloading_count == 1
|
|
assert result.box_net_count == 3
|
|
|
|
|
|
def test_run_merged_pipeline_propagates_box_counts(tmp_path, monkeypatch):
|
|
"""run_merged_pipeline passes counter box counts into MergedPipelineResult."""
|
|
from src.pipeline import run_merged_pipeline
|
|
|
|
_patch_pipeline_fakes(monkeypatch)
|
|
video_path = str(tmp_path / "t.mp4")
|
|
_write_video(video_path)
|
|
cfgs = [
|
|
ModelConfig(filename="a.pt", path=str(tmp_path / "a.pt"), stem="a", known_classes=["sack", "box"]),
|
|
ModelConfig(filename="b.pt", path=str(tmp_path / "b.pt"), stem="b", known_classes=["sack", "box"]),
|
|
]
|
|
|
|
result = run_merged_pipeline(
|
|
video_path=video_path,
|
|
model_configs=cfgs,
|
|
output_path=str(tmp_path / "merged.mp4"),
|
|
preview_enabled=False,
|
|
)
|
|
assert result.loading_count == 7
|
|
assert result.box_loading_count == 4
|
|
assert result.box_unloading_count == 1
|
|
assert result.box_net_count == 3
|
|
|
|
|
|
# ── Perf/output tuning params (signature contracts) ────────────────────
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"fn_name", ["run_pipeline", "run_merged_pipeline"]
|
|
)
|
|
def test_pipeline_perf_params_signature(fn_name):
|
|
"""Both runners expose detect_interval=2, output_max_height=None, preview_every_n=5."""
|
|
import inspect
|
|
|
|
from src import pipeline as pipe_mod
|
|
|
|
sig = inspect.signature(getattr(pipe_mod, fn_name))
|
|
assert sig.parameters["detect_interval"].default == 2
|
|
assert sig.parameters["output_max_height"].default is None
|
|
assert sig.parameters["preview_every_n"].default == 5
|