feat: sync live-count, exemplar annotation modules, and update .gitignore

This commit is contained in:
ervanfahriaw committed 2026-08-24 14:57:42 +07:00
1 parent b6624eeff9
commit ac95674c07
39 files changed
+3979 -444

No files matched your search

+14 -3
View File
@@ -36,6 +36,11 @@ class AutolabelRequest(BaseModel):
append: bool = False
custom_model_path: Optional[str] = None
class Exemplar(BaseModel):
box: list[float] # [cx, cy, w, h], normalized 0..1
positive: bool = True
class PreviewRequest(BaseModel):
frame_id: int
engine: str
@@ -44,6 +49,10 @@ class PreviewRequest(BaseModel):
min_box_frac: float = 0.0
target_class_names: Optional[list[str]] = None
custom_model_path: Optional[str] = None
# Drawn box exemplars (REQ-172): normalized cxcywh, positive or negative.
# Preview-only — `autolabel.start` deliberately has no equivalent.
exemplars: Optional[list[Exemplar]] = None
exemplar_class_name: Optional[str] = None
@router.post("/api/projects/{project_id}/batches")
@@ -162,14 +171,14 @@ async def autolabel_with_model(
@router.post("/api/batches/{batch_id}/preview")
def preview_autolabel(batch_id: int, request: PreviewRequest) -> dict:
from backend import autolabel, jobs
from backend import jobs, preview
if not jobs.gpu_lock.acquire(timeout=20):
busy = jobs.running_types()
kind = busy[0] if busy else "background"
raise HTTPException(409, f"The GPU is busy with a {kind} job — wait for it to finish")
try:
shapes = autolabel.preview_frame(
shapes = preview.preview_frame(
batch_id=batch_id,
frame_id=request.frame_id,
engine=request.engine,
@@ -177,7 +186,9 @@ def preview_autolabel(batch_id: int, request: PreviewRequest) -> dict:
iou_threshold=request.iou_threshold,
min_box_frac=request.min_box_frac,
target_class_names=request.target_class_names,
custom_model_path=request.custom_model_path
custom_model_path=request.custom_model_path,
exemplars=[e.model_dump() for e in request.exemplars or []],
exemplar_class_name=request.exemplar_class_name,
)
return {"shapes": shapes}
except Exception as exc:
+30 -6
View File
@@ -15,9 +15,9 @@ router = APIRouter(tags=["live-count"])
class StartRequest(BaseModel):
# Either a raw source (RTSP URL or absolute path) or an archive-relative
# path like "2026-08-13/batch001.mp4", which the backend resolves — the
# frontend never needs to know where the archive is mounted.
# Either a live stream — which must be a WebRTC (WHEP) URL, REQ-176 — or an
# archive-relative path like "2026-08-13/batch001.mp4", which the backend
# resolves so the frontend never needs to know where the archive is mounted.
source: str = ""
source_rel: Optional[str] = None
model_path: Optional[str] = None
@@ -60,14 +60,22 @@ def available_models(project_id: int) -> dict:
def start(project_id: int, request: StartRequest) -> dict:
project = project_or_404(project_id)
source = request.source
source, whep = request.source, ""
if request.source_rel:
try:
source = library.resolve(project["video_root"], request.source_rel)
except library.LibraryError as exc:
raise HTTPException(400, str(exc))
elif source:
# A live source is a WebRTC URL and nothing else. The browser watches it
# over WebRTC; the counter decodes the RTSP leg of the same MediaMTX
# path, derived here (REQ-176).
try:
whep, source = live_count.whep_url(source), live_count.whep_to_rtsp(source)
except live_count.LiveCountError as exc:
raise HTTPException(400, str(exc))
if not source:
raise HTTPException(400, "Pick a video or enter a stream URL")
raise HTTPException(400, "Pick a video or enter a WebRTC stream URL")
path = request.model_path
if not path and request.model_version_id is not None:
@@ -88,6 +96,7 @@ def start(project_id: int, request: StartRequest) -> dict:
unload_confirm_frames=request.unload_confirm_frames,
min_area_scale=request.min_area_scale,
spatial_dedup=request.spatial_dedup,
whep=whep,
)
except live_count.LiveCountError as exc:
raise HTTPException(400, str(exc))
@@ -118,9 +127,24 @@ def status() -> dict:
return live_count.status()
@router.get("/api/live-count/overlay")
def overlay() -> dict:
"""Geometry only — boxes, line and counts for the frame just processed.
What the WebRTC preview draws over the video, instead of the server
encoding a JPEG per frame for it (REQ-177).
"""
return live_count.overlay()
@router.get("/api/live-count/stream")
def stream():
"""MJPEG of the annotated frames. Ends when the session does."""
"""MJPEG of the annotated frames — the archive-file preview. Ends with the session."""
if live_count.status().get("preview") == "webrtc":
# Nothing is encoding JPEGs for this session; without this the generator
# would sit on a worker thread for ten seconds producing nothing.
raise HTTPException(409, "This session is watched over WebRTC, not MJPEG")
def frames():
blank_streak = 0
while True:
+37
View File
@@ -40,6 +40,25 @@ class AssistRequest(BaseModel):
threshold: float = 0.5
class PoolExemplar(BaseModel):
# Normalized xyxy against the frame, as drawn on the review canvas.
box: List[float]
positive: bool = True
class ExemplarLabelRequest(BaseModel):
# The whole frame-local pool, newest last (REQ-173/174).
exemplars: List[PoolExemplar]
class_id: int = 0
# The filter panel (REQ-175). Defaults mirror `exemplar.DEFAULTS`.
threshold: float = 0.5
iou_threshold: float = 0.8
min_box_frac: float = 0.002
max_detections: int = 100
# Off by default: a drag previews, only Apply writes.
apply: bool = False
@router.get("/api/frames/{frame_id}/annotations")
def list_annotations(frame_id: int) -> dict:
target = review_store.frame(frame_id)
@@ -99,6 +118,24 @@ def assist(frame_id: int, request: AssistRequest) -> dict:
raise HTTPException(400, str(exc))
@router.post("/api/frames/{frame_id}/exemplar-label")
def exemplar_label(frame_id: int, request: ExemplarLabelRequest) -> dict:
from backend import exemplar as exemplar_store
try:
return exemplar_store.label(
frame_id, request.class_id,
[item.model_dump() for item in request.exemplars],
threshold=request.threshold,
iou_threshold=request.iou_threshold,
min_box_frac=request.min_box_frac,
max_detections=request.max_detections,
apply=request.apply,
)
except review_store.ReviewError as exc:
raise HTTPException(400, str(exc))
@router.post("/api/frames/{frame_id}/status")
def set_status(frame_id: int, request: StatusRequest) -> dict:
try:
Binary file not shown.
-122
View File
@@ -259,125 +259,3 @@ def _reset_reviewed(batch_id: int) -> None:
"AND review_status = 'approved'",
(batch_id,),
)
def preview_frame(
batch_id: int,
frame_id: int,
engine: str,
threshold: float = DEFAULT_THRESHOLD,
iou_threshold: float = DEFAULT_IOU,
min_box_frac: float = 0.0,
target_class_names: Optional[List[str]] = None,
custom_model_path: Optional[str] = None
) -> List[dict]:
batch = batches.get(batch_id)
if not batch:
raise ValueError("No such batch")
project = projects.get(batch["project_id"])
frame = next((f for f in batches.frames(batch_id) if f["id"] == frame_id), None)
if not frame:
raise ValueError("Frame not found")
directory = batches.frames_dir(batch["project_slug"], batch_id)
frame_file = os.path.join(directory, frame["filename"])
fw = max(1, frame.get("width") or 1)
fh = max(1, frame.get("height") or 1)
yolo_model = None
sam3_target_classes = []
if engine == "sam3" and not custom_model_path:
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
if allowed_classes_set:
sam3_target_classes = [c for c in project["classes"] if c["name"].strip().lower() in allowed_classes_set or c["prompt"].strip().lower() in allowed_classes_set]
else:
sam3_target_classes = [c for c in project["classes"]]
prompts = [c["prompt"] for c in sam3_target_classes]
if prompts:
from backend.sam3_engine import get_engine
get_engine()
else:
from ultralytics import YOLO
if custom_model_path and os.path.isfile(custom_model_path):
m_path = custom_model_path
else:
m_path = projects.training_start_point(project)
with db.cursor() as cur:
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
row = cur.fetchone()
if row and os.path.isfile(row[0]):
m_path = row[0]
yolo_model = YOLO(m_path)
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
all_raw_detections = []
if yolo_model is not None:
results = yolo_model.predict(frame_file, conf=threshold, verbose=False)
if results and len(results) > 0:
model_names = results[0].names
for box in results[0].boxes:
cls_idx = int(box.cls[0].item())
raw_cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
target_class_id = name_to_class_id.get(raw_cls_name)
if target_class_id is None:
for item in project["classes"]:
if item["class_id"] == cls_idx:
target_class_id = item["class_id"]
break
if target_class_id is None and 0 <= cls_idx < len(project["classes"]):
target_class_id = project["classes"][cls_idx]["class_id"]
if target_class_id is None:
continue
target_cls_obj = next((c for c in project["classes"] if c["class_id"] == target_class_id), None)
proj_cls_name = target_cls_obj["name"].strip().lower() if target_cls_obj else ""
if allowed_classes_set is not None:
if (raw_cls_name not in allowed_classes_set and
proj_cls_name not in allowed_classes_set and
str(target_class_id) not in allowed_classes_set):
continue
score = float(box.conf[0].item())
xyxyn = box.xyxyn[0].tolist()
all_raw_detections.append(labeling.Detection(
class_id=target_class_id,
class_name=proj_cls_name or raw_cls_name,
box=[xyxyn[0]*fw, xyxyn[1]*fh, xyxyn[2]*fw, xyxyn[3]*fh],
score=score,
mask=None
))
if engine == "sam3" and sam3_target_classes:
prompts = [(c.get("prompt") or c["name"]).strip() for c in sam3_target_classes]
res = labeling.label_image(
frame_file, frame["filename"], prompts, threshold,
iou_threshold=iou_threshold, min_box_frac=min_box_frac
)
if not res.error and res.detections:
for det in res.detections:
if 0 <= det.class_id < len(sam3_target_classes):
real_cls = sam3_target_classes[det.class_id]
det.class_id = real_cls["class_id"]
det.class_name = real_cls["name"]
all_raw_detections.append(det)
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_threshold)
items = []
for det in kept:
if project["label_type"] == "bbox" or det.mask is None:
geom = review.bbox(det.box[0]/fw, det.box[1]/fh, det.box[2]/fw, det.box[3]/fh)
items.append({"class_id": det.class_id, "geometry": geom, "score": det.score})
else:
for geometry in _geometries(det, fw, fh, project["label_type"]):
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
return items
+270
View File
@@ -0,0 +1,270 @@
"""Exemplar-driven manual labeling in the review editor (REQ-173/174/175).
A drag on the review canvas is not just a rectangle: it is a visual prompt.
The drawn box joins a frame-local pool, the pool is replayed against SAM3
together with the class's text prompt, and the whole class is re-detected on
that frame from the result.
The pool lives in the editor, not in the database, and is sent whole on every
call. That keeps this module stateless and matches REQ-172's reasoning: SAM3's
geometric prompts pool features from *this* image, so a pool only means
anything for as long as the user is looking at the frame it was drawn on.
Split out of `review.py` because that file is already at the 400-line limit.
"""
import json
import time
from typing import List, Optional
from backend import db, projects, review
# A detection this close to a box the user drew is the same object: the user's
# own shape wins, so the detection is dropped rather than stacked on top of it.
DUPLICATE_IOU = 0.6
# A detection overlapping a negative box by this much is what the user pointed
# at when they said "not this" (REQ-174). Lower than DUPLICATE_IOU because a
# negative is drawn roughly, around something the user wants gone.
NEGATIVE_IOU = 0.3
# Below this, no detection is really "inside" a drawn box, so a polygon project
# keeps the rectangle rather than snapping to an unrelated mask.
SNAP_IOU = 0.1
# What the filter panel opens with (REQ-175). Measured on a dense `sack` frame:
# NMS at 0.8 only removes near-duplicates and a 0.002 area floor only removes
# specks, where the aggressive-looking values delete real, touching objects.
DEFAULTS = {
"threshold": 0.5,
"iou_threshold": 0.8,
"min_box_frac": 0.002,
"max_detections": 100,
}
def _class_prompt(project_id: int, class_id: int) -> str:
project = projects.get(project_id)
for item in project["classes"]:
if item["class_id"] == class_id:
return (item.get("prompt") or item["name"]).strip()
raise review.ReviewError(f"Class {class_id} does not exist in this project")
def _rect(points: List[float], label_type: str) -> dict:
"""The drawn rectangle as a storable shape for this project."""
x0, y0, x1, y1 = points
if label_type == "bbox":
return review.bbox(x0, y0, x1, y1)
return review.polygon([(x0, y0), (x1, y0), (x1, y1), (x0, y1)])
def _cxcywh(points: List[float]) -> List[float]:
x0, y0, x1, y1 = points
return [(x0 + x1) / 2, (y0 + y1) / 2, x1 - x0, y1 - y0]
def _replace_class(frame_id: int, class_id: int, items: List[dict]) -> None:
"""Swap every shape of one class on one frame for a fresh set.
"Replace everything, re-add drawn": the user's exemplar shapes are part of
`items`, so they come back verbatim in the same transaction.
"""
now = time.time()
with db.cursor() as cur:
cur.execute("DELETE FROM annotations WHERE frame_id = ? AND class_id = ?",
(frame_id, class_id))
cur.executemany(
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
VALUES (?, ?, ?, ?, ?, ?)""",
[(frame_id, class_id, json.dumps(item["geometry"]), item.get("score", 1.0),
item.get("source", "auto"), now) for item in items],
)
def _append_drawn(frame_id: int, class_id: int, drawn: List[dict]) -> int:
"""Add drawn shapes the frame does not already carry, leaving the rest alone.
The pool is re-sent whole on every call, so most of it is usually already
stored; only what is genuinely new gets inserted.
"""
from backend.labeling import _iou
existing = [review.to_box(row["geometry"]) for row in review.listing(frame_id)
if row["class_id"] == class_id]
fresh = [item for item in drawn
if not any(_iou(review.to_box(item["geometry"]), box) >= 0.9
for box in existing)]
for item in fresh:
review.add(frame_id, class_id, item["geometry"], source="manual")
return len(fresh)
def _drop_negative_overlaps(frame_id: int, class_id: int,
negatives: List[List[float]]) -> int:
"""Delete shapes of this class the user shift-dragged over (REQ-174).
Used on the path where SAM3 never runs; the re-detect path filters the
detections instead, which has the same effect on what ends up stored.
"""
from backend.labeling import _iou
doomed = [row["id"] for row in review.listing(frame_id)
if row["class_id"] == class_id
and any(_iou(review.to_box(row["geometry"]), box) >= NEGATIVE_IOU
for box in negatives)]
return review.delete_many(doomed)
def label(frame_id: int, class_id: int, exemplars: List[dict],
threshold: float = 0.5, iou_threshold: float = 0.8,
min_box_frac: float = 0.002, max_detections: int = 100,
apply: bool = False) -> dict:
"""Detect one class on one frame from the frame's exemplar pool.
`exemplars` is the whole pool, newest last, each
`{"box": [x0, y0, x1, y1], "positive": bool}` normalized to the frame.
Positive boxes are both prompts and labels; negative boxes are prompts and
deletions, never labels.
Nothing is written unless `apply` is set (REQ-175): a drag previews, the
filter panel re-previews, and only Apply touches the frame. Apply re-runs
rather than trusting shapes sent back from the browser — SAM3 is
deterministic for a given pool and threshold, so the second pass reproduces
what was previewed.
"""
from backend import batches, jobs
from backend.labeling import _iou, deduplicate
from PIL import Image
target = review.frame(frame_id)
if target is None:
raise review.ReviewError("No such frame")
label_type = target["label_type"]
prompt = _class_prompt(target["project_id"], class_id)
positives, negatives = [], []
for item in exemplars:
box = review.validate({"type": "bbox", "points": item["box"]}, "bbox")["points"]
(positives if item.get("positive", True) else negatives).append(box)
if not positives and not negatives:
raise review.ReviewError("No exemplars to run")
# The GPU lock is shared with background jobs. Shorter than `review.assist`'s
# 20s on purpose: this fires from a mouse gesture, so a long stall would feel
# like a hung editor — and the fallback keeps the drawing rather than failing.
if not jobs.gpu_lock.acquire(timeout=5):
busy = jobs.running_types()
kind = busy[0] if busy else "background"
drawn = [{"geometry": _rect(box, label_type), "score": 1.0, "source": "manual"}
for box in positives]
if apply:
# Never the replace path here: with no detections to put back, it
# would wipe the class and leave only the drawings. Applying a
# detection-less run just files the drawings and honours the
# negatives.
_append_drawn(frame_id, class_id, drawn)
_drop_negative_overlaps(frame_id, class_id, negatives)
return _result(frame_id, drawn, apply, redetected=False,
message=f"The GPU is busy with a {kind} job — this is your drawing "
"only, nothing was detected")
try:
from backend.sam3_engine import get_engine
path = batches.frame_path(frame_id)
with Image.open(path) as handle:
image = handle.convert("RGB")
width, height = image.size
engine = get_engine()
state = engine.open_state(image)
found = engine.apply_prompts(
state, threshold=threshold, text=prompt,
exemplars=[{"box": _cxcywh(box), "positive": True} for box in positives]
+ [{"box": _cxcywh(box), "positive": False} for box in negatives],
)
finally:
jobs.gpu_lock.release()
# The panel's filters, in the order the batch job applies them (REQ-175):
# area floor, then NMS, then the cap on how many survive.
if min_box_frac > 0:
floor = width * height * min_box_frac
found = [d for d in found
if (d.box[2] - d.box[0]) * (d.box[3] - d.box[1]) >= floor]
found = deduplicate(found, iou_threshold)
found.sort(key=lambda d: d.score, reverse=True)
if max_detections > 0:
found = found[:max_detections]
detections = [(_norm_box(d.box, width, height), d) for d in found]
items: List[dict] = []
# The user's own boxes first, so the duplicate check below measures against
# what they drew rather than the other way round.
for box in positives:
geometry = _rect(box, label_type)
if label_type != "bbox":
snapped = _snap(box, detections, width, height)
if snapped is not None:
geometry = snapped
items.append({"geometry": geometry, "score": 1.0, "source": "manual"})
for norm, detection in detections:
if any(_iou(norm, box) >= NEGATIVE_IOU for box in negatives):
continue
if any(_iou(norm, box) >= DUPLICATE_IOU for box in positives):
continue
for geometry in _detection_shapes(detection, norm, width, height, label_type):
items.append({"geometry": geometry, "score": detection.score, "source": "auto"})
if apply:
_replace_class(frame_id, class_id, items)
return _result(frame_id, items, apply, redetected=True, message=None)
def _result(frame_id: int, items: List[dict], applied: bool,
redetected: bool, message: Optional[str]) -> dict:
"""A preview carries the shapes; an apply also carries the frame as stored."""
return {
"shapes": items,
"applied": applied,
"redetected": redetected,
"message": message,
"annotations": review.listing(frame_id) if applied else None,
}
def _norm_box(box: List[float], width: int, height: int) -> List[float]:
return [box[0] / width, box[1] / height, box[2] / width, box[3] / height]
def _snap(drawn: List[float], detections, width: int, height: int) -> Optional[dict]:
"""The mask polygon of whatever SAM3 found inside a drawn box.
A rectangle is a bad polygon label, so in a polygon project the drag is a
prompt for the shape rather than the shape itself (REQ-173).
"""
from backend.labeling import _iou
best = None
best_iou = SNAP_IOU
for norm, detection in detections:
if detection.mask is None:
continue
overlap = _iou(norm, drawn)
if overlap >= best_iou:
best, best_iou = detection, overlap
if best is None:
return None
points = review.mask_to_polygons(best.mask)
if not points or len(points[0]) < 3:
return None
return review.polygon([(x / width, y / height) for x, y in points[0]])
def _detection_shapes(detection, norm: List[float], width: int, height: int,
label_type: str) -> List[dict]:
if label_type == "bbox" or detection.mask is None:
return [review.bbox(*norm)]
return [review.polygon([(x / width, y / height) for x, y in points])
for points in review.mask_to_polygons(detection.mask)
if len(points) >= 3]
+12 -2
View File
@@ -68,8 +68,13 @@ def label_image(
threshold: float,
iou_threshold: float = 0.8,
min_box_frac: float = 0.0,
exemplar_index: int = -1,
exemplars: Optional[List[dict]] = None,
) -> ImageResult:
"""Detect every prompt in one image and return the surviving instances."""
"""Detect every prompt in one image and return the surviving instances.
When `exemplars` are given, the prompt at `exemplar_index` also carries them
as drawn box exemplars (REQ-172); every other prompt runs on text alone."""
try:
image = Image.open(image_path).convert("RGB")
except Exception as exc: # unreadable/corrupt frame: report, don't abort the job
@@ -77,7 +82,12 @@ def label_image(
width, height = image.size
try:
detections = get_engine().detect(image, prompts, threshold)
if exemplars and 0 <= exemplar_index < len(prompts):
detections = get_engine().detect_with_exemplars(
image, prompts, threshold, exemplar_index, exemplars
)
else:
detections = get_engine().detect(image, prompts, threshold)
except Exception as exc:
return ImageResult(image_path, rel_path, width, height, error=str(exc))
+54 -128
View File
@@ -1,4 +1,4 @@
"""Live counting test bench: point a trained model at an RTSP stream and watch it count.
"""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
@@ -22,6 +22,10 @@ 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.*`.
@@ -31,10 +35,6 @@ for candidate in ("/app", os.path.join(_REPO, "algoritma-batch")):
sys.path.insert(0, candidate)
class LiveCountError(Exception):
pass
class Session:
"""One running counter. Owns a capture thread and the latest rendered frame."""
@@ -43,8 +43,11 @@ class Session:
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):
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
@@ -75,6 +78,7 @@ class Session:
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)
@@ -92,6 +96,10 @@ class Session:
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.
@@ -119,6 +127,8 @@ class Session:
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,
@@ -227,7 +237,11 @@ class Session:
self.fps = since / (now - tick)
tick, since = now, 0
self._render(frame, inside, outside, counter)
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:
@@ -263,59 +277,33 @@ class Session:
if not self.error:
self.error = f"trace write failed: {exc}"
def _render(self, frame, detections, ignored, counter) -> None:
height, width = frame.shape[:2]
# Shade what the region excludes. Without this the neighbouring truck's
# sacks simply vanish from the overlay, and "are they being ignored?"
# looks identical to "is the model missing them?".
if self.line_x_start > 0 or self.line_x_end < width:
shade = frame.copy()
if self.line_x_start > 0:
cv2.rectangle(shade, (0, 0), (self.line_x_start, height), (0, 0, 0), -1)
if self.line_x_end < width:
cv2.rectangle(shade, (self.line_x_end, 0), (width, height), (0, 0, 0), -1)
cv2.addWeighted(shade, 0.55, frame, 0.45, 0, frame)
# Ignored detections stay visible, in grey, so the region can be judged.
for det in ignored:
x1, y1, x2, y2 = (int(v) for v in det.bbox)
cv2.rectangle(frame, (x1, y1), (x2, y2), (130, 130, 130), 1)
for edge in (self.line_x_start, self.line_x_end):
if 0 < edge < width:
cv2.line(frame, (edge, 0), (edge, height), (255, 0, 255), 2)
cv2.putText(frame, "IGNORED", (max(4, self.line_x_start - 92), height - 14),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 0, 255), 1, cv2.LINE_AA)
cv2.putText(frame, "IGNORED", (min(width - 88, self.line_x_end + 8), height - 14),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 0, 255), 1, cv2.LINE_AA)
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:
x1, y1, x2, y2 = (int(v) for v in det.bbox)
counted = bool(counter.counted_tracks.get(det.track_id))
colour = (74, 222, 128) if counted else (248, 191, 113)
cv2.rectangle(frame, (x1, y1), (x2, y2), colour, 2)
cv2.putText(frame, f"#{det.track_id} {det.confidence:.2f}", (x1, max(14, y1 - 6)),
cv2.FONT_HERSHEY_SIMPLEX, 0.45, colour, 1, cv2.LINE_AA)
cv2.line(frame, (self.line_x_start, self.line_y), (self.line_x_end, self.line_y),
(0, 255, 255), 2)
for edge in (self.line_y - self.margin, self.line_y + self.margin):
cv2.line(frame, (self.line_x_start, edge), (self.line_x_end, edge),
(0, 160, 160), 1)
panel = f"IN {self.loading} OUT {self.unloading} NET {self.loading - self.unloading}"
cv2.rectangle(frame, (12, 12), (12 + 9 * len(panel) + 20, 84), (0, 0, 0), -1)
cv2.putText(frame, panel, (24, 46), cv2.FONT_HERSHEY_SIMPLEX, 0.8,
(74, 222, 128), 2, cv2.LINE_AA)
cv2.putText(frame, f"{self.fps:.1f} fps {self.tracked} tracked {self.ignored} ignored",
(24, 72), cv2.FONT_HERSHEY_SIMPLEX, 0.55, (200, 200, 200), 1, cv2.LINE_AA)
ok, buffer = cv2.imencode(".jpg", frame, [cv2.IMWRITE_JPEG_QUALITY, 75])
if ok:
with self._lock:
self._jpeg = buffer.tobytes()
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
@@ -339,74 +327,6 @@ def _too_small(bbox, scale: float) -> bool:
return (x2 - x1) * (y2 - y1) < minimum * scale
def _is_stream(source: str) -> bool:
return str(source).startswith(("rtsp://", "rtmp://", "http://", "https://"))
class _ThreadedStream:
"""Decode in a background thread and always hand out the newest frame.
A plain VideoCapture.read() on RTSP is blocking, and decoding 1080p costs
more than inference does — measured here at 6.7 fps end-to-end against 200
fps for the model itself. Worse, reading slower than the camera sends builds
a backlog, so the picture drifts further behind real time the longer it
runs. Dropping stale frames keeps latency flat, which is what a counting
test needs to mean anything. `predict.py` does the same thing.
"""
def __init__(self, source: str):
self._capture = cv2.VideoCapture(source)
try:
self._capture.set(cv2.CAP_PROP_BUFFERSIZE, 1)
except Exception:
pass
self._frame = None
self._lock = threading.Lock()
self._running = True
self._thread = threading.Thread(target=self._pump, daemon=True)
if self._capture.isOpened():
self._thread.start()
def _pump(self) -> None:
while self._running:
ok, frame = self._capture.read()
if not ok:
time.sleep(0.01)
continue
with self._lock:
self._frame = frame
def isOpened(self) -> bool:
return self._capture.isOpened()
def read(self):
with self._lock:
if self._frame is None:
return False, None
frame, self._frame = self._frame, None
return True, frame
def release(self) -> None:
self._running = False
# The thread is only started when the capture opened, so a failed
# source would otherwise raise "cannot join thread before it is
# started" here — inside the caller's finally, skipping the GPU lock
# release and wedging every later session on "GPU busy".
if self._thread.is_alive():
self._thread.join(timeout=2)
self._capture.release()
def _open(source: str):
if _is_stream(source):
os.environ.setdefault(
"OPENCV_FFMPEG_CAPTURE_OPTIONS",
"rtsp_transport;tcp|buffer_size;20480000|max_delay;500000",
)
return _ThreadedStream(source)
return cv2.VideoCapture(source)
# ---- module-level single session ----------------------------------------
_session: Optional[Session] = None
@@ -417,7 +337,8 @@ def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_e
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) -> dict:
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"]:
@@ -429,7 +350,7 @@ def start(source: str, model_path: str, line_y: int, line_x_start: int, line_x_e
_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)
spatial_dedup, whep)
_session.start()
time.sleep(0.4) # let an immediate failure surface in the response
return _session.status()
@@ -459,9 +380,14 @@ def status() -> dict:
"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 {}
+60
View File
@@ -0,0 +1,60 @@
"""Burning the counting overlay into the frame, as an MJPEG.
Split out of `live_count.py` to keep it inside the 400-line limit. This is the
fallback preview, used for archive files; a WebRTC session draws the same
geometry on a canvas in the browser instead (REQ-177).
"""
import cv2
def render(session, frame, detections, ignored, counter):
height, width = frame.shape[:2]
# Shade what the region excludes. Without this the neighbouring truck's
# sacks simply vanish from the overlay, and "are they being ignored?"
# looks identical to "is the model missing them?".
if session.line_x_start > 0 or session.line_x_end < width:
shade = frame.copy()
if session.line_x_start > 0:
cv2.rectangle(shade, (0, 0), (session.line_x_start, height), (0, 0, 0), -1)
if session.line_x_end < width:
cv2.rectangle(shade, (session.line_x_end, 0), (width, height), (0, 0, 0), -1)
cv2.addWeighted(shade, 0.55, frame, 0.45, 0, frame)
# Ignored detections stay visible, in grey, so the region can be judged.
for det in ignored:
x1, y1, x2, y2 = (int(v) for v in det.bbox)
cv2.rectangle(frame, (x1, y1), (x2, y2), (130, 130, 130), 1)
for edge in (session.line_x_start, session.line_x_end):
if 0 < edge < width:
cv2.line(frame, (edge, 0), (edge, height), (255, 0, 255), 2)
cv2.putText(frame, "IGNORED", (max(4, session.line_x_start - 92), height - 14),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 0, 255), 1, cv2.LINE_AA)
cv2.putText(frame, "IGNORED", (min(width - 88, session.line_x_end + 8), height - 14),
cv2.FONT_HERSHEY_SIMPLEX, 0.5, (255, 0, 255), 1, cv2.LINE_AA)
for det in detections:
x1, y1, x2, y2 = (int(v) for v in det.bbox)
counted = bool(counter.counted_tracks.get(det.track_id))
colour = (74, 222, 128) if counted else (248, 191, 113)
cv2.rectangle(frame, (x1, y1), (x2, y2), colour, 2)
cv2.putText(frame, f"#{det.track_id} {det.confidence:.2f}", (x1, max(14, y1 - 6)),
cv2.FONT_HERSHEY_SIMPLEX, 0.45, colour, 1, cv2.LINE_AA)
cv2.line(frame, (session.line_x_start, session.line_y), (session.line_x_end, session.line_y),
(0, 255, 255), 2)
for edge in (session.line_y - session.margin, session.line_y + session.margin):
cv2.line(frame, (session.line_x_start, edge), (session.line_x_end, edge),
(0, 160, 160), 1)
panel = f"IN {session.loading} OUT {session.unloading} NET {session.loading - session.unloading}"
cv2.rectangle(frame, (12, 12), (12 + 9 * len(panel) + 20, 84), (0, 0, 0), -1)
cv2.putText(frame, panel, (24, 46), cv2.FONT_HERSHEY_SIMPLEX, 0.8,
(74, 222, 128), 2, cv2.LINE_AA)
cv2.putText(frame, f"{session.fps:.1f} fps {session.tracked} tracked {session.ignored} ignored",
(24, 72), cv2.FONT_HERSHEY_SIMPLEX, 0.55, (200, 200, 200), 1, cv2.LINE_AA)
ok, buffer = cv2.imencode(".jpg", frame, [cv2.IMWRITE_JPEG_QUALITY, 75])
return buffer.tobytes() if ok else None
+131
View File
@@ -0,0 +1,131 @@
"""Opening a live source, and the URL algebra around it.
Split out of `live_count.py` so that file stays inside the 400-line limit: this
is the transport layer (what a "source" is, how it is opened, how a WebRTC URL
maps to the RTSP leg of the same stream), not the counting logic.
"""
import os
import threading
import time
import cv2
class LiveCountError(Exception):
pass
def _is_stream(source: str) -> bool:
return str(source).startswith(("rtsp://", "rtmp://", "http://", "https://"))
# MediaMTX fans one camera out to several protocols on one host: WHEP for
# browsers, RTSP for decoders. The ports are the server's, not ours, so they
# come from the environment rather than the code (REQ-176).
WHEP_PATH = os.environ.get("MEDIAMTX_WHEP_PATH", "/whep")
RTSP_PORT = os.environ.get("MEDIAMTX_RTSP_PORT", "8554")
def whep_url(source: str) -> str:
"""The URL a browser POSTs its SDP offer to, for a WebRTC stream URL."""
trimmed = source.rstrip("/")
return trimmed if trimmed.endswith(WHEP_PATH) else trimmed + WHEP_PATH
def whep_to_rtsp(source: str) -> str:
"""`http://host:8889/cam` -> `rtsp://host:8554/cam`.
The browser and the counter watch the same camera, but not over the same
protocol: WebRTC is what makes the *preview* cheap, while pulling it into
Python would add ICE and a jitter buffer on top of exactly the same H.264
decode RTSP already does. So the AI counts from the RTSP leg of the same
MediaMTX path — one ingest, two consumers.
"""
from urllib.parse import urlparse
parsed = urlparse(source)
if parsed.scheme not in ("http", "https") or not parsed.hostname:
raise LiveCountError(
"A live source must be a WebRTC (WHEP) URL, e.g. http://host:8889/cam")
path = parsed.path.rstrip("/")
if path.endswith(WHEP_PATH):
path = path[: -len(WHEP_PATH)]
if not path.strip("/"):
raise LiveCountError(f"No stream path in {source} — expected e.g. http://host:8889/cam")
return f"rtsp://{parsed.hostname}:{RTSP_PORT}{path}"
class _ThreadedStream:
"""Decode in a background thread and always hand out the newest frame.
A plain VideoCapture.read() on RTSP is blocking, and getting the frames
there costs far more than looking at them does — measured on this camera at
8.3 fps for decode alone against 205 fps for tracking and inference. Worse, reading slower than the camera sends builds
a backlog, so the picture drifts further behind real time the longer it
runs. Dropping stale frames keeps latency flat, which is what a counting
test needs to mean anything. `predict.py` does the same thing.
"""
def __init__(self, source: str):
self._capture = cv2.VideoCapture(source)
try:
self._capture.set(cv2.CAP_PROP_BUFFERSIZE, 1)
except Exception:
pass
self._frame = None
self._lock = threading.Lock()
self._running = True
self._thread = threading.Thread(target=self._pump, daemon=True)
if self._capture.isOpened():
self._thread.start()
def _pump(self) -> None:
while self._running:
ok, frame = self._capture.read()
if not ok:
time.sleep(0.01)
continue
with self._lock:
self._frame = frame
def isOpened(self) -> bool:
return self._capture.isOpened()
def read(self):
with self._lock:
if self._frame is None:
return False, None
frame, self._frame = self._frame, None
return True, frame
def release(self) -> None:
self._running = False
# The thread is only started when the capture opened, so a failed
# source would otherwise raise "cannot join thread before it is
# started" here — inside the caller's finally, skipping the GPU lock
# release and wedging every later session on "GPU busy".
if self._thread.is_alive():
self._thread.join(timeout=2)
self._capture.release()
# TCP, and it is worth writing down why, because the obvious reasoning gives the
# wrong answer here. The link to the streaming server is a ZeroTier VPN with 15%
# packet loss and a 41-104 ms round trip, which is exactly the case where UDP is
# supposed to win — and with the ffmpeg CLI it does, 16 fps against 6.8. Through
# OpenCV it loses badly: measured over 30 s on this camera, tcp 6.4 fps, udp with
# a socket buffer 4.4, bare udp 2.4. OpenCV drops what it cannot reassemble
# rather than showing it, so the loss becomes missing frames. Left overridable
# because on a clean LAN the answer flips back.
RTSP_TRANSPORT = os.environ.get("RTSP_TRANSPORT", "tcp")
def _open(source: str):
if _is_stream(source):
os.environ.setdefault(
"OPENCV_FFMPEG_CAPTURE_OPTIONS",
f"rtsp_transport;{RTSP_TRANSPORT}|buffer_size;1048576|max_delay;200000",
)
return _ThreadedStream(source)
return cv2.VideoCapture(source)
+151
View File
@@ -0,0 +1,151 @@
"""The auto-annotate preview: one frame, run now, nothing written (REQ-171/172).
Split out of `autolabel.py` to keep that file under the 400-line limit. The job
path and this path share `labeling.label_image`, so what the preview shows is
what a batch run would write — with one deliberate exception: drawn box
exemplars (REQ-172) only ever apply here. SAM3's geometric prompts pool features
from the current image, so replaying them on another frame would ask about
whatever happens to sit at those coordinates there.
"""
import os
from typing import List, Optional
from backend import batches, db, labeling, projects, review
from backend.autolabel import DEFAULT_IOU, DEFAULT_THRESHOLD, _geometries
def preview_frame(
batch_id: int,
frame_id: int,
engine: str,
threshold: float = DEFAULT_THRESHOLD,
iou_threshold: float = DEFAULT_IOU,
min_box_frac: float = 0.0,
target_class_names: Optional[List[str]] = None,
custom_model_path: Optional[str] = None,
exemplars: Optional[List[dict]] = None,
exemplar_class_name: Optional[str] = None,
) -> List[dict]:
batch = batches.get(batch_id)
if not batch:
raise ValueError("No such batch")
project = projects.get(batch["project_id"])
frame = next((f for f in batches.frames(batch_id) if f["id"] == frame_id), None)
if not frame:
raise ValueError("Frame not found")
directory = batches.frames_dir(batch["project_slug"], batch_id)
frame_file = os.path.join(directory, frame["filename"])
fw = max(1, frame.get("width") or 1)
fh = max(1, frame.get("height") or 1)
yolo_model = None
sam3_target_classes = []
if engine == "sam3" and not custom_model_path:
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
if allowed_classes_set:
sam3_target_classes = [c for c in project["classes"] if c["name"].strip().lower() in allowed_classes_set or c["prompt"].strip().lower() in allowed_classes_set]
else:
sam3_target_classes = [c for c in project["classes"]]
prompts = [c["prompt"] for c in sam3_target_classes]
if prompts:
from backend.sam3_engine import get_engine
get_engine()
else:
from ultralytics import YOLO
if custom_model_path and os.path.isfile(custom_model_path):
m_path = custom_model_path
else:
m_path = projects.training_start_point(project)
with db.cursor() as cur:
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
row = cur.fetchone()
if row and os.path.isfile(row[0]):
m_path = row[0]
yolo_model = YOLO(m_path)
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
all_raw_detections = []
if yolo_model is not None:
results = yolo_model.predict(frame_file, conf=threshold, verbose=False)
if results and len(results) > 0:
model_names = results[0].names
for box in results[0].boxes:
cls_idx = int(box.cls[0].item())
raw_cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
target_class_id = name_to_class_id.get(raw_cls_name)
if target_class_id is None:
for item in project["classes"]:
if item["class_id"] == cls_idx:
target_class_id = item["class_id"]
break
if target_class_id is None and 0 <= cls_idx < len(project["classes"]):
target_class_id = project["classes"][cls_idx]["class_id"]
if target_class_id is None:
continue
target_cls_obj = next((c for c in project["classes"] if c["class_id"] == target_class_id), None)
proj_cls_name = target_cls_obj["name"].strip().lower() if target_cls_obj else ""
if allowed_classes_set is not None:
if (raw_cls_name not in allowed_classes_set and
proj_cls_name not in allowed_classes_set and
str(target_class_id) not in allowed_classes_set):
continue
score = float(box.conf[0].item())
xyxyn = box.xyxyn[0].tolist()
all_raw_detections.append(labeling.Detection(
class_id=target_class_id,
class_name=proj_cls_name or raw_cls_name,
box=[xyxyn[0]*fw, xyxyn[1]*fh, xyxyn[2]*fw, xyxyn[3]*fh],
score=score,
mask=None
))
if engine == "sam3" and sam3_target_classes:
prompts = [(c.get("prompt") or c["name"]).strip() for c in sam3_target_classes]
# Exemplars belong to exactly one class — the chip that was active when
# they were drawn. An unknown name means no exemplar class, so the run
# falls back to plain text rather than silently attaching the boxes to
# whichever class happens to be first.
exemplar_index = -1
if exemplars and exemplar_class_name:
wanted = exemplar_class_name.strip().lower()
exemplar_index = next(
(i for i, c in enumerate(sam3_target_classes)
if c["name"].strip().lower() == wanted),
-1,
)
res = labeling.label_image(
frame_file, frame["filename"], prompts, threshold,
iou_threshold=iou_threshold, min_box_frac=min_box_frac,
exemplar_index=exemplar_index, exemplars=exemplars
)
if not res.error and res.detections:
for det in res.detections:
if 0 <= det.class_id < len(sam3_target_classes):
real_cls = sam3_target_classes[det.class_id]
det.class_id = real_cls["class_id"]
det.class_name = real_cls["name"]
all_raw_detections.append(det)
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_threshold)
items = []
for det in kept:
if project["label_type"] == "bbox" or det.mask is None:
geom = review.bbox(det.box[0]/fw, det.box[1]/fh, det.box[2]/fw, det.box[3]/fh)
items.append({"class_id": det.class_id, "geometry": geom, "score": det.score})
else:
for geometry in _geometries(det, fw, fh, project["label_type"]):
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
return items
+35
View File
@@ -103,6 +103,41 @@ class Sam3Engine:
def detect_with_exemplars(
self,
image: Image.Image,
prompts: List[str],
threshold: float,
exemplar_index: int,
exemplars: List[dict],
) -> List[Detection]:
"""`detect()`, but one prompt also carries drawn box exemplars (REQ-172).
Still one `set_image` for the whole call. The prompt set is reset before
every class because `state["geometric_prompt"]` survives `set_text_prompt`
— without the reset, one class's boxes would leak into the next class.
"""
processor = Sam3Processor(self.model, device=self.device)
processor.confidence_threshold = threshold
detections: List[Detection] = []
with torch.autocast(self.device, dtype=self.autocast_dtype):
state = processor.set_image(image)
for class_id, prompt in enumerate(prompts):
processor.reset_all_prompts(state)
output = processor.set_text_prompt(prompt=prompt, state=state)
if class_id == exemplar_index:
for exemplar in exemplars:
output = processor.add_geometric_prompt(
box=exemplar["box"],
label=bool(exemplar.get("positive", True)),
state=state,
)
detections.extend(self._collect(output, class_id, prompt))
del state
return detections
# ---- interactive / exemplar prompting ------------------------------