- 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
447 lines
17 KiB
Python
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}")
|
|
|