- config.py: resolve_data_path (legacy abs + rel) + rel_data_path - all file-opening reads wrapped: preview, autolabel, training, model download, live count, projects.get; training_start_point hack replaced - new writes store paths relative to data/ - legacy stale rows (/home/asus/reTraining/...) resolve without migration - requirements: REQ-187 added; REQ-188 (per-class max box) + REQ-186 copy-line amendment drafted for the next task
324 lines
15 KiB
Python
324 lines
15 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"):
|
|
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
|
|
|
|
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 = 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]
|
|
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,
|
|
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,),
|
|
)
|