Files
asus d3a6aa49b6 feat: per-class max box fraction (REQ-188)
- labeling/preview/autolabel: max_box_frac + per-class max_box_fracs
  ceiling filter (0=none, 1=off) before NMS, mirrors min_box_frac
- exemplar review-assist: ceiling before max_detections truncation;
  ExemplarLabelRequest + filter panel 'Max box size' slider
- ClassParamsTable: MaxBox column; copy line now REQ-186 amended
  order: class <name> conf <v> iou <v> minbox <v> maxbox <v> container
- docs: design/ui-spec/tasks updated (incl. stale 4-slider narrative)
2026-10-02 17:27:11 +07:00

330 lines
16 KiB
Python

"""The auto-annotation job: SAM3 over every frame of a batch (REQ-030…034).
Detection itself is `labeling.label_image`, unchanged — one `set_image` per
frame with the class prompts looped over that cached state, then greedy IoU NMS
across prompts. This module's job is only to turn its output into rows and to
keep the user's own corrections out of the way.
"""
import os
from typing import List, Optional
from PIL import Image
from backend import batches, config, db, jobs, labeling, projects, review
from backend.batches import BatchError
DEFAULT_THRESHOLD = 0.35
DEFAULT_IOU = 0.0
def _parse_class_params(raw) -> dict:
"""Normalize `class_params` (REQ-181): lowercased class name → overrides.
Unknown keys and non-numeric values are dropped, so a malformed payload
degrades to the global values instead of failing the job."""
out: dict = {}
for name, values in (raw or {}).items():
if not isinstance(values, dict):
continue
entry = {}
for key in ("threshold", "iou_threshold", "min_box_frac", "max_box_frac"):
if key in values:
try:
entry[key] = float(values[key])
except (TypeError, ValueError):
pass
if entry:
out[str(name).strip().lower()] = entry
return out
def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
iou_threshold: float = DEFAULT_IOU, min_box_frac: float = 0.0,
resume: bool = False, append: bool = False, engine: str = "base_model",
engines: Optional[List[str]] = None,
class_ids: Optional[List[int]] = None,
engine_classes: Optional[dict[str, List[str]]] = None,
custom_model_path: Optional[str] = None,
target_class_names: Optional[List[str]] = None,
class_params: Optional[dict] = None) -> dict:
batch = batches.get(batch_id)
if batch is None:
raise batches.BatchError("No such batch")
if batch["frame_count"] == 0:
raise batches.BatchError("This batch has no frames yet")
active_engines = engines if (engines and len(engines) > 0) else [engine]
job = jobs.create(
"autolabel",
params={"batch_id": batch_id, "threshold": threshold,
"iou_threshold": iou_threshold, "min_box_frac": min_box_frac,
"resume": resume, "append": append, "engine": active_engines[0], "engines": active_engines,
"class_ids": class_ids, "engine_classes": engine_classes,
"custom_model_path": custom_model_path, "target_class_names": target_class_names,
"class_params": class_params},
project_id=batch["project_id"],
batch_id=batch_id,
message=f"{batch['date_label']}/{batch['batch_label']} ({'+'.join(e.upper() for e in active_engines)})",
)
return job.to_dict()
def _geometries(detection, width: int, height: int, label_type: str) -> List[dict]:
if label_type == "bbox":
x0, y0, x1, y1 = detection.box
return [review.bbox(x0 / width, y0 / height, x1 / width, y1 / height)]
shapes = []
for points in review.mask_to_polygons(detection.mask):
if len(points) >= 3:
shapes.append(review.polygon([(x / width, y / height) for x, y in points]))
return shapes
@jobs.handler("autolabel")
def _run_autolabel(job) -> None:
batch = batches.get(job.params["batch_id"])
if batch is None:
raise batches.BatchError("The batch disappeared before labeling started")
project = projects.get(batch["project_id"])
selected_engine = job.params.get("engine", "base_model")
frames = batches.frames(batch["id"])
batches.set_status(batch["id"], "labeling")
job.progress(0, len(frames))
skip = review.frames_with_auto(batch["id"]) if job.params.get("resume") else set()
if skip:
job.log(f"Resuming: skipping {len(skip)} frame(s) that already have automatic shapes")
directory = batches.frames_dir(batch["project_slug"], batch["id"])
written = 0
attempted = 0
failures = []
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
per_class = _parse_class_params(job.params.get("class_params"))
# Predict at the LOWEST threshold in play so a class with a lower override
# can still see its boxes; each box then passes its own class gate below.
predict_conf = min([conf] + [v["threshold"] for v in per_class.values() if "threshold" in v])
yolo_model = None
sam3_target_classes = []
custom_path = job.params.get("custom_model_path")
target_class_names = job.params.get("target_class_names")
engine_classes = job.params.get("engine_classes")
if not target_class_names and isinstance(engine_classes, dict):
c_names = engine_classes.get(selected_engine) or engine_classes.get("sam3") or []
if isinstance(c_names, list) and len(c_names) > 0:
target_class_names = c_names
if selected_engine == "sam3" and not custom_path:
allowed_classes_set = {c.strip().lower() for c in target_class_names} if target_class_names else None
if allowed_classes_set:
# Add any new target class names that aren't in project classes yet
existing_names = {c["name"].strip().lower() for c in project["classes"]}
for name in target_class_names:
if name.strip().lower() not in existing_names:
try:
updated_proj = projects.add_class(project["id"], name=name.strip(), prompt=name.strip())
project["classes"] = updated_proj["classes"]
except Exception as exc:
job.log(f"Warning adding class '{name}': {exc}")
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 engine_is_loaded, get_engine
if not engine_is_loaded():
job.log("Loading SAM3 (the first run downloads ~3.4 GB from HuggingFace)…")
engine = get_engine()
job.log(f"SAM3 ready on {engine.device}; prompts: {', '.join(prompts)}")
else:
job.log("SAM3 selected but 0 prompts match project classes.")
else:
from ultralytics import YOLO
if custom_path and os.path.isfile(custom_path):
m_path = custom_path
job.log(f"Loading Custom Model: {os.path.basename(m_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
job.log(f"Loading Base Model: {os.path.basename(m_path)}...")
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
# REQ-184: classes marked container keep boxes that sit inside them. The
# prompt pass above 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")}
job.log(f"Starting auto-labeling with {selected_engine}...")
for index, frame in enumerate(frames):
if job.cancelled:
job.log(f"Cancelled after {index} frame(s)")
break
if frame["id"] in skip:
job.progress(index + 1, len(frames))
continue
attempted += 1
try:
frame_file = os.path.join(directory, frame["filename"])
fw = max(1, frame.get("width") or 1)
fh = max(1, frame.get("height") or 1)
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", conf):
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 selected_engine == "sam3" and sam3_target_classes:
prompts = [(c.get("prompt") or c["name"]).strip() for c in sam3_target_classes]
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", conf) 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_thresh) 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", job.params.get("min_box_frac", 0.0)) 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, conf,
iou_threshold=iou_thresh, min_box_frac=job.params.get("min_box_frac", 0.0),
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)
elif res.error:
job.log(f"[SAM3 ERROR] {frame['filename']}: {res.error}")
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_thresh,
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})
if job.params.get("append"):
review.append_auto(frame["id"], items)
else:
review.replace_auto(frame["id"], items)
written += len(items)
job.progress(index + 1, len(frames), f"{frame['filename']}: {len(items)} shape(s)")
except Exception as exc:
failures.append(str(exc))
job.log(f"[ERROR] {frame['filename']}: {exc}")
job.progress(index + 1, len(frames))
# "Every frame failed" is not a finished job with no findings — it is a
# broken run, and reporting `done` for it would be the system lying about
# its own state. An empty frame is fine (REQ-033); an errored one is not.
if attempted and len(failures) == attempted:
batches.set_status(batch["id"], "failed")
raise BatchError(f"All {attempted} frame(s) failed. First error: {failures[0]}")
batches.set_status(batch["id"], "reviewing")
_reset_reviewed(batch["id"])
if failures:
job.log(f"{len(failures)} of {attempted} frame(s) failed — see the errors above")
job.log(f"Wrote {written} shape(s) across {attempted - len(failures)} frame(s)")
def _reset_reviewed(batch_id: int) -> None:
"""Approvals were given against the previous labels, so a re-run puts those
frames back in the queue. Manual shapes stay; the sign-off does not."""
with db.cursor() as cur:
cur.execute(
"UPDATE frames SET review_status = 'pending' WHERE batch_id = ? "
"AND review_status = 'approved'",
(batch_id,),
)