189 lines
8.9 KiB
Python
189 lines
8.9 KiB
Python
"""The auto-annotate preview: one frame, run now, nothing written (REQ-171/172).
|
|
|
|
Split out of `autolabel.py` to keep that file under the 400-line limit. The job
|
|
path and this path share `labeling.label_image`, so what the preview shows is
|
|
what a batch run would write — with one deliberate exception: drawn box
|
|
exemplars (REQ-172) only ever apply here. SAM3's geometric prompts pool features
|
|
from the current image, so replaying them on another frame would ask about
|
|
whatever happens to sit at those coordinates there.
|
|
"""
|
|
|
|
import os
|
|
from typing import List, Optional
|
|
|
|
from backend import batches, db, labeling, projects, review
|
|
from backend.autolabel import DEFAULT_IOU, DEFAULT_THRESHOLD, _geometries, _parse_class_params
|
|
|
|
def preview_frame(
|
|
batch_id: int,
|
|
frame_id: int,
|
|
engine: str,
|
|
threshold: float = DEFAULT_THRESHOLD,
|
|
iou_threshold: float = DEFAULT_IOU,
|
|
min_box_frac: float = 0.0,
|
|
target_class_names: Optional[List[str]] = None,
|
|
custom_model_path: Optional[str] = None,
|
|
exemplars: Optional[List[dict]] = None,
|
|
exemplar_class_name: Optional[str] = None,
|
|
class_params: Optional[dict] = None,
|
|
) -> List[dict]:
|
|
batch = batches.get(batch_id)
|
|
if not batch:
|
|
raise ValueError("No such batch")
|
|
project = projects.get(batch["project_id"])
|
|
|
|
frame = next((f for f in batches.frames(batch_id) if f["id"] == frame_id), None)
|
|
if not frame:
|
|
raise ValueError("Frame not found")
|
|
|
|
directory = batches.frames_dir(batch["project_slug"], batch_id)
|
|
frame_file = os.path.join(directory, frame["filename"])
|
|
fw = max(1, frame.get("width") or 1)
|
|
fh = max(1, frame.get("height") or 1)
|
|
|
|
yolo_model = None
|
|
sam3_target_classes = []
|
|
|
|
if engine == "sam3" and not custom_model_path:
|
|
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
|
if allowed_classes_set:
|
|
sam3_target_classes = [c for c in project["classes"] if c["name"].strip().lower() in allowed_classes_set or c["prompt"].strip().lower() in allowed_classes_set]
|
|
else:
|
|
sam3_target_classes = [c for c in project["classes"]]
|
|
|
|
prompts = [c["prompt"] for c in sam3_target_classes]
|
|
if prompts:
|
|
from backend.sam3_engine import get_engine
|
|
get_engine()
|
|
else:
|
|
from ultralytics import YOLO
|
|
if custom_model_path and os.path.isfile(custom_model_path):
|
|
m_path = custom_model_path
|
|
else:
|
|
m_path = projects.training_start_point(project)
|
|
with db.cursor() as cur:
|
|
cur.execute("SELECT weights_path FROM model_versions WHERE project_id = ? ORDER BY version DESC LIMIT 1", (project["id"],))
|
|
row = cur.fetchone()
|
|
if row and os.path.isfile(row[0]):
|
|
m_path = row[0]
|
|
yolo_model = YOLO(m_path)
|
|
|
|
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
|
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
|
|
per_class = _parse_class_params(class_params)
|
|
predict_conf = min([threshold] + [v["threshold"] for v in per_class.values() if "threshold" in v])
|
|
# REQ-184: classes marked container keep boxes that sit inside them. The
|
|
# prompt pass below still sees prompt-index ids, so it gets the index-aligned set.
|
|
container_ids = {c["class_id"] for c in project["classes"] if c.get("container")}
|
|
prompt_container_ids = {i for i, c in enumerate(sam3_target_classes) if c.get("container")}
|
|
|
|
all_raw_detections = []
|
|
|
|
if yolo_model is not None:
|
|
results = yolo_model.predict(frame_file, conf=predict_conf, verbose=False)
|
|
if results and len(results) > 0:
|
|
model_names = results[0].names
|
|
for box in results[0].boxes:
|
|
cls_idx = int(box.cls[0].item())
|
|
raw_cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
|
|
|
|
target_class_id = name_to_class_id.get(raw_cls_name)
|
|
if target_class_id is None:
|
|
for item in project["classes"]:
|
|
if item["class_id"] == cls_idx:
|
|
target_class_id = item["class_id"]
|
|
break
|
|
if target_class_id is None and 0 <= cls_idx < len(project["classes"]):
|
|
target_class_id = project["classes"][cls_idx]["class_id"]
|
|
|
|
if target_class_id is None:
|
|
continue
|
|
|
|
target_cls_obj = next((c for c in project["classes"] if c["class_id"] == target_class_id), None)
|
|
proj_cls_name = target_cls_obj["name"].strip().lower() if target_cls_obj else ""
|
|
|
|
if allowed_classes_set is not None:
|
|
if (raw_cls_name not in allowed_classes_set and
|
|
proj_cls_name not in allowed_classes_set and
|
|
str(target_class_id) not in allowed_classes_set):
|
|
continue
|
|
|
|
score = float(box.conf[0].item())
|
|
xyxyn = box.xyxyn[0].tolist()
|
|
|
|
cp = per_class.get(proj_cls_name) or per_class.get(raw_cls_name)
|
|
if cp:
|
|
if score < cp.get("threshold", threshold):
|
|
continue
|
|
mb = cp.get("min_box_frac")
|
|
if mb and mb > 0 and (xyxyn[2]-xyxyn[0])*(xyxyn[3]-xyxyn[1]) < mb:
|
|
continue
|
|
|
|
all_raw_detections.append(labeling.Detection(
|
|
class_id=target_class_id,
|
|
class_name=proj_cls_name or raw_cls_name,
|
|
box=[xyxyn[0]*fw, xyxyn[1]*fh, xyxyn[2]*fw, xyxyn[3]*fh],
|
|
score=score,
|
|
mask=None
|
|
))
|
|
|
|
if engine == "sam3" and sam3_target_classes:
|
|
prompts = [(c.get("prompt") or c["name"]).strip() for c in sam3_target_classes]
|
|
# Exemplars belong to exactly one class — the chip that was active when
|
|
# they were drawn. An unknown name means no exemplar class, so the run
|
|
# falls back to plain text rather than silently attaching the boxes to
|
|
# whichever class happens to be first.
|
|
exemplar_index = -1
|
|
if exemplars and exemplar_class_name:
|
|
wanted = exemplar_class_name.strip().lower()
|
|
exemplar_index = next(
|
|
(i for i, c in enumerate(sam3_target_classes)
|
|
if c["name"].strip().lower() == wanted),
|
|
-1,
|
|
)
|
|
thr_list = iou_list = mb_list = None
|
|
if per_class:
|
|
names = [c["name"].strip().lower() for c in sam3_target_classes]
|
|
if any("threshold" in v for v in per_class.values()):
|
|
thr_list = [per_class.get(n, {}).get("threshold", threshold) for n in names]
|
|
if any("iou_threshold" in v for v in per_class.values()):
|
|
iou_list = [per_class.get(n, {}).get("iou_threshold", iou_threshold) for n in names]
|
|
if any("min_box_frac" in v for v in per_class.values()):
|
|
mb_list = [per_class.get(n, {}).get("min_box_frac", min_box_frac) for n in names]
|
|
res = labeling.label_image(
|
|
frame_file, frame["filename"], prompts, threshold,
|
|
iou_threshold=iou_threshold, min_box_frac=min_box_frac,
|
|
exemplar_index=exemplar_index, exemplars=exemplars,
|
|
thresholds=thr_list,
|
|
iou_by_class=dict(enumerate(iou_list)) if iou_list else None,
|
|
min_box_fracs=mb_list,
|
|
container_ids=prompt_container_ids,
|
|
)
|
|
if not res.error and res.detections:
|
|
for det in res.detections:
|
|
if 0 <= det.class_id < len(sam3_target_classes):
|
|
real_cls = sam3_target_classes[det.class_id]
|
|
det.class_id = real_cls["class_id"]
|
|
det.class_name = real_cls["name"]
|
|
all_raw_detections.append(det)
|
|
|
|
iou_by_class_proj = {
|
|
c["class_id"]: per_class[c["name"].strip().lower()]["iou_threshold"]
|
|
for c in project["classes"]
|
|
if c["name"].strip().lower() in per_class
|
|
and "iou_threshold" in per_class[c["name"].strip().lower()]
|
|
}
|
|
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_threshold,
|
|
iou_by_class=iou_by_class_proj or None,
|
|
container_ids=container_ids)
|
|
items = []
|
|
for det in kept:
|
|
if project["label_type"] == "bbox" or det.mask is None:
|
|
geom = review.bbox(det.box[0]/fw, det.box[1]/fh, det.box[2]/fw, det.box[3]/fh)
|
|
items.append({"class_id": det.class_id, "geometry": geom, "score": det.score})
|
|
else:
|
|
for geometry in _geometries(det, fw, fh, project["label_type"]):
|
|
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
|
|
|
|
return items
|