diff --git a/.gitignore b/.gitignore index 3117ab7..d82e562 100644 --- a/.gitignore +++ b/.gitignore @@ -73,6 +73,11 @@ weights/ *.sqlite *.sqlite3 +# Asset exceptions +!backend/assets/ +!backend/assets/** + + # Training Logs & Caches *.log *.tfevents* diff --git a/backend/api/batches.py b/backend/api/batches.py index 9dfc5cd..392ec0e 100644 --- a/backend/api/batches.py +++ b/backend/api/batches.py @@ -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: diff --git a/backend/api/live_count.py b/backend/api/live_count.py index 243ef20..e8656d3 100644 --- a/backend/api/live_count.py +++ b/backend/api/live_count.py @@ -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: diff --git a/backend/api/review.py b/backend/api/review.py index 007901c..0969dec 100644 --- a/backend/api/review.py +++ b/backend/api/review.py @@ -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: diff --git a/backend/assets/clock_glyphs.npz b/backend/assets/clock_glyphs.npz new file mode 100644 index 0000000..d11a1ff Binary files /dev/null and b/backend/assets/clock_glyphs.npz differ diff --git a/backend/autolabel.py b/backend/autolabel.py index 2657481..0324e8a 100644 --- a/backend/autolabel.py +++ b/backend/autolabel.py @@ -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 diff --git a/backend/exemplar.py b/backend/exemplar.py new file mode 100644 index 0000000..eeae72b --- /dev/null +++ b/backend/exemplar.py @@ -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] diff --git a/backend/labeling.py b/backend/labeling.py index f0ff939..04f7563 100644 --- a/backend/labeling.py +++ b/backend/labeling.py @@ -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)) diff --git a/backend/live_count.py b/backend/live_count.py index 17cf188..d73682d 100644 --- a/backend/live_count.py +++ b/backend/live_count.py @@ -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 {} diff --git a/backend/live_render.py b/backend/live_render.py new file mode 100644 index 0000000..23b9dac --- /dev/null +++ b/backend/live_render.py @@ -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 diff --git a/backend/live_source.py b/backend/live_source.py new file mode 100644 index 0000000..662c3e7 --- /dev/null +++ b/backend/live_source.py @@ -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) diff --git a/backend/preview.py b/backend/preview.py new file mode 100644 index 0000000..70430bb --- /dev/null +++ b/backend/preview.py @@ -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 diff --git a/backend/sam3_engine.py b/backend/sam3_engine.py index f735171..d4e060f 100644 --- a/backend/sam3_engine.py +++ b/backend/sam3_engine.py @@ -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 ------------------------------ diff --git a/docs/GT.xlsx b/docs/GT.xlsx new file mode 100644 index 0000000..0463bbe Binary files /dev/null and b/docs/GT.xlsx differ diff --git a/docs/design.md b/docs/design.md index 9ce6ef0..bec76c9 100644 --- a/docs/design.md +++ b/docs/design.md @@ -123,9 +123,14 @@ box. | `batches.py` | batch lifecycle | **new** | | `review.py` | annotation CRUD, frame status, click-assist | **new** | | `autolabel.py` | the SAM3 job over a whole batch | **new** | +| `preview.py` | one-frame preview for the auto-annotate modal (REQ-171,172) | **new** | +| `exemplar.py` | exemplar-driven labeling in the review editor (REQ-173,174) | **new** | | `dataset.py` | merge into the master dataset, stable split | **new** | | `evaluate.py` | validate base vs new model | **new** | | `hardware.py` | VRAM detection → training defaults | **new** | +| `live_count.py` | the live counting session: capture → track → count | **new** | +| `live_source.py` | what a source is, how it opens, WHEP↔RTSP (REQ-176) | **new** | +| `live_render.py` | the MJPEG overlay, archive-file preview only (REQ-177) | **new** | | `api/` | the FastAPI routes, one module per domain | **new** | Removed: `uploads.py`, `static/index.html`, and the old flow's endpoints. @@ -167,6 +172,9 @@ POST /api/projects/{id}/batches # {rel, start_sec, end_sec, fps} GET /api/batches/{id} # status + review progress (REQ-045) GET /api/batches/{id}/frames # frames + statuses POST /api/batches/{id}/autolabel # {threshold} → job (REQ-030,032,034) +POST /api/batches/{id}/preview # one frame, run now, nothing written; + # +exemplars[] {box:[cx,cy,w,h], positive} + # +exemplar_class_name (REQ-171,172) DELETE /api/batches/{id}/classes/{class_id}/annotations # clear all shapes of class in batch (REQ-046) POST /api/batches/{ids}/approve # one or many, comma-separated → one merge job (REQ-131) GET /api/batches/{ids}/triage/summary # one or many, comma-separated (REQ-130) @@ -180,6 +188,7 @@ POST /api/frames/{id}/annotations # add a manual shape (REQ-042) PATCH /api/annotations/{id} # move/resize/reclass DELETE /api/annotations/{id} POST /api/frames/{id}/assist # click/box → SAM3 shape (REQ-043) +POST /api/frames/{id}/exemplar-label # drawn pool → re-detect one class (REQ-173,174) POST /api/frames/{id}/status # approved | rejected | pending (REQ-041) POST /api/projects/{id}/train # → train job (REQ-060,061,062) @@ -187,11 +196,36 @@ GET /api/projects/{id}/models # versions + metrics (REQ-063,06 GET /api/models/{id}/weights # download best.pt POST /api/models/{id}/promote # make it the project's base model (REQ-064) +GET /api/projects/{id}/live-count/models # weights this project can count with +POST /api/projects/{id}/live-count/start # {source|source_rel, model_path, dials} + # source must be a WHEP URL (REQ-176) +POST /api/live-count/stop +PATCH /api/live-count/line # move the line mid-session +GET /api/live-count/status # counts + preview: "webrtc" | "mjpeg" +GET /api/live-count/overlay # boxes/line/counts for the canvas (REQ-177) +GET /api/live-count/stream # MJPEG; 409 on a WebRTC session (REQ-177) + GET /api/jobs?project_id=… REQ-070,071 GET /api/jobs/{id} POST /api/jobs/{id}/cancel ``` +### Live counting (REQ-176, REQ-177) + +One camera, one ingest on the streaming server, two consumers: + +``` +camera ──▶ MediaMTX ──┬── WHEP :8889/cam/whep ──▶ browser