- config.yaml gains counting.max_reid_frames (200) + debounce_frames (8); CountingConfig + parsing + pinned-value test added - predict.py load_zones() delegates to read_zone_polygons() (geometry only); legacy zones.json knob reads removed (config.yaml is canonical) - run_prediction wires 5 more live knobs from cfg.counting.* (entry/exit overlap, tolerance_missing_frames, max_reid_frames, debounce_frames) - delete deprecated src/config.py (v3 keys) + tests/test_config.py; src/main.py bridges onto unified config so the deprecated entrypoint keeps working - 8 never-read globals (JARAK_ABSORBSI_GHOST, CAMERA_NOISE_DEADBAND, ...) left untouched — flagged for Phase 2 dead-code cleanup Verified: pytest 21 passed; full pipeline runs on synthetic video (engines load, 30 frames, clean finish).
288 lines
9.6 KiB
Python
288 lines
9.6 KiB
Python
"""Main pipeline — wires all components together via dependency injection.
|
|
|
|
Pipeline: Frame → TruckDetect → ROI → SackTrack → Stabilize → Count
|
|
The stabilizer sits between the tracker and counter, smoothing bbox
|
|
coordinates and holding lost tracks through momentary dropouts.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import time
|
|
|
|
import cv2
|
|
|
|
from src.batch import BatchLifecycleManager
|
|
from src.config_loader import load_config as _load_unified_config
|
|
from src.counting import LineCrossCounter
|
|
from src.dashboard import DashboardOverlay
|
|
from src.detection import TruckDetector
|
|
from src.logger import CSVLogger
|
|
from src.stabilizer import BboxStabilizer
|
|
from src.streaming import RTSPSource, VideoFileSource
|
|
from src.tracking import ByteTrackTracker
|
|
from src.truck_roi import TruckROITracker
|
|
|
|
|
|
def _legacy_cfg_bridge(env_path: str = ".env"):
|
|
"""Bridge the deprecated v3 Config API onto the unified config_loader.Config.
|
|
|
|
`python -m src.main` is deprecated (use predict.py), but must keep working.
|
|
Only the fields actually used below are bridged.
|
|
"""
|
|
import os
|
|
from types import SimpleNamespace
|
|
|
|
try:
|
|
from dotenv import load_dotenv
|
|
load_dotenv(env_path)
|
|
except ImportError:
|
|
pass
|
|
cfg = _load_unified_config("config.yaml")
|
|
mode = cfg.get_active_mode()
|
|
truck_key = next((e.path for e in mode.engines if "truck" in e.classes), "combined")
|
|
sack_key = next((e.path for e in mode.engines if "sack" in e.classes), "combined")
|
|
base = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
return SimpleNamespace(
|
|
local_rtsp=cfg.stream.rtsp_url,
|
|
jetson_rtsp="",
|
|
sack_model_path=cfg.engine_path(sack_key, base),
|
|
truck_model_path=cfg.engine_path(truck_key, base),
|
|
counting_line_y=0.60,
|
|
counting_line_x_start=0.38,
|
|
counting_line_x_end=0.72,
|
|
sack_conf=cfg.detection_params_for("sack").conf,
|
|
truck_conf=cfg.detection_params_for("truck").conf,
|
|
batch_timeout_seconds=cfg.batch.timeout_seconds,
|
|
csv_output_dir="./output",
|
|
seed=42,
|
|
)
|
|
|
|
|
|
def load_config(env_path: str = ".env"):
|
|
"""Deprecated shim — use src.config_loader.load_config instead."""
|
|
return _legacy_cfg_bridge(env_path)
|
|
|
|
|
|
def build_pipeline(cfg, source_path: str | None = None):
|
|
"""Construct all components from config."""
|
|
|
|
# ── Stream source ────────────────────────────────────────────
|
|
if source_path:
|
|
stream = VideoFileSource(source_path)
|
|
elif cfg.local_rtsp:
|
|
stream = RTSPSource(cfg.local_rtsp)
|
|
else:
|
|
raise ValueError("No video source: pass --source or set LOCAL_RTSP")
|
|
|
|
if not stream.open():
|
|
raise RuntimeError(
|
|
f"Cannot open stream: {source_path or cfg.local_rtsp}"
|
|
)
|
|
|
|
w, h = stream.frame_size
|
|
print(f"Stream opened: {w}x{h} @ {stream.fps:.1f} FPS")
|
|
|
|
# ── Components ───────────────────────────────────────────────
|
|
truck_detector = TruckDetector(cfg.truck_model_path, cfg.truck_conf)
|
|
tracker = ByteTrackTracker(cfg.sack_model_path, cfg.sack_conf)
|
|
stabilizer = BboxStabilizer(
|
|
ema_alpha=0.35,
|
|
max_hold_frames=10,
|
|
max_height_ratio=1.5,
|
|
min_height_ratio=0.70,
|
|
)
|
|
|
|
roi_tracker = TruckROITracker(frame_width=w, frame_height=h)
|
|
|
|
# Initial line position (updated dynamically by ROI tracker)
|
|
counter = LineCrossCounter(
|
|
line_y=int(h * 0.50),
|
|
line_x_start=int(w * 0.38),
|
|
line_x_end=int(w * 0.72),
|
|
margin=20,
|
|
)
|
|
|
|
batch_mgr = BatchLifecycleManager(cfg.batch_timeout_seconds)
|
|
dashboard = DashboardOverlay()
|
|
logger = CSVLogger(cfg.csv_output_dir)
|
|
|
|
# ── Wire callbacks ───────────────────────────────────────────
|
|
def on_batch_start(batch_id: int, timestamp: float) -> None:
|
|
print(f"\n>>> BATCH #{batch_id} STARTED")
|
|
counter.reset()
|
|
stabilizer.reset()
|
|
roi_tracker.reset()
|
|
|
|
def on_batch_end(record) -> None:
|
|
print(
|
|
f"\n>>> BATCH #{record.batch_id} ENDED — "
|
|
f"L={record.loading_count} U={record.unloading_count} "
|
|
f"Net={record.net_count}"
|
|
)
|
|
logger.log_batch(record)
|
|
|
|
batch_mgr.on_batch_start(on_batch_start)
|
|
batch_mgr.on_batch_end(on_batch_end)
|
|
|
|
return {
|
|
"stream": stream,
|
|
"truck_detector": truck_detector,
|
|
"tracker": tracker,
|
|
"stabilizer": stabilizer,
|
|
"roi_tracker": roi_tracker,
|
|
"counter": counter,
|
|
"batch_mgr": batch_mgr,
|
|
"dashboard": dashboard,
|
|
"logger": logger,
|
|
}
|
|
|
|
|
|
def _filter_sacks_in_roi(detections, roi):
|
|
"""Keep only sacks whose centroid X falls within the truck ROI."""
|
|
if roi is None:
|
|
return []
|
|
return [
|
|
d for d in detections
|
|
if roi.contains_x((d.bbox[0] + d.bbox[2]) / 2.0)
|
|
]
|
|
|
|
|
|
def run(cfg, source_path: str | None = None) -> None:
|
|
"""Main processing loop — processes EVERY frame for accuracy."""
|
|
p = build_pipeline(cfg, source_path)
|
|
|
|
stream = p["stream"]
|
|
tracker = p["tracker"]
|
|
stabilizer = p["stabilizer"]
|
|
truck_det = p["truck_detector"]
|
|
roi_tracker = p["roi_tracker"]
|
|
counter = p["counter"]
|
|
batch_mgr = p["batch_mgr"]
|
|
dashboard = p["dashboard"]
|
|
logger = p["logger"]
|
|
|
|
frame_idx = 0
|
|
prev_loading = 0
|
|
prev_unloading = 0
|
|
|
|
# Truck detection runs every N frames (heavy model, truck moves slow)
|
|
TRUCK_DET_INTERVAL = 15
|
|
|
|
try:
|
|
while True:
|
|
ok, frame = stream.read()
|
|
if not ok:
|
|
break
|
|
|
|
timestamp = time.time()
|
|
frame_idx += 1
|
|
|
|
# ── Truck detection (every N frames — truck is slow) ─
|
|
roi = roi_tracker.roi
|
|
if 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
|
|
|
|
# Sync counter line to ROI
|
|
if roi is not None:
|
|
counter.line_y = roi.line_y
|
|
counter.line_x_start = roi.x1
|
|
counter.line_x_end = roi.x2
|
|
|
|
if frame_idx % TRUCK_DET_INTERVAL == 0:
|
|
batch_mgr.update(
|
|
truck_detected=truck_present,
|
|
timestamp=timestamp,
|
|
loading_count=counter.loading_count,
|
|
unloading_count=counter.unloading_count,
|
|
)
|
|
|
|
# ── Track → Stabilize → Filter → Count ──────────────
|
|
tracked_sacks = []
|
|
if batch_mgr.is_active:
|
|
raw_tracked = tracker.update(frame, [])
|
|
stable = stabilizer.update(raw_tracked)
|
|
tracked_sacks = _filter_sacks_in_roi(stable, roi)
|
|
events = counter.update(tracked_sacks)
|
|
|
|
# Log crossing events
|
|
for ev in events:
|
|
logger.log_event(
|
|
batch_mgr.current_batch_id or 0,
|
|
ev["track_id"],
|
|
ev["direction"],
|
|
timestamp,
|
|
)
|
|
|
|
if counter.loading_count != prev_loading:
|
|
print(
|
|
f" [F{frame_idx}] Loading: "
|
|
f"{counter.loading_count}"
|
|
)
|
|
if counter.unloading_count != prev_unloading:
|
|
print(
|
|
f" [F{frame_idx}] Unloading: "
|
|
f"{counter.unloading_count}"
|
|
)
|
|
prev_loading = counter.loading_count
|
|
prev_unloading = counter.unloading_count
|
|
|
|
# ── Dashboard ────────────────────────────────────────
|
|
viz = dashboard.draw(
|
|
frame=frame,
|
|
detections=tracked_sacks,
|
|
roi=roi,
|
|
loading_count=counter.loading_count,
|
|
unloading_count=counter.unloading_count,
|
|
batch_id=batch_mgr.current_batch_id,
|
|
history=batch_mgr.history,
|
|
)
|
|
|
|
cv2.imshow("Sack Counter", viz)
|
|
key = cv2.waitKey(1) & 0xFF
|
|
if key == ord("q"):
|
|
break
|
|
elif key == ord("r"):
|
|
counter.reset()
|
|
stabilizer.reset()
|
|
print("Counter reset manually")
|
|
|
|
finally:
|
|
stream.release()
|
|
cv2.destroyAllWindows()
|
|
print(f"\nProcessed {frame_idx} frames")
|
|
print(
|
|
f"Final — Loading: {counter.loading_count} "
|
|
f"Unloading: {counter.unloading_count} "
|
|
f"Net: {counter.net_count}"
|
|
)
|
|
|
|
|
|
def main() -> None:
|
|
import sys
|
|
print(
|
|
"DEPRECATED: `python -m src.main` hanya untuk eksperimen. "
|
|
"Untuk pipeline produksi gunakan: python predict.py --source VIDEO --env .env",
|
|
file=sys.stderr,
|
|
)
|
|
parser = argparse.ArgumentParser(description="Sack Counting System v3 (deprecated entrypoint)")
|
|
parser.add_argument(
|
|
"--source", type=str, default=None,
|
|
help="Video file path (overrides RTSP)",
|
|
)
|
|
parser.add_argument(
|
|
"--env", type=str, default=".env",
|
|
help="Path to .env file",
|
|
)
|
|
args = parser.parse_args()
|
|
|
|
cfg = load_config(args.env)
|
|
run(cfg, args.source)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|