166 lines
6.4 KiB
Python
166 lines
6.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 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))
|