This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
93 lines
3.0 KiB
Python
93 lines
3.0 KiB
Python
"""Run every class prompt against one frame and return the surviving instances.
|
|
|
|
Each class is its own prompt, so prompt index is class id. Prompts overlap in
|
|
practice ("sack" and "woven plastic sack" both fire on the same object), so
|
|
detections are deduplicated across prompts by IoU, keeping the higher-scoring
|
|
one (REQ-031).
|
|
|
|
The set_image-once-per-image rule lives in `sam3_engine.detect`, which this
|
|
calls — see the domain invariants in `../AGENTS.md`.
|
|
"""
|
|
|
|
from dataclasses import dataclass, field
|
|
from typing import List, Optional
|
|
|
|
from PIL import Image
|
|
|
|
from backend.sam3_engine import Detection, get_engine
|
|
|
|
|
|
@dataclass
|
|
class ImageResult:
|
|
image_path: str
|
|
rel_path: str
|
|
width: int
|
|
height: int
|
|
detections: List[Detection] = field(default_factory=list)
|
|
error: Optional[str] = None
|
|
|
|
|
|
def _iou(box_a: List[float], box_b: List[float]) -> float:
|
|
ax0, ay0, ax1, ay1 = box_a
|
|
bx0, by0, bx1, by1 = box_b
|
|
inter_w = max(0.0, min(ax1, bx1) - max(ax0, bx0))
|
|
inter_h = max(0.0, min(ay1, by1) - max(ay0, by0))
|
|
inter = inter_w * inter_h
|
|
if inter <= 0:
|
|
return 0.0
|
|
area_a = max(0.0, ax1 - ax0) * max(0.0, ay1 - ay0)
|
|
area_b = max(0.0, bx1 - bx0) * max(0.0, by1 - by0)
|
|
union = area_a + area_b - inter
|
|
return inter / union if union > 0 else 0.0
|
|
|
|
|
|
def deduplicate(detections: List[Detection], iou_threshold: float = 0.8) -> List[Detection]:
|
|
"""Greedy NMS per class: highest score wins within the SAME class."""
|
|
if iou_threshold <= 0.0:
|
|
return detections
|
|
|
|
by_class: dict[int, List[Detection]] = {}
|
|
|
|
for det in detections:
|
|
by_class.setdefault(det.class_id, []).append(det)
|
|
|
|
kept: List[Detection] = []
|
|
for cls_dets in by_class.values():
|
|
cls_kept: List[Detection] = []
|
|
for det in sorted(cls_dets, key=lambda d: d.score, reverse=True):
|
|
if all(_iou(det.box, k.box) < iou_threshold for k in cls_kept):
|
|
cls_kept.append(det)
|
|
kept.extend(cls_kept)
|
|
return kept
|
|
|
|
|
|
def label_image(
|
|
image_path: str,
|
|
rel_path: str,
|
|
prompts: List[str],
|
|
threshold: float,
|
|
iou_threshold: float = 0.8,
|
|
min_box_frac: float = 0.0,
|
|
) -> ImageResult:
|
|
"""Detect every prompt in one image and return the surviving instances."""
|
|
try:
|
|
image = Image.open(image_path).convert("RGB")
|
|
except Exception as exc: # unreadable/corrupt frame: report, don't abort the job
|
|
return ImageResult(image_path, rel_path, 0, 0, error=str(exc))
|
|
|
|
width, height = image.size
|
|
try:
|
|
detections = get_engine().detect(image, prompts, threshold)
|
|
except Exception as exc:
|
|
return ImageResult(image_path, rel_path, width, height, error=str(exc))
|
|
|
|
if min_box_frac > 0:
|
|
floor = width * height * min_box_frac
|
|
detections = [
|
|
d for d in detections
|
|
if (d.box[2] - d.box[0]) * (d.box[3] - d.box[1]) >= floor
|
|
]
|
|
|
|
return ImageResult(image_path, rel_path, width, height,
|
|
deduplicate(detections, iou_threshold))
|