Files
feedmill-auto-label/backend/labeling.py
T
asus 5c7c122105 feat: add counting bench, triage, and dataset modules
This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
2026-08-14 16:28:52 +07:00

93 lines
3.0 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,
) -> ImageResult:
"""Detect every prompt in one image and return the surviving instances."""
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:
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))