Files
reTraining/backend/projects.py
T
asus 5c7c122105 feat: add counting bench, triage, and dataset modules
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.
2026-08-14 16:28:52 +07:00

440 lines
16 KiB
Python

"""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()),
)
# No data.yaml to refresh here any more: it is assembled per training run
# from the datasets that run picks, so it always reflects the current classes.
return get(project_id)
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)
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 not os.path.isfile(path):
if path.startswith("/data/"):
alt_path = os.path.join(config.DATA_DIR, path[6:])
if os.path.isfile(alt_path):
path = alt_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}")