feat: setup dataset enrichment app codebase and scripts
This commit is contained in:
1 parent
b5c28cc98a
commit
d07578462e
72 files changed
+11370
No files matched your search
@@ -0,0 +1,438 @@
|
||||
"""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()),
|
||||
)
|
||||
|
||||
updated = get(project_id)
|
||||
if updated["dataset"]["train"] + updated["dataset"]["val"] > 0:
|
||||
from backend import dataset
|
||||
|
||||
dataset.write_data_yaml(updated)
|
||||
return updated
|
||||
|
||||
|
||||
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)
|
||||
dataset.write_data_yaml(updated)
|
||||
|
||||
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 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}")
|
||||
|
||||
Reference in new issue
Block a user