From d07578462ecfbcd42d7c45cc53e4dbe7ebbb68d0 Mon Sep 17 00:00:00 2001 From: asus Date: Wed, 5 Aug 2026 11:52:27 +0700 Subject: [PATCH] feat: setup dataset enrichment app codebase and scripts --- .gitignore | 85 + .gitmodules | 3 + AGENTS.md | 129 ++ CLAUDE.md | 1 + Dockerfile | 35 + backend/__init__.py | 0 backend/api/__init__.py | 0 backend/api/batches.py | 152 ++ backend/api/common.py | 35 + backend/api/jobs.py | 29 + backend/api/models.py | 63 + backend/api/projects.py | 207 +++ backend/api/review.py | 91 + backend/autolabel.py | 241 +++ backend/batches.py | 253 +++ backend/config.py | 39 + backend/dataset.py | 269 +++ backend/db.py | 178 ++ backend/evaluate.py | 64 + backend/hardware.py | 66 + backend/jobs.py | 255 +++ backend/labeling.py | 81 + backend/library.py | 117 ++ backend/main.py | 75 + backend/projects.py | 438 +++++ backend/review.py | 350 ++++ backend/sam3_engine.py | 242 +++ backend/training.py | 197 +++ backend/video.py | 139 ++ docker-compose.override.yml | 4 + docker-compose.yml | 30 + docs/design.md | 285 +++ docs/requirements.md | 162 ++ docs/tasks.md | 840 +++++++++ frontend/.dockerignore | 2 + frontend/.gitignore | 24 + frontend/.oxlintrc.json | 8 + frontend/Dockerfile | 12 + frontend/README.md | 16 + frontend/index.html | 17 + frontend/nginx.conf | 29 + frontend/package-lock.json | 1665 ++++++++++++++++++ frontend/package.json | 19 + frontend/public/favicon.svg | 1 + frontend/public/icons.svg | 24 + frontend/src/App.jsx | 143 ++ frontend/src/api.js | 107 ++ frontend/src/app.css | 681 +++++++ frontend/src/components/AnnotationCanvas.jsx | 240 +++ frontend/src/components/Filmstrip.jsx | 21 + frontend/src/components/Icons.jsx | 138 ++ frontend/src/components/QuickReclassBar.jsx | 28 + frontend/src/components/ReviewSidebar.jsx | 102 ++ frontend/src/components/Shape.jsx | 123 ++ frontend/src/components/ShortcutsPanel.jsx | 37 + frontend/src/components/Sidebar.jsx | 104 ++ frontend/src/main.jsx | 11 + frontend/src/pages/LibraryPage.jsx | 591 +++++++ frontend/src/pages/ModelsPage.jsx | 378 ++++ frontend/src/pages/ProjectsPage.jsx | 377 ++++ frontend/src/pages/ReviewPage.jsx | 338 ++++ frontend/src/pages/TrimPage.jsx | 195 ++ frontend/src/roboflow.css | 266 +++ frontend/src/theme.css | 255 +++ frontend/vite.config.js | 17 + install_nvidia.sh | 24 + requirements.txt | 20 + sam3 | 1 + scripts/auto_pull.py | 46 + scripts/restart_app.sh | 23 + scripts/webhook.py | 67 + start.sh | 65 + 72 files changed, 11370 insertions(+) create mode 100644 .gitignore create mode 100644 .gitmodules create mode 100644 AGENTS.md create mode 120000 CLAUDE.md create mode 100644 Dockerfile create mode 100644 backend/__init__.py create mode 100644 backend/api/__init__.py create mode 100644 backend/api/batches.py create mode 100644 backend/api/common.py create mode 100644 backend/api/jobs.py create mode 100644 backend/api/models.py create mode 100644 backend/api/projects.py create mode 100644 backend/api/review.py create mode 100644 backend/autolabel.py create mode 100644 backend/batches.py create mode 100644 backend/config.py create mode 100644 backend/dataset.py create mode 100644 backend/db.py create mode 100644 backend/evaluate.py create mode 100644 backend/hardware.py create mode 100644 backend/jobs.py create mode 100644 backend/labeling.py create mode 100644 backend/library.py create mode 100644 backend/main.py create mode 100644 backend/projects.py create mode 100644 backend/review.py create mode 100644 backend/sam3_engine.py create mode 100644 backend/training.py create mode 100644 backend/video.py create mode 100644 docker-compose.override.yml create mode 100644 docker-compose.yml create mode 100644 docs/design.md create mode 100644 docs/requirements.md create mode 100644 docs/tasks.md create mode 100644 frontend/.dockerignore create mode 100644 frontend/.gitignore create mode 100644 frontend/.oxlintrc.json create mode 100644 frontend/Dockerfile create mode 100644 frontend/README.md create mode 100644 frontend/index.html create mode 100644 frontend/nginx.conf create mode 100644 frontend/package-lock.json create mode 100644 frontend/package.json create mode 100644 frontend/public/favicon.svg create mode 100644 frontend/public/icons.svg create mode 100644 frontend/src/App.jsx create mode 100644 frontend/src/api.js create mode 100644 frontend/src/app.css create mode 100644 frontend/src/components/AnnotationCanvas.jsx create mode 100644 frontend/src/components/Filmstrip.jsx create mode 100644 frontend/src/components/Icons.jsx create mode 100644 frontend/src/components/QuickReclassBar.jsx create mode 100644 frontend/src/components/ReviewSidebar.jsx create mode 100644 frontend/src/components/Shape.jsx create mode 100644 frontend/src/components/ShortcutsPanel.jsx create mode 100644 frontend/src/components/Sidebar.jsx create mode 100644 frontend/src/main.jsx create mode 100644 frontend/src/pages/LibraryPage.jsx create mode 100644 frontend/src/pages/ModelsPage.jsx create mode 100644 frontend/src/pages/ProjectsPage.jsx create mode 100644 frontend/src/pages/ReviewPage.jsx create mode 100644 frontend/src/pages/TrimPage.jsx create mode 100644 frontend/src/roboflow.css create mode 100644 frontend/src/theme.css create mode 100644 frontend/vite.config.js create mode 100644 install_nvidia.sh create mode 100644 requirements.txt create mode 160000 sam3 create mode 100755 scripts/auto_pull.py create mode 100755 scripts/restart_app.sh create mode 100755 scripts/webhook.py create mode 100755 start.sh diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..3117ab7 --- /dev/null +++ b/.gitignore @@ -0,0 +1,85 @@ +# Python +__pycache__/ +*.pyc +*.pyo +*.pyd +.venv/ +venv/ +*.egg-info/ +build/ +dist/ + +# Application Data & Environment +data/ +.env +.env.local + +# Frontend +frontend/node_modules/ +frontend/dist/ +frontend/*.local + +# OS / IDE +.DS_Store +.vscode/ +.idea/ + +# Video formats +*.mp4 +*.avi +*.mov +*.mkv +*.webm +*.flv +*.wmv +*.m4v +*.mpg +*.mpeg +# Image formats & Datasets +*.jpg +*.jpeg +*.png +*.webp +*.bmp +*.tiff +*.gif +datasets/ +dataset/ + +# ML Models & Checkpoints +*.pt +*.pth +*.bin +*.safetensors +*.ckpt +*.onnx +*.engine +*.plan +*.trt +*.h5 +runs/ +checkpoints/ +weights/ + +# Binary & Data Arrays +*.npy +*.npz +*.parquet +*.feather +*.pkl +*.pickle +*.joblib +*.db +*.sqlite +*.sqlite3 + +# Training Logs & Caches +*.log +*.tfevents* +wandb/ +mlruns/ +.cache/ +.pytest_cache/ +.mypy_cache/ +.ruff_cache/ + diff --git a/.gitmodules b/.gitmodules new file mode 100644 index 0000000..6fa4a2d --- /dev/null +++ b/.gitmodules @@ -0,0 +1,3 @@ +[submodule "sam3"] + path = sam3 + url = https://github.com/facebookresearch/sam3.git diff --git a/AGENTS.md b/AGENTS.md new file mode 100644 index 0000000..bd63075 --- /dev/null +++ b/AGENTS.md @@ -0,0 +1,129 @@ +# Agent Instructions + +Working rules for agents in this repo. A merge of the [Karpathy coding +guidelines](https://github.com/multica-ai/andrej-karpathy-skills) and the +[Chain of Truth](https://faridsurya-dev.github.io/Vibe-Coding-Research/en/welcome) +method: **validated artifacts are the source of truth, AI is a generator and accelerator.** + +Please Response in extremely concise and precise. + +## 1. Think Before Coding + +**Don't assume. Don't hide confusion. Surface tradeoffs.** + +- State your assumptions explicitly. If uncertain, ask. +- If multiple interpretations exist, present them — don't pick silently. +- If a simpler approach exists, say so. Push back when warranted. +- If something is unclear, stop. Name what's confusing. Ask. + +## 2. Simplicity First + +**Minimum code that solves the problem. Nothing speculative.** + +- No features beyond what was asked. +- No abstractions for single-use code. +- No "flexibility" or "configurability" that wasn't requested. +- No error handling for impossible scenarios. +- If you write 200 lines and it could be 50, rewrite it. + +## 3. Surgical Changes + +**Touch only what you must. Clean up only your own mess.** + +- Don't "improve" adjacent code, comments, or formatting. +- Don't refactor things that aren't broken. +- Match existing style, even if you'd do it differently. +- If you notice unrelated dead code, mention it — don't delete it. +- Remove imports/variables/functions that *your* changes made unused. + +The test: every changed line should trace directly to the user's request. + +## 4. Goal-Driven Execution + +**Define success criteria. Loop until verified.** + +Turn tasks into verifiable goals, and for multi-step work state a brief plan: + +``` +1. [Step] → verify: [check] +2. [Step] → verify: [check] +``` + +This repo has no automated tests, so verification means **running something**: hit the +endpoint, run the job, look at the files it produced. + +## 5. Chain of Truth — documents first, then code + +`docs/` is the source of truth, not the chat prompt. + + +| Document | Contents | +| ---------------------- | ----------------------------------------------------------------------------------- | +| `docs/requirements.md` | Numbered `REQ-xxx` requirements. Changes only with the user's approval. | +| `docs/design.md` | Data schema, API contract, disk layout. Each section names the `REQ-xxx` it serves. | +| `docs/tasks.md` | Implementation steps + verification criteria, status `[TODO]`/`[DONE]`. | + + +The rules: + +- Before writing feature code, make sure a `REQ-xxx` covers it. If none does, propose +adding one to the user first. +- Once a task is finished **and verified**, flip its status in `docs/tasks.md` to `[DONE]` +in the same commit. +- If the implementation diverges from `docs/design.md`, update the design — never let a +document lie. +- Use relative paths in markdown (`./`, `../`), not absolute ones. + +## 6. Repo rules + +- **Package manager: `uv`.** No `pip`, `poetry`, or bare `python`/`python3`. Dependencies +live in `requirements.txt`; install them with `uv pip install -r requirements.txt` and run +scripts with `uv run`. The Docker image installs the same file, so the two environments +cannot drift. +- **File size limit: 400 lines.** Any new or refactored file that exceeds it must be split +into smaller, logical modules. +- `**sam3/` is a vendor copy** of Meta's library. It's a dependency, not app code — don't +add scripts there or edit anything inside it. +- **Never write into the user's video archive.** All output goes under `data/`. +- Secrets (`HF_TOKEN`) come from `.env` only; they never belong in code or docs. + +## 7. UI/UX + +Frontend work follows [ui-ux-pro-max](https://github.com/nextlevelbuilder/ui-ux-pro-max-skill): +generate the design system first (style, palette, typography), then build against it, then +validate before delivering. Consistency across pages beats per-page cleverness. + +This app is a **dense internal tool**, not a landing page. Its screens are for long review +sessions in front of a screen: the video frame and the annotation canvas are the content, +everything else is chrome and stays quiet. No decorative gradients, no marketing motion. + +Pre-delivery checklist — a UI task is not `[DONE]` until all of it passes: + +- [ ] No emoji as icons (SVG only: Heroicons/Lucide) +- [ ] `cursor: pointer` on every clickable element +- [ ] Hover states with smooth transitions (150–300 ms) +- [ ] Text contrast at least 4.5:1 +- [ ] Focus states visible for keyboard navigation +- [ ] `prefers-reduced-motion` respected +- [ ] Responsive at 375 / 768 / 1024 / 1440 px + +The review editor also has to survive keyboard-only use — see `docs/design.md`, "Frontend". + +## 8. Domain invariants + +Two things are easy to break without noticing, and breaking either makes the whole system +lie: + +1. **Stable val split.** Once a frame lands in `val`, it stays in `val` forever. Otherwise +he base-vs-new mAP comparison is meaningless. +2. **One `set_image` per image.** `Sam3Processor.set_image()` runs the vision backbone; +set_text_prompt()`only re-runs the grounding head against the cached`backbone_out`. n N-prompt job calls` set_image` **once per image** and loops prompts over that same +tate. Don't restructure this into set_image-per-prompt. + +## 9. Scalability & Portability + +**Never hardcode something that will change across environments.** + +- **Hardware Agnosticism:** Do not hardcode hardware requirements (e.g., GPU configurations in `docker-compose.yml`) directly into base configuration files. Instead, use dynamic startup scripts (like `start.sh`) or environment overrides to detect the host's capabilities and inject the appropriate settings automatically. +- **Portability:** The app must be fully deployable and scalable on any device (from a CPU-only laptop to a massive multi-GPU rig) without requiring manual code edits to run. +- **Dynamic Configuration:** Do not hardcode absolute IP addresses, local network paths, or machine-specific environment variables in code. Rely on relative paths and configuration files to ensure maximum scalability. diff --git a/CLAUDE.md b/CLAUDE.md new file mode 120000 index 0000000..47dc3e3 --- /dev/null +++ b/CLAUDE.md @@ -0,0 +1 @@ +AGENTS.md \ No newline at end of file diff --git a/Dockerfile b/Dockerfile new file mode 100644 index 0000000..6a1a7d7 --- /dev/null +++ b/Dockerfile @@ -0,0 +1,35 @@ +FROM python:3.12-slim + +ENV PYTHONUNBUFFERED=1 \ + UV_SYSTEM_PYTHON=1 \ + APP_DATA_DIR=/data \ + VIDEO_ARCHIVE=/videos + +# ffmpeg does the trimming and frame extraction; libgl1/libglib2.0-0 are what +# OpenCV needs once ultralytics pulls in the non-headless build. +RUN apt-get update && apt-get install -y --no-install-recommends \ + ffmpeg libgl1 libglib2.0-0 \ + && rm -rf /var/lib/apt/lists/* + +# uv comes from PyPI rather than `COPY --from=ghcr.io/astral-sh/uv`: pulling it +# from a second registry made every build fail whenever ghcr.io was unreachable, +# even with all the dependency layers already cached. This is the only pip call +# in the project — bootstrapping the tool that installs everything else. +RUN pip install --no-cache-dir uv + +WORKDIR /app + +# Dependencies first so code edits don't re-download ~3 GB of CUDA wheels. +COPY requirements.txt ./ +RUN uv pip install -r requirements.txt + +# The vendored SAM3 declares timm/ftfy/regex/tqdm itself, so this install keeps +# its dependencies (unlike einops + pycocotools, which it needs but never +# declares — those are pinned in requirements.txt). +COPY sam3/ ./sam3/ +RUN uv pip install -e ./sam3 + +COPY backend/ ./backend/ + +EXPOSE 8000 +CMD ["uvicorn", "backend.main:app", "--host", "0.0.0.0", "--port", "8000"] diff --git a/backend/__init__.py b/backend/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/api/__init__.py b/backend/api/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/backend/api/batches.py b/backend/api/batches.py new file mode 100644 index 0000000..f2a3233 --- /dev/null +++ b/backend/api/batches.py @@ -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} + diff --git a/backend/api/common.py b/backend/api/common.py new file mode 100644 index 0000000..40902d4 --- /dev/null +++ b/backend/api/common.py @@ -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)) diff --git a/backend/api/jobs.py b/backend/api/jobs.py new file mode 100644 index 0000000..4a108fb --- /dev/null +++ b/backend/api/jobs.py @@ -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} diff --git a/backend/api/models.py b/backend/api/models.py new file mode 100644 index 0000000..7a81e72 --- /dev/null +++ b/backend/api/models.py @@ -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)) diff --git a/backend/api/projects.py b/backend/api/projects.py new file mode 100644 index 0000000..e966851 --- /dev/null +++ b/backend/api/projects.py @@ -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), + }, + ) diff --git a/backend/api/review.py b/backend/api/review.py new file mode 100644 index 0000000..076b638 --- /dev/null +++ b/backend/api/review.py @@ -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)} diff --git a/backend/autolabel.py b/backend/autolabel.py new file mode 100644 index 0000000..30f0e8e --- /dev/null +++ b/backend/autolabel.py @@ -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,), + ) diff --git a/backend/batches.py b/backend/batches.py new file mode 100644 index 0000000..9ee76fc --- /dev/null +++ b/backend/batches.py @@ -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 + + diff --git a/backend/config.py b/backend/config.py new file mode 100644 index 0000000..ca36c8f --- /dev/null +++ b/backend/config.py @@ -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) diff --git a/backend/dataset.py b/backend/dataset.py new file mode 100644 index 0000000..3c72a3b --- /dev/null +++ b/backend/dataset.py @@ -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}") diff --git a/backend/db.py b/backend/db.py new file mode 100644 index 0000000..215d3af --- /dev/null +++ b/backend/db.py @@ -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 diff --git a/backend/evaluate.py b/backend/evaluate.py new file mode 100644 index 0000000..1f5bee0 --- /dev/null +++ b/backend/evaluate.py @@ -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} diff --git a/backend/hardware.py b/backend/hardware.py new file mode 100644 index 0000000..9f4d8d2 --- /dev/null +++ b/backend/hardware.py @@ -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 diff --git a/backend/jobs.py b/backend/jobs.py new file mode 100644 index 0000000..18e2827 --- /dev/null +++ b/backend/jobs.py @@ -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() diff --git a/backend/labeling.py b/backend/labeling.py new file mode 100644 index 0000000..cc4fa3c --- /dev/null +++ b/backend/labeling.py @@ -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)) diff --git a/backend/library.py b/backend/library.py new file mode 100644 index 0000000..8b4f3f2 --- /dev/null +++ b/backend/library.py @@ -0,0 +1,117 @@ +"""The video archive, read as `//.` (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 `/` 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 diff --git a/backend/main.py b/backend/main.py new file mode 100644 index 0000000..12439f9 --- /dev/null +++ b/backend/main.py @@ -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() diff --git a/backend/projects.py b/backend/projects.py new file mode 100644 index 0000000..407f8b0 --- /dev/null +++ b/backend/projects.py @@ -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}") + diff --git a/backend/review.py b/backend/review.py new file mode 100644 index 0000000..2364050 --- /dev/null +++ b/backend/review.py @@ -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") + diff --git a/backend/sam3_engine.py b/backend/sam3_engine.py new file mode 100644 index 0000000..869a166 --- /dev/null +++ b/backend/sam3_engine.py @@ -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 diff --git a/backend/training.py b/backend/training.py new file mode 100644 index 0000000..5d187a4 --- /dev/null +++ b/backend/training.py @@ -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) + diff --git a/backend/video.py b/backend/video.py new file mode 100644 index 0000000..32eacc0 --- /dev/null +++ b/backend/video.py @@ -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 diff --git a/docker-compose.override.yml b/docker-compose.override.yml new file mode 100644 index 0000000..f784924 --- /dev/null +++ b/docker-compose.override.yml @@ -0,0 +1,4 @@ +services: + backend: + devices: + - nvidia.com/gpu=all diff --git a/docker-compose.yml b/docker-compose.yml new file mode 100644 index 0000000..cad395c --- /dev/null +++ b/docker-compose.yml @@ -0,0 +1,30 @@ +services: + backend: + build: . + + environment: + HF_TOKEN: ${HF_TOKEN:-} + HF_HUB_DISABLE_XET: "1" + APP_DATA_DIR: /data + VIDEO_ARCHIVE: /videos + CORS_ORIGINS: ${CORS_ORIGINS:-http://localhost:5173} + volumes: + - ./data:/data + - ${VIDEO_ARCHIVE_HOST:-./data/archive}:/videos:ro + - hf-cache:/root/.cache/huggingface + ports: + - "8000:8000" + shm_size: '8gb' + ipc: host + restart: unless-stopped + + frontend: + build: ./frontend + ports: + - "${WEB_PORT:-8080}:80" + depends_on: + - backend + restart: unless-stopped + +volumes: + hf-cache: diff --git a/docs/design.md b/docs/design.md new file mode 100644 index 0000000..72948e1 --- /dev/null +++ b/docs/design.md @@ -0,0 +1,285 @@ +# Design + +Serves `./requirements.md`. Each section names the REQs it fulfils. + +## Overview + +``` +Browser (React + Vite, served by nginx) + │ /api → proxy + ▼ +FastAPI ──► jobs (1 worker thread, 1 GPU) + │ ├─ extract : ffmpeg (CPU) + │ ├─ autolabel: SAM3 (GPU) + │ ├─ merge : copy + write labels (CPU) + │ └─ train : Ultralytics + eval (GPU) + ├──► SQLite (metadata & status) + └──► data/ (frames, master dataset, weights) +``` + +Storage split rule: **SQLite holds metadata and status; the disk holds pixels, final labels, +and weights.** The master dataset must stay useful even if the database is lost +(REQ-006, REQ-054). + +## Disk layout (REQ-006) + +``` +data/ # Docker volume + app.db # SQLite (WAL) + projects// + base/model.pt # the project's active base model (REQ-003) + dataset/ # MASTER, accumulative (REQ-050…052) + images/{train,val}/…jpg + labels/{train,val}/…txt + data.yaml + batches// + frames/000001.jpg … # extraction output (REQ-022) + models// + best.pt + metrics.json # base vs new metrics (REQ-063) + runs/ # Ultralytics run directory +``` + +Master dataset filenames: `__.jpg` — unique across batches and +self-documenting about where each image came from. The user's video archive is read-only +(REQ-074). + +## SQLite schema + +Created by an idempotent migration in `backend/db.py` at startup. + +```sql +projects( + id, slug UNIQUE, name, label_type CHECK(bbox|polygon), + base_model_path, base_model_kind CHECK(uploaded|pretrained|trained), + video_root, val_every DEFAULT 5, created_at) + +project_classes( + id, project_id → projects, class_id INT, name, prompt, + UNIQUE(project_id, class_id)) -- class_id = the YOLO class index (REQ-003/005) + +batches( + id, project_id → projects, video_path, date_label, batch_label, + start_sec REAL, end_sec REAL, fps REAL, + status CHECK(extracting|extracted|labeling|reviewing|approved|merged|failed), + frame_count INT, created_at, merged_at) + +frames( + id, batch_id → batches, idx INT, filename, width INT, height INT, + review_status CHECK(pending|approved|rejected) DEFAULT 'pending', + UNIQUE(batch_id, idx)) + +annotations( + id, frame_id → frames, class_id INT, + geometry TEXT, -- JSON; see "Geometry format" + score REAL, source CHECK(auto|manual), created_at) + +dataset_items( -- master dataset membership (REQ-052) + id, project_id → projects, frame_id → frames UNIQUE, + split CHECK(train|val), image_rel, label_rel, added_at) + +model_versions( + id, project_id → projects, version INT, weights_path, + parent_model_path, metrics TEXT, base_metrics TEXT, created_at, + UNIQUE(project_id, version)) + +jobs( -- persistent (REQ-071) + id, project_id, batch_id, type CHECK(extract|autolabel|merge|train), + status CHECK(queued|running|done|failed|cancelled), + progress INT, total INT, message, error, log TEXT, + created_at, started_at, finished_at) +``` + +**The stable val split (REQ-052)** is enforced by `dataset_items`: an existing row never +changes its `split`. On merge, only frames without a row are assigned, using a per-project +round-robin counter (`val_every`) that continues from the previous count. + +**Geometry format.** One JSON column covers both label types (REQ-002): + +- `bbox` → `{"type":"bbox","points":[x0,y0,x1,y1]}` +- `polygon` → `{"type":"polygon","points":[[x,y], …]}` + +Coordinates are stored **normalized 0–1** against the frame size, so neither the editor nor +the exporter needs to know the display size. SAM3 mask → polygon conversion is +`review.mask_to_polygons()`; for `bbox` projects the mask is only used to take its bounding +box. + +## Backend modules + +`app/` moves to `backend/`. Reuse existing code wherever possible: + +| Module | Role | Status | +|---|---|---| +| `sam3_engine.py` | SAM3 singleton, `open_state`/`apply_prompts`/`segment_at` | reused, plus a `release()` for REQ-065 | +| `labeling.py` | per-frame detection + cross-prompt NMS (REQ-031) | reused; the folder-walking half went with the old flow | +| `exporters.py` | ~~YOLO label writing~~ | **deleted** — `dataset.py` writes labels, `mask_to_polygons` moved to `review.py` | +| `sessions.py` | ~~exemplar/tap interaction~~ | **deleted** — see below | +| `jobs.py` | single-worker queue | extended: job types + persistence | +| `training.py` | Ultralytics fine-tune | changed: starts from the base model, args from `hardware.py` | +| `db.py` | connection + migration | **new** | +| `projects.py` | project CRUD, reads classes from a `.pt` | **new** | +| `library.py` | scans `//` | **new** | +| `video.py` | `ffprobe`, Range streaming, `ffmpeg` extraction | **new** | +| `batches.py` | batch lifecycle | **new** | +| `review.py` | annotation CRUD, frame status, click-assist | **new** | +| `autolabel.py` | the SAM3 job over a whole batch | **new** | +| `dataset.py` | merge into the master dataset, stable split | **new** | +| `evaluate.py` | validate base vs new model | **new** | +| `hardware.py` | VRAM detection → training defaults | **new** | +| `api/` | the FastAPI routes, one module per domain | **new** | + +Removed: `uploads.py`, `static/index.html`, and the old flow's endpoints. + +`sessions.py` was meant to be reused for click-assist, but it existed to hold GPU-resident +state for an interactive session — a whole eviction policy, an undo stack, and a per-session +annotation store, all of which the database and the stateless `review.assist` now cover. +Adapting 357 lines to do what 60 lines do was not worth it, so the module is gone. The one +thing it knew that mattered — SAM3 wants exemplar boxes as normalized centre-x, centre-y, +width, height — moved with it. + +The 400-line file limit (see `../AGENTS.md`) applies to all of the above. It is why the +routes live in `backend/api/{projects,batches,review,models,jobs}.py` rather than in +`main.py`, which now only builds the app and owns startup. Route modules import their +domain module under an alias (`from backend import projects as project_store`) so the two +namespaces stay distinguishable. + +## API contract + +``` +GET /api/health REQ-073 + +GET /api/projects REQ-001 +POST /api/projects REQ-001,002,004,005 +GET /api/projects/{id} # includes label_type_locked: bool (REQ-002) +DELETE /api/projects/{id} +PATCH /api/projects/{id} # class prompts, val_every (REQ-005) +POST /api/projects/{id}/classes # add new class {name, prompt} (REQ-008) +DELETE /api/projects/{id}/classes/{class_id} # delete class, delete shapes, reindex classes (REQ-007) +POST /api/projects/{id}/base-model # upload .pt, read classes (REQ-003) +GET /api/projects/{id}/dataset # master dataset summary (REQ-053) +GET /api/projects/{id}/dataset/download # zip (REQ-054) + +GET /api/projects/{id}/library # list of dates (REQ-011) +GET /api/projects/{id}/library/{date} # videos + duration/resolution (REQ-012) +GET /api/projects/{id}/video?rel=… # Range streaming (REQ-013) + +POST /api/projects/{id}/batches # {rel, start_sec, end_sec, fps} → extract job +GET /api/batches/{id} # status + review progress (REQ-045) +GET /api/batches/{id}/frames # frames + statuses +POST /api/batches/{id}/autolabel # {threshold} → job (REQ-030,032,034) +DELETE /api/batches/{id}/classes/{class_id}/annotations # clear all shapes of class in batch (REQ-046) +POST /api/batches/{id}/approve # → merge job (REQ-050) + +GET /api/frames/{id}/image?w=… # frame image / thumbnail +GET /api/frames/{id}/annotations +POST /api/frames/{id}/annotations # add a manual shape (REQ-042) +PATCH /api/annotations/{id} # move/resize/reclass +DELETE /api/annotations/{id} +POST /api/frames/{id}/assist # click/box → SAM3 shape (REQ-043) +POST /api/frames/{id}/status # approved | rejected | pending (REQ-041) + +POST /api/projects/{id}/train # → train job (REQ-060,061,062) +GET /api/projects/{id}/models # versions + metrics (REQ-063,064) +GET /api/models/{id}/weights # download best.pt +POST /api/models/{id}/promote # make it the project's base model (REQ-064) + +GET /api/jobs?project_id=… REQ-070,071 +GET /api/jobs/{id} +POST /api/jobs/{id}/cancel +``` + +## Job flows + +**extract (REQ-020…023).** `ffmpeg -ss -to -i