Files

103 lines
3.4 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,
exemplar_index: int = -1,
exemplars: Optional[List[dict]] = None,
) -> ImageResult:
"""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
return ImageResult(image_path, rel_path, 0, 0, error=str(exc))
width, height = image.size
try:
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))
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))