Files
feedmill-auto-label/backend/batches.py
T

254 lines
8.4 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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