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.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:
|
||||||
|
|||||||
@@ -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