Files
reTraining/backend/projects.py
T
asus dee58e4ae5 fix: resolve model weight paths against DATA_DIR (REQ-187)
- config.py: resolve_data_path (legacy abs + rel) + rel_data_path
- all file-opening reads wrapped: preview, autolabel, training, model
  download, live count, projects.get; training_start_point hack replaced
- new writes store paths relative to data/
- legacy stale rows (/home/asus/reTraining/...) resolve without migration
- requirements: REQ-187 added; REQ-188 (per-class max box) + REQ-186
  copy-line amendment drafted for the next task
2026-10-02 16:59:25 +07:00

447 lines
17 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 Dict, 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, config.rel_data_path(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, container) "
"VALUES (?, ?, ?, ?, ?)",
[(project_id, index, item["name"], item["prompt"], item.get("container", 0))
for index, item in enumerate(classes)],
)
def _row_to_dict(cur, row) -> dict:
cur.execute(
"SELECT class_id, name, prompt, container 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": config.resolve_data_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"],
"video_root_linux": config.archive_host_paths()["linux"],
"video_root_windows": config.archive_host_paths()["windows"],
"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,
containers: Optional[Dict[int, bool]] = None) -> dict:
"""Edit the things that are safe to change: prompts, container flags,
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)),
)
for class_id, flag in (containers or {}).items():
cur.execute(
"UPDATE project_classes SET container = ? WHERE project_id = ? AND class_id = ?",
(1 if flag else 0, 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 = ?",
(config.rel_data_path(stored), project_id),
)
_write_classes(cur, project_id, _kept_classes(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_classes(project: dict, names: List[str]) -> List[dict]:
"""Keep the stored per-class attributes — prompt (REQ-171) and container
flag (REQ-184) — for a class that survives a base-model swap; new classes
fall back to the class name as prompt and no flag."""
known = {item["name"]: item for item in project["classes"]}
return [{"name": name,
"prompt": known[name]["prompt"] if name in known else name,
"container": known[name].get("container", 0) if name in known else 0}
for name in names]
def training_start_point(project: dict) -> str:
"""The weights a training run should start from (REQ-060, REQ-004)."""
path = config.resolve_data_path(project["base_model_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}")