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
+66 -29

No files matched your search

+39 -29
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,17 +144,18 @@ 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 truck_model_config is not None: if fixed_zone is None:
# Load a separate truck detection model if truck_model_config is not None:
truck_shared = YOLO(truck_model_config.path) # Load a separate truck detection model
truck_detector = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) truck_shared = YOLO(truck_model_config.path)
elif "truck" in (model_config.known_classes or []): truck_detector = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",))
# Use the same model for truck detection elif "truck" in (model_config.known_classes or []):
truck_detector = BaseDetector( # Use the same model for truck detection
shared_model, conf=truck_conf, class_filter=("truck",) truck_detector = BaseDetector(
) shared_model, conf=truck_conf, class_filter=("truck",)
)
tracker = ByteTrackTracker(shared_model, conf=sack_conf) tracker = ByteTrackTracker(shared_model, conf=sack_conf)
stabilizer = BboxStabilizer() stabilizer = BboxStabilizer()
@@ -199,10 +201,13 @@ def run_pipeline(
break break
# Truck detection # Truck detection
roi = roi_tracker.roi if fixed_zone is not None:
if truck_detector is not None and frame_idx % truck_det_interval == 0: roi = fixed_zone
trucks = truck_detector.detect(frame) else:
roi = roi_tracker.update(trucks) 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 truck_present = roi is not None and roi.confidence > 0
@@ -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,17 +405,18 @@ 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 truck_model_config is not None: if fixed_zone is None:
truck_shared = YOLO(truck_model_config.path) if truck_model_config is not None:
truck_det = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",)) truck_shared = YOLO(truck_model_config.path)
# ponytail: first-match truck detector; could prefer dedicated truck model over sack model w/ truck class truck_det = BaseDetector(truck_shared, conf=truck_conf, class_filter=("truck",))
else: # ponytail: first-match truck detector; could prefer dedicated truck model over sack model w/ truck class
for i, cfg in enumerate(model_configs): else:
if "truck" in (cfg.known_classes or []): for i, cfg in enumerate(model_configs):
truck_det = detectors[i] if "truck" in (cfg.known_classes or []):
break truck_det = detectors[i]
break
# Single tracker using first model's weights # Single tracker using first model's weights
tracker_model = YOLO(model_configs[0].path) tracker_model = YOLO(model_configs[0].path)
@@ -453,10 +460,13 @@ def run_merged_pipeline(
break break
# Truck detection # Truck detection
roi = roi_tracker.roi if fixed_zone is not None:
if truck_det is not None and frame_idx % truck_det_interval == 0: roi = fixed_zone
trucks = truck_det.detect(frame) else:
roi = roi_tracker.update(trucks) 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 truck_present = roi is not None and roi.confidence > 0
if roi is not None: if roi is not None:
+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