Files
reTraining/backend/labeling.py
T

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