diff --git a/labeling/backend/README.md b/labeling/backend/README.md index d659d75..2152ed2 100644 --- a/labeling/backend/README.md +++ b/labeling/backend/README.md @@ -4,13 +4,10 @@ API + data model for the bounding-box labeling tool. See CLAUDE.md §3.2 for the full requirements this is built against. **Status:** core data model + CRUD API implemented (sets, hierarchical classes, -frames, labels, per-frame/per-set status), frame image serving, and bulk frame -ingest from a recording session (see `labeling/frontend/` for the UI). **Not yet -implemented:** dataset versioning/promotion (`enemy-v1`, `enemy-v2`, ...) — -deferred since it needs its own design pass (snapshot semantics: labels are -editable at any time, but a promoted dataset version must stay reproducible). -Also not yet implemented: auth/login (see below) or the active-learning -auto-label workflow. +frames, labels, per-frame/per-set status), frame image serving, bulk frame +ingest from a recording session, and dataset version promotion (see +`labeling/frontend/` for the UI). **Not yet implemented:** auth/login (see +below) or the active-learning auto-label workflow. ## Stack @@ -32,6 +29,11 @@ auto-label workflow. frame + set + sub-class. - `FrameSetStatus` — per-frame, per-set label state (`unlabeled` / `auto_labeled` / `reviewed`). +- `DatasetVersion` / `DatasetVersionFrame` / `DatasetVersionLabel` — a named, + immutable snapshot of a set's frames+labels as of promotion time (`enemy-v1`, + `enemy-v2`, ...). Labels stay live-editable at any time; a promoted version + copies the label data at that moment, so it stays reproducible regardless of + later edits or deletions to the live labels. Frame images are served from `FRAMES_ROOT` (env var, default `./frames`) at `/images/`; `Frame.image_path` is always relative to that root. @@ -71,6 +73,23 @@ Safe to re-run: frames already registered for that session (by frame index) are skipped. `--width`/`--height` default to 1280x720 (Spelunky Classic HD's fixed capture resolution). +## Promoting a dataset version + +Once frames are marked `reviewed` (see the frontend, or `PUT +/frames/{id}/sets/{id}/status`), freeze them into a named, reproducible +version: + +``` +POST /sets/{set_id}/dataset-versions +{"name": "enemy-v1", "description": "first reviewed batch"} +``` + +Defaults to every currently-`reviewed` frame in that set; pass an explicit +`frame_ids` list to promote a different selection instead. Fetch the frozen +result (what training will eventually consume) via `GET +/dataset-versions/{id}` — it returns each frame plus the exact labels that +existed at promotion time, unaffected by any later edits to the live labels. + ## Testing ``` diff --git a/labeling/backend/src/spelunkai_labeling_backend/main.py b/labeling/backend/src/spelunkai_labeling_backend/main.py index 3d77d29..aea78b0 100644 --- a/labeling/backend/src/spelunkai_labeling_backend/main.py +++ b/labeling/backend/src/spelunkai_labeling_backend/main.py @@ -17,6 +17,7 @@ from fastapi.middleware.cors import CORSMiddleware from fastapi.staticfiles import StaticFiles from .db import Database, get_db +from .routers.dataset_versions import router as dataset_versions_router from .routers.frames import router as frames_router from .routers.labels import router as labels_router from .routers.sets import main_classes_router, sets_router @@ -49,6 +50,7 @@ def create_app(database_url: Optional[str] = None, frames_root: Optional[str] = app.include_router(main_classes_router) app.include_router(frames_router) app.include_router(labels_router) + app.include_router(dataset_versions_router) app.mount("/images", StaticFiles(directory=frames_dir), name="images") diff --git a/labeling/backend/src/spelunkai_labeling_backend/models.py b/labeling/backend/src/spelunkai_labeling_backend/models.py index 692c86c..e65b3a3 100644 --- a/labeling/backend/src/spelunkai_labeling_backend/models.py +++ b/labeling/backend/src/spelunkai_labeling_backend/models.py @@ -119,3 +119,46 @@ class Label(Base): updated_at = Column(DateTime, default=_utcnow, onupdate=_utcnow) sub_class = relationship("SubClass") + + +class DatasetVersion(Base): + """A named, immutable snapshot of a set's reviewed frames+labels, for + reproducible training runs (labels stay editable at any time, but a + promoted version's contents never change).""" + + __tablename__ = "dataset_versions" + __table_args__ = (UniqueConstraint("set_id", "name", name="uq_dataset_version_set_name"),) + + id = Column(Integer, primary_key=True) + set_id = Column(Integer, ForeignKey("label_sets.id"), nullable=False) + name = Column(String, nullable=False) + description = Column(String, nullable=True) + created_at = Column(DateTime, default=_utcnow) + created_by_id = Column(Integer, ForeignKey("users.id"), nullable=True) + + +class DatasetVersionFrame(Base): + """Which frames are included in a dataset version.""" + + __tablename__ = "dataset_version_frames" + __table_args__ = (UniqueConstraint("dataset_version_id", "frame_id", name="uq_dataset_version_frame"),) + + id = Column(Integer, primary_key=True) + dataset_version_id = Column(Integer, ForeignKey("dataset_versions.id"), nullable=False) + frame_id = Column(Integer, ForeignKey("frames.id"), nullable=False) + + +class DatasetVersionLabel(Base): + """A frozen copy of a Label as of promotion time - the live Label row it + was copied from may later be edited or deleted without affecting this.""" + + __tablename__ = "dataset_version_labels" + + id = Column(Integer, primary_key=True) + dataset_version_id = Column(Integer, ForeignKey("dataset_versions.id"), nullable=False) + frame_id = Column(Integer, ForeignKey("frames.id"), nullable=False) + sub_class_id = Column(Integer, ForeignKey("sub_classes.id"), nullable=False) + x = Column(Float, nullable=False) + y = Column(Float, nullable=False) + width = Column(Float, nullable=False) + height = Column(Float, nullable=False) diff --git a/labeling/backend/src/spelunkai_labeling_backend/routers/dataset_versions.py b/labeling/backend/src/spelunkai_labeling_backend/routers/dataset_versions.py new file mode 100644 index 0000000..024eaf1 --- /dev/null +++ b/labeling/backend/src/spelunkai_labeling_backend/routers/dataset_versions.py @@ -0,0 +1,129 @@ +"""Dataset version promotion: frozen, reproducible snapshots of a set's +reviewed frames and labels, for training to consume without being affected +by later labeling edits. +""" +from __future__ import annotations + +from fastapi import APIRouter, Depends, HTTPException +from sqlalchemy.orm import Session + +from .. import models, schemas +from ..db import get_db +from ..users import get_or_create_user + +router = APIRouter(tags=["dataset-versions"]) + + +@router.post("/sets/{set_id}/dataset-versions", response_model=schemas.DatasetVersionSummary, status_code=201) +def create_dataset_version(set_id: int, payload: schemas.DatasetVersionCreate, db: Session = Depends(get_db)): + if db.get(models.LabelSet, set_id) is None: + raise HTTPException(404, "set not found") + if db.query(models.DatasetVersion).filter_by(set_id=set_id, name=payload.name).first(): + raise HTTPException(409, f"dataset version '{payload.name}' already exists for this set") + + if payload.frame_ids is not None: + frame_ids = payload.frame_ids + found = { + row.id for row in db.query(models.Frame.id).filter(models.Frame.id.in_(frame_ids)) + } + missing = sorted(set(frame_ids) - found) + if missing: + raise HTTPException(422, f"unknown frame ids: {missing}") + else: + frame_ids = [ + row.frame_id + for row in db.query(models.FrameSetStatus.frame_id).filter_by( + set_id=set_id, status=models.LabelStatus.REVIEWED + ) + ] + if not frame_ids: + raise HTTPException(422, "no frames to promote (no reviewed frames, and no frame_ids given)") + + version = models.DatasetVersion( + set_id=set_id, + name=payload.name, + description=payload.description, + created_by_id=get_or_create_user(db, payload.created_by).id if payload.created_by else None, + ) + db.add(version) + db.flush() # assigns version.id without committing yet + + for frame_id in frame_ids: + db.add(models.DatasetVersionFrame(dataset_version_id=version.id, frame_id=frame_id)) + for label in db.query(models.Label).filter_by(frame_id=frame_id, set_id=set_id).all(): + db.add(models.DatasetVersionLabel( + dataset_version_id=version.id, + frame_id=frame_id, + sub_class_id=label.sub_class_id, + x=label.x, + y=label.y, + width=label.width, + height=label.height, + )) + + db.commit() + return schemas.DatasetVersionSummary( + id=version.id, + set_id=version.set_id, + name=version.name, + description=version.description, + frame_count=len(frame_ids), + ) + + +@router.get("/sets/{set_id}/dataset-versions", response_model=list[schemas.DatasetVersionSummary]) +def list_dataset_versions(set_id: int, db: Session = Depends(get_db)): + if db.get(models.LabelSet, set_id) is None: + raise HTTPException(404, "set not found") + + results = [] + for version in db.query(models.DatasetVersion).filter_by(set_id=set_id).all(): + frame_count = db.query(models.DatasetVersionFrame).filter_by(dataset_version_id=version.id).count() + results.append(schemas.DatasetVersionSummary( + id=version.id, + set_id=version.set_id, + name=version.name, + description=version.description, + frame_count=frame_count, + )) + return results + + +@router.get("/dataset-versions/{version_id}", response_model=schemas.DatasetVersionDetail) +def get_dataset_version(version_id: int, db: Session = Depends(get_db)): + version = db.get(models.DatasetVersion, version_id) + if version is None: + raise HTTPException(404, "dataset version not found") + + frame_ids = [ + row.frame_id + for row in db.query(models.DatasetVersionFrame.frame_id).filter_by(dataset_version_id=version_id) + ] + frames_by_id = { + f.id: f for f in db.query(models.Frame).filter(models.Frame.id.in_(frame_ids)) + } if frame_ids else {} + + labels_by_frame: dict[int, list[models.DatasetVersionLabel]] = {} + for label in db.query(models.DatasetVersionLabel).filter_by(dataset_version_id=version_id).all(): + labels_by_frame.setdefault(label.frame_id, []).append(label) + + frames_out = [ + schemas.DatasetVersionFrameRead( + frame=schemas.FrameRead.model_validate(frames_by_id[frame_id]), + labels=[ + schemas.DatasetVersionLabelRead( + sub_class_id=label.sub_class_id, x=label.x, y=label.y, width=label.width, height=label.height, + ) + for label in labels_by_frame.get(frame_id, []) + ], + ) + for frame_id in frame_ids + ] + + return schemas.DatasetVersionDetail( + id=version.id, + set_id=version.set_id, + name=version.name, + description=version.description, + frames=frames_out, + ) diff --git a/labeling/backend/src/spelunkai_labeling_backend/schemas.py b/labeling/backend/src/spelunkai_labeling_backend/schemas.py index 166f2d5..8d2fd69 100644 --- a/labeling/backend/src/spelunkai_labeling_backend/schemas.py +++ b/labeling/backend/src/spelunkai_labeling_backend/schemas.py @@ -112,3 +112,41 @@ class FrameSetStatusUpdate(BaseModel): class FrameWithStatus(BaseModel): frame: FrameRead status: LabelStatus + + +class DatasetVersionCreate(BaseModel): + name: str + description: Optional[str] = None + created_by: Optional[str] = None + # Which frames to freeze into the version. Default (None): every frame + # currently `reviewed` for this set - the normal "promote what's done" case. + frame_ids: Optional[list[int]] = None + + +class DatasetVersionSummary(BaseModel): + id: int + set_id: int + name: str + description: Optional[str] = None + frame_count: int + + +class DatasetVersionLabelRead(BaseModel): + sub_class_id: int + x: float + y: float + width: float + height: float + + +class DatasetVersionFrameRead(BaseModel): + frame: FrameRead + labels: list[DatasetVersionLabelRead] + + +class DatasetVersionDetail(BaseModel): + id: int + set_id: int + name: str + description: Optional[str] = None + frames: list[DatasetVersionFrameRead] diff --git a/labeling/backend/tests/test_dataset_versions.py b/labeling/backend/tests/test_dataset_versions.py new file mode 100644 index 0000000..56f882b --- /dev/null +++ b/labeling/backend/tests/test_dataset_versions.py @@ -0,0 +1,111 @@ +def _make_set_with_subclass(client, name="Enemy", sub_name="Bat"): + set_id = client.post("/sets", json={"name": name}).json()["id"] + main_class_id = client.post(f"/sets/{set_id}/main-classes", json={"name": name}).json()["id"] + sub_class_id = client.post(f"/main-classes/{main_class_id}/sub-classes", json={"name": sub_name}).json()["id"] + return set_id, sub_class_id + + +def _make_frame(client, session_name="run01", frame_index=0): + payload = {"session_name": session_name, "frame_index": frame_index, "image_path": "x.png", "width": 1280, "height": 720} + return client.post("/frames", json=payload).json()["id"] + + +def _make_reviewed_frame_with_label(client, set_id, sub_class_id, frame_index): + frame_id = _make_frame(client, frame_index=frame_index) + client.post( + f"/frames/{frame_id}/sets/{set_id}/labels", + json={"sub_class_id": sub_class_id, "x": 1, "y": 2, "width": 3, "height": 4}, + ) + client.put(f"/frames/{frame_id}/sets/{set_id}/status", json={"status": "reviewed"}) + return frame_id + + +def test_promote_defaults_to_reviewed_frames(client): + set_id, sub_class_id = _make_set_with_subclass(client) + reviewed_frame = _make_reviewed_frame_with_label(client, set_id, sub_class_id, frame_index=0) + _make_frame(client, frame_index=1) # stays unlabeled, should not be promoted + + resp = client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1"}) + assert resp.status_code == 201 + assert resp.json()["frame_count"] == 1 + + detail = client.get(f"/dataset-versions/{resp.json()['id']}").json() + assert [f["frame"]["id"] for f in detail["frames"]] == [reviewed_frame] + assert detail["frames"][0]["labels"][0]["sub_class_id"] == sub_class_id + + +def test_promote_rejects_duplicate_name(client): + set_id, sub_class_id = _make_set_with_subclass(client) + _make_reviewed_frame_with_label(client, set_id, sub_class_id, frame_index=0) + + client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1"}) + resp = client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1"}) + assert resp.status_code == 409 + + +def test_promote_rejects_empty_selection(client): + set_id, _ = _make_set_with_subclass(client) + resp = client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1"}) + assert resp.status_code == 422 + + +def test_promote_with_explicit_frame_ids(client): + set_id, sub_class_id = _make_set_with_subclass(client) + frame_id = _make_frame(client, frame_index=0) # not reviewed, but explicitly selected + client.post( + f"/frames/{frame_id}/sets/{set_id}/labels", + json={"sub_class_id": sub_class_id, "x": 0, "y": 0, "width": 1, "height": 1}, + ) + + resp = client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1", "frame_ids": [frame_id]}) + assert resp.status_code == 201 + assert resp.json()["frame_count"] == 1 + + +def test_promote_rejects_unknown_frame_ids(client): + set_id, _ = _make_set_with_subclass(client) + resp = client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1", "frame_ids": [999]}) + assert resp.status_code == 422 + + +def test_promote_requires_existing_set(client): + resp = client.post("/sets/999/dataset-versions", json={"name": "v1"}) + assert resp.status_code == 404 + + +def test_dataset_version_is_frozen_after_label_edits(client): + set_id, sub_class_id = _make_set_with_subclass(client) + frame_id = _make_reviewed_frame_with_label(client, set_id, sub_class_id, frame_index=0) + + version_id = client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1"}).json()["id"] + original_labels = client.get(f"/dataset-versions/{version_id}").json()["frames"][0]["labels"] + assert original_labels[0]["x"] == 1 + + # edit the live label after promotion + label_id = client.get(f"/frames/{frame_id}/sets/{set_id}/labels").json()[0]["id"] + client.patch(f"/labels/{label_id}", json={"x": 999}) + + # the promoted snapshot must be unaffected + frozen_labels = client.get(f"/dataset-versions/{version_id}").json()["frames"][0]["labels"] + assert frozen_labels[0]["x"] == 1 + + # deleting the live label must not touch the snapshot either + client.delete(f"/labels/{label_id}") + frozen_labels_after_delete = client.get(f"/dataset-versions/{version_id}").json()["frames"][0]["labels"] + assert len(frozen_labels_after_delete) == 1 + + +def test_list_dataset_versions(client): + set_id, sub_class_id = _make_set_with_subclass(client) + _make_reviewed_frame_with_label(client, set_id, sub_class_id, frame_index=0) + client.post(f"/sets/{set_id}/dataset-versions", json={"name": "v1"}) + + resp = client.get(f"/sets/{set_id}/dataset-versions") + assert resp.status_code == 200 + assert [v["name"] for v in resp.json()] == ["v1"] + assert resp.json()[0]["frame_count"] == 1 + + +def test_get_missing_dataset_version(client): + resp = client.get("/dataset-versions/999") + assert resp.status_code == 404