"""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 Dict, List, Optional, Set 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 _containment_fraction(box_a: List[float], box_b: List[float]) -> float: """Intersection over the smaller box's area: how much of the smaller box sits inside the other one (REQ-184's containment test).""" 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)) area_a = max(0.0, ax1 - ax0) * max(0.0, ay1 - ay0) area_b = max(0.0, bx1 - bx0) * max(0.0, by1 - by0) small = min(area_a, area_b) return inter_w * inter_h / small if small > 0 else 0.0 def deduplicate(detections: List[Detection], iou_threshold: float = 0.8, iou_by_class: Optional[Dict[int, float]] = None, container_ids: Optional[Set[int]] = None) -> List[Detection]: """Greedy NMS: highest score wins within the SAME class, then across classes (REQ-031) — unless the kept box is a container class and the candidate is at least 90% inside it, which is containment, not overlap (REQ-184). `iou_by_class` overrides the threshold per class id (REQ-181); `container_ids` are the class ids marked container.""" if not iou_by_class and 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, cls_dets in by_class.items(): iou = (iou_by_class or {}).get(cls, iou_threshold) if iou <= 0.0: kept.extend(cls_dets) continue 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 for k in cls_kept): cls_kept.append(det) kept.extend(cls_kept) # Cross-class pass (REQ-031), greedy score-desc. A pair whose resolved # threshold is <= 0 is never suppressed — zero means NMS off for that # class, mirroring the within-class pass above. cross_kept: List[Detection] = [] for det in sorted(kept, key=lambda d: d.score, reverse=True): drop = False for k in cross_kept: if det.class_id == k.class_id: continue if (container_ids and k.class_id in container_ids and _containment_fraction(det.box, k.box) >= 0.9): continue iou = (iou_by_class or {}).get(det.class_id, iou_threshold) if iou > 0.0 and _iou(det.box, k.box) >= iou: drop = True break if not drop: cross_kept.append(det) return cross_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, thresholds: Optional[List[float]] = None, iou_by_class: Optional[Dict[int, float]] = None, min_box_fracs: Optional[List[float]] = None, container_ids: Optional[Set[int]] = 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. `thresholds`, `iou_by_class` and `min_box_fracs` are per-prompt overrides aligned with `prompts` (REQ-181); classes without one use the global values. `container_ids` are prompt indices marked container (REQ-184) — the caller maps them, because detections still carry prompt-index class ids here.""" 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, thresholds=thresholds ) else: detections = get_engine().detect(image, prompts, threshold, thresholds=thresholds) except Exception as exc: return ImageResult(image_path, rel_path, width, height, error=str(exc)) if min_box_fracs is not None: def _keep(det) -> bool: frac = min_box_fracs[det.class_id] if det.class_id < len(min_box_fracs) else min_box_frac if frac <= 0: return True floor = width * height * frac return (det.box[2] - det.box[0]) * (det.box[3] - det.box[1]) >= floor detections = [d for d in detections if _keep(d)] elif 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, iou_by_class=iou_by_class, container_ids=container_ids))