394 lines
16 KiB
Python
394 lines
16 KiB
Python
"""Live counting test bench: point a trained model at a live stream and watch it count.
|
|
|
|
This is a **test harness**, not the production counter. It reuses the real
|
|
pipeline pieces from `algoritma-batch` — ByteTrack, the bbox stabiliser and the
|
|
line-cross counter with its spatial dedup — so what you see here is what
|
|
`predict.py` would do. What it deliberately leaves out is everything stateful:
|
|
no batch lifecycle, no SQLite, no truck-presence state machine. The question it
|
|
answers is "does this model count correctly on this camera", and those parts
|
|
only get in the way of answering it.
|
|
|
|
One session at a time, holding the GPU lock, because the GPU is shared with
|
|
training and auto-annotation (REQ-065, REQ-070).
|
|
"""
|
|
|
|
import os
|
|
import sys
|
|
import threading
|
|
import time
|
|
from typing import Optional
|
|
|
|
import cv2
|
|
import numpy as np
|
|
|
|
from backend import config, jobs
|
|
from backend import live_render
|
|
from backend.live_source import ( # re-exported: callers catch live_count.LiveCountError
|
|
LiveCountError, _is_stream, _open, whep_to_rtsp, whep_url,
|
|
)
|
|
|
|
# `algoritma-batch/src` is copied to /app/src in the image; in a source checkout
|
|
# it still lives under algoritma-batch/. Both are made importable as `src.*`.
|
|
_REPO = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
|
for candidate in ("/app", os.path.join(_REPO, "algoritma-batch")):
|
|
if os.path.isdir(os.path.join(candidate, "src")) and candidate not in sys.path:
|
|
sys.path.insert(0, candidate)
|
|
|
|
|
|
class Session:
|
|
"""One running counter. Owns a capture thread and the latest rendered frame."""
|
|
|
|
def __init__(self, source: str, model_path: str, line_y: int,
|
|
line_x_start: int, line_x_end: int, conf: float,
|
|
dedup_radius: float, margin: int, imgsz: int,
|
|
entry_travel_min: float, handoff_radius: float,
|
|
unload_confirm_frames: int, min_area_scale: float,
|
|
spatial_dedup: bool, whep_url: str = ""):
|
|
self.source = source
|
|
# When the browser watches the camera over WebRTC it never asks for the
|
|
# MJPEG, so encoding a JPEG per frame would be pure waste (REQ-177).
|
|
self.whep_url = whep_url
|
|
self.model_path = model_path
|
|
self.line_y = line_y
|
|
self.line_x_start = line_x_start
|
|
self.line_x_end = line_x_end
|
|
self.conf = conf
|
|
self.dedup_radius = dedup_radius
|
|
self.margin = margin
|
|
self.imgsz = imgsz
|
|
self.entry_travel_min = entry_travel_min
|
|
self.handoff_radius = handoff_radius
|
|
self.unload_confirm_frames = unload_confirm_frames
|
|
self.min_area_scale = min_area_scale
|
|
self.spatial_dedup = spatial_dedup
|
|
|
|
self.started_at = time.time()
|
|
self.error = ""
|
|
self.stopping = False
|
|
self.frames = 0
|
|
self.fps = 0.0
|
|
self.loading = 0
|
|
self.unloading = 0
|
|
self.tracked = 0
|
|
self.ignored = 0
|
|
self.too_small = 0
|
|
self.traced = 0
|
|
self.trace_path = os.path.join(
|
|
config.DATA_DIR, "live-count", f"session-{int(self.started_at)}.jsonl")
|
|
self.events: list[dict] = []
|
|
|
|
self._jpeg: Optional[bytes] = None
|
|
self._overlay: dict = {}
|
|
self._counter = None # set once the worker builds it
|
|
self._lock = threading.Lock()
|
|
self._thread = threading.Thread(target=self._run, name="live-count", daemon=True)
|
|
|
|
# -- public -----------------------------------------------------------
|
|
|
|
def start(self) -> None:
|
|
self._thread.start()
|
|
|
|
def stop(self) -> None:
|
|
self.stopping = True
|
|
self._thread.join(timeout=15)
|
|
|
|
def snapshot(self) -> Optional[bytes]:
|
|
with self._lock:
|
|
return self._jpeg
|
|
|
|
def overlay(self) -> dict:
|
|
with self._lock:
|
|
return dict(self._overlay)
|
|
|
|
def move_line(self, line_y: Optional[int] = None, line_x_start: Optional[int] = None,
|
|
line_x_end: Optional[int] = None) -> dict:
|
|
"""Reposition the counting line without restarting.
|
|
|
|
Placing a line correctly means watching the stream while you move it,
|
|
which is impossible if moving it requires a restart — and a restart
|
|
throws the counts away. Already-counted tracks keep their verdict; the
|
|
counter only decides for a track the first time it crosses.
|
|
"""
|
|
if line_y is not None:
|
|
self.line_y = max(0, min(720, int(line_y)))
|
|
if line_x_start is not None:
|
|
self.line_x_start = max(0, min(1280, int(line_x_start)))
|
|
if line_x_end is not None:
|
|
self.line_x_end = max(0, min(1280, int(line_x_end)))
|
|
if self.line_x_end < self.line_x_start:
|
|
self.line_x_start, self.line_x_end = self.line_x_end, self.line_x_start
|
|
if self._counter is not None:
|
|
self._counter.line_y = self.line_y
|
|
self._counter.line_x_start = self.line_x_start
|
|
self._counter.line_x_end = self.line_x_end
|
|
return {"y": self.line_y, "x_start": self.line_x_start, "x_end": self.line_x_end}
|
|
|
|
def status(self) -> dict:
|
|
return {
|
|
"running": self._thread.is_alive(),
|
|
"source": self.source,
|
|
"whep_url": self.whep_url,
|
|
"preview": "webrtc" if self.whep_url else "mjpeg",
|
|
"model_path": self.model_path,
|
|
"error": self.error,
|
|
"frames": self.frames,
|
|
"fps": round(self.fps, 1),
|
|
"loading": self.loading,
|
|
"unloading": self.unloading,
|
|
"net": self.loading - self.unloading,
|
|
"tracked": self.tracked,
|
|
"ignored": self.ignored,
|
|
"elapsed": round(time.time() - self.started_at, 1),
|
|
"line": {"y": self.line_y, "x_start": self.line_x_start, "x_end": self.line_x_end},
|
|
"too_small": self.too_small,
|
|
"traced": self.traced,
|
|
"trace_path": self.trace_path,
|
|
"events": self.events[-25:],
|
|
}
|
|
|
|
# -- worker -----------------------------------------------------------
|
|
|
|
def _run(self) -> None:
|
|
# The GPU is shared. Waiting here rather than failing means "start" is
|
|
# safe to press while a training run is finishing.
|
|
if not jobs.gpu_lock.acquire(timeout=30):
|
|
busy = jobs.running_types()
|
|
self.error = f"GPU busy with a {busy[0] if busy else 'background'} job"
|
|
return
|
|
capture = None
|
|
try:
|
|
from ultralytics import YOLO
|
|
from src.counting import LineCrossCounter
|
|
from src.stabilizer import BboxStabilizer
|
|
from src.tracking import ByteTrackTracker
|
|
|
|
model = YOLO(self.model_path)
|
|
# Warm-up: the first CUDA call inside the tracker has been seen to
|
|
# segfault without it (same reason predict.py does this).
|
|
model(np.zeros((720, 1280, 3), dtype=np.uint8), imgsz=self.imgsz, verbose=False)
|
|
|
|
tracker = ByteTrackTracker(model, self.conf)
|
|
stabilizer = BboxStabilizer(ema_alpha=0.35, max_hold_frames=10,
|
|
max_height_ratio=1.5, min_height_ratio=0.70)
|
|
counter = LineCrossCounter(
|
|
line_y=self.line_y, line_x_start=self.line_x_start,
|
|
line_x_end=self.line_x_end, margin=self.margin,
|
|
dedup_radius=self.dedup_radius,
|
|
entry_travel_min=self.entry_travel_min,
|
|
handoff_radius=self.handoff_radius,
|
|
unload_confirm_frames=self.unload_confirm_frames,
|
|
spatial_dedup=self.spatial_dedup,
|
|
)
|
|
self._counter = counter
|
|
|
|
capture = _open(self.source)
|
|
if capture is None or not capture.isOpened():
|
|
raise LiveCountError(f"Could not open source: {self.source}")
|
|
|
|
tick = time.time()
|
|
since = 0
|
|
while not self.stopping:
|
|
ok, frame = capture.read()
|
|
if not ok or frame is None:
|
|
# A file simply ends. On a stream this only means the
|
|
# decoder has not produced a new frame yet, so wait briefly
|
|
# — long enough not to spin, short enough not to become the
|
|
# new frame-rate ceiling.
|
|
if _is_stream(self.source):
|
|
time.sleep(0.005)
|
|
continue
|
|
break
|
|
|
|
frame = cv2.resize(frame, (1280, 720))
|
|
detections = [d for d in tracker.update(frame, []) if d.class_name == "sack"]
|
|
stable = stabilizer.update(detections)
|
|
inside, outside = [], []
|
|
small = 0
|
|
for det in stable:
|
|
x1, y1, x2, y2 = det.bbox
|
|
centre_x = (x1 + x2) / 2
|
|
# Perspective-aware area gate, the curve `predict.py` uses:
|
|
# a box that small at that depth is a fragment, not a sack.
|
|
if _too_small(det.bbox, self.min_area_scale):
|
|
small += 1
|
|
outside.append(det)
|
|
elif self.line_x_start <= centre_x <= self.line_x_end:
|
|
inside.append(det)
|
|
else:
|
|
outside.append(det)
|
|
self.too_small = small
|
|
for event in counter.update(inside):
|
|
self.events.append({
|
|
"track_id": event.get("track_id"),
|
|
"direction": event.get("direction", "loading"),
|
|
"at": round(time.time() - self.started_at, 1),
|
|
})
|
|
|
|
self._write_traces(counter.drain_traces())
|
|
self.loading = counter.loading_count
|
|
self.unloading = counter.unloading_count
|
|
self.tracked = len(inside)
|
|
self.ignored = len(outside)
|
|
self.frames += 1
|
|
since += 1
|
|
|
|
now = time.time()
|
|
if now - tick >= 1.0:
|
|
self.fps = since / (now - tick)
|
|
tick, since = now, 0
|
|
|
|
self._set_overlay(inside, outside, counter)
|
|
if not self.whep_url:
|
|
jpeg = live_render.render(self, frame, inside, outside, counter)
|
|
with self._lock:
|
|
self._jpeg = jpeg
|
|
except Exception as exc: # surfaced in status(), not swallowed
|
|
self.error = f"{type(exc).__name__}: {exc}"
|
|
finally:
|
|
# The lock is released no matter what tearing down the capture does.
|
|
# It was the other order once, and one exception in release() leaked
|
|
# the GPU for the lifetime of the process.
|
|
try:
|
|
if capture is not None:
|
|
capture.release()
|
|
except Exception as exc:
|
|
if not self.error:
|
|
self.error = f"capture release failed: {exc}"
|
|
finally:
|
|
jobs.gpu_lock.release()
|
|
|
|
def _write_traces(self, traces: list) -> None:
|
|
"""Append finished tracks to a JSONL, one object per track.
|
|
|
|
This is the file that answers "was that the model, the tracker or the
|
|
counter" on a clip with a known count: every track that ever existed
|
|
lands here with its trajectory and the reason it did or did not count.
|
|
"""
|
|
if not traces:
|
|
return
|
|
import json
|
|
try:
|
|
os.makedirs(os.path.dirname(self.trace_path), exist_ok=True)
|
|
with open(self.trace_path, "a", encoding="utf-8") as handle:
|
|
for record in traces:
|
|
handle.write(json.dumps(record, default=str) + "\n")
|
|
self.traced += len(traces)
|
|
except OSError as exc:
|
|
if not self.error:
|
|
self.error = f"trace write failed: {exc}"
|
|
|
|
def _set_overlay(self, detections, ignored, counter) -> None:
|
|
"""The same drawing `_render` burns into the JPEG, as geometry.
|
|
|
|
The browser draws it on a canvas over the WebRTC video, so the frames
|
|
themselves never pass through this process. Coordinates are in the
|
|
1280x720 space the pipeline works in; the canvas scales them.
|
|
"""
|
|
boxes = []
|
|
for det in detections:
|
|
state = "counted" if counter.counted_tracks.get(det.track_id) else "tracked"
|
|
boxes.append({"id": det.track_id, "b": [int(v) for v in det.bbox],
|
|
"c": round(float(det.confidence), 2), "s": state})
|
|
for det in ignored:
|
|
boxes.append({"id": det.track_id, "b": [int(v) for v in det.bbox],
|
|
"c": round(float(det.confidence), 2), "s": "ignored"})
|
|
payload = {
|
|
"frame": self.frames,
|
|
"line": {"y": self.line_y, "x_start": self.line_x_start,
|
|
"x_end": self.line_x_end, "margin": self.margin},
|
|
"boxes": boxes,
|
|
"loading": self.loading,
|
|
"unloading": self.unloading,
|
|
"net": self.loading - self.unloading,
|
|
"fps": round(self.fps, 1),
|
|
}
|
|
with self._lock:
|
|
self._overlay = payload
|
|
|
|
def _too_small(bbox, scale: float) -> bool:
|
|
"""`predict.py`'s perspective curve: a sack at the top of the frame is
|
|
genuinely smaller in pixels than the same sack at the bottom, so one flat
|
|
threshold either lets fragments through up close or discards real sacks far
|
|
away. Interpolated in the 1280x720 space the pipeline works in. `scale` of 0
|
|
turns the gate off."""
|
|
if scale <= 0:
|
|
return False
|
|
x1, y1, x2, y2 = bbox
|
|
centre_y = (y1 + y2) / 2.0
|
|
top_y, bottom_y = 133.0, 720.0
|
|
top_area, bottom_area = 3556.0, 11111.0
|
|
if centre_y <= top_y:
|
|
minimum = top_area
|
|
elif centre_y >= bottom_y:
|
|
minimum = bottom_area
|
|
else:
|
|
ratio = (centre_y - top_y) / (bottom_y - top_y)
|
|
minimum = top_area + ratio * (bottom_area - top_area)
|
|
return (x2 - x1) * (y2 - y1) < minimum * scale
|
|
|
|
|
|
# ---- module-level single session ----------------------------------------
|
|
|
|
_session: Optional[Session] = None
|
|
_guard = threading.Lock()
|
|
|
|
|
|
def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_end: int,
|
|
conf: float = 0.35, dedup_radius: float = 60.0, margin: int = 5,
|
|
imgsz: int = 640, entry_travel_min: float = 60.0,
|
|
handoff_radius: float = 100.0, unload_confirm_frames: int = 3,
|
|
min_area_scale: float = 1.0, spatial_dedup: bool = False,
|
|
whep: str = "") -> dict:
|
|
global _session
|
|
with _guard:
|
|
if _session is not None and _session.status()["running"]:
|
|
raise LiveCountError("A counting session is already running — stop it first")
|
|
if not os.path.isfile(model_path):
|
|
raise LiveCountError(f"Model not found: {model_path}")
|
|
if not _is_stream(source) and not os.path.isfile(source):
|
|
raise LiveCountError(f"Source not found: {source}")
|
|
_session = Session(source, model_path, line_y, line_x_start, line_x_end,
|
|
conf, dedup_radius, margin, imgsz, entry_travel_min,
|
|
handoff_radius, unload_confirm_frames, min_area_scale,
|
|
spatial_dedup, whep)
|
|
_session.start()
|
|
time.sleep(0.4) # let an immediate failure surface in the response
|
|
return _session.status()
|
|
|
|
|
|
def stop() -> dict:
|
|
global _session
|
|
with _guard:
|
|
if _session is None:
|
|
return {"running": False}
|
|
_session.stop()
|
|
report = _session.status()
|
|
_session = None
|
|
return report
|
|
|
|
|
|
def move_line(line_y=None, line_x_start=None, line_x_end=None) -> dict:
|
|
if _session is None:
|
|
raise LiveCountError("No counting session is running")
|
|
return _session.move_line(line_y, line_x_start, line_x_end)
|
|
|
|
|
|
def status() -> dict:
|
|
if _session is None:
|
|
# Same shape as a live session, so callers never branch on presence.
|
|
return {"running": False, "loading": 0, "unloading": 0, "net": 0, "frames": 0,
|
|
"fps": 0.0, "tracked": 0, "ignored": 0, "too_small": 0, "traced": 0,
|
|
"trace_path": "", "elapsed": 0.0, "events": [],
|
|
"error": "", "source": "", "model_path": "",
|
|
"whep_url": "", "preview": "mjpeg",
|
|
"line": {"y": 0, "x_start": 0, "x_end": 1280}}
|
|
return _session.status()
|
|
|
|
|
|
def snapshot() -> Optional[bytes]:
|
|
return _session.snapshot() if _session is not None else None
|
|
|
|
|
|
def overlay() -> dict:
|
|
return _session.overlay() if _session is not None else {}
|