Files
reTraining/backend/api/batches.py
T

322 lines
11 KiB
Python

import json
import os
import shutil
import tempfile
from typing import Optional
from fastapi import APIRouter, File, Form, HTTPException, Response, UploadFile
from fastapi.responses import FileResponse
from pydantic import BaseModel
from backend import autolabel, dataset, library, projects
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
target_class_names: Optional[list[str]] = None
threshold: float = autolabel.DEFAULT_THRESHOLD
iou_threshold: float = autolabel.DEFAULT_IOU
min_box_frac: float = 0.0
resume: bool = False
append: bool = False
custom_model_path: Optional[str] = None
class Exemplar(BaseModel):
box: list[float] # [cx, cy, w, h], normalized 0..1
positive: bool = True
class PreviewRequest(BaseModel):
frame_id: int
engine: str
threshold: float = autolabel.DEFAULT_THRESHOLD
iou_threshold: float = autolabel.DEFAULT_IOU
min_box_frac: float = 0.0
target_class_names: Optional[list[str]] = None
custom_model_path: Optional[str] = None
# Drawn box exemplars (REQ-172): normalized cxcywh, positive or negative.
# Preview-only — `autolabel.start` deliberately has no equivalent.
exemplars: Optional[list[Exemplar]] = None
exemplar_class_name: Optional[str] = None
@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, append=request.append,
engine=request.engine, engines=engine_list, class_ids=request.class_ids,
engine_classes=request.engine_classes,
target_class_names=request.target_class_names,
custom_model_path=request.custom_model_path)
except batch_store.BatchError as exc:
raise HTTPException(400, str(exc))
@router.post("/api/batches/inspect-model")
async def inspect_model(file: UploadFile = File(...)) -> dict:
if not (file.filename or "").endswith(".pt"):
raise HTTPException(400, "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:
classes = projects.read_model_classes(staged_path)
return {"filename": file.filename, "classes": classes, "staged_path": staged_path}
except Exception as exc:
if os.path.exists(staged_path):
os.unlink(staged_path)
raise HTTPException(400, f"Invalid model: {exc}")
@router.post("/api/batches/{batch_id}/autolabel-with-model")
async def autolabel_with_model(
batch_id: int,
file: UploadFile = File(...),
threshold: float = Form(0.35),
iou_threshold: float = Form(0.0),
selected_classes: str = Form("[]"),
append: bool = Form(True),
) -> dict:
if not (file.filename or "").endswith(".pt"):
raise HTTPException(400, "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:
target_classes = json.loads(selected_classes) if selected_classes else None
return autolabel.start(
batch_id,
threshold=threshold,
iou_threshold=iou_threshold,
append=append,
engine="custom",
custom_model_path=staged_path,
target_class_names=target_classes,
)
except Exception as exc:
if os.path.exists(staged_path):
os.unlink(staged_path)
raise HTTPException(400, f"Auto-annotation failed to start: {exc}")
@router.post("/api/batches/{batch_id}/preview")
def preview_autolabel(batch_id: int, request: PreviewRequest) -> dict:
from backend import jobs, preview
if not jobs.gpu_lock.acquire(timeout=20):
busy = jobs.running_types()
kind = busy[0] if busy else "background"
raise HTTPException(409, f"The GPU is busy with a {kind} job — wait for it to finish")
try:
shapes = preview.preview_frame(
batch_id=batch_id,
frame_id=request.frame_id,
engine=request.engine,
threshold=request.threshold,
iou_threshold=request.iou_threshold,
min_box_frac=request.min_box_frac,
target_class_names=request.target_class_names,
custom_model_path=request.custom_model_path,
exemplars=[e.model_dump() for e in request.exemplars or []],
exemplar_class_name=request.exemplar_class_name,
)
return {"shapes": shapes}
except Exception as exc:
raise HTTPException(400, str(exc))
finally:
jobs.gpu_lock.release()
@router.post("/api/sam3/playground-test")
async def sam3_playground_test(
file: UploadFile = File(...),
prompts: str = Form(...),
threshold: float = Form(0.35),
iou_threshold: float = Form(0.0),
) -> dict:
from PIL import Image
from backend import labeling
from backend.sam3_engine import get_engine
try:
image = Image.open(file.file).convert("RGB")
except Exception as exc:
raise HTTPException(400, f"Could not read image: {exc}")
width, height = image.size
prompt_list = [p.strip() for p in prompts.split(",") if p.strip()]
if not prompt_list:
raise HTTPException(400, "At least one text prompt is required")
from backend import jobs
if not jobs.gpu_lock.acquire(timeout=20):
busy = jobs.running_types()
kind = busy[0] if busy else "background"
raise HTTPException(409, f"The GPU is busy with a {kind} job — wait for it to finish")
try:
engine = get_engine()
raw_dets = engine.detect(image, prompt_list, threshold)
kept_dets = labeling.deduplicate(raw_dets, iou_threshold=iou_threshold)
except Exception as exc:
raise HTTPException(500, f"SAM3 inference failed: {exc}")
finally:
jobs.gpu_lock.release()
results = []
for det in kept_dets:
norm_box = [
det.box[0] / width,
det.box[1] / height,
det.box[2] / width,
det.box[3] / height,
]
polys = []
if det.mask is not None:
raw_polys = review_store.mask_to_polygons(det.mask)
polys = [[[float(pt[0]), float(pt[1])] for pt in poly] for poly in raw_polys]
results.append({
"class_id": det.class_id,
"prompt": det.class_name,
"score": round(float(det.score), 4),
"box": [round(v, 5) for v in norm_box],
"polygons": polys,
})
return {
"width": width,
"height": height,
"detections": results,
}
@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}
class ApproveRequest(BaseModel):
dataset_id: Optional[int] = None
dataset_name: str = ""
@router.post("/api/batches/{batch_ids}/approve")
def approve_batch(batch_ids: str, request: ApproveRequest = ApproveRequest()) -> dict:
"""`batch_ids` is one id or a comma-separated selection — one merge, one
dataset, however many batches Data Prep was tuned against (REQ-131)."""
try:
return dataset.approve(batch_ids, dataset_id=request.dataset_id,
dataset_name=request.dataset_name)
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/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}
@router.post("/api/batches/{batch_id}/reset-auto-annotations")
def reset_batch_auto_annotations(batch_id: int) -> dict:
if batch_store.get(batch_id) is None:
raise HTTPException(404, "No such batch")
deleted = review_store.clear_batch_auto_annotations(batch_id)
return {"deleted": deleted}