feat: add fixed_zone param to pipelines, skip truck detection when set

This commit is contained in:
jetson committed 2026-09-22 13:01:40 +07:00
1 parent 4ad4284d9a
commit ad3e7a7e9e
2 files changed
+40 -3

No files matched your search

+13 -3
View File
@@ -20,7 +20,7 @@ from src.interfaces import Detection
from src.model_registry import ModelConfig from src.model_registry import ModelConfig
from src.stabilizer import BboxStabilizer from src.stabilizer import BboxStabilizer
from src.tracking import ByteTrackTracker from src.tracking import ByteTrackTracker
from src.truck_roi import TruckROITracker from src.truck_roi import TruckROI, TruckROITracker
from src.video_writer import AnnotatedVideoWriter from src.video_writer import AnnotatedVideoWriter
@@ -78,6 +78,7 @@ def run_pipeline(
output_path: str, output_path: str,
class_filter: list[str] | None = None, class_filter: list[str] | None = None,
truck_model_config: ModelConfig | None = None, truck_model_config: ModelConfig | None = None,
fixed_zone: TruckROI | None = None,
sack_conf: float = 0.4, sack_conf: float = 0.4,
truck_conf: float = 0.5, truck_conf: float = 0.5,
truck_det_interval: int = 15, truck_det_interval: int = 15,
@@ -143,8 +144,9 @@ def run_pipeline(
shared_model, conf=sack_conf, class_filter=effective_filter 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 truck_detector = None
if fixed_zone is None:
if truck_model_config is not None: if truck_model_config is not None:
# Load a separate truck detection model # Load a separate truck detection model
truck_shared = YOLO(truck_model_config.path) truck_shared = YOLO(truck_model_config.path)
@@ -199,6 +201,9 @@ def run_pipeline(
break break
# Truck detection # Truck detection
if fixed_zone is not None:
roi = fixed_zone
else:
roi = roi_tracker.roi roi = roi_tracker.roi
if truck_detector is not None and frame_idx % truck_det_interval == 0: if truck_detector is not None and frame_idx % truck_det_interval == 0:
trucks = truck_detector.detect(frame) trucks = truck_detector.detect(frame)
@@ -362,6 +367,7 @@ def run_merged_pipeline(
output_path: str, output_path: str,
class_filters: dict[str, list[str] | None] | None = None, class_filters: dict[str, list[str] | None] | None = None,
truck_model_config: ModelConfig | None = None, truck_model_config: ModelConfig | None = None,
fixed_zone: TruckROI | None = None,
sack_conf: float = 0.4, sack_conf: float = 0.4,
truck_conf: float = 0.5, truck_conf: float = 0.5,
truck_det_interval: int = 15, truck_det_interval: int = 15,
@@ -399,8 +405,9 @@ def run_merged_pipeline(
eff_filter = (class_filters or {}).get(cfg.stem) or cfg.known_classes or None 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)) detectors.append(BaseDetector(model, conf=sack_conf, class_filter=eff_filter))
# Truck detector # Truck detector — skip if fixed zone
truck_det = None truck_det = None
if fixed_zone is None:
if truck_model_config is not None: if truck_model_config is not None:
truck_shared = YOLO(truck_model_config.path) truck_shared = YOLO(truck_model_config.path)
truck_det = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) truck_det = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",))
@@ -453,6 +460,9 @@ def run_merged_pipeline(
break break
# Truck detection # Truck detection
if fixed_zone is not None:
roi = fixed_zone
else:
roi = roi_tracker.roi roi = roi_tracker.roi
if truck_det is not None and frame_idx % truck_det_interval == 0: if truck_det is not None and frame_idx % truck_det_interval == 0:
trucks = truck_det.detect(frame) trucks = truck_det.detect(frame)
+27
View File
@@ -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