From ad3e7a7e9e3b57c4d60ed7015ce433e59740a2e4 Mon Sep 17 00:00:00 2001 From: jetson Date: Tue, 22 Sep 2026 13:01:40 +0700 Subject: [PATCH] feat: add fixed_zone param to pipelines, skip truck detection when set --- src/pipeline.py | 68 ++++++++++++++++++------------- tests/test_fixed_zone_pipeline.py | 27 ++++++++++++ 2 files changed, 66 insertions(+), 29 deletions(-) create mode 100644 tests/test_fixed_zone_pipeline.py diff --git a/src/pipeline.py b/src/pipeline.py index ccd3b54..604d043 100644 --- a/src/pipeline.py +++ b/src/pipeline.py @@ -20,7 +20,7 @@ from src.interfaces import Detection from src.model_registry import ModelConfig from src.stabilizer import BboxStabilizer from src.tracking import ByteTrackTracker -from src.truck_roi import TruckROITracker +from src.truck_roi import TruckROI, TruckROITracker from src.video_writer import AnnotatedVideoWriter @@ -78,6 +78,7 @@ def run_pipeline( output_path: str, class_filter: list[str] | None = None, truck_model_config: ModelConfig | None = None, + fixed_zone: TruckROI | None = None, sack_conf: float = 0.4, truck_conf: float = 0.5, truck_det_interval: int = 15, @@ -143,17 +144,18 @@ def run_pipeline( shared_model, conf=sack_conf, class_filter=effective_filter ) - # Truck detector: use explicit truck_model_config, or same model if it has "truck" class + # Truck detector — skip if fixed zone truck_detector = None - if truck_model_config is not None: - # Load a separate truck detection model - truck_shared = YOLO(truck_model_config.path) - truck_detector = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) - elif "truck" in (model_config.known_classes or []): - # Use the same model for truck detection - truck_detector = BaseDetector( - shared_model, conf=truck_conf, class_filter=("truck",) - ) + if fixed_zone is None: + if truck_model_config is not None: + # Load a separate truck detection model + truck_shared = YOLO(truck_model_config.path) + truck_detector = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) + elif "truck" in (model_config.known_classes or []): + # Use the same model for truck detection + truck_detector = BaseDetector( + shared_model, conf=truck_conf, class_filter=("truck",) + ) tracker = ByteTrackTracker(shared_model, conf=sack_conf) stabilizer = BboxStabilizer() @@ -199,10 +201,13 @@ def run_pipeline( break # Truck detection - roi = roi_tracker.roi - if truck_detector is not None and frame_idx % truck_det_interval == 0: - trucks = truck_detector.detect(frame) - roi = roi_tracker.update(trucks) + if fixed_zone is not None: + roi = fixed_zone + else: + roi = roi_tracker.roi + if truck_detector is not None and frame_idx % truck_det_interval == 0: + trucks = truck_detector.detect(frame) + roi = roi_tracker.update(trucks) truck_present = roi is not None and roi.confidence > 0 @@ -362,6 +367,7 @@ def run_merged_pipeline( output_path: str, class_filters: dict[str, list[str] | None] | None = None, truck_model_config: ModelConfig | None = None, + fixed_zone: TruckROI | None = None, sack_conf: float = 0.4, truck_conf: float = 0.5, truck_det_interval: int = 15, @@ -399,17 +405,18 @@ def run_merged_pipeline( eff_filter = (class_filters or {}).get(cfg.stem) or cfg.known_classes or None detectors.append(BaseDetector(model, conf=sack_conf, class_filter=eff_filter)) - # Truck detector + # Truck detector — skip if fixed zone truck_det = None - if truck_model_config is not None: - truck_shared = YOLO(truck_model_config.path) - truck_det = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) - # ponytail: first-match truck detector; could prefer dedicated truck model over sack model w/ truck class - else: - for i, cfg in enumerate(model_configs): - if "truck" in (cfg.known_classes or []): - truck_det = detectors[i] - break + if fixed_zone is None: + if truck_model_config is not None: + truck_shared = YOLO(truck_model_config.path) + truck_det = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) + # ponytail: first-match truck detector; could prefer dedicated truck model over sack model w/ truck class + else: + for i, cfg in enumerate(model_configs): + if "truck" in (cfg.known_classes or []): + truck_det = detectors[i] + break # Single tracker using first model's weights tracker_model = YOLO(model_configs[0].path) @@ -453,10 +460,13 @@ def run_merged_pipeline( break # Truck detection - roi = roi_tracker.roi - if truck_det is not None and frame_idx % truck_det_interval == 0: - trucks = truck_det.detect(frame) - roi = roi_tracker.update(trucks) + if fixed_zone is not None: + roi = fixed_zone + else: + roi = roi_tracker.roi + if truck_det is not None and frame_idx % truck_det_interval == 0: + trucks = truck_det.detect(frame) + roi = roi_tracker.update(trucks) truck_present = roi is not None and roi.confidence > 0 if roi is not None: diff --git a/tests/test_fixed_zone_pipeline.py b/tests/test_fixed_zone_pipeline.py new file mode 100644 index 0000000..c38bca9 --- /dev/null +++ b/tests/test_fixed_zone_pipeline.py @@ -0,0 +1,27 @@ +"""Tests for fixed_zone parameter in pipeline functions.""" + +import inspect + +from src.pipeline import run_pipeline, run_merged_pipeline + + +def test_run_pipeline_accepts_fixed_zone(): + sig = inspect.signature(run_pipeline) + assert "fixed_zone" in sig.parameters + + +def test_run_merged_pipeline_accepts_fixed_zone(): + sig = inspect.signature(run_merged_pipeline) + assert "fixed_zone" in sig.parameters + + +def test_fixed_zone_defaults_to_none(): + sig = inspect.signature(run_pipeline) + param = sig.parameters["fixed_zone"] + assert param.default is None + + +def test_merged_fixed_zone_defaults_to_none(): + sig = inspect.signature(run_merged_pipeline) + param = sig.parameters["fixed_zone"] + assert param.default is None