This commit includes major additions and updates to the frontend and backend architectures, introducing new dataset management, live counting features, batch processing, and triage logic. Includes new UI pages, components, and API routes.
221 lines
8.2 KiB
Python
221 lines
8.2 KiB
Python
"""Named master datasets — several per project, each a full standalone copy.
|
|
|
|
One project used to have exactly one master dataset, so merging a batch twice
|
|
was a conflict. Now a merge targets a *named* dataset, and the same batch can go
|
|
into as many as you like: "batch7+8 strict rules" and "batch7+8 after I fixed
|
|
the annotations" are two datasets holding the same frames with different labels.
|
|
|
|
Each dataset owns its files under
|
|
|
|
data/projects/<slug>/datasets/<id>/{images,labels}/{train,val}
|
|
|
|
Images are copied rather than shared, so a dataset folder can be moved or
|
|
archived on its own without silently losing pixels.
|
|
|
|
Combining datasets for a training run is "newest wins": if a frame appears in
|
|
two of them, the one created later is treated as the correction of the earlier,
|
|
and the frame is emitted once. Emitting it twice would hand the model two
|
|
contradictory labels for the same image.
|
|
"""
|
|
|
|
import json
|
|
import os
|
|
import shutil
|
|
import time
|
|
from typing import List, Optional
|
|
|
|
from backend import config, db
|
|
|
|
|
|
class DatasetError(Exception):
|
|
pass
|
|
|
|
|
|
def dataset_root(project_slug: str, dataset_id: int) -> str:
|
|
return os.path.join(config.project_dir(project_slug), "datasets", str(dataset_id))
|
|
|
|
|
|
def adopt_legacy_tree() -> int:
|
|
"""Move a pre-rename project's files under the dataset that adopted its rows.
|
|
|
|
`_migrate_dataset_items` gave the old rows a home in the `datasets` table but
|
|
left the pixels at `<project>/dataset/`, so the adopting dataset points at a
|
|
directory that does not exist and a training run would find no images.
|
|
Idempotent: a dataset whose root already exists is left alone.
|
|
"""
|
|
moved = 0
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""SELECT s.id, p.slug FROM datasets s
|
|
JOIN projects p ON p.id = s.project_id
|
|
WHERE s.note = 'Adopted from the original single dataset'""",
|
|
)
|
|
adopted = [(row[0], row[1]) for row in cur.fetchall()]
|
|
|
|
for dataset_id, slug in adopted:
|
|
root = dataset_root(slug, dataset_id)
|
|
legacy = os.path.join(config.project_dir(slug), "dataset")
|
|
if os.path.isdir(root) or not os.path.isdir(os.path.join(legacy, "images")):
|
|
continue
|
|
os.makedirs(os.path.dirname(root), exist_ok=True)
|
|
for name in ("images", "labels"):
|
|
source = os.path.join(legacy, name)
|
|
if os.path.isdir(source):
|
|
os.makedirs(root, exist_ok=True)
|
|
shutil.move(source, os.path.join(root, name))
|
|
# The rest of the legacy tree is a stale data.yaml and the old
|
|
# `selected/` symlinks, both rebuilt per run now.
|
|
shutil.rmtree(legacy, ignore_errors=True)
|
|
moved += 1
|
|
return moved
|
|
|
|
|
|
def create(project_id: int, name: str = "", note: str = "",
|
|
rule_version: Optional[str] = None, rules: Optional[List[dict]] = None) -> dict:
|
|
label = (name or "").strip() or time.strftime("Master Dataset %Y-%m-%d %H:%M")
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"""INSERT INTO datasets (project_id, name, note, rule_version, rules_json,
|
|
created_at)
|
|
VALUES (?, ?, ?, ?, ?, ?)""",
|
|
(project_id, label, note, rule_version,
|
|
json.dumps(rules) if rules is not None else None, time.time()),
|
|
)
|
|
dataset_id = cur.lastrowid
|
|
return get(dataset_id)
|
|
|
|
|
|
def get(dataset_id: int) -> Optional[dict]:
|
|
with db.cursor() as cur:
|
|
cur.execute("SELECT * FROM datasets WHERE id = ?", (dataset_id,))
|
|
row = cur.fetchone()
|
|
if row is None:
|
|
return None
|
|
return _with_counts(cur, dict(row))
|
|
|
|
|
|
def listing(project_id: int) -> List[dict]:
|
|
with db.cursor() as cur:
|
|
cur.execute(
|
|
"SELECT * FROM datasets WHERE project_id = ? ORDER BY created_at DESC",
|
|
(project_id,),
|
|
)
|
|
rows = [dict(row) for row in cur.fetchall()]
|
|
return [_with_counts(cur, row) for row in rows]
|
|
|
|
|
|
def _with_counts(cur, row: dict) -> dict:
|
|
cur.execute(
|
|
"""SELECT split, COUNT(*) FROM dataset_items
|
|
WHERE dataset_id = ? GROUP BY split""",
|
|
(row["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, COUNT(d.id)
|
|
FROM dataset_items d
|
|
JOIN frames f ON f.id = d.frame_id
|
|
JOIN batches b ON b.id = f.batch_id
|
|
WHERE d.dataset_id = ? GROUP BY b.id ORDER BY b.id""",
|
|
(row["id"],),
|
|
)
|
|
batches = [{"id": r[0], "date_label": r[1], "batch_label": r[2], "images": r[3]}
|
|
for r in cur.fetchall()]
|
|
row["splits"] = splits
|
|
row["total"] = splits["train"] + splits["val"]
|
|
row["batches"] = batches
|
|
row["rules"] = json.loads(row.pop("rules_json") or "[]")
|
|
return row
|
|
|
|
|
|
def snapshot_rules(dataset_id: int, rules: List[dict], rule_version: str) -> None:
|
|
"""Freeze the rules a dataset was cut under (REQ-132)."""
|
|
with db.cursor() as cur:
|
|
cur.execute("UPDATE datasets SET rules_json = ?, rule_version = ? WHERE id = ?",
|
|
(json.dumps(rules), rule_version, dataset_id))
|
|
|
|
|
|
def rename(dataset_id: int, name: str = None, note: str = None) -> dict:
|
|
fields, args = [], []
|
|
if name is not None:
|
|
fields.append("name = ?")
|
|
args.append(name.strip())
|
|
if note is not None:
|
|
fields.append("note = ?")
|
|
args.append(note)
|
|
if fields:
|
|
args.append(dataset_id)
|
|
with db.cursor() as cur:
|
|
cur.execute(f"UPDATE datasets SET {', '.join(fields)} WHERE id = ?", args)
|
|
return get(dataset_id)
|
|
|
|
|
|
def delete(dataset_id: int) -> bool:
|
|
dataset = get(dataset_id)
|
|
if dataset is None:
|
|
return False
|
|
with db.cursor() as cur:
|
|
cur.execute("SELECT slug FROM projects WHERE id = ?", (dataset["project_id"],))
|
|
row = cur.fetchone()
|
|
if row is not None:
|
|
shutil.rmtree(dataset_root(row["slug"], dataset_id), ignore_errors=True)
|
|
with db.cursor() as cur:
|
|
cur.execute("DELETE FROM datasets WHERE id = ?", (dataset_id,))
|
|
return True
|
|
|
|
|
|
def combined_items(project_id: int, dataset_ids: List[int]) -> List[dict]:
|
|
"""Frames from these datasets, newest dataset winning on a repeated frame.
|
|
|
|
A frame in two datasets means the later one is a correction — a rule change
|
|
or a fixed annotation. Training on both copies would teach the model that
|
|
the same pixels are two different things.
|
|
"""
|
|
if not dataset_ids:
|
|
return []
|
|
with db.cursor() as cur:
|
|
placeholders = ",".join("?" for _ in dataset_ids)
|
|
cur.execute(
|
|
f"""SELECT d.frame_id, d.split, d.image_rel, d.label_rel, d.dataset_id,
|
|
s.created_at, s.name
|
|
FROM dataset_items d
|
|
JOIN datasets s ON s.id = d.dataset_id
|
|
WHERE d.project_id = ? AND d.dataset_id IN ({placeholders})
|
|
ORDER BY s.created_at ASC, d.id ASC""",
|
|
[project_id] + list(dataset_ids),
|
|
)
|
|
rows = cur.fetchall()
|
|
|
|
# Ordered oldest first, so a later dataset simply overwrites the entry.
|
|
winner = {}
|
|
for row in rows:
|
|
winner[row["frame_id"]] = {
|
|
"frame_id": row["frame_id"],
|
|
"split": row["split"],
|
|
"image_rel": row["image_rel"],
|
|
"label_rel": row["label_rel"],
|
|
"dataset_id": row["dataset_id"],
|
|
"dataset_name": row["name"],
|
|
}
|
|
return list(winner.values())
|
|
|
|
|
|
def overlap_report(project_id: int, dataset_ids: List[int], total_unique: int) -> dict:
|
|
"""How many frames the chosen datasets share, so the user is told rather
|
|
than quietly given fewer images than the totals suggest."""
|
|
if len(dataset_ids) < 2:
|
|
return {"shared_frames": 0, "total_unique": total_unique}
|
|
with db.cursor() as cur:
|
|
placeholders = ",".join("?" for _ in dataset_ids)
|
|
cur.execute(
|
|
f"""SELECT COUNT(*) FROM (
|
|
SELECT frame_id FROM dataset_items
|
|
WHERE project_id = ? AND dataset_id IN ({placeholders})
|
|
GROUP BY frame_id HAVING COUNT(DISTINCT dataset_id) > 1)""",
|
|
[project_id] + list(dataset_ids),
|
|
)
|
|
shared = cur.fetchone()[0]
|
|
return {"shared_frames": shared, "total_unique": total_unique}
|