feat: setup dataset enrichment app codebase and scripts
This commit is contained in:
1 parent
b5c28cc98a
commit
d07578462e
72 files changed
+11370
No files matched your search
@@ -0,0 +1,81 @@
|
||||
"""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 across all prompts: highest score wins an overlapping region."""
|
||||
kept: List[Detection] = []
|
||||
for det in sorted(detections, key=lambda d: d.score, reverse=True):
|
||||
if all(_iou(det.box, k.box) < iou_threshold for k in kept):
|
||||
kept.append(det)
|
||||
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))
|
||||
Reference in new issue
Block a user