feat: add fixed_zone param to pipelines, skip truck detection when set
This commit is contained in:
1 parent
4ad4284d9a
commit
ad3e7a7e9e
2 files changed
+66
-29
No files matched your search
+39
-29
@@ -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:
|
||||
|
||||
@@ -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
|
||||
Reference in new issue
Block a user