"""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, config, 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: resolved = config.resolve_data_path(row[0]) if os.path.isfile(resolved): m_path = resolved 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 xb = cp.get("max_box_frac") if xb and xb < 1 and (xyxyn[2]-xyxyn[0])*(xyxyn[3]-xyxyn[1]) > xb: 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 = mx_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] if any("max_box_frac" in v for v in per_class.values()): mx_list = [per_class.get(n, {}).get("max_box_frac", 1.0) 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, max_box_fracs=mx_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