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,241 @@
|
||||
"""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 backend import batches, db, jobs, labeling, projects, review
|
||||
from backend.batches import BatchError
|
||||
|
||||
DEFAULT_THRESHOLD = 0.35
|
||||
DEFAULT_IOU = 0.8
|
||||
|
||||
|
||||
def start(batch_id: int, threshold: float = DEFAULT_THRESHOLD,
|
||||
iou_threshold: float = DEFAULT_IOU, min_box_frac: float = 0.0,
|
||||
resume: bool = False, engine: str = "sam3",
|
||||
engines: Optional[List[str]] = None,
|
||||
class_ids: Optional[List[int]] = None,
|
||||
engine_classes: Optional[dict[str, List[str]]] = 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, "engine": active_engines[0], "engines": active_engines,
|
||||
"class_ids": class_ids, "engine_classes": engine_classes},
|
||||
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"])
|
||||
|
||||
raw_active = job.params.get("engines") or [job.params.get("engine", "sam3")]
|
||||
expanded_engines = []
|
||||
for eng in raw_active:
|
||||
if eng == "both":
|
||||
expanded_engines.extend(["base_model", "secondary_model"])
|
||||
elif eng == "sam3+model1":
|
||||
expanded_engines.extend(["sam3", "base_model"])
|
||||
else:
|
||||
expanded_engines.append(eng)
|
||||
expanded_engines = list(dict.fromkeys(expanded_engines))
|
||||
|
||||
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 = []
|
||||
|
||||
from ultralytics import YOLO
|
||||
m1_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]):
|
||||
m1_path = row[0]
|
||||
|
||||
yolo_models = {}
|
||||
if "base_model" in expanded_engines or "yolo" in expanded_engines:
|
||||
job.log(f"Loading Base/Trained Model: {os.path.basename(m1_path)}...")
|
||||
yolo_models["base_model"] = YOLO(m1_path)
|
||||
|
||||
if "secondary_model" in expanded_engines:
|
||||
m2_path = project["secondary_model_path"] if (project.get("secondary_model_path") and os.path.isfile(project["secondary_model_path"])) else m1_path
|
||||
label_name = project.get("secondary_model_name") or os.path.basename(m2_path)
|
||||
job.log(f"Loading Secondary Model: {label_name}...")
|
||||
yolo_models["secondary_model"] = YOLO(m2_path)
|
||||
|
||||
allowed_class_ids = set(job.params["class_ids"]) if job.params.get("class_ids") is not None else None
|
||||
engine_classes = job.params.get("engine_classes") or {}
|
||||
|
||||
sam3_target_classes = []
|
||||
if "sam3" in expanded_engines:
|
||||
sam3_classes = engine_classes.get("sam3")
|
||||
if sam3_classes is not None:
|
||||
allowed_set = {c.strip().lower() for c in sam3_classes}
|
||||
sam3_target_classes = [
|
||||
c for c in project["classes"]
|
||||
if c["name"].strip().lower() in allowed_set or c["prompt"].strip().lower() in allowed_set
|
||||
]
|
||||
else:
|
||||
sam3_target_classes = [
|
||||
c for c in project["classes"]
|
||||
if (allowed_class_ids is None or c["class_id"] in allowed_class_ids)
|
||||
]
|
||||
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 class filter.")
|
||||
|
||||
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
||||
conf = job.params.get("threshold", DEFAULT_THRESHOLD)
|
||||
iou_thresh = job.params.get("iou_threshold", DEFAULT_IOU)
|
||||
|
||||
job.log(f"Starting multi-engine auto-labeling ({', '.join(expanded_engines)})...")
|
||||
|
||||
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"])
|
||||
all_raw_detections = []
|
||||
|
||||
for eng_key, y_model in yolo_models.items():
|
||||
allowed_for_eng = engine_classes.get(eng_key)
|
||||
if allowed_for_eng is not None and len(allowed_for_eng) == 0:
|
||||
continue
|
||||
results = y_model.predict(frame_file, conf=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())
|
||||
cls_name = str(model_names.get(cls_idx, cls_idx)).strip().lower()
|
||||
if allowed_for_eng is not None and cls_name not in [c.strip().lower() for c in allowed_for_eng]:
|
||||
continue
|
||||
if cls_name not in name_to_class_id:
|
||||
try:
|
||||
updated_proj = projects.add_class(project["id"], {"name": cls_name, "prompt": cls_name})
|
||||
project["classes"] = updated_proj["classes"]
|
||||
name_to_class_id = {item["name"].strip().lower(): item["class_id"] for item in project["classes"]}
|
||||
except Exception:
|
||||
pass
|
||||
if cls_name in name_to_class_id:
|
||||
target_class_id = name_to_class_id[cls_name]
|
||||
else:
|
||||
continue
|
||||
score = float(box.conf[0].item())
|
||||
xyxyn = box.xyxyn[0].tolist()
|
||||
all_raw_detections.append(labeling.Detection(
|
||||
class_id=target_class_id,
|
||||
class_name=cls_name,
|
||||
box=[xyxyn[0]*frame["width"], xyxyn[1]*frame["height"], xyxyn[2]*frame["width"], xyxyn[3]*frame["height"]],
|
||||
score=score,
|
||||
mask=None
|
||||
))
|
||||
|
||||
if "sam3" in expanded_engines and sam3_target_classes:
|
||||
prompts = [c["prompt"] for c in sam3_target_classes]
|
||||
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)
|
||||
)
|
||||
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)
|
||||
|
||||
kept = labeling.deduplicate(all_raw_detections, iou_threshold=iou_thresh)
|
||||
items = []
|
||||
for det in kept:
|
||||
if project["label_type"] == "bbox" or det.mask is None:
|
||||
geom = review.bbox(det.box[0]/frame["width"], det.box[1]/frame["height"], det.box[2]/frame["width"], det.box[3]/frame["height"])
|
||||
items.append({"class_id": det.class_id, "geometry": geom, "score": det.score})
|
||||
else:
|
||||
for geometry in _geometries(det, frame["width"], frame["height"], project["label_type"]):
|
||||
items.append({"class_id": det.class_id, "geometry": geometry, "score": det.score})
|
||||
|
||||
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,),
|
||||
)
|
||||
Reference in new issue
Block a user