feat: sync live-count, exemplar annotation modules, and update .gitignore
This commit is contained in:
1 parent
b6624eeff9
commit
ac95674c07
39 files changed
+3979
-444
No files matched your search
+14
-3
@@ -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:
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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.
@@ -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
|
||||
@@ -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
@@ -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
@@ -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 {}
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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 ------------------------------
|
||||
|
||||
|
||||
Reference in new issue
Block a user