Add dataset version promotion (enemy-v1, enemy-v2, ...)
Add DatasetVersion/DatasetVersionFrame/DatasetVersionLabel and a
promote endpoint (POST /sets/{id}/dataset-versions) that freezes a
set's currently-reviewed frames (or an explicit frame_ids selection)
into a named, immutable snapshot: it copies each label's data at
promotion time rather than referencing the live rows, so later edits
or deletes to those labels can't retroactively change an already
-promoted version. GET /dataset-versions/{id} returns the frozen
frames+labels - this is what the training pipeline will eventually
pull from.
This was the labeling backend's last deliberately-deferred piece from
the original data model (needed its own design pass for snapshot
semantics). 31/31 backend tests pass, including one that promotes a
version, edits and deletes the live label afterward, and asserts the
snapshot is untouched.
Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
parent
96b7d42c34
commit
4e993ccf52
@ -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/<image_path>`; `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
|
||||
|
||||
```
|
||||
|
||||
@ -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")
|
||||
|
||||
|
||||
@ -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)
|
||||
|
||||
@ -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,
|
||||
)
|
||||
@ -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]
|
||||
|
||||
111
labeling/backend/tests/test_dataset_versions.py
Normal file
111
labeling/backend/tests/test_dataset_versions.py
Normal file
@ -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
|
||||
Loading…
x
Reference in New Issue
Block a user