"""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))