feat: setup dataset enrichment app codebase and scripts

This commit is contained in:
asus committed 2026-08-05 11:52:27 +07:00
1 parent b5c28cc98a
commit d07578462e
72 files changed
+11370

No files matched your search

View File
Whitespace-only changes.
View File
Whitespace-only changes.
+152
View File
@@ -0,0 +1,152 @@
"""Batch, frame and auto-annotation routes (REQ-020…034)."""
import os
from typing import Optional
from fastapi import APIRouter, HTTPException, Response
from fastapi.responses import FileResponse
from pydantic import BaseModel
from backend import autolabel, dataset, library
from backend import batches as batch_store
from backend import review as review_store
from backend.api.common import project_or_404, thumbnail
router = APIRouter(tags=["batches"])
class BatchRequest(BaseModel):
rel: str
start_sec: float = 0.0
end_sec: float
fps: float = 1.0
class AutolabelRequest(BaseModel):
engine: str = "sam3"
engines: Optional[list[str]] = None
class_ids: Optional[list[int]] = None
engine_classes: Optional[dict[str, list[str]]] = None
threshold: float = autolabel.DEFAULT_THRESHOLD
iou_threshold: float = autolabel.DEFAULT_IOU
min_box_frac: float = 0.0
resume: bool = False
@router.post("/api/projects/{project_id}/batches")
def create_batch(project_id: int, request: BatchRequest) -> dict:
try:
return batch_store.create(project_id, request.rel, request.start_sec,
request.end_sec, request.fps)
except (batch_store.BatchError, library.LibraryError) as exc:
raise HTTPException(400, str(exc))
@router.get("/api/projects/{project_id}/batches")
def list_batches(project_id: int) -> dict:
project_or_404(project_id)
return {"batches": batch_store.listing(project_id)}
class BatchPatch(BaseModel):
batch_label: Optional[str] = None
date_label: Optional[str] = None
status: Optional[str] = None
@router.get("/api/batches/{batch_id}")
def read_batch(batch_id: int) -> dict:
batch = batch_store.get(batch_id)
if batch is None:
raise HTTPException(404, "No such batch")
return batch
@router.patch("/api/batches/{batch_id}")
def update_batch(batch_id: int, request: BatchPatch) -> dict:
try:
return batch_store.update(batch_id, request.model_dump(exclude_unset=True))
except batch_store.BatchError as exc:
raise HTTPException(400, str(exc))
@router.delete("/api/batches/{batch_id}")
def delete_batch(batch_id: int) -> dict:
if not batch_store.delete(batch_id):
raise HTTPException(404, "No such batch")
return {"deleted": True}
@router.get("/api/batches/{batch_id}/frames")
def list_frames(batch_id: int) -> dict:
if batch_store.get(batch_id) is None:
raise HTTPException(404, "No such batch")
return {"frames": batch_store.frames(batch_id)}
@router.post("/api/batches/{batch_id}/autolabel")
def start_autolabel(batch_id: int, request: AutolabelRequest) -> dict:
try:
engine_list = request.engines if (request.engines and len(request.engines) > 0) else [request.engine]
return autolabel.start(batch_id, request.threshold, request.iou_threshold,
request.min_box_frac, resume=request.resume,
engines=engine_list, class_ids=request.class_ids,
engine_classes=request.engine_classes)
except batch_store.BatchError as exc:
raise HTTPException(400, str(exc))
@router.post("/api/batches/{batch_id}/approve-all")
def approve_all_batch_frames(batch_id: int) -> dict:
if batch_store.get(batch_id) is None:
raise HTTPException(404, "No such batch")
updated = batch_store.approve_all_frames(batch_id)
return {"approved_count": updated}
@router.post("/api/batches/{batch_id}/approve")
def approve_batch(batch_id: int) -> dict:
try:
return dataset.approve(batch_id)
except dataset.DatasetError as exc:
raise HTTPException(400, str(exc))
@router.get("/api/projects/{project_id}/dataset")
def dataset_summary(project_id: int) -> dict:
project_or_404(project_id)
return dataset.summary(project_id)
@router.get("/api/projects/{project_id}/dataset/download")
def dataset_download(project_id: int):
project = project_or_404(project_id)
try:
path = dataset.zip_path(project)
except dataset.DatasetError as exc:
raise HTTPException(400, str(exc))
return FileResponse(path, media_type="application/zip",
filename=f"{project['slug']}-dataset.zip")
@router.get("/api/frames/{frame_id}/image")
def frame_image(frame_id: int, w: int = 0):
path = batch_store.frame_path(frame_id)
if path is None or not os.path.isfile(path):
raise HTTPException(404, "No such frame")
if w and 16 <= w <= 2048:
return Response(content=thumbnail(path, w), media_type="image/jpeg",
headers={"Cache-Control": "public, max-age=3600"})
return FileResponse(path, media_type="image/jpeg",
headers={"Cache-Control": "public, max-age=3600"})
@router.delete("/api/batches/{batch_id}/classes/{class_id}/annotations")
def clear_batch_class_annotations(batch_id: int, class_id: int) -> dict:
if batch_store.get(batch_id) is None:
raise HTTPException(404, "No such batch")
deleted = review_store.clear_batch_class_annotations(batch_id, class_id)
return {"deleted": deleted}
+35
View File
@@ -0,0 +1,35 @@
"""Helpers shared by the route modules."""
import io
import os
from functools import lru_cache
from fastapi import HTTPException
from backend import projects as project_store
def project_or_404(project_id: int) -> dict:
project = project_store.get(project_id)
if project is None:
raise HTTPException(404, "No such project")
return project
@lru_cache(maxsize=512)
def _thumbnail_cached(path: str, width: int, mtime: float) -> bytes:
from PIL import Image
with Image.open(path) as image:
image = image.convert("RGB")
if image.width > width:
height = max(1, round(image.height * width / image.width))
image = image.resize((width, height), Image.BILINEAR)
buffer = io.BytesIO()
image.save(buffer, "JPEG", quality=72)
return buffer.getvalue()
def thumbnail(path: str, width: int) -> bytes:
# mtime is part of the key so a re-extracted frame never serves a stale one.
return _thumbnail_cached(path, width, os.path.getmtime(path))
+29
View File
@@ -0,0 +1,29 @@
"""Job routes (REQ-070, REQ-071)."""
from typing import Optional
from fastapi import APIRouter, HTTPException
from backend import jobs as job_store
router = APIRouter(prefix="/api/jobs", tags=["jobs"])
@router.get("")
def list_jobs(project_id: Optional[int] = None, limit: int = 50) -> dict:
return {"jobs": [job.to_dict() for job in job_store.listing(project_id, limit)]}
@router.get("/{job_id}")
def get_job(job_id: int) -> dict:
job = job_store.get(job_id)
if job is None:
raise HTTPException(404, "No such job")
return job.to_dict()
@router.post("/{job_id}/cancel")
def cancel_job(job_id: int) -> dict:
if not job_store.cancel(job_id):
raise HTTPException(400, "Job is not cancellable")
return {"ok": True}
+63
View File
@@ -0,0 +1,63 @@
"""Training and model-version routes (REQ-060…065)."""
import os
from typing import Optional, Union
from fastapi import APIRouter, HTTPException
from fastapi.responses import FileResponse
from pydantic import BaseModel
from backend import hardware, training
from backend.api.common import project_or_404
router = APIRouter(tags=["models"])
class TrainRequest(BaseModel):
epochs: int = 50
batch: Optional[int] = None
imgsz: Optional[int] = None
device: Optional[Union[int, str]] = None
batch_ids: Optional[list] = None
@router.get("/api/hardware")
def read_hardware() -> dict:
return hardware.defaults()
@router.post("/api/projects/{project_id}/train")
def start_training(project_id: int, request: TrainRequest) -> dict:
project_or_404(project_id)
try:
return training.start(
project_id, request.epochs,
{"batch": request.batch, "imgsz": request.imgsz, "device": request.device},
batch_ids=request.batch_ids,
)
except training.TrainingError as exc:
raise HTTPException(400, str(exc))
@router.get("/api/projects/{project_id}/models")
def list_models(project_id: int) -> dict:
project = project_or_404(project_id)
return {"models": training.listing(project_id),
"base_model_kind": project["base_model_kind"]}
@router.get("/api/models/{model_id}/weights")
def download_weights(model_id: int):
version = training.get_version(model_id)
if version is None or not os.path.isfile(version["weights_path"]):
raise HTTPException(404, "No weights for that version")
return FileResponse(version["weights_path"], media_type="application/octet-stream",
filename=f"v{version['version']}-best.pt")
@router.post("/api/models/{model_id}/promote")
def promote_model(model_id: int) -> dict:
try:
return training.promote(model_id)
except training.TrainingError as exc:
raise HTTPException(400, str(exc))
+207
View File
@@ -0,0 +1,207 @@
"""Project, library and video-streaming routes (REQ-001…013)."""
import os
import shutil
import tempfile
from typing import Dict, List, Optional
from fastapi import APIRouter, File, HTTPException, Request, UploadFile
from fastapi.responses import FileResponse, StreamingResponse
from pydantic import BaseModel, Field
from backend import config, library, video
from backend import projects as project_store
from backend.api.common import project_or_404
router = APIRouter(prefix="/api/projects", tags=["projects"])
VIDEO_MEDIA = {".mp4": "video/mp4", ".m4v": "video/mp4", ".mkv": "video/x-matroska",
".mov": "video/quicktime", ".webm": "video/webm", ".avi": "video/x-msvideo"}
STREAM_CHUNK = 1024 * 512
class ClassSpec(BaseModel):
name: str
prompt: Optional[str] = None
class ProjectRequest(BaseModel):
name: str
label_type: str = "bbox"
video_root: str
classes: List[ClassSpec] = Field(default_factory=list)
val_every: int = 5
class ProjectPatch(BaseModel):
prompts: Optional[Dict[int, str]] = None
val_every: Optional[int] = None
video_root: Optional[str] = None
@router.get("")
def list_projects() -> dict:
return {"projects": project_store.listing(), "video_root_default": config.VIDEO_ROOT}
@router.post("")
def create_project(request: ProjectRequest) -> dict:
try:
return project_store.create(
name=request.name,
label_type=request.label_type,
video_root=request.video_root,
classes=[item.model_dump() for item in request.classes],
val_every=request.val_every,
)
except project_store.ProjectError as exc:
raise HTTPException(400, str(exc))
@router.get("/{project_id}")
def read_project(project_id: int) -> dict:
return project_or_404(project_id)
@router.patch("/{project_id}")
def patch_project(project_id: int, request: ProjectPatch) -> dict:
try:
return project_store.update(project_id, prompts=request.prompts,
val_every=request.val_every,
video_root=request.video_root)
except project_store.ProjectError as exc:
raise HTTPException(400, str(exc))
@router.post("/{project_id}/base-model")
async def upload_base_model(project_id: int, file: UploadFile = File(...)) -> dict:
if not (file.filename or "").endswith(".pt"):
raise HTTPException(400, "Base model must be a .pt checkpoint")
# Staged to a temp file first: reading the classes can fail, and a rejected
# upload must not leave a broken model.pt in the project.
with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
shutil.copyfileobj(file.file, staged)
staged_path = staged.name
await file.close()
try:
return project_store.set_base_model(project_id, staged_path)
except project_store.ProjectError as exc:
raise HTTPException(400, str(exc))
finally:
os.unlink(staged_path)
@router.post("/{project_id}/secondary-model")
async def upload_secondary_model(project_id: int, file: UploadFile = File(...)) -> dict:
project_or_404(project_id)
if not file.filename.endswith(".pt"):
raise HTTPException(400, "The secondary model must be a .pt file")
with tempfile.NamedTemporaryFile(suffix=".pt", delete=False) as staged:
shutil.copyfileobj(file.file, staged)
staged_path = staged.name
await file.close()
try:
return project_store.set_secondary_model(project_id, staged_path, name=file.filename)
except project_store.ProjectError as exc:
raise HTTPException(400, str(exc))
finally:
os.unlink(staged_path)
@router.post("/{project_id}/classes")
def add_class(project_id: int, request: ClassSpec) -> dict:
try:
return project_store.add_class(project_id, request.name, request.prompt)
except project_store.ProjectError as exc:
raise HTTPException(400, str(exc))
@router.delete("/{project_id}/classes/{class_id}")
def delete_class(project_id: int, class_id: int) -> dict:
try:
return project_store.delete_class(project_id, class_id)
except project_store.ProjectError as exc:
raise HTTPException(400, str(exc))
@router.delete("/{project_id}")
def delete_project(project_id: int) -> dict:
return {"deleted": project_store.delete(project_id)}
@router.get("/{project_id}/library")
def list_library(project_id: int) -> dict:
project = project_or_404(project_id)
try:
return {"video_root": project["video_root"],
"dates": library.list_dates(project["video_root"])}
except library.LibraryError as exc:
raise HTTPException(400, str(exc))
@router.get("/{project_id}/library/{date}")
def list_library_date(project_id: int, date: str) -> dict:
project = project_or_404(project_id)
try:
return {"date": date,
"videos": library.list_videos(project["video_root"], date, project_id)}
except library.LibraryError as exc:
raise HTTPException(400, str(exc))
@router.get("/{project_id}/video/info")
def video_info(project_id: int, rel: str) -> dict:
project = project_or_404(project_id)
try:
path = library.resolve(project["video_root"], rel)
info = dict(video.probe(path))
except (library.LibraryError, video.VideoError) as exc:
raise HTTPException(400, str(exc))
info.update({"rel": rel, "batch_label": library.batch_label(os.path.basename(path)),
"date_label": rel.split("/", 1)[0]})
return info
@router.get("/{project_id}/video")
def stream_video(project_id: int, rel: str, request: Request):
"""Serve a video with Range support so the player can seek (REQ-013)."""
project = project_or_404(project_id)
try:
path = library.resolve(project["video_root"], rel)
except library.LibraryError as exc:
raise HTTPException(404, str(exc))
media = VIDEO_MEDIA.get(os.path.splitext(path)[1].lower(), "application/octet-stream")
size = os.path.getsize(path)
header = request.headers.get("range")
if not header or not header.startswith("bytes="):
return FileResponse(path, media_type=media, headers={"Accept-Ranges": "bytes"})
start_text, _, end_text = header[6:].partition("-")
start = int(start_text) if start_text else 0
end = int(end_text) if end_text else size - 1
start, end = max(0, start), min(end, size - 1)
if start > end:
raise HTTPException(416, "Requested range is not satisfiable")
def chunks():
with open(path, "rb") as handle:
handle.seek(start)
remaining = end - start + 1
while remaining > 0:
block = handle.read(min(STREAM_CHUNK, remaining))
if not block:
break
remaining -= len(block)
yield block
return StreamingResponse(
chunks(), status_code=206, media_type=media,
headers={
"Content-Range": f"bytes {start}-{end}/{size}",
"Accept-Ranges": "bytes",
"Content-Length": str(end - start + 1),
},
)
+91
View File
@@ -0,0 +1,91 @@
"""Annotation and review-state routes (REQ-040…045)."""
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, HTTPException
from pydantic import BaseModel
from backend import batches as batch_store
from backend import review as review_store
router = APIRouter(tags=["review"])
class AnnotationRequest(BaseModel):
class_id: int = 0
geometry: Dict[str, Any]
class AnnotationPatch(BaseModel):
class_id: Optional[int] = None
geometry: Optional[Dict[str, Any]] = None
class StatusRequest(BaseModel):
status: str
class AssistRequest(BaseModel):
box: List[float]
class_id: int = 0
threshold: float = 0.5
@router.get("/api/frames/{frame_id}/annotations")
def list_annotations(frame_id: int) -> dict:
target = review_store.frame(frame_id)
if target is None:
raise HTTPException(404, "No such frame")
return {
"frame": {
"id": target["id"], "batch_id": target["batch_id"], "idx": target["idx"],
"width": target["width"], "height": target["height"],
"review_status": target["review_status"], "label_type": target["label_type"],
},
"annotations": review_store.listing(frame_id),
}
@router.post("/api/frames/{frame_id}/annotations")
def add_annotation(frame_id: int, request: AnnotationRequest) -> dict:
try:
return review_store.add(frame_id, request.class_id, request.geometry)
except review_store.ReviewError as exc:
raise HTTPException(400, str(exc))
@router.patch("/api/annotations/{annotation_id}")
def patch_annotation(annotation_id: int, request: AnnotationPatch) -> dict:
try:
return review_store.update(annotation_id, request.class_id, request.geometry)
except review_store.ReviewError as exc:
raise HTTPException(400, str(exc))
@router.delete("/api/annotations/{annotation_id}")
def delete_annotation(annotation_id: int) -> dict:
return {"deleted": review_store.delete(annotation_id)}
@router.post("/api/frames/{frame_id}/assist")
def assist(frame_id: int, request: AssistRequest) -> dict:
try:
return review_store.assist(frame_id, request.box, request.class_id,
request.threshold)
except review_store.ReviewError as exc:
raise HTTPException(400, str(exc))
@router.post("/api/frames/{frame_id}/status")
def set_status(frame_id: int, request: StatusRequest) -> dict:
try:
return review_store.set_status(frame_id, request.status)
except review_store.ReviewError as exc:
raise HTTPException(400, str(exc))
@router.get("/api/batches/{batch_id}/next-pending")
def next_pending(batch_id: int, after_idx: int = -1) -> dict:
if batch_store.get(batch_id) is None:
raise HTTPException(404, "No such batch")
return {"frame_id": review_store.next_pending(batch_id, after_idx)}
+241
View File
@@ -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,),
)
+253
View File
@@ -0,0 +1,253 @@
"""Batch lifecycle: one trimmed range of one video, turned into frames.
A batch is the unit of work everything downstream hangs off — auto-annotation,
review, and the merge into the master dataset all address a batch. The same
video can produce many batches with different ranges (REQ-023).
extracting -> extracted -> labeling -> reviewing -> approved -> merged
\\-> failed
"""
import os
import time
from typing import List, Optional
from PIL import Image
from backend import config, db, jobs, library, projects, video
class BatchError(Exception):
pass
def batch_dir(project_slug: str, batch_id: int) -> str:
return os.path.join(config.project_dir(project_slug), "batches", str(batch_id))
def frames_dir(project_slug: str, batch_id: int) -> str:
return os.path.join(batch_dir(project_slug, batch_id), "frames")
def create(project_id: int, rel: str, start_sec: float, end_sec: float, fps: float) -> dict:
"""Register a batch and queue its extraction job (REQ-020…022)."""
project = projects.get(project_id)
if project is None:
raise BatchError("No such project")
try:
video_path = library.resolve(project["video_root"], rel)
except library.LibraryError as exc:
raise BatchError(str(exc))
try:
info = video.probe(video_path)
except video.VideoError as exc:
raise BatchError(str(exc))
end_sec = min(end_sec, info["duration"]) if info["duration"] else end_sec
if end_sec <= start_sec:
raise BatchError("The end of the range must be after its start")
if fps <= 0:
raise BatchError("fps must be greater than 0")
date_label, filename = rel.split("/", 1)
with db.cursor() as cur:
cur.execute(
"""INSERT INTO batches (project_id, video_path, date_label, batch_label,
start_sec, end_sec, fps, status, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, 'extracting', ?)""",
(project_id, video_path, date_label, library.batch_label(filename),
float(start_sec), float(end_sec), float(fps), time.time()),
)
batch_id = cur.lastrowid
jobs.create(
"extract",
params={"batch_id": batch_id},
project_id=project_id,
batch_id=batch_id,
message=f"{date_label}/{library.batch_label(filename)}",
)
return get(batch_id)
def get(batch_id: int) -> Optional[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT b.*, p.slug AS project_slug, p.name AS project_name
FROM batches b JOIN projects p ON p.id = b.project_id WHERE b.id = ?""",
(batch_id,),
)
row = cur.fetchone()
if row is None:
return None
return _row_to_dict(cur, row)
def listing(project_id: int) -> List[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT b.*, p.slug AS project_slug, p.name AS project_name
FROM batches b JOIN projects p ON p.id = b.project_id
WHERE b.project_id = ? ORDER BY b.created_at DESC""",
(project_id,),
)
return [_row_to_dict(cur, row) for row in cur.fetchall()]
def _row_to_dict(cur, row) -> dict:
cur.execute(
"""SELECT review_status, COUNT(*) FROM frames WHERE batch_id = ?
GROUP BY review_status""",
(row["id"],),
)
review = {"pending": 0, "approved": 0, "rejected": 0}
for status, count in cur.fetchall():
review[status] = count
cur.execute(
"""SELECT COUNT(*) FROM annotations a JOIN frames f ON f.id = a.frame_id
WHERE f.batch_id = ?""",
(row["id"],),
)
annotation_count = cur.fetchone()[0]
return {
"id": row["id"],
"project_id": row["project_id"],
"project_slug": row["project_slug"],
"project_name": row["project_name"],
"video_path": row["video_path"],
"date_label": row["date_label"],
"batch_label": row["batch_label"],
"start_sec": row["start_sec"],
"end_sec": row["end_sec"],
"fps": row["fps"],
"status": row["status"],
"frame_count": row["frame_count"],
"created_at": row["created_at"],
"merged_at": row["merged_at"],
"review": review,
"reviewed": review["approved"] + review["rejected"],
"annotation_count": annotation_count,
}
def frames(batch_id: int) -> List[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT f.*, (SELECT COUNT(*) FROM annotations a WHERE a.frame_id = f.id)
AS annotation_count
FROM frames f WHERE f.batch_id = ? ORDER BY f.idx""",
(batch_id,),
)
return [dict(row) for row in cur.fetchall()]
def frame_path(frame_id: int) -> Optional[str]:
with db.cursor() as cur:
cur.execute(
"""SELECT f.filename, b.id AS batch_id, p.slug
FROM frames f
JOIN batches b ON b.id = f.batch_id
JOIN projects p ON p.id = b.project_id
WHERE f.id = ?""",
(frame_id,),
)
row = cur.fetchone()
if row is None:
return None
return os.path.join(frames_dir(row["slug"], row["batch_id"]), row["filename"])
def set_status(batch_id: int, status: str) -> None:
with db.cursor() as cur:
cur.execute("UPDATE batches SET status = ? WHERE id = ?", (status, batch_id))
@jobs.handler("extract")
def _run_extract(job) -> None:
batch = get(job.params["batch_id"])
if batch is None:
raise BatchError("The batch disappeared before extraction started")
out_dir = frames_dir(batch["project_slug"], batch["id"])
expected = video.frame_count(batch["start_sec"], batch["end_sec"], batch["fps"])
job.log(f"Extracting {expected} frame(s) at {batch['fps']} fps from "
f"{batch['date_label']}/{batch['batch_label']} "
f"[{batch['start_sec']:.1f}s – {batch['end_sec']:.1f}s]")
job.progress(0, expected)
try:
names = video.extract_frames(
batch["video_path"], out_dir,
batch["start_sec"], batch["end_sec"], batch["fps"],
on_progress=lambda written: job.progress(written, expected),
should_stop=lambda: job.cancelled,
)
except video.VideoError as exc:
set_status(batch["id"], "failed")
raise BatchError(str(exc))
if not names:
set_status(batch["id"], "failed")
raise BatchError("ffmpeg produced no frames for that range")
with Image.open(os.path.join(out_dir, names[0])) as first:
width, height = first.size
with db.cursor() as cur:
cur.executemany(
"INSERT OR IGNORE INTO frames (batch_id, idx, filename, width, height) "
"VALUES (?, ?, ?, ?, ?)",
[(batch["id"], index, name, width, height) for index, name in enumerate(names)],
)
cur.execute("UPDATE batches SET frame_count = ?, status = 'extracted' WHERE id = ?",
(len(names), batch["id"]))
job.progress(len(names), len(names))
job.log(f"Extracted {len(names)} frame(s) at {width}×{height}")
def update(batch_id: int, patch: dict) -> dict:
batch = get(batch_id)
if batch is None:
raise BatchError("No such batch")
fields = []
args = []
if "batch_label" in patch and patch["batch_label"] is not None:
fields.append("batch_label = ?")
args.append(str(patch["batch_label"]).strip())
if "date_label" in patch and patch["date_label"] is not None:
fields.append("date_label = ?")
args.append(str(patch["date_label"]).strip())
if "status" in patch and patch["status"] is not None:
fields.append("status = ?")
args.append(str(patch["status"]).strip())
if fields:
args.append(batch_id)
with db.cursor() as cur:
cur.execute(f"UPDATE batches SET {', '.join(fields)} WHERE id = ?", args)
return get(batch_id)
def delete(batch_id: int) -> bool:
import shutil
batch = get(batch_id)
if batch is None:
return False
with db.cursor() as cur:
cur.execute("DELETE FROM batches WHERE id = ?", (batch_id,))
shutil.rmtree(batch_dir(batch["project_slug"], batch_id), ignore_errors=True)
return True
def approve_all_frames(batch_id: int) -> int:
with db.cursor() as cur:
cur.execute("UPDATE frames SET review_status = 'approved' WHERE batch_id = ?", (batch_id,))
return cur.rowcount
+39
View File
@@ -0,0 +1,39 @@
"""Paths and settings. Everything is environment-driven so the same image runs
locally and in Docker without code changes (REQ-072)."""
import os
from dotenv import load_dotenv
load_dotenv()
# huggingface_hub reads HF_TOKEN / HUGGING_FACE_HUB_TOKEN; accept either name in
# .env so pasting a token under the obvious name just works.
if os.environ.get("HF_TOKEN") and not os.environ.get("HUGGING_FACE_HUB_TOKEN"):
os.environ["HUGGING_FACE_HUB_TOKEN"] = os.environ["HF_TOKEN"]
REPO_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
DATA_DIR = os.path.abspath(os.environ.get("APP_DATA_DIR", os.path.join(REPO_ROOT, "data")))
PROJECTS_DIR = os.path.join(DATA_DIR, "projects")
DB_PATH = os.path.join(DATA_DIR, "app.db")
# Where the video archive is mounted. Projects store a path relative to nothing —
# they store an absolute one — but this is the default the UI starts browsing from.
VIDEO_ROOT = os.path.abspath(os.environ.get("VIDEO_ARCHIVE", os.path.join(DATA_DIR, "archive")))
# Vite dev server needs cross-origin access; in Docker nginx proxies /api and
# this is irrelevant.
CORS_ORIGINS = [
origin.strip()
for origin in os.environ.get("CORS_ORIGINS", "http://localhost:5173").split(",")
if origin.strip()
]
def ensure_dirs() -> None:
os.makedirs(PROJECTS_DIR, exist_ok=True)
def project_dir(slug: str) -> str:
return os.path.join(PROJECTS_DIR, slug)
+269
View File
@@ -0,0 +1,269 @@
"""The master dataset: approved frames merged in, batch after batch (REQ-050…054).
The one rule that matters here is the stable val split. A frame's membership is
recorded once in `dataset_items` and never revised, so an image that was in
`val` for the last comparison is still in `val` for the next one. Without that,
a rising mAP could just mean an easier val set.
Label files are plain YOLO:
detect class_id cx cy w h (normalized)
segment class_id x1 y1 x2 y2 … (normalized polygon)
"""
import os
import shutil
import time
from backend import batches, config, db, jobs, projects, review
class DatasetError(Exception):
pass
def dataset_dir(project_slug: str) -> str:
return os.path.join(config.project_dir(project_slug), "dataset")
def approve(batch_id: int) -> dict:
"""Sign a batch off and queue its merge (REQ-045, REQ-050)."""
batch = batches.get(batch_id)
if batch is None:
raise DatasetError("No such batch")
if batch["status"] == "merged":
raise DatasetError("This batch is already in the master dataset")
if batch["review"]["pending"] > 0:
raise DatasetError(
f"{batch['review']['pending']} frame(s) still need a decision before this "
"batch can be approved"
)
if batch["review"]["approved"] == 0:
raise DatasetError("Every frame was rejected — there is nothing to merge")
batches.set_status(batch_id, "approved")
job = jobs.create(
"merge",
params={"batch_id": batch_id},
project_id=batch["project_id"],
batch_id=batch_id,
message=f"{batch['date_label']}/{batch['batch_label']}",
)
return job.to_dict()
def _label_line(class_id: int, geometry: dict, label_type: str) -> str:
if label_type == "bbox":
x0, y0, x1, y1 = review.to_box(geometry)
return (f"{class_id} {(x0 + x1) / 2:.6f} {(y0 + y1) / 2:.6f} "
f"{x1 - x0:.6f} {y1 - y0:.6f}")
points = geometry["points"]
if geometry["type"] == "bbox":
x0, y0, x1, y1 = geometry["points"]
points = [[x0, y0], [x1, y0], [x1, y1], [x0, y1]]
coords = " ".join(f"{value:.6f}" for point in points for value in point)
return f"{class_id} {coords}"
def _next_split(cur, project_id: int, val_every: int) -> str:
"""Continue the every-Nth pattern from wherever the last merge left off."""
if val_every <= 0:
return "train"
cur.execute("SELECT COUNT(*) FROM dataset_items WHERE project_id = ?", (project_id,))
position = cur.fetchone()[0]
return "val" if position % val_every == val_every - 1 else "train"
def write_data_yaml(project: dict, batch_ids: list = None) -> str:
"""Rebuild data.yaml from the project's classes (REQ-051)."""
root = dataset_dir(project["slug"])
os.makedirs(root, exist_ok=True)
counts = summary(project["id"])["splits"]
names = ", ".join(f"'{item['name']}'" for item in project["classes"])
if batch_ids:
with db.cursor() as cur:
placeholders = ",".join("?" for _ in batch_ids)
cur.execute(
f"""SELECT d.image_rel, d.split FROM dataset_items d
JOIN frames f ON f.id = d.frame_id
WHERE d.project_id = ? AND f.batch_id IN ({placeholders})""",
[project["id"]] + list(batch_ids),
)
rows = cur.fetchall()
train_files = [row[0] for row in rows if row[1] == "train"]
val_files = [row[0] for row in rows if row[1] == "val"] or train_files
train_txt = os.path.join(root, "selected_train.txt")
val_txt = os.path.join(root, "selected_val.txt")
with open(train_txt, "w", encoding="utf-8") as handle:
handle.write("\n".join(os.path.join(root, rel) for rel in train_files) + "\n")
with open(val_txt, "w", encoding="utf-8") as handle:
handle.write("\n".join(os.path.join(root, rel) for rel in val_files) + "\n")
path = os.path.join(root, "selected_data.yaml")
with open(path, "w", encoding="utf-8") as handle:
handle.write(f"path: {root}\n")
handle.write(f"train: {train_txt}\n")
handle.write(f"val: {val_txt}\n\n")
handle.write(f"nc: {len(project['classes'])}\n")
handle.write(f"names: [{names}]\n")
return path
path = os.path.join(root, "data.yaml")
with open(path, "w", encoding="utf-8") as handle:
handle.write(f"path: {root}\n")
handle.write("train: images/train\n")
handle.write(f"val: images/{'val' if counts['val'] > 0 else 'train'}\n\n")
handle.write(f"nc: {len(project['classes'])}\n")
handle.write(f"names: [{names}]\n")
return path
def summary(project_id: int) -> dict:
with db.cursor() as cur:
cur.execute(
"SELECT split, COUNT(*) FROM dataset_items WHERE project_id = ? GROUP BY split",
(project_id,),
)
splits = {"train": 0, "val": 0}
for split, count in cur.fetchall():
splits[split] = count
cur.execute(
"""SELECT b.id, b.date_label, b.batch_label, b.merged_at,
COUNT(d.id) AS images
FROM batches b
LEFT JOIN frames f ON f.batch_id = b.id
LEFT JOIN dataset_items d ON d.frame_id = f.id
WHERE b.project_id = ? AND b.status = 'merged'
GROUP BY b.id ORDER BY b.merged_at""",
(project_id,),
)
merged = [dict(row) for row in cur.fetchall()]
return {"splits": splits, "total": splits["train"] + splits["val"], "batches": merged}
def drop_class_from_labels(project: dict, class_id: int) -> dict:
"""Rewrite every label file on disk after a class is deleted (REQ-007).
Two edits per file: lines of the deleted class go, and every id above it
comes down by one. Skipping this would leave `2` in old files meaning a
class that is now `1` — labels that quietly name the wrong thing are worse
than labels that are missing.
"""
root = dataset_dir(project["slug"])
with db.cursor() as cur:
cur.execute("SELECT label_rel FROM dataset_items WHERE project_id = ?",
(project["id"],))
label_files = [row[0] for row in cur.fetchall()]
rewritten = 0
dropped = 0
for rel in label_files:
path = os.path.join(root, rel)
if not os.path.isfile(path):
continue
with open(path, encoding="utf-8") as handle:
lines = handle.read().splitlines()
kept, touched = [], False
for line in lines:
if not line.strip():
continue
head, _, rest = line.partition(" ")
try:
current = int(head)
except ValueError:
kept.append(line)
continue
if current == class_id:
dropped += 1
touched = True
continue
if current > class_id:
current -= 1
touched = True
kept.append(f"{current} {rest}")
if touched:
# An emptied file stays as an empty file: the image is still a valid
# negative sample (REQ-033), it just has nothing on it any more.
with open(path, "w", encoding="utf-8") as handle:
handle.write("\n".join(kept) + ("\n" if kept else ""))
rewritten += 1
return {"label_files_rewritten": rewritten, "dataset_lines_removed": dropped}
def zip_path(project: dict) -> str:
"""Zip the master dataset for download (REQ-054)."""
root = dataset_dir(project["slug"])
if not os.path.isdir(os.path.join(root, "images")):
raise DatasetError("This project's dataset is still empty")
archive = os.path.join(config.project_dir(project["slug"]), "dataset")
return shutil.make_archive(archive, "zip", root)
@jobs.handler("merge")
def _run_merge(job) -> None:
batch = batches.get(job.params["batch_id"])
if batch is None:
raise DatasetError("The batch disappeared before the merge started")
project = projects.get(batch["project_id"])
root = dataset_dir(project["slug"])
for split in ("train", "val"):
os.makedirs(os.path.join(root, "images", split), exist_ok=True)
os.makedirs(os.path.join(root, "labels", split), exist_ok=True)
frames = [f for f in batches.frames(batch["id"]) if f["review_status"] == "approved"]
source_dir = batches.frames_dir(project["slug"], batch["id"])
job.progress(0, len(frames))
job.log(f"Merging {len(frames)} approved frame(s) into the master dataset")
added = {"train": 0, "val": 0}
skipped = 0
for index, frame in enumerate(frames):
if job.cancelled:
job.log(f"Cancelled after {index} frame(s)")
break
with db.cursor() as cur:
cur.execute("SELECT 1 FROM dataset_items WHERE frame_id = ?", (frame["id"],))
if cur.fetchone() is not None:
skipped += 1
job.progress(index + 1, len(frames))
continue
split = _next_split(cur, project["id"], project["val_every"])
stem = f"{batch['id']}__{os.path.splitext(frame['filename'])[0]}"
image_rel = f"images/{split}/{stem}.jpg"
label_rel = f"labels/{split}/{stem}.txt"
shutil.copyfile(os.path.join(source_dir, frame["filename"]),
os.path.join(root, image_rel))
lines = [_label_line(item["class_id"], item["geometry"], project["label_type"])
for item in review.listing(frame["id"])]
# An approved frame with nothing on it is a negative sample, and an
# empty .txt is how YOLO spells that (REQ-033).
with open(os.path.join(root, label_rel), "w", encoding="utf-8") as handle:
handle.write("\n".join(lines) + ("\n" if lines else ""))
cur.execute(
"""INSERT INTO dataset_items (project_id, frame_id, split, image_rel,
label_rel, added_at)
VALUES (?, ?, ?, ?, ?, ?)""",
(project["id"], frame["id"], split, image_rel, label_rel, time.time()),
)
added[split] += 1
job.progress(index + 1, len(frames))
with db.cursor() as cur:
cur.execute("UPDATE batches SET status = 'merged', merged_at = ? WHERE id = ?",
(time.time(), batch["id"]))
path = write_data_yaml(projects.get(project["id"]))
totals = summary(project["id"])["splits"]
job.log(f"Added {added['train']} train / {added['val']} val"
+ (f", skipped {skipped} already merged" if skipped else ""))
job.log(f"Master dataset now {totals['train']} train / {totals['val']} val — {path}")
+178
View File
@@ -0,0 +1,178 @@
"""SQLite storage for metadata and status.
The split is deliberate: this database holds *what* and *where*, the disk holds
the pixels, the final YOLO labels, and the weights. A master dataset stays
trainable even if this file is deleted (REQ-006, REQ-054).
Connections are per-call rather than shared, because the job worker runs on its
own thread and SQLite connections are not safely shared across threads. WAL mode
lets that worker write while requests read.
"""
import os
import sqlite3
from contextlib import contextmanager
from typing import Iterator
from backend import config
SCHEMA = [
"""
CREATE TABLE IF NOT EXISTS projects (
id INTEGER PRIMARY KEY AUTOINCREMENT,
slug TEXT NOT NULL UNIQUE,
name TEXT NOT NULL,
label_type TEXT NOT NULL CHECK (label_type IN ('bbox', 'polygon')),
base_model_path TEXT,
base_model_kind TEXT CHECK (base_model_kind IN ('uploaded', 'pretrained', 'trained')),
video_root TEXT NOT NULL,
val_every INTEGER NOT NULL DEFAULT 5,
created_at REAL NOT NULL
)
""",
"""
CREATE TABLE IF NOT EXISTS project_classes (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
class_id INTEGER NOT NULL,
name TEXT NOT NULL,
prompt TEXT NOT NULL,
UNIQUE (project_id, class_id)
)
""",
"""
CREATE TABLE IF NOT EXISTS batches (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
video_path TEXT NOT NULL,
date_label TEXT NOT NULL,
batch_label TEXT NOT NULL,
start_sec REAL NOT NULL,
end_sec REAL NOT NULL,
fps REAL NOT NULL,
status TEXT NOT NULL CHECK (status IN (
'extracting', 'extracted', 'labeling', 'reviewing',
'approved', 'merged', 'failed')),
frame_count INTEGER NOT NULL DEFAULT 0,
created_at REAL NOT NULL,
merged_at REAL
)
""",
"""
CREATE TABLE IF NOT EXISTS frames (
id INTEGER PRIMARY KEY AUTOINCREMENT,
batch_id INTEGER NOT NULL REFERENCES batches(id) ON DELETE CASCADE,
idx INTEGER NOT NULL,
filename TEXT NOT NULL,
width INTEGER NOT NULL,
height INTEGER NOT NULL,
review_status TEXT NOT NULL DEFAULT 'pending'
CHECK (review_status IN ('pending', 'approved', 'rejected')),
UNIQUE (batch_id, idx)
)
""",
"""
CREATE TABLE IF NOT EXISTS annotations (
id INTEGER PRIMARY KEY AUTOINCREMENT,
frame_id INTEGER NOT NULL REFERENCES frames(id) ON DELETE CASCADE,
class_id INTEGER NOT NULL,
geometry TEXT NOT NULL,
score REAL NOT NULL DEFAULT 1.0,
source TEXT NOT NULL CHECK (source IN ('auto', 'manual')),
created_at REAL NOT NULL
)
""",
"""
CREATE TABLE IF NOT EXISTS dataset_items (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
frame_id INTEGER NOT NULL REFERENCES frames(id) ON DELETE CASCADE UNIQUE,
split TEXT NOT NULL CHECK (split IN ('train', 'val')),
image_rel TEXT NOT NULL,
label_rel TEXT NOT NULL,
added_at REAL NOT NULL
)
""",
"""
CREATE TABLE IF NOT EXISTS model_versions (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER NOT NULL REFERENCES projects(id) ON DELETE CASCADE,
version INTEGER NOT NULL,
weights_path TEXT NOT NULL,
parent_model_path TEXT,
metrics TEXT,
base_metrics TEXT,
created_at REAL NOT NULL,
UNIQUE (project_id, version)
)
""",
"""
CREATE TABLE IF NOT EXISTS jobs (
id INTEGER PRIMARY KEY AUTOINCREMENT,
project_id INTEGER REFERENCES projects(id) ON DELETE CASCADE,
batch_id INTEGER REFERENCES batches(id) ON DELETE CASCADE,
type TEXT NOT NULL CHECK (type IN ('extract', 'autolabel', 'merge', 'train')),
status TEXT NOT NULL CHECK (status IN (
'queued', 'running', 'done', 'failed', 'cancelled')),
params TEXT NOT NULL DEFAULT '{}',
progress INTEGER NOT NULL DEFAULT 0,
total INTEGER NOT NULL DEFAULT 0,
message TEXT NOT NULL DEFAULT '',
error TEXT NOT NULL DEFAULT '',
log TEXT NOT NULL DEFAULT '',
created_at REAL NOT NULL,
started_at REAL,
finished_at REAL
)
""",
"CREATE INDEX IF NOT EXISTS idx_frames_batch ON frames(batch_id, idx)",
"CREATE INDEX IF NOT EXISTS idx_annotations_frame ON annotations(frame_id)",
"CREATE INDEX IF NOT EXISTS idx_batches_project ON batches(project_id)",
"CREATE INDEX IF NOT EXISTS idx_jobs_project ON jobs(project_id, created_at)",
"CREATE INDEX IF NOT EXISTS idx_dataset_items_project ON dataset_items(project_id)",
]
def connect() -> sqlite3.Connection:
os.makedirs(os.path.dirname(config.DB_PATH), exist_ok=True)
connection = sqlite3.connect(config.DB_PATH, timeout=30.0)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA journal_mode = WAL")
connection.execute("PRAGMA foreign_keys = ON")
connection.execute("PRAGMA busy_timeout = 30000")
return connection
@contextmanager
def cursor() -> Iterator[sqlite3.Cursor]:
"""Transactional cursor: commits on success, rolls back on exception."""
connection = connect()
try:
with connection:
yield connection.cursor()
finally:
connection.close()
def migrate() -> None:
"""Create every table and index. Idempotent — safe on every startup."""
with cursor() as cur:
for statement in SCHEMA:
cur.execute(statement)
cur.execute("PRAGMA table_info(projects)")
cols = [column[1] for column in cur.fetchall()]
if "secondary_model_path" not in cols:
cur.execute("ALTER TABLE projects ADD COLUMN secondary_model_path TEXT")
if "secondary_model_name" not in cols:
cur.execute("ALTER TABLE projects ADD COLUMN secondary_model_name TEXT")
if "secondary_model_classes" not in cols:
cur.execute("ALTER TABLE projects ADD COLUMN secondary_model_classes TEXT")
def healthy() -> bool:
try:
with cursor() as cur:
cur.execute("SELECT 1")
return True
except sqlite3.Error:
return False
+64
View File
@@ -0,0 +1,64 @@
"""Base model versus new model, on the same val set (REQ-063).
Both models are validated against one `data.yaml`, so the numbers differ only
because the weights differ. Combined with the stable val split in `dataset.py`,
that is what makes "the model improved" a claim rather than a hope.
"""
from typing import Optional
class EvaluateError(Exception):
pass
def class_names(weights: str) -> list:
from ultralytics import YOLO
names = YOLO(weights).names
return [names[key] for key in sorted(names)] if isinstance(names, dict) else list(names)
def validate(weights: str, data_yaml: str, imgsz: int = 640, device=0,
batch: int = 8) -> dict:
from ultralytics import YOLO
metrics = YOLO(weights).val(
data=data_yaml, imgsz=imgsz, device=device, batch=batch,
split="val", plots=False, verbose=False,
)
box = metrics.box
return {
"map50": round(float(box.map50), 4),
"map50_95": round(float(box.map), 4),
"precision": round(float(box.mp), 4),
"recall": round(float(box.mr), 4),
}
def compare(base_weights: Optional[str], new_weights: str, data_yaml: str,
expected_classes: list, imgsz: int = 640, device=0,
batch: int = 8) -> dict:
"""Validate both models where that is meaningful, and say so when it is not.
A base model whose class list differs from the project's cannot be scored on
this dataset — its class ids mean something else. Reporting nothing beats
reporting a number that looks like a regression but is a mismatch.
"""
new_metrics = validate(new_weights, data_yaml, imgsz, device, batch)
base_metrics = None
skipped = None
if not base_weights:
skipped = "This project has no base model yet — nothing to compare against."
else:
try:
base_metrics = validate(base_weights, data_yaml, imgsz, device, batch)
except Exception as exc:
base_metrics, skipped = None, f"Could not evaluate base model: {exc}"
delta = None
if base_metrics:
delta = {key: round(new_metrics[key] - base_metrics[key], 4) for key in new_metrics}
return {"base": base_metrics, "new": new_metrics, "delta": delta, "skipped": skipped}
+66
View File
@@ -0,0 +1,66 @@
"""Training defaults derived from the machine this happens to be running on.
The point is REQ-062: moving to a bigger GPU should change the numbers in the
form, not the code. Everything here is a *default* — the user can override all
of it per run.
"""
from typing import Optional
SAM3_RESIDENT_GB = 3.9
SAM3_HEADROOM_GB = 0.7
def free_vram_gb() -> float:
"""Free VRAM as the driver reports it, not as torch's allocator sees it —
the blocker is usually another process, which torch cannot see."""
import torch
if not torch.cuda.is_available():
return 0.0
free, _total = torch.cuda.mem_get_info()
return round(free / (1024 ** 3), 2)
def detect() -> dict:
import torch
if not torch.cuda.is_available():
return {"device": "cpu", "gpu": None, "vram_gb": 0.0}
properties = torch.cuda.get_device_properties(0)
return {
"device": "cuda",
"gpu": properties.name,
"vram_gb": round(properties.total_memory / (1024 ** 3), 1),
}
def defaults(epochs: int = 50) -> dict:
"""Batch size and image size that should fit, given the VRAM we can see."""
info = detect()
vram = info["vram_gb"]
if info["device"] == "cpu":
settings = {"batch": 4, "imgsz": 512, "device": "cpu", "workers": 2}
note = "No GPU visible — training on CPU will be very slow."
elif vram < 8:
settings = {"batch": 8, "imgsz": 640, "device": 0, "workers": 2}
note = f"{vram} GB of VRAM: small batches, 640 px."
elif vram <= 16:
settings = {"batch": 32, "imgsz": 640, "device": 0, "workers": 8}
note = f"{vram} GB of VRAM: optimized batch 32, 640 px."
else:
settings = {"batch": 32, "imgsz": 768, "device": 0, "workers": 8}
note = f"{vram} GB of VRAM: room for larger batches and 768 px."
return {**info, **settings, "epochs": epochs, "note": note}
def resolve(overrides: Optional[dict] = None, epochs: int = 50) -> dict:
"""Defaults with the user's overrides applied on top."""
settings = defaults(epochs)
for key, value in (overrides or {}).items():
if value is not None and key in ("batch", "imgsz", "device", "epochs", "workers"):
settings[key] = value
return settings
+255
View File
@@ -0,0 +1,255 @@
"""Job queue: one worker thread, because there is one GPU (REQ-070).
Jobs are rows in SQLite, so the list survives a restart (REQ-071). A job that was
still running when the process died is marked failed at the next startup — it
cannot be resumed, and pretending otherwise would be worse than saying so.
Handlers register themselves by job type:
@jobs.handler("extract")
def _extract(job: Job) -> None:
...
A handler reports progress with `job.progress(i, n, "...")`, writes user-facing
lines with `job.log("...")`, and checks `job.cancelled` between units of work.
Raising anything marks the job failed with that exception on it.
"""
import json
import queue
import threading
import time
import traceback
from typing import Callable, Dict, List, Optional
from backend import db
MAX_LOG_LINES = 500
PROGRESS_FLUSH_SECONDS = 0.5
JOB_TYPES = ("extract", "autolabel", "merge", "train")
GPU_JOB_TYPES = ("autolabel", "train")
"""`extract` is ffmpeg and `merge` is file copying — neither touches the card,
so neither should be able to block an interactive assist."""
gpu_lock = threading.Lock()
"""Held for the duration of any GPU work. The job worker takes it around a
handler; the interactive assist route takes it around one SAM3 call. One card,
one holder (REQ-065)."""
class Job:
"""One queued unit of work. The database row is the source of truth; this
object is the handle a handler writes through."""
def __init__(self, row):
self.id: int = row["id"]
self.type: str = row["type"]
self.project_id: Optional[int] = row["project_id"]
self.batch_id: Optional[int] = row["batch_id"]
self.params: dict = json.loads(row["params"])
self.status: str = row["status"]
self.current: int = row["progress"]
self.total: int = row["total"]
self.message: str = row["message"]
self.error: str = row["error"]
self.lines: List[str] = row["log"].splitlines() if row["log"] else []
self.created_at: float = row["created_at"]
self.started_at: Optional[float] = row["started_at"]
self.finished_at: Optional[float] = row["finished_at"]
self._flushed_at = 0.0
# ---- what handlers call ------------------------------------------------
@property
def cancelled(self) -> bool:
return self.id in _cancelled
def progress(self, current: int, total: Optional[int] = None,
message: Optional[str] = None) -> None:
self.current = current
if total is not None:
self.total = total
if message is not None:
self.message = message
# Throttled: a 3000-frame job would otherwise write 3000 times.
if time.time() - self._flushed_at >= PROGRESS_FLUSH_SECONDS:
self.flush()
def log(self, message: str) -> None:
self.lines.append(f"[{time.strftime('%H:%M:%S')}] {message}")
if len(self.lines) > MAX_LOG_LINES:
del self.lines[: len(self.lines) - MAX_LOG_LINES]
self.flush()
def flush(self) -> None:
self._flushed_at = time.time()
with db.cursor() as cur:
cur.execute(
"""UPDATE jobs SET status = ?, progress = ?, total = ?, message = ?,
error = ?, log = ?, started_at = ?, finished_at = ?
WHERE id = ?""",
(self.status, self.current, self.total, self.message, self.error,
"\n".join(self.lines), self.started_at, self.finished_at, self.id),
)
# ---- serialisation -----------------------------------------------------
def to_dict(self) -> dict:
end = self.finished_at or time.time()
return {
"id": self.id,
"type": self.type,
"project_id": self.project_id,
"batch_id": self.batch_id,
"params": self.params,
"status": self.status,
"progress": self.current,
"total": self.total,
"message": self.message,
"error": self.error,
"log": self.lines,
"created_at": self.created_at,
"started_at": self.started_at,
"finished_at": self.finished_at,
"elapsed": end - (self.started_at or self.created_at),
}
_handlers: Dict[str, Callable[[Job], None]] = {}
_queue: "queue.Queue[int]" = queue.Queue()
_cancelled: set = set()
_worker: Optional[threading.Thread] = None
_worker_lock = threading.Lock()
def handler(job_type: str):
"""Register the function that runs jobs of this type."""
if job_type not in JOB_TYPES:
raise ValueError(f"Unknown job type: {job_type}")
def decorate(function: Callable[[Job], None]) -> Callable[[Job], None]:
_handlers[job_type] = function
return function
return decorate
def create(job_type: str, params: Optional[dict] = None, project_id: Optional[int] = None,
batch_id: Optional[int] = None, message: str = "") -> Job:
if job_type not in _handlers:
raise ValueError(f"No handler registered for job type: {job_type}")
with db.cursor() as cur:
cur.execute(
"""INSERT INTO jobs (project_id, batch_id, type, status, params, message, created_at)
VALUES (?, ?, ?, 'queued', ?, ?, ?)""",
(project_id, batch_id, job_type, json.dumps(params or {}), message, time.time()),
)
job_id = cur.lastrowid
job = get(job_id)
assert job is not None
_queue.put(job_id)
_ensure_worker()
return job
def get(job_id: int) -> Optional[Job]:
with db.cursor() as cur:
cur.execute("SELECT * FROM jobs WHERE id = ?", (job_id,))
row = cur.fetchone()
return Job(row) if row else None
def listing(project_id: Optional[int] = None, limit: int = 50) -> List[Job]:
with db.cursor() as cur:
if project_id is None:
cur.execute("SELECT * FROM jobs ORDER BY created_at DESC LIMIT ?", (limit,))
else:
cur.execute(
"SELECT * FROM jobs WHERE project_id = ? ORDER BY created_at DESC LIMIT ?",
(project_id, limit),
)
return [Job(row) for row in cur.fetchall()]
def running_types() -> List[str]:
"""Job types occupying the worker right now — the GPU is shared with the
interactive SAM3 calls the review editor makes."""
with db.cursor() as cur:
cur.execute("SELECT DISTINCT type FROM jobs WHERE status = 'running'")
return [row[0] for row in cur.fetchall()]
def cancel(job_id: int) -> bool:
job = get(job_id)
if job is None or job.status in ("done", "failed", "cancelled"):
return False
_cancelled.add(job_id)
if job.status == "queued":
# Never started, so no handler will notice the flag — close it out here.
job.status = "cancelled"
job.finished_at = time.time()
job.log("Cancelled before it started")
else:
job.log("Cancellation requested")
return True
def recover() -> int:
"""Close out jobs left behind by a previous process (REQ-071)."""
with db.cursor() as cur:
cur.execute(
"""UPDATE jobs SET status = 'failed', error = 'interrupted by a server restart',
finished_at = ?
WHERE status IN ('queued', 'running')""",
(time.time(),),
)
return cur.rowcount
def _ensure_worker() -> None:
global _worker
with _worker_lock:
if _worker is None or not _worker.is_alive():
_worker = threading.Thread(target=_worker_loop, name="job-worker", daemon=True)
_worker.start()
def _worker_loop() -> None:
while True:
job_id = _queue.get()
job = get(job_id)
if job is None:
continue
if job.id in _cancelled:
_finish(job, "cancelled")
continue
_run(job)
def _run(job: Job) -> None:
job.status = "running"
job.started_at = time.time()
job.flush()
try:
if job.type in GPU_JOB_TYPES:
with gpu_lock:
_handlers[job.type](job)
else:
_handlers[job.type](job)
except Exception as exc:
job.error = f"{type(exc).__name__}: {exc}"
job.log(f"FAILED: {job.error}")
job.log(traceback.format_exc().strip().splitlines()[-1])
_finish(job, "failed")
return
_finish(job, "cancelled" if job.cancelled else "done")
def _finish(job: Job, status: str) -> None:
_cancelled.discard(job.id)
job.status = status
job.finished_at = time.time()
job.message = ""
job.flush()
+81
View File
@@ -0,0 +1,81 @@
"""Run every class prompt against one frame and return the surviving instances.
Each class is its own prompt, so prompt index is class id. Prompts overlap in
practice ("sack" and "woven plastic sack" both fire on the same object), so
detections are deduplicated across prompts by IoU, keeping the higher-scoring
one (REQ-031).
The set_image-once-per-image rule lives in `sam3_engine.detect`, which this
calls — see the domain invariants in `../AGENTS.md`.
"""
from dataclasses import dataclass, field
from typing import List, Optional
from PIL import Image
from backend.sam3_engine import Detection, get_engine
@dataclass
class ImageResult:
image_path: str
rel_path: str
width: int
height: int
detections: List[Detection] = field(default_factory=list)
error: Optional[str] = None
def _iou(box_a: List[float], box_b: List[float]) -> float:
ax0, ay0, ax1, ay1 = box_a
bx0, by0, bx1, by1 = box_b
inter_w = max(0.0, min(ax1, bx1) - max(ax0, bx0))
inter_h = max(0.0, min(ay1, by1) - max(ay0, by0))
inter = inter_w * inter_h
if inter <= 0:
return 0.0
area_a = max(0.0, ax1 - ax0) * max(0.0, ay1 - ay0)
area_b = max(0.0, bx1 - bx0) * max(0.0, by1 - by0)
union = area_a + area_b - inter
return inter / union if union > 0 else 0.0
def deduplicate(detections: List[Detection], iou_threshold: float = 0.8) -> List[Detection]:
"""Greedy NMS across all prompts: highest score wins an overlapping region."""
kept: List[Detection] = []
for det in sorted(detections, key=lambda d: d.score, reverse=True):
if all(_iou(det.box, k.box) < iou_threshold for k in kept):
kept.append(det)
return kept
def label_image(
image_path: str,
rel_path: str,
prompts: List[str],
threshold: float,
iou_threshold: float = 0.8,
min_box_frac: float = 0.0,
) -> ImageResult:
"""Detect every prompt in one image and return the surviving instances."""
try:
image = Image.open(image_path).convert("RGB")
except Exception as exc: # unreadable/corrupt frame: report, don't abort the job
return ImageResult(image_path, rel_path, 0, 0, error=str(exc))
width, height = image.size
try:
detections = get_engine().detect(image, prompts, threshold)
except Exception as exc:
return ImageResult(image_path, rel_path, width, height, error=str(exc))
if min_box_frac > 0:
floor = width * height * min_box_frac
detections = [
d for d in detections
if (d.box[2] - d.box[0]) * (d.box[3] - d.box[1]) >= floor
]
return ImageResult(image_path, rel_path, width, height,
deduplicate(detections, iou_threshold))
+117
View File
@@ -0,0 +1,117 @@
"""The video archive, read as `<video_root>/<date>/<batch>.<ext>` (REQ-010…012).
Read-only by construction: this module only ever lists and stats files, never
writes into the user's recordings (REQ-074).
"""
import os
import re
from typing import List, Optional
from backend import config, db, video
class LibraryError(Exception):
pass
def _safe_join(video_root: str, *parts: str) -> str:
"""Join under the archive root, refusing anything that escapes it."""
root = os.path.realpath(video_root)
target = os.path.realpath(os.path.join(root, *parts))
if target != root and not target.startswith(root + os.sep):
raise LibraryError("Path is outside the video archive")
return target
def batch_label(filename: str) -> str:
"""`batch-4.mp4` -> `batch-4`. The filename is the batch's identity."""
return os.path.splitext(filename)[0]
def _batch_sort_key(filename: str):
"""Sort batch-2 before batch-10, and keep unnumbered names after them."""
numbers = re.findall(r"\d+", batch_label(filename))
return (0, int(numbers[0])) if numbers else (1, 0), filename.lower()
def _effective_root(video_root: str) -> str:
"""Return video_root if it exists, otherwise fall back to config.VIDEO_ROOT."""
if os.path.isdir(video_root):
return video_root
if os.path.isdir(config.VIDEO_ROOT):
return config.VIDEO_ROOT
return video_root
def list_dates(video_root: str) -> List[dict]:
"""Date folders, newest name first, with how many videos each holds."""
video_root = _effective_root(video_root)
if not os.path.isdir(video_root):
raise LibraryError(f"Video archive folder not found: {video_root}")
dates = []
for name in sorted(os.listdir(video_root), reverse=True):
path = os.path.join(video_root, name)
if not os.path.isdir(path) or name.startswith("."):
continue
try:
count = sum(1 for f in os.listdir(path) if f.lower().endswith(video.VIDEO_EXTS))
except OSError:
continue
dates.append({"date": name, "video_count": count})
return dates
def list_videos(video_root: str, date: str, project_id: Optional[int] = None) -> List[dict]:
"""Videos in one date folder, with metadata and how often each was used."""
video_root = _effective_root(video_root)
folder = _safe_join(video_root, date)
if not os.path.isdir(folder):
raise LibraryError(f"No such date in the archive: {date}")
used = _usage(project_id)
videos = []
for filename in sorted(
(f for f in os.listdir(folder) if f.lower().endswith(video.VIDEO_EXTS)),
key=_batch_sort_key,
):
path = os.path.join(folder, filename)
entry = {
"rel": f"{date}/{filename}",
"filename": filename,
"batch_label": batch_label(filename),
"used_count": used.get(os.path.realpath(path), 0),
}
try:
entry.update(video.probe(path))
except video.VideoError as exc:
# A file ffprobe cannot read still belongs in the list, flagged —
# hiding it would look like the archive is missing recordings.
entry.update({"duration": 0.0, "width": 0, "height": 0, "fps": 0.0,
"error": str(exc)})
videos.append(entry)
return videos
def _usage(project_id: Optional[int]) -> dict:
"""How many batches already came out of each video path (REQ-012)."""
if project_id is None:
return {}
with db.cursor() as cur:
cur.execute(
"SELECT video_path, COUNT(*) FROM batches WHERE project_id = ? GROUP BY video_path",
(project_id,),
)
return {os.path.realpath(row[0]): row[1] for row in cur.fetchall()}
def resolve(video_root: str, rel: str) -> str:
"""Turn a `<date>/<file>` reference into an absolute path inside the archive."""
video_root = _effective_root(video_root)
path = _safe_join(video_root, rel)
if not os.path.isfile(path):
raise LibraryError(f"No such video: {rel}")
if not path.lower().endswith(video.VIDEO_EXTS):
raise LibraryError("That file is not a video")
return path
+75
View File
@@ -0,0 +1,75 @@
"""FastAPI server for the dataset enrichment platform.
Run it with: .venv/bin/uvicorn backend.main:app --port 8000
or: docker compose up
Routes live in `backend/api/`, one module per domain; this file only wires them
together and owns startup.
"""
import os
import shutil
from contextlib import asynccontextmanager
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from backend import config, db, jobs
from backend.api import batches, jobs as job_routes, models, projects, review
@asynccontextmanager
async def lifespan(_app: FastAPI):
config.ensure_dirs()
db.migrate()
from backend import projects as project_store
project_store.ensure_seed_project()
interrupted = jobs.recover()
if interrupted:
print(f"[startup] closed {interrupted} job(s) interrupted by the last restart")
yield
app = FastAPI(title="Dataset Enrichment", lifespan=lifespan)
app.add_middleware(
CORSMiddleware,
allow_origins=config.CORS_ORIGINS,
allow_methods=["*"],
allow_headers=["*"],
)
app.include_router(projects.router)
app.include_router(batches.router)
app.include_router(review.router)
app.include_router(models.router)
app.include_router(job_routes.router)
@app.get("/api/health")
def health() -> dict:
import torch
from backend import hardware
free_vram = hardware.free_vram_gb()
needed = hardware.SAM3_RESIDENT_GB + hardware.SAM3_HEADROOM_GB
return {
"device": "cuda" if torch.cuda.is_available() else "cpu",
"gpu": torch.cuda.get_device_name(0) if torch.cuda.is_available() else None,
"vram_free_gb": free_vram,
"sam3_ready": free_vram >= needed or _engine_loaded(),
"ffmpeg": shutil.which("ffmpeg") is not None,
"ffprobe": shutil.which("ffprobe") is not None,
"hf_token": bool(os.environ.get("HUGGING_FACE_HUB_TOKEN")),
"db": db.healthy(),
"data_dir": config.DATA_DIR,
"video_root": config.VIDEO_ROOT,
"model_loaded": _engine_loaded(),
}
def _engine_loaded() -> bool:
from backend.sam3_engine import engine_is_loaded
return engine_is_loaded()
+438
View File
@@ -0,0 +1,438 @@
"""Projects: the unit that makes this system reusable (REQ-001…006).
A project owns a base model, a locked class list, a video archive root, and its
own accumulating master dataset. Everything it produces lives under one folder,
so a project can be copied or backed up whole.
Classes come from the base model whenever there is one — `model.names` is the
only thing that keeps the master dataset, the auto-annotation prompts, and the
fine-tune consistent with each other (REQ-003).
"""
import os
import re
import shutil
import time
from typing import List, Optional
from backend import config, db
LABEL_TYPES = ("bbox", "polygon")
# Starting points when a project has no base model of its own (REQ-004).
PRETRAINED = {"bbox": "yolo11n.pt", "polygon": "yolo11n-seg.pt"}
class ProjectError(Exception):
"""Something the user can fix: a bad name, a missing folder, a locked field."""
def slugify(name: str) -> str:
slug = re.sub(r"[^a-z0-9]+", "-", name.strip().lower()).strip("-")
return slug or "project"
def _unique_slug(cur, name: str) -> str:
base = slugify(name)
slug, suffix = base, 2
while True:
cur.execute("SELECT 1 FROM projects WHERE slug = ?", (slug,))
if cur.fetchone() is None:
return slug
slug, suffix = f"{base}-{suffix}", suffix + 1
def read_model_classes(weights_path: str) -> List[str]:
"""Class names in a YOLO checkpoint, in class-id order."""
from ultralytics import YOLO
try:
names = YOLO(weights_path).names
except Exception as exc:
raise ProjectError(f"Could not read classes from that model: {exc}")
if isinstance(names, dict):
return [names[key] for key in sorted(names)]
return list(names)
def _project_paths(slug: str) -> dict:
root = config.project_dir(slug)
return {
"root": root,
"base": os.path.join(root, "base"),
"dataset": os.path.join(root, "dataset"),
"batches": os.path.join(root, "batches"),
"models": os.path.join(root, "models"),
}
def create(name: str, label_type: str, video_root: str, classes: Optional[List[dict]] = None,
base_model_path: Optional[str] = None, val_every: int = 5) -> dict:
"""Create a project. `classes` is [{"name": ..., "prompt": ...}, …] and is
ignored when a base model is given — that model's names win."""
if not name.strip():
raise ProjectError("A project name is required")
if label_type not in LABEL_TYPES:
raise ProjectError(f"label_type must be one of {LABEL_TYPES}")
video_root = os.path.abspath(os.path.expanduser(video_root))
if not os.path.isdir(video_root):
raise ProjectError(f"Video archive folder not found: {video_root}")
if base_model_path:
names = read_model_classes(base_model_path)
classes = [{"name": n, "prompt": n} for n in names]
if not classes:
raise ProjectError("Give a base model to read classes from, or list the classes")
cleaned = []
for index, item in enumerate(classes):
class_name = str(item.get("name", "")).strip()
if not class_name:
raise ProjectError(f"Class {index} has no name")
cleaned.append({"name": class_name,
"prompt": str(item.get("prompt") or class_name).strip()})
if len({c["name"] for c in cleaned}) != len(cleaned):
raise ProjectError("Class names must be unique")
with db.cursor() as cur:
slug = _unique_slug(cur, name)
paths = _project_paths(slug)
for path in paths.values():
os.makedirs(path, exist_ok=True)
stored_model = ""
kind = "pretrained"
if base_model_path:
stored_model = os.path.join(paths["base"], "model.pt")
shutil.copyfile(base_model_path, stored_model)
kind = "uploaded"
cur.execute(
"""INSERT INTO projects (slug, name, label_type, base_model_path,
base_model_kind, video_root, val_every, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
(slug, name.strip(), label_type, stored_model, kind, video_root,
max(0, val_every), time.time()),
)
project_id = cur.lastrowid
_write_classes(cur, project_id, cleaned)
return get(project_id)
def _write_classes(cur, project_id: int, classes: List[dict]) -> None:
cur.execute("DELETE FROM project_classes WHERE project_id = ?", (project_id,))
cur.executemany(
"INSERT INTO project_classes (project_id, class_id, name, prompt) VALUES (?, ?, ?, ?)",
[(project_id, index, item["name"], item["prompt"])
for index, item in enumerate(classes)],
)
def _row_to_dict(cur, row) -> dict:
cur.execute(
"SELECT class_id, name, prompt FROM project_classes WHERE project_id = ? ORDER BY class_id",
(row["id"],),
)
classes = [dict(item) for item in cur.fetchall()]
# How many shapes hang off each class — the number the user needs before
# agreeing to delete one (REQ-007).
cur.execute(
"""SELECT a.class_id, COUNT(*) FROM annotations a
JOIN frames f ON f.id = a.frame_id
JOIN batches b ON b.id = f.batch_id
WHERE b.project_id = ? GROUP BY a.class_id""",
(row["id"],),
)
usage = dict(cur.fetchall())
for item in classes:
item["annotation_count"] = usage.get(item["class_id"], 0)
cur.execute("SELECT COUNT(*) FROM batches WHERE project_id = ?", (row["id"],))
batch_count = cur.fetchone()[0]
cur.execute(
"SELECT split, COUNT(*) FROM dataset_items WHERE project_id = ? GROUP BY split",
(row["id"],),
)
dataset = {"train": 0, "val": 0}
for split, count in cur.fetchall():
dataset[split] = count
sec_classes = []
if "secondary_model_classes" in row.keys() and row["secondary_model_classes"]:
try:
sec_classes = json.loads(row["secondary_model_classes"])
except Exception:
sec_classes = []
paths = _project_paths(row["slug"])
return {
"id": row["id"],
"slug": row["slug"],
"name": row["name"],
"label_type": row["label_type"],
"base_model_path": row["base_model_path"],
"base_model_kind": row["base_model_kind"],
"secondary_model_path": row["secondary_model_path"] if "secondary_model_path" in row.keys() else None,
"secondary_model_name": row["secondary_model_name"] if "secondary_model_name" in row.keys() else None,
"secondary_model_classes": sec_classes,
"base_model_fallback": PRETRAINED[row["label_type"]],
"video_root": row["video_root"],
"val_every": row["val_every"],
"created_at": row["created_at"],
"classes": classes,
"batch_count": batch_count,
"dataset": dataset,
# Once anything has been merged the label type is settled (REQ-002).
"label_type_locked": (dataset["train"] + dataset["val"]) > 0,
"paths": paths,
}
def get(project_id: int) -> Optional[dict]:
with db.cursor() as cur:
cur.execute("SELECT * FROM projects WHERE id = ?", (project_id,))
row = cur.fetchone()
return _row_to_dict(cur, row) if row else None
def listing() -> List[dict]:
with db.cursor() as cur:
cur.execute("SELECT * FROM projects ORDER BY created_at DESC")
return [_row_to_dict(cur, row) for row in cur.fetchall()]
def update(project_id: int, prompts: Optional[dict] = None, val_every: Optional[int] = None,
video_root: Optional[str] = None) -> dict:
"""Edit the things that are safe to change: prompts, split ratio, archive root.
Class names and label type are not among them."""
if get(project_id) is None:
raise ProjectError("No such project")
with db.cursor() as cur:
if val_every is not None:
cur.execute("UPDATE projects SET val_every = ? WHERE id = ?",
(max(0, val_every), project_id))
if video_root is not None:
resolved = os.path.abspath(os.path.expanduser(video_root))
if not os.path.isdir(resolved):
raise ProjectError(f"Video archive folder not found: {resolved}")
cur.execute("UPDATE projects SET video_root = ? WHERE id = ?",
(resolved, project_id))
for class_id, prompt in (prompts or {}).items():
cur.execute(
"UPDATE project_classes SET prompt = ? WHERE project_id = ? AND class_id = ?",
(str(prompt).strip(), project_id, int(class_id)),
)
return get(project_id)
def add_class(project_id: int, name: str, prompt: Optional[str] = None) -> dict:
"""Append a class to an existing project (REQ-008).
Appending is the easy direction: the new class takes the next id, so no
existing annotation or label file means anything different afterwards. Only
`data.yaml` has to be rewritten, because it carries `nc`.
"""
project = get(project_id)
if project is None:
raise ProjectError("No such project")
clean = name.strip()
if not clean:
raise ProjectError("A class needs a name")
if any(item["name"] == clean for item in project["classes"]):
raise ProjectError(f'This project already has a class called "{clean}"')
next_id = max((item["class_id"] for item in project["classes"]), default=-1) + 1
with db.cursor() as cur:
cur.execute(
"INSERT INTO project_classes (project_id, class_id, name, prompt) "
"VALUES (?, ?, ?, ?)",
(project_id, next_id, clean, (prompt or clean).strip()),
)
updated = get(project_id)
if updated["dataset"]["train"] + updated["dataset"]["val"] > 0:
from backend import dataset
dataset.write_data_yaml(updated)
return updated
def delete_class(project_id: int, class_id: int) -> dict:
"""Remove a class and renumber the ones above it, everywhere (REQ-007).
"Everywhere" is the whole point: the database rows, and the label files
already written into the master dataset. A YOLO label is an integer index,
so a class list and a set of label files that disagree do not fail loudly —
they train a model on the wrong names.
"""
project = get(project_id)
if project is None:
raise ProjectError("No such project")
target = next((c for c in project["classes"] if c["class_id"] == class_id), None)
if target is None:
raise ProjectError(f"This project has no class {class_id}")
if len(project["classes"]) == 1:
raise ProjectError("A project needs at least one class")
from backend import dataset
with db.cursor() as cur:
frames_of_project = """
SELECT f.id FROM frames f
JOIN batches b ON b.id = f.batch_id
WHERE b.project_id = ?
"""
cur.execute(
f"DELETE FROM annotations WHERE class_id = ? AND frame_id IN ({frames_of_project})",
(class_id, project_id),
)
removed = cur.rowcount
cur.execute(
f"""UPDATE annotations SET class_id = class_id - 1
WHERE class_id > ? AND frame_id IN ({frames_of_project})""",
(class_id, project_id),
)
cur.execute("DELETE FROM project_classes WHERE project_id = ? AND class_id = ?",
(project_id, class_id))
cur.execute(
"UPDATE project_classes SET class_id = class_id - 1 "
"WHERE project_id = ? AND class_id > ?",
(project_id, class_id),
)
report = dataset.drop_class_from_labels(project, class_id)
updated = get(project_id)
dataset.write_data_yaml(updated)
return {
"project": updated,
"removed": {
"class": target["name"],
"annotations": removed,
**report,
},
}
def set_base_model(project_id: int, weights_path: str) -> dict:
"""Point the project at a new base model and re-read its classes (REQ-003).
Refused once the master dataset exists and the new model's classes differ —
a dataset labelled against one class list cannot be trained against another.
"""
project = get(project_id)
if project is None:
raise ProjectError("No such project")
names = read_model_classes(weights_path)
existing = [item["name"] for item in project["classes"]]
if project["dataset"]["train"] + project["dataset"]["val"] > 0 and set(names) != set(existing):
raise ProjectError(
"That model's classes differ from the ones this project's dataset was "
f"labelled with ({existing} vs {names}). Create a new project for it."
)
paths = _project_paths(project["slug"])
os.makedirs(paths["base"], exist_ok=True)
stored = os.path.join(paths["base"], "model.pt")
if os.path.abspath(weights_path) != os.path.abspath(stored):
shutil.copyfile(weights_path, stored)
with db.cursor() as cur:
cur.execute(
"UPDATE projects SET base_model_path = ?, base_model_kind = 'uploaded' WHERE id = ?",
(stored, project_id),
)
_write_classes(cur, project_id,
[{"name": n, "prompt": p}
for n, p in zip(names, _kept_prompts(project, names))])
return get(project_id)
def set_secondary_model(project_id: int, weights_path: str, name: str = "") -> dict:
"""Point the project at a secondary model for auto-annotation."""
project = get(project_id)
if project is None:
raise ProjectError("No such project")
paths = _project_paths(project["slug"])
os.makedirs(paths["base"], exist_ok=True)
stored = os.path.join(paths["base"], "secondary_model.pt")
if os.path.abspath(weights_path) != os.path.abspath(stored):
shutil.copyfile(weights_path, stored)
names = read_model_classes(weights_path)
model_label = name.strip() or os.path.basename(weights_path)
with db.cursor() as cur:
cur.execute(
"UPDATE projects SET secondary_model_path = ?, secondary_model_name = ?, secondary_model_classes = ? WHERE id = ?",
(stored, model_label, json.dumps(names), project_id),
)
return get(project_id)
def _kept_prompts(project: dict, names: List[str]) -> List[str]:
"""Keep the prompt the user already wrote for a class that survives a
base-model swap; fall back to the class name for new ones."""
known = {item["name"]: item["prompt"] for item in project["classes"]}
return [known.get(name, name) for name in names]
def training_start_point(project: dict) -> str:
"""The weights a training run should start from (REQ-060, REQ-004)."""
path = project["base_model_path"]
if path and os.path.isfile(path):
return path
return PRETRAINED[project["label_type"]]
def delete(project_id: int) -> bool:
project = get(project_id)
if project is None:
return False
with db.cursor() as cur:
cur.execute("DELETE FROM projects WHERE id = ?", (project_id,))
shutil.rmtree(project["paths"]["root"], ignore_errors=True)
return True
def ensure_seed_project() -> None:
"""Ensure at least one project exists on startup using legacy data if available."""
with db.cursor() as cur:
cur.execute("SELECT COUNT(*) FROM projects")
if cur.fetchone()[0] > 0:
return
legacy_model = "/data/base_models/best.pt"
if not os.path.exists(legacy_model):
legacy_model = "/videos/model/best.pt"
if not os.path.exists(legacy_model):
legacy_model = os.path.join(config.VIDEO_ROOT, "model", "best.pt")
if os.path.exists(legacy_model):
try:
create(
name="Cargo & Sack Detection",
label_type="bbox",
video_root=config.VIDEO_ROOT,
base_model_path=legacy_model,
)
print("[startup] Seeded default project 'Cargo & Sack Detection' from legacy model")
return
except Exception as exc:
print(f"[startup] Failed to seed project from legacy model: {exc}")
try:
create(
name="Default Detection Project",
label_type="bbox",
video_root=config.VIDEO_ROOT,
classes=[{"name": "sack", "prompt": "sack"}, {"name": "truck", "prompt": "truck"}],
)
print("[startup] Seeded 'Default Detection Project'")
except Exception as exc:
print(f"[startup] Seed project creation skipped: {exc}")
+350
View File
@@ -0,0 +1,350 @@
"""Annotations and per-frame review state (REQ-040…045).
Geometry is stored normalized 0–1 against the frame, as JSON:
bbox {"type": "bbox", "points": [x0, y0, x1, y1]}
polygon {"type": "polygon", "points": [[x, y], …]}
Normalized because the editor scales the frame to whatever the window allows,
and the exporter needs the same numbers YOLO wants — neither should care about
the display size.
`source` separates what SAM3 produced from what the user drew. Re-running
auto-annotation replaces only the former (REQ-034).
"""
import json
import time
from typing import List, Optional
from backend import db
STATUSES = ("pending", "approved", "rejected")
class ReviewError(Exception):
pass
# ---- geometry ----------------------------------------------------------
def _clamp(value: float) -> float:
return max(0.0, min(1.0, float(value)))
def bbox(x0: float, y0: float, x1: float, y1: float) -> dict:
left, right = sorted((_clamp(x0), _clamp(x1)))
top, bottom = sorted((_clamp(y0), _clamp(y1)))
return {"type": "bbox", "points": [left, top, right, bottom]}
def polygon(points) -> dict:
return {"type": "polygon", "points": [[_clamp(x), _clamp(y)] for x, y in points]}
def mask_to_polygons(mask, min_area_px: int = 24, max_polygons: int = 1) -> list:
"""Contour a boolean SAM3 mask into polygon point arrays, largest first.
YOLO-seg expects one polygon per instance, so only the largest connected
component is kept by default — SAM3 masks are occasionally speckled. The raw
contour is one point per pixel step, which would make label files enormous
for no accuracy gain, so it is simplified first.
"""
import cv2
import numpy as np
contours, _ = cv2.findContours((mask.astype(np.uint8)) * 255,
cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)
polygons = []
for contour in sorted(contours, key=cv2.contourArea, reverse=True)[:max_polygons]:
if cv2.contourArea(contour) < min_area_px:
continue
epsilon = 0.002 * cv2.arcLength(contour, True)
approx = cv2.approxPolyDP(contour, epsilon, True).reshape(-1, 2)
if approx.shape[0] >= 3:
polygons.append(approx.astype(np.float32))
return polygons
def validate(geometry: dict, label_type: str) -> dict:
"""Reject shapes that would export as broken labels."""
if not isinstance(geometry, dict):
raise ReviewError("geometry must be an object")
kind = geometry.get("type")
points = geometry.get("points") or []
if kind == "bbox":
if len(points) != 4:
raise ReviewError("a bbox needs [x0, y0, x1, y1]")
shape = bbox(*points)
left, top, right, bottom = shape["points"]
if right - left < 0.002 or bottom - top < 0.002:
raise ReviewError("that box is too small to be a label")
return shape
if kind == "polygon":
if len(points) < 3:
raise ReviewError("a polygon needs at least 3 points")
if label_type == "bbox":
raise ReviewError("this project stores boxes, not polygons")
return polygon(points)
raise ReviewError(f"unknown geometry type: {kind}")
def to_box(geometry: dict) -> List[float]:
"""The bounding box of any shape, normalized — used for the bbox export
and for deduplicating polygon detections against each other."""
points = geometry["points"]
if geometry["type"] == "bbox":
return list(points)
xs = [point[0] for point in points]
ys = [point[1] for point in points]
return [min(xs), min(ys), max(xs), max(ys)]
# ---- frames ------------------------------------------------------------
def frame(frame_id: int) -> Optional[dict]:
with db.cursor() as cur:
cur.execute(
"""SELECT f.*, b.id AS batch_id, b.project_id, p.slug AS project_slug,
p.label_type
FROM frames f
JOIN batches b ON b.id = f.batch_id
JOIN projects p ON p.id = b.project_id
WHERE f.id = ?""",
(frame_id,),
)
row = cur.fetchone()
return dict(row) if row else None
def set_status(frame_id: int, status: str) -> dict:
if status not in STATUSES:
raise ReviewError(f"status must be one of {STATUSES}")
if frame(frame_id) is None:
raise ReviewError("No such frame")
with db.cursor() as cur:
cur.execute("UPDATE frames SET review_status = ? WHERE id = ?", (status, frame_id))
return {"frame_id": frame_id, "review_status": status}
def next_pending(batch_id: int, after_idx: int = -1) -> Optional[int]:
"""The next frame still needing a decision, for the 'jump to unreviewed' key."""
with db.cursor() as cur:
cur.execute(
"""SELECT id FROM frames
WHERE batch_id = ? AND review_status = 'pending' AND idx > ?
ORDER BY idx LIMIT 1""",
(batch_id, after_idx),
)
row = cur.fetchone()
if row:
return row["id"]
cur.execute(
"""SELECT id FROM frames WHERE batch_id = ? AND review_status = 'pending'
ORDER BY idx LIMIT 1""",
(batch_id,),
)
row = cur.fetchone()
return row["id"] if row else None
# ---- annotations -------------------------------------------------------
def _row_to_dict(row) -> dict:
return {
"id": row["id"],
"frame_id": row["frame_id"],
"class_id": row["class_id"],
"geometry": json.loads(row["geometry"]),
"score": row["score"],
"source": row["source"],
}
def listing(frame_id: int) -> List[dict]:
with db.cursor() as cur:
cur.execute("SELECT * FROM annotations WHERE frame_id = ? ORDER BY id", (frame_id,))
return [_row_to_dict(row) for row in cur.fetchall()]
def add(frame_id: int, class_id: int, geometry: dict, source: str = "manual",
score: float = 1.0) -> dict:
target = frame(frame_id)
if target is None:
raise ReviewError("No such frame")
shape = validate(geometry, target["label_type"])
_check_class(target["project_id"], class_id)
with db.cursor() as cur:
cur.execute(
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
VALUES (?, ?, ?, ?, ?, ?)""",
(frame_id, class_id, json.dumps(shape), score, source, time.time()),
)
cur.execute("SELECT * FROM annotations WHERE id = ?", (cur.lastrowid,))
return _row_to_dict(cur.fetchone())
def update(annotation_id: int, class_id: Optional[int] = None,
geometry: Optional[dict] = None) -> dict:
with db.cursor() as cur:
cur.execute("SELECT * FROM annotations WHERE id = ?", (annotation_id,))
row = cur.fetchone()
if row is None:
raise ReviewError("No such annotation")
target = frame(row["frame_id"])
new_geometry = row["geometry"]
if geometry is not None:
new_geometry = json.dumps(validate(geometry, target["label_type"]))
new_class = row["class_id"] if class_id is None else class_id
_check_class(target["project_id"], new_class)
with db.cursor() as cur:
# Any edit makes it the user's shape, so it survives a re-run of
# auto-annotation (REQ-034).
cur.execute(
"UPDATE annotations SET class_id = ?, geometry = ?, source = 'manual' WHERE id = ?",
(new_class, new_geometry, annotation_id),
)
cur.execute("SELECT * FROM annotations WHERE id = ?", (annotation_id,))
return _row_to_dict(cur.fetchone())
def delete(annotation_id: int) -> bool:
with db.cursor() as cur:
cur.execute("DELETE FROM annotations WHERE id = ?", (annotation_id,))
return cur.rowcount > 0
def replace_auto(frame_id: int, items: List[dict]) -> int:
"""Swap this frame's automatic shapes for a fresh set, leaving manual ones."""
with db.cursor() as cur:
cur.execute("DELETE FROM annotations WHERE frame_id = ? AND source = 'auto'",
(frame_id,))
cur.executemany(
"""INSERT INTO annotations (frame_id, class_id, geometry, score, source, created_at)
VALUES (?, ?, ?, ?, 'auto', ?)""",
[(frame_id, item["class_id"], json.dumps(item["geometry"]),
item.get("score", 1.0), time.time()) for item in items],
)
return len(items)
def frames_with_auto(batch_id: int) -> set:
"""Frame ids that already carry automatic shapes — the resume skip-list
for REQ-035."""
with db.cursor() as cur:
cur.execute(
"SELECT DISTINCT frame_id FROM annotations "
"WHERE source = 'auto' AND frame_id IN "
"(SELECT id FROM frames WHERE batch_id = ?)",
(batch_id,),
)
return {row[0] for row in cur.fetchall()}
def assist(frame_id: int, box: List[float], class_id: int = 0,
threshold: float = 0.5) -> dict:
"""Drag a rough box, get SAM3's shape for the object inside it (REQ-043).
The box is a visual exemplar rather than a crop: SAM3 may return several
matches, so the one overlapping what the user drew is the one kept.
"""
from backend import batches, jobs
from backend.sam3_engine import get_engine
from PIL import Image
target = frame(frame_id)
if target is None:
raise ReviewError("No such frame")
_check_class(target["project_id"], class_id)
# Acquire the shared GPU lock with a 20s timeout. 20 seconds is chosen so that
# short CPU/ffmpeg jobs let assist through, while long GPU jobs fail fast with
# a legible message (REQ-065, REQ-070).
if not jobs.gpu_lock.acquire(timeout=20):
busy = jobs.running_types()
kind = busy[0] if busy else "background"
raise ReviewError(
f"The GPU is busy with a {kind} job — wait for it to finish, or draw the "
"shape by hand"
)
try:
drawn = validate({"type": "bbox", "points": box}, "bbox")["points"]
x0, y0, x1, y1 = drawn
exemplar = [(x0 + x1) / 2, (y0 + y1) / 2, x1 - x0, y1 - y0]
path = batches.frame_path(frame_id)
with Image.open(path) as handle:
image = handle.convert("RGB")
width, height = image.size
engine = get_engine()
state = engine.open_state(image)
found = engine.apply_prompts(
state, threshold=threshold,
exemplars=[{"box": exemplar, "positive": True}],
)
if not found:
raise ReviewError("SAM3 found nothing in that box — draw it tighter, or add the "
"shape by hand")
detection = max(found, key=lambda d: _overlap(d.box, drawn, width, height))
if target["label_type"] == "bbox":
bx0, by0, bx1, by1 = detection.box
geometry = bbox(bx0 / width, by0 / height, bx1 / width, by1 / height)
else:
polygons = mask_to_polygons(detection.mask)
if not polygons:
raise ReviewError("SAM3's mask was too small to turn into a polygon")
geometry = polygon([(x / width, y / height) for x, y in polygons[0]])
finally:
jobs.gpu_lock.release()
return add(frame_id, class_id, geometry, source="manual", score=detection.score)
def _overlap(detection_box: List[float], drawn: List[float],
width: int, height: int) -> float:
"""IoU between a pixel-space detection and the normalized box drawn."""
box = [detection_box[0] / width, detection_box[1] / height,
detection_box[2] / width, detection_box[3] / height]
ix0, iy0 = max(box[0], drawn[0]), max(box[1], drawn[1])
ix1, iy1 = min(box[2], drawn[2]), min(box[3], drawn[3])
inter = max(0.0, ix1 - ix0) * max(0.0, iy1 - iy0)
if inter <= 0:
return 0.0
area_box = (box[2] - box[0]) * (box[3] - box[1])
area_drawn = (drawn[2] - drawn[0]) * (drawn[3] - drawn[1])
return inter / (area_box + area_drawn - inter)
def clear_batch_class_annotations(batch_id: int, class_id: int) -> int:
"""Delete all annotations matching class_id across all frames in a batch (REQ-046)."""
with db.cursor() as cur:
cur.execute(
"""DELETE FROM annotations
WHERE class_id = ? AND frame_id IN (
SELECT id FROM frames WHERE batch_id = ?
)""",
(class_id, batch_id),
)
return cur.rowcount
def _check_class(project_id: int, class_id: int) -> None:
with db.cursor() as cur:
cur.execute(
"SELECT 1 FROM project_classes WHERE project_id = ? AND class_id = ?",
(project_id, class_id),
)
if cur.fetchone() is None:
raise ReviewError(f"Class {class_id} does not exist in this project")
+242
View File
@@ -0,0 +1,242 @@
"""SAM3 text-prompted detection, wrapped for reuse across labeling jobs.
The model is expensive to build (weights come from the gated HuggingFace repo
`facebook/sam3`), so it is loaded once per process and kept resident.
The important performance detail: `Sam3Processor.set_image()` runs the vision
backbone, while `set_text_prompt()` only runs the (much cheaper) grounding head
against the cached `backbone_out`. So for an N-prompt job we call `set_image`
once per image and loop the prompts over that same state.
"""
import os
import sys
import threading
from dataclasses import dataclass, field
from typing import List, Optional
import numpy as np
import torch
from PIL import Image
# Fix python import path masking issue where sam3 is imported as a namespace package
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "sam3")))
from sam3.model.sam3_image_processor import Sam3Processor
from sam3.model_builder import build_sam3_image_model
@dataclass
class Detection:
"""One labeled instance in one image."""
class_id: int
class_name: str = ""
score: float = 0.0
box: List[float] = field(default_factory=list) # xyxy in pixels
mask: Optional[np.ndarray] = None # bool array, (H, W) at original image size
class Sam3Engine:
def __init__(self, checkpoint_path: Optional[str] = None):
# SAM3 is CUDA-only in practice: `PositionEmbeddingSine` precomputes its
# tables with a hardcoded `device="cuda"`, so a CPU run dies deep inside
# the backbone with an unrelated-looking error. Fail here instead, where
# the message can say something useful.
if not torch.cuda.is_available():
raise RuntimeError(
"SAM3 requires a CUDA GPU. No GPU is visible to torch — check "
"`nvidia-smi`, that CUDA_VISIBLE_DEVICES isn't set to empty, and "
"that this venv has a CUDA build of torch installed."
)
self.device = "cuda"
self.autocast_dtype = torch.float16
# `enable_inst_interactivity=True` builds a SAM1-style click predictor,
# but in this vendored copy its `image_encoder` is None and its expected
# feature sizes (288/144/72) don't match the 1008px image pipeline, so
# `predictor.set_image()` always fails. It costs ~0.4 GB for nothing, so
# it stays off. Box exemplars cover the interactive use case instead.
self.supports_tap = False
self.model = build_sam3_image_model(
device=self.device,
checkpoint_path=checkpoint_path,
load_from_HF=checkpoint_path is None,
)
self.processor = Sam3Processor(self.model, device=self.device)
def detect(self, image: Image.Image, prompts: List[str], threshold: float) -> List[Detection]:
"""Run every prompt against one image; prompt index becomes the class id."""
self.processor.confidence_threshold = threshold
detections: List[Detection] = []
with torch.autocast(self.device, dtype=self.autocast_dtype):
state = self.processor.set_image(image)
for class_id, prompt in enumerate(prompts):
output = self.processor.set_text_prompt(prompt=prompt, state=state)
masks, boxes, scores = output["masks"], output["boxes"], output["scores"]
if masks.shape[0] == 0:
continue
# Pull off the GPU immediately: the next prompt overwrites these
# tensors, and full-resolution masks are the memory hog here.
masks_np = masks.squeeze(1).to(torch.uint8).cpu().numpy().astype(bool)
boxes_np = boxes.float().cpu().numpy()
scores_np = scores.float().cpu().numpy()
for i in range(masks_np.shape[0]):
detections.append(
Detection(
class_id=class_id,
class_name=prompt,
score=float(scores_np[i]),
box=[float(v) for v in boxes_np[i]],
mask=masks_np[i],
)
)
del state
if self.device == "cuda":
torch.cuda.empty_cache()
return detections
# ---- interactive / exemplar prompting ------------------------------
def open_state(self, image: Image.Image):
"""Run the vision backbone once and hand back the reusable state."""
with torch.autocast(self.device, dtype=self.autocast_dtype):
return self.processor.set_image(image)
def apply_prompts(
self,
state,
threshold: float,
text: Optional[str] = None,
exemplars: Optional[List[dict]] = None,
) -> List[Detection]:
"""Re-run grounding for this image from scratch with the given prompts.
Exemplars are boxes in normalized cxcywh with a positive/negative flag.
The prompt set is always replayed from empty because SAM3 only supports
appending geometric prompts — that's how undo is implemented.
"""
self.processor.confidence_threshold = threshold
exemplars = exemplars or []
with torch.autocast(self.device, dtype=self.autocast_dtype):
self.processor.reset_all_prompts(state)
output = None
if text:
output = self.processor.set_text_prompt(prompt=text, state=state)
for exemplar in exemplars:
output = self.processor.add_geometric_prompt(
box=exemplar["box"], label=bool(exemplar.get("positive", True)),
state=state,
)
if output is None:
return []
return self._collect(output, class_id=0, class_name=text or "visual")
def segment_at(
self,
image: Image.Image,
points: Optional[List[List[float]]] = None,
labels: Optional[List[int]] = None,
box: Optional[List[float]] = None,
) -> Optional[Detection]:
"""Tap-to-segment: one point (or box) in pixels -> that object's mask."""
if not self.supports_tap:
return None
predictor = self.model.inst_interactive_predictor
predictor.set_image(np.array(image))
masks, scores, _ = predictor.predict(
point_coords=np.array(points, dtype=np.float32) if points else None,
point_labels=np.array(labels, dtype=np.int32) if labels else None,
box=np.array(box, dtype=np.float32) if box else None,
multimask_output=True,
)
if masks.shape[0] == 0:
return None
best = int(np.argmax(scores))
mask = masks[best].astype(bool)
ys, xs = np.where(mask)
if xs.size == 0:
return None
return Detection(
class_id=0,
class_name="tap",
score=float(scores[best]),
box=[float(xs.min()), float(ys.min()), float(xs.max()), float(ys.max())],
mask=mask,
)
def _collect(self, output, class_id: int, class_name: str) -> List[Detection]:
masks, boxes, scores = output["masks"], output["boxes"], output["scores"]
if masks.shape[0] == 0:
return []
masks_np = masks.squeeze(1).to(torch.uint8).cpu().numpy().astype(bool)
boxes_np = boxes.float().cpu().numpy()
scores_np = scores.float().cpu().numpy()
return [
Detection(
class_id=class_id,
class_name=class_name,
score=float(scores_np[i]),
box=[float(v) for v in boxes_np[i]],
mask=masks_np[i],
)
for i in range(masks_np.shape[0])
]
_engine: Optional[Sam3Engine] = None
_engine_lock = threading.Lock()
def get_engine() -> Sam3Engine:
"""Build the model on first use, then hand out the same instance."""
global _engine
from backend import hardware
with _engine_lock:
if _engine is None:
needed = hardware.SAM3_RESIDENT_GB + hardware.SAM3_HEADROOM_GB
free = hardware.free_vram_gb()
if free < needed:
raise RuntimeError(
f"SAM3 needs ~{needed:.1f} GB free but only {free:.1f} GB is available. "
"Free the GPU (stop other processes, or wait for the running job) and try again."
)
try:
_engine = Sam3Engine()
except (ImportError, RuntimeError) as exc:
curr_free = hardware.free_vram_gb()
raise RuntimeError(
f"{exc} (Available VRAM: {curr_free:.1f} GB)"
) from exc
return _engine
def engine_is_loaded() -> bool:
return _engine is not None
def release_engine() -> bool:
"""Drop the model and free its VRAM (REQ-065).
SAM3 holds ~3.4 GB resident. On a 6 GB card that is most of the memory a
training run needs, so the two must never be loaded at once. The next job
that needs SAM3 rebuilds it from the local cache in about 12 seconds.
"""
global _engine
import gc
with _engine_lock:
if _engine is None:
return False
_engine = None
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
return True
+197
View File
@@ -0,0 +1,197 @@
"""Fine-tune the project's base model on its master dataset (REQ-060…065).
The default is old + new together: the master dataset already accumulates every
merged batch, so a run sees the whole history. Training on the newest batch
alone is what makes a model quietly forget what it used to know, so it is not
what happens here.
"""
import json
import os
import shutil
import time
from typing import Optional
from backend import config, dataset, db, evaluate, hardware, jobs, projects
PRETRAINED = {"bbox": "yolo11n.pt", "polygon": "yolo11n-seg.pt"}
class TrainingError(Exception):
pass
def models_dir(project_slug: str) -> str:
return os.path.join(config.project_dir(project_slug), "models")
def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, batch_ids: Optional[list] = None) -> dict:
project = projects.get(project_id)
if project is None:
raise TrainingError("No such project")
counts = dataset.summary(project_id)["splits"]
if counts["train"] == 0:
raise TrainingError(
"The master dataset is empty — approve and merge a batch before training"
)
settings = hardware.resolve(overrides, epochs)
job = jobs.create(
"train",
params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids},
project_id=project_id,
message=f"{counts['train']} train / {counts['val']} val",
)
return job.to_dict()
def listing(project_id: int) -> list:
with db.cursor() as cur:
cur.execute(
"SELECT * FROM model_versions WHERE project_id = ? ORDER BY version DESC",
(project_id,),
)
rows = []
for row in cur.fetchall():
item = dict(row)
item["metrics"] = json.loads(item["metrics"] or "null")
item["base_metrics"] = json.loads(item["base_metrics"] or "null")
rows.append(item)
return rows
def get_version(model_id: int) -> Optional[dict]:
with db.cursor() as cur:
cur.execute("SELECT * FROM model_versions WHERE id = ?", (model_id,))
row = cur.fetchone()
return dict(row) if row else None
def promote(model_id: int) -> dict:
"""Make a trained version the project's base model for the next round (REQ-064)."""
version = get_version(model_id)
if version is None:
raise TrainingError("No such model version")
project = projects.get(version["project_id"])
base_path = os.path.join(config.project_dir(project["slug"]), "base", "model.pt")
os.makedirs(os.path.dirname(base_path), exist_ok=True)
shutil.copyfile(version["weights_path"], base_path)
with db.cursor() as cur:
cur.execute(
"UPDATE projects SET base_model_path = ?, base_model_kind = 'trained' WHERE id = ?",
(base_path, project["id"]),
)
return projects.get(project["id"])
def _next_version(cur, project_id: int) -> int:
cur.execute(
"SELECT COALESCE(MAX(version), 0) + 1 FROM model_versions WHERE project_id = ?",
(project_id,),
)
return cur.fetchone()[0]
@jobs.handler("train")
def _run_train(job) -> None:
from ultralytics import YOLO
os.environ["ULTRALYTICS_OFFLINE"] = "true"
os.environ["YOLO_OFFLINE"] = "true"
project = projects.get(job.params["project_id"])
settings = job.params["settings"]
batch_ids = job.params.get("batch_ids")
data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids)
# SAM3 and a training run must not hold VRAM at the same time (REQ-065).
from backend.sam3_engine import release_engine
if release_engine():
job.log("Released SAM3 from VRAM before training")
start_point = project["base_model_path"] or PRETRAINED[project["label_type"]]
if not project["base_model_path"]:
job.log(f"No base model on this project — starting from {start_point}")
job.log(f"Fine-tuning {os.path.basename(start_point)} for {settings['epochs']} epoch(s) "
f"(batch={settings['batch']}, imgsz={settings['imgsz']}, "
f"device={settings['device']})")
with db.cursor() as cur:
version = _next_version(cur, project["id"])
out_dir = os.path.join(models_dir(project["slug"]), str(version))
os.makedirs(out_dir, exist_ok=True)
model = YOLO(start_point)
def on_epoch(trainer):
# trainer.epoch is 0-based; report a human-facing 1-based count.
epoch = getattr(trainer, 'epoch', 0) + 1
total = getattr(trainer, 'epochs', settings["epochs"])
job.progress(epoch, total, f"epoch {epoch}/{total}")
model.add_callback("on_fit_epoch_end", on_epoch)
job.progress(0, settings["epochs"])
keep_run_dir = False
try:
model.train(
data=data_yaml,
epochs=settings["epochs"],
imgsz=settings["imgsz"],
batch=settings["batch"],
device=settings["device"],
workers=settings.get("workers", 8),
cache=False,
project=os.path.join(out_dir, "runs"),
name="train",
exist_ok=True,
amp=True,
plots=False,
verbose=False,
)
produced = os.path.join(out_dir, "runs", "train", "weights", "best.pt")
if not os.path.isfile(produced):
raise TrainingError("Training finished without producing best.pt")
weights = os.path.join(out_dir, "best.pt")
shutil.copyfile(produced, weights)
job.log("Validating the base model and the new one on the same val set…")
comparison = evaluate.compare(
project["base_model_path"] or None, weights, data_yaml,
[item["name"] for item in project["classes"]],
imgsz=settings["imgsz"], device=settings["device"], batch=settings["batch"],
)
with open(os.path.join(out_dir, "metrics.json"), "w", encoding="utf-8") as handle:
json.dump(comparison, handle, indent=2)
with db.cursor() as cur:
cur.execute(
"""INSERT INTO model_versions (project_id, version, weights_path,
parent_model_path, metrics, base_metrics,
created_at)
VALUES (?, ?, ?, ?, ?, ?, ?)""",
(project["id"], version, weights, project["base_model_path"],
json.dumps(comparison["new"]), json.dumps(comparison["base"]), time.time()),
)
new = comparison["new"]
if comparison["delta"]:
delta = comparison["delta"]
job.log(f"v{version}: mAP50 {new['map50']:.4f} ({delta['map50']:+.4f} vs base), "
f"mAP50-95 {new['map50_95']:.4f} ({delta['map50_95']:+.4f})")
else:
job.log(f"v{version}: mAP50 {new['map50']:.4f}, mAP50-95 {new['map50_95']:.4f} "
f"— {comparison['skipped']}")
except Exception:
# Keep the runs directory on failure: results.csv and logs are the only record
# of why training failed (REQ-006, REQ-064).
keep_run_dir = True
raise
finally:
# Delete run directory on success or cancellation. job.cancelled finishes training
# without raising an exception, so it takes this delete path.
if not keep_run_dir:
shutil.rmtree(os.path.join(out_dir, "runs"), ignore_errors=True)
+139
View File
@@ -0,0 +1,139 @@
"""ffmpeg/ffprobe wrappers: read a video's metadata, cut a range into frames.
ffmpeg is a hard dependency rather than an OpenCV fallback because seeking to a
timestamp in a long recording has to be exact — an off-by-a-few-seconds trim
silently produces frames of the wrong thing (REQ-020…022).
"""
import json
import os
import shutil
import subprocess
from typing import Callable, List, Optional
VIDEO_EXTS = (".mp4", ".mkv", ".mov", ".avi", ".webm", ".m4v")
_probe_cache: dict = {}
class VideoError(Exception):
pass
def available() -> bool:
return shutil.which("ffmpeg") is not None and shutil.which("ffprobe") is not None
def probe(path: str) -> dict:
"""Duration, resolution and fps of a video, cached by (path, mtime, size).
The cache matters: the Library page probes every file in a date folder, and
ffprobe on a cold cache over dozens of long recordings is slow (REQ-012).
"""
try:
stat = os.stat(path)
except OSError as exc:
raise VideoError(str(exc))
key = (os.path.realpath(path), stat.st_mtime, stat.st_size)
if key in _probe_cache:
return _probe_cache[key]
if not available():
raise VideoError("ffprobe is not installed in this environment")
result = subprocess.run(
["ffprobe", "-v", "error", "-select_streams", "v:0",
"-show_entries", "stream=width,height,avg_frame_rate",
"-show_entries", "format=duration",
"-of", "json", path],
capture_output=True, text=True,
)
if result.returncode != 0:
raise VideoError(result.stderr.strip().splitlines()[-1] if result.stderr else "ffprobe failed")
payload = json.loads(result.stdout or "{}")
streams = payload.get("streams") or [{}]
info = {
"duration": float(payload.get("format", {}).get("duration") or 0.0),
"width": int(streams[0].get("width") or 0),
"height": int(streams[0].get("height") or 0),
"fps": _parse_fps(streams[0].get("avg_frame_rate")),
"size": stat.st_size,
}
_probe_cache[key] = info
return info
def _parse_fps(value: Optional[str]) -> float:
# ffprobe reports "30000/1001", and "0/0" for streams it cannot work out.
if not value or "/" not in value:
return 0.0
numerator, denominator = value.split("/", 1)
try:
return round(float(numerator) / float(denominator), 3) if float(denominator) else 0.0
except ValueError:
return 0.0
def frame_count(start_sec: float, end_sec: float, fps: float) -> int:
"""How many frames a trim will produce — shown before extraction (REQ-021)."""
span = max(0.0, end_sec - start_sec)
return max(0, int(span * fps))
def extract_frames(
video_path: str,
out_dir: str,
start_sec: float,
end_sec: float,
fps: float,
on_progress: Optional[Callable[[int], None]] = None,
should_stop: Optional[Callable[[], bool]] = None,
) -> List[str]:
"""Cut [start, end] at `fps` into numbered JPEGs. Returns their filenames.
`-ss` before `-i` seeks by keyframe (fast) and ffmpeg then decodes accurately
from there, which is what makes a two-minute cut out of a two-hour recording
quick instead of a full decode.
"""
if not available():
raise VideoError("ffmpeg is not installed in this environment")
if end_sec <= start_sec:
raise VideoError("The end of the range must be after its start")
if fps <= 0:
raise VideoError("fps must be greater than 0")
os.makedirs(out_dir, exist_ok=True)
command = [
"ffmpeg", "-hide_banner", "-loglevel", "error", "-y",
"-ss", f"{start_sec:.3f}", "-to", f"{end_sec:.3f}", "-i", video_path,
"-vf", f"fps={fps}", "-q:v", "2",
os.path.join(out_dir, "%06d.jpg"),
]
process = subprocess.Popen(command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True)
expected = frame_count(start_sec, end_sec, fps)
while process.poll() is None:
if should_stop is not None and should_stop():
process.terminate()
process.wait(timeout=10)
raise VideoError("cancelled")
if on_progress is not None:
on_progress(min(_written(out_dir), expected))
try:
process.wait(timeout=1)
except subprocess.TimeoutExpired:
pass
if process.returncode != 0:
raise VideoError((process.stderr.read() or "ffmpeg failed").strip().splitlines()[-1])
return sorted(name for name in os.listdir(out_dir) if name.endswith(".jpg"))
def _written(out_dir: str) -> int:
try:
return sum(1 for name in os.listdir(out_dir) if name.endswith(".jpg"))
except OSError:
return 0