feat: setup dataset enrichment app codebase and scripts
This commit is contained in:
1 parent
b5c28cc98a
commit
d07578462e
72 files changed
+11370
No files matched your search
Whitespace-only changes.
Whitespace-only changes.
@@ -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}
|
||||
|
||||
@@ -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))
|
||||
@@ -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}
|
||||
@@ -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))
|
||||
@@ -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),
|
||||
},
|
||||
)
|
||||
@@ -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)}
|
||||
@@ -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,),
|
||||
)
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
@@ -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
@@ -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
|
||||
@@ -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}
|
||||
@@ -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
@@ -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()
|
||||
@@ -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))
|
||||
@@ -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
|
||||
@@ -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()
|
||||
@@ -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}")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in new issue
Block a user