401 lines
14 KiB
Python
401 lines
14 KiB
Python
"""Annotations and per-frame review state (REQ-040…045).
|
||
|
||
Geometry is stored normalized 0–1 against the frame, as JSON:
|
||
|
||
bbox {"type": "bbox", "points": [x0, y0, x1, y1]}
|
||
polygon {"type": "polygon", "points": [[x, y], …]}
|
||
|
||
Normalized because the editor scales the frame to whatever the window allows,
|
||
and the exporter needs the same numbers YOLO wants — neither should care about
|
||
the display size.
|
||
|
||
`source` separates what SAM3 produced from what the user drew. Re-running
|
||
auto-annotation replaces only the former (REQ-034).
|
||
"""
|
||
|
||
import json
|
||
import time
|
||
from typing import List, Optional
|
||
|
||
from backend import db
|
||
|
||
STATUSES = ("pending", "approved", "rejected")
|
||
|
||
|
||
class ReviewError(Exception):
|
||
pass
|
||
|
||
|
||
# ---- geometry ----------------------------------------------------------
|
||
|
||
def _clamp(value: float) -> float:
|
||
return max(0.0, min(1.0, float(value)))
|
||
|
||
|
||
def bbox(x0: float, y0: float, x1: float, y1: float) -> dict:
|
||
left, right = sorted((_clamp(x0), _clamp(x1)))
|
||
top, bottom = sorted((_clamp(y0), _clamp(y1)))
|
||
return {"type": "bbox", "points": [left, top, right, bottom]}
|
||
|
||
|
||
def polygon(points) -> dict:
|
||
return {"type": "polygon", "points": [[_clamp(x), _clamp(y)] for x, y in points]}
|
||
|
||
|
||
def mask_to_polygons(mask, min_area_px: int = 24, max_polygons: int = 1) -> list:
|
||
"""Contour a boolean SAM3 mask into polygon point arrays, largest first.
|
||
|
||
YOLO-seg expects one polygon per instance, so only the largest connected
|
||
component is kept by default — SAM3 masks are occasionally speckled. The raw
|
||
contour is one point per pixel step, which would make label files enormous
|
||
for no accuracy gain, so it is simplified first.
|
||
"""
|
||
import cv2
|
||
import numpy as np
|
||
|
||
contours, _ = cv2.findContours((mask.astype(np.uint8)) * 255,
|
||
cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
|
||
polygons = []
|
||
for contour in sorted(contours, key=cv2.contourArea, reverse=True)[:max_polygons]:
|
||
if cv2.contourArea(contour) < min_area_px:
|
||
continue
|
||
epsilon = 0.002 * cv2.arcLength(contour, True)
|
||
approx = cv2.approxPolyDP(contour, epsilon, True).reshape(-1, 2)
|
||
if approx.shape[0] >= 3:
|
||
polygons.append(approx.astype(np.float32))
|
||
return polygons
|
||
|
||
|
||
def validate(geometry: dict, label_type: str) -> dict:
|
||
"""Reject shapes that would export as broken labels."""
|
||
if not isinstance(geometry, dict):
|
||
raise ReviewError("geometry must be an object")
|
||
kind = geometry.get("type")
|
||
points = geometry.get("points") or []
|
||
|
||
if kind == "bbox":
|
||
if len(points) != 4:
|
||
raise ReviewError("a bbox needs [x0, y0, x1, y1]")
|
||
shape = bbox(*points)
|
||
left, top, right, bottom = shape["points"]
|
||
if right - left < 0.002 or bottom - top < 0.002:
|
||
raise ReviewError("that box is too small to be a label")
|
||
return shape
|
||
|
||
if kind == "polygon":
|
||
if len(points) < 3:
|
||
raise ReviewError("a polygon needs at least 3 points")
|
||
if label_type == "bbox":
|
||
raise ReviewError("this project stores boxes, not polygons")
|
||
return polygon(points)
|
||
|
||
raise ReviewError(f"unknown geometry type: {kind}")
|
||
|
||
|
||
def to_box(geometry: dict) -> List[float]:
|
||
"""The bounding box of any shape, normalized — used for the bbox export
|
||
and for deduplicating polygon detections against each other."""
|
||
points = geometry["points"]
|
||
if geometry["type"] == "bbox":
|
||
return list(points)
|
||
xs = [point[0] for point in points]
|
||
ys = [point[1] for point in points]
|
||
return [min(xs), min(ys), max(xs), max(ys)]
|
||
|
||
|
||
# ---- frames ------------------------------------------------------------
|
||
|
||
def frame(frame_id: int) -> Optional[dict]:
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""SELECT f.*, b.id AS batch_id, b.project_id, p.slug AS project_slug,
|
||
p.label_type
|
||
FROM frames f
|
||
JOIN batches b ON b.id = f.batch_id
|
||
JOIN projects p ON p.id = b.project_id
|
||
WHERE f.id = ?""",
|
||
(frame_id,),
|
||
)
|
||
row = cur.fetchone()
|
||
return dict(row) if row else None
|
||
|
||
|
||
def set_status(frame_id: int, status: str) -> dict:
|
||
if status not in STATUSES:
|
||
raise ReviewError(f"status must be one of {STATUSES}")
|
||
if frame(frame_id) is None:
|
||
raise ReviewError("No such frame")
|
||
with db.cursor() as cur:
|
||
cur.execute("UPDATE frames SET review_status = ? WHERE id = ?", (status, frame_id))
|
||
return {"frame_id": frame_id, "review_status": status}
|
||
|
||
|
||
def next_pending(batch_id: int, after_idx: int = -1) -> Optional[int]:
|
||
"""The next frame still needing a decision, for the 'jump to unreviewed' key."""
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""SELECT id FROM frames
|
||
WHERE batch_id = ? AND review_status = 'pending' AND idx > ?
|
||
ORDER BY idx LIMIT 1""",
|
||
(batch_id, after_idx),
|
||
)
|
||
row = cur.fetchone()
|
||
if row:
|
||
return row["id"]
|
||
cur.execute(
|
||
"""SELECT id FROM frames WHERE batch_id = ? AND review_status = 'pending'
|
||
ORDER BY idx LIMIT 1""",
|
||
(batch_id,),
|
||
)
|
||
row = cur.fetchone()
|
||
return row["id"] if row else None
|
||
|
||
|
||
# ---- annotations -------------------------------------------------------
|
||
|
||
def _row_to_dict(row) -> dict:
|
||
return {
|
||
"id": row["id"],
|
||
"frame_id": row["frame_id"],
|
||
"class_id": row["class_id"],
|
||
"geometry": json.loads(row["geometry"]),
|
||
"score": row["score"],
|
||
"source": row["source"],
|
||
}
|
||
|
||
|
||
def listing(frame_id: int) -> List[dict]:
|
||
with db.cursor() as cur:
|
||
cur.execute("SELECT * FROM annotations WHERE frame_id = ? ORDER BY id", (frame_id,))
|
||
return [_row_to_dict(row) for row in cur.fetchall()]
|
||
|
||
|
||
def add(frame_id: int, class_id: int, geometry: dict, source: str = "manual",
|
||
score: float = 1.0) -> dict:
|
||
target = frame(frame_id)
|
||
if target is None:
|
||
raise ReviewError("No such frame")
|
||
shape = validate(geometry, target["label_type"])
|
||
_check_class(target["project_id"], class_id)
|
||
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
|
||
VALUES (?, ?, ?, ?, ?, ?)""",
|
||
(frame_id, class_id, json.dumps(shape), score, source, time.time()),
|
||
)
|
||
cur.execute("SELECT * FROM annotations WHERE id = ?", (cur.lastrowid,))
|
||
return _row_to_dict(cur.fetchone())
|
||
|
||
|
||
def update(annotation_id: int, class_id: Optional[int] = None,
|
||
geometry: Optional[dict] = None) -> dict:
|
||
with db.cursor() as cur:
|
||
cur.execute("SELECT * FROM annotations WHERE id = ?", (annotation_id,))
|
||
row = cur.fetchone()
|
||
if row is None:
|
||
raise ReviewError("No such annotation")
|
||
target = frame(row["frame_id"])
|
||
|
||
new_geometry = row["geometry"]
|
||
if geometry is not None:
|
||
new_geometry = json.dumps(validate(geometry, target["label_type"]))
|
||
new_class = row["class_id"] if class_id is None else class_id
|
||
_check_class(target["project_id"], new_class)
|
||
|
||
with db.cursor() as cur:
|
||
# Any edit makes it the user's shape, so it survives a re-run of
|
||
# auto-annotation (REQ-034).
|
||
cur.execute(
|
||
"UPDATE annotations SET class_id = ?, geometry = ?, source = 'manual' WHERE id = ?",
|
||
(new_class, new_geometry, annotation_id),
|
||
)
|
||
cur.execute("SELECT * FROM annotations WHERE id = ?", (annotation_id,))
|
||
return _row_to_dict(cur.fetchone())
|
||
|
||
|
||
def delete(annotation_id: int) -> bool:
|
||
with db.cursor() as cur:
|
||
cur.execute("DELETE FROM annotations WHERE id = ?", (annotation_id,))
|
||
return cur.rowcount > 0
|
||
|
||
|
||
def replace_auto(frame_id: int, items: List[dict]) -> int:
|
||
"""Swap this frame's automatic shapes for a fresh set, leaving manual ones."""
|
||
with db.cursor() as cur:
|
||
cur.execute("DELETE FROM annotations WHERE frame_id = ? AND source = 'auto'",
|
||
(frame_id,))
|
||
cur.executemany(
|
||
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
|
||
VALUES (?, ?, ?, ?, 'auto', ?)""",
|
||
[(frame_id, item["class_id"], json.dumps(item["geometry"]),
|
||
item.get("score", 1.0), time.time()) for item in items],
|
||
)
|
||
return len(items)
|
||
|
||
|
||
def append_auto(frame_id: int, items: List[dict]) -> int:
|
||
"""Add new automatic shapes to this frame without duplicating existing ones."""
|
||
if not items:
|
||
return 0
|
||
existing = listing(frame_id)
|
||
filtered_items = []
|
||
for item in items:
|
||
is_dup = False
|
||
item_box = to_box(item["geometry"])
|
||
for ex in existing:
|
||
if ex["class_id"] == item["class_id"]:
|
||
ex_box = to_box(ex["geometry"])
|
||
from backend.labeling import _iou
|
||
if _iou(item_box, ex_box) >= 0.85:
|
||
is_dup = True
|
||
break
|
||
if not is_dup:
|
||
filtered_items.append(item)
|
||
|
||
if not filtered_items:
|
||
return 0
|
||
|
||
with db.cursor() as cur:
|
||
cur.executemany(
|
||
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
|
||
VALUES (?, ?, ?, ?, 'auto', ?)""",
|
||
[(frame_id, item["class_id"], json.dumps(item["geometry"]),
|
||
item.get("score", 1.0), time.time()) for item in filtered_items],
|
||
)
|
||
return len(filtered_items)
|
||
|
||
|
||
def frames_with_auto(batch_id: int) -> set:
|
||
"""Frame ids that already carry automatic shapes — the resume skip-list
|
||
for REQ-035."""
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"SELECT DISTINCT frame_id FROM annotations "
|
||
"WHERE source = 'auto' AND frame_id IN "
|
||
"(SELECT id FROM frames WHERE batch_id = ?)",
|
||
(batch_id,),
|
||
)
|
||
return {row[0] for row in cur.fetchall()}
|
||
|
||
|
||
|
||
def assist(frame_id: int, box: List[float], class_id: int = 0,
|
||
threshold: float = 0.5) -> dict:
|
||
"""Drag a rough box, get SAM3's shape for the object inside it (REQ-043).
|
||
|
||
The box is a visual exemplar rather than a crop: SAM3 may return several
|
||
matches, so the one overlapping what the user drew is the one kept.
|
||
"""
|
||
from backend import batches, jobs
|
||
from backend.sam3_engine import get_engine
|
||
from PIL import Image
|
||
|
||
target = frame(frame_id)
|
||
if target is None:
|
||
raise ReviewError("No such frame")
|
||
_check_class(target["project_id"], class_id)
|
||
|
||
# Acquire the shared GPU lock with a 20s timeout. 20 seconds is chosen so that
|
||
# short CPU/ffmpeg jobs let assist through, while long GPU jobs fail fast with
|
||
# a legible message (REQ-065, REQ-070).
|
||
if not jobs.gpu_lock.acquire(timeout=20):
|
||
busy = jobs.running_types()
|
||
kind = busy[0] if busy else "background"
|
||
raise ReviewError(
|
||
f"The GPU is busy with a {kind} job — wait for it to finish, or draw the "
|
||
"shape by hand"
|
||
)
|
||
|
||
try:
|
||
drawn = validate({"type": "bbox", "points": box}, "bbox")["points"]
|
||
x0, y0, x1, y1 = drawn
|
||
exemplar = [(x0 + x1) / 2, (y0 + y1) / 2, x1 - x0, y1 - y0]
|
||
|
||
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,
|
||
exemplars=[{"box": exemplar, "positive": True}],
|
||
)
|
||
|
||
if not found:
|
||
raise ReviewError("SAM3 found nothing in that box — draw it tighter, or add the "
|
||
"shape by hand")
|
||
|
||
detection = max(found, key=lambda d: _overlap(d.box, drawn, width, height))
|
||
if target["label_type"] == "bbox":
|
||
bx0, by0, bx1, by1 = detection.box
|
||
geometry = bbox(bx0 / width, by0 / height, bx1 / width, by1 / height)
|
||
else:
|
||
polygons = mask_to_polygons(detection.mask)
|
||
if not polygons:
|
||
raise ReviewError("SAM3's mask was too small to turn into a polygon")
|
||
geometry = polygon([(x / width, y / height) for x, y in polygons[0]])
|
||
finally:
|
||
jobs.gpu_lock.release()
|
||
|
||
return add(frame_id, class_id, geometry, source="manual", score=detection.score)
|
||
|
||
|
||
|
||
def _overlap(detection_box: List[float], drawn: List[float],
|
||
width: int, height: int) -> float:
|
||
"""IoU between a pixel-space detection and the normalized box drawn."""
|
||
box = [detection_box[0] / width, detection_box[1] / height,
|
||
detection_box[2] / width, detection_box[3] / height]
|
||
ix0, iy0 = max(box[0], drawn[0]), max(box[1], drawn[1])
|
||
ix1, iy1 = min(box[2], drawn[2]), min(box[3], drawn[3])
|
||
inter = max(0.0, ix1 - ix0) * max(0.0, iy1 - iy0)
|
||
if inter <= 0:
|
||
return 0.0
|
||
area_box = (box[2] - box[0]) * (box[3] - box[1])
|
||
area_drawn = (drawn[2] - drawn[0]) * (drawn[3] - drawn[1])
|
||
return inter / (area_box + area_drawn - inter)
|
||
|
||
|
||
def clear_batch_class_annotations(batch_id: int, class_id: int) -> int:
|
||
"""Delete all annotations matching class_id across all frames in a batch (REQ-046)."""
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""DELETE FROM annotations
|
||
WHERE class_id = ? AND frame_id IN (
|
||
SELECT id FROM frames WHERE batch_id = ?
|
||
)""",
|
||
(class_id, batch_id),
|
||
)
|
||
return cur.rowcount
|
||
|
||
|
||
def clear_batch_auto_annotations(batch_id: int) -> int:
|
||
"""Delete all automatic annotations (source = 'auto') for a batch and reset frame statuses."""
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"""DELETE FROM annotations
|
||
WHERE source = 'auto' AND frame_id IN (
|
||
SELECT id FROM frames WHERE batch_id = ?
|
||
)""",
|
||
(batch_id,),
|
||
)
|
||
deleted = cur.rowcount
|
||
cur.execute(
|
||
"UPDATE frames SET review_status = 'pending' WHERE batch_id = ?",
|
||
(batch_id,),
|
||
)
|
||
return deleted
|
||
|
||
|
||
def _check_class(project_id: int, class_id: int) -> None:
|
||
with db.cursor() as cur:
|
||
cur.execute(
|
||
"SELECT 1 FROM project_classes WHERE project_id = ? AND class_id = ?",
|
||
(project_id, class_id),
|
||
)
|
||
if cur.fetchone() is None:
|
||
raise ReviewError(f"Class {class_id} does not exist in this project")
|
||
|