Files

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 {}