103 lines
3.4 KiB
Python
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))
|