Files
reTraining/backend/exemplar.py
T

284 lines
12 KiB
Python

"""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:
is_pos = bool(item.get("positive", True))
target_list = positives if is_pos else negatives
if item.get("point") is not None:
px, py = float(item["point"][0]), float(item["point"][1])
target_list.append({"kind": "point", "cxcywh": [px, py, 0.03, 0.03], "point": [px, py]})
elif item.get("box") is not None:
box = review.validate({"type": "bbox", "points": item["box"]}, "bbox")["points"]
target_list.append({"kind": "box", "cxcywh": _cxcywh(box), "box": box})
elif "points" in item and len(item["points"]) == 2:
px, py = float(item["points"][0]), float(item["points"][1])
target_list.append({"kind": "point", "cxcywh": [px, py, 0.03, 0.03], "point": [px, py]})
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(p["box"], label_type), "score": 1.0, "source": "manual"}
for p in positives if p["kind"] == "box"]
if apply:
_append_drawn(frame_id, class_id, drawn)
_drop_negative_overlaps(frame_id, class_id, [n["box"] for n in negatives if n["kind"] == "box"])
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": p["cxcywh"], "positive": True} for p in positives]
+ [{"box": n["cxcywh"], "positive": False} for n 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] = []
# Manual drawn boxes first:
for item in positives:
if item["kind"] != "box":
continue
box = item["box"]
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"})
manual_boxes = [p["box"] for p in positives if p["kind"] == "box"]
negative_boxes = [n["box"] for n in negatives if n["kind"] == "box"]
negative_points = [n["point"] for n in negatives if n["kind"] == "point"]
for norm, detection in detections:
if any(_iou(norm, box) >= NEGATIVE_IOU for box in negative_boxes):
continue
if any(norm[0] <= pt[0] <= norm[2] and norm[1] <= pt[1] <= norm[3] for pt in negative_points):
continue
if any(_iou(norm, box) >= DUPLICATE_IOU for box in manual_boxes):
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]