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,197 @@
|
||||
"""Fine-tune the project's base model on its master dataset (REQ-060…065).
|
||||
|
||||
The default is old + new together: the master dataset already accumulates every
|
||||
merged batch, so a run sees the whole history. Training on the newest batch
|
||||
alone is what makes a model quietly forget what it used to know, so it is not
|
||||
what happens here.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import shutil
|
||||
import time
|
||||
from typing import Optional
|
||||
|
||||
from backend import config, dataset, db, evaluate, hardware, jobs, projects
|
||||
|
||||
PRETRAINED = {"bbox": "yolo11n.pt", "polygon": "yolo11n-seg.pt"}
|
||||
|
||||
|
||||
class TrainingError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def models_dir(project_slug: str) -> str:
|
||||
return os.path.join(config.project_dir(project_slug), "models")
|
||||
|
||||
|
||||
def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None, batch_ids: Optional[list] = None) -> dict:
|
||||
project = projects.get(project_id)
|
||||
if project is None:
|
||||
raise TrainingError("No such project")
|
||||
counts = dataset.summary(project_id)["splits"]
|
||||
if counts["train"] == 0:
|
||||
raise TrainingError(
|
||||
"The master dataset is empty — approve and merge a batch before training"
|
||||
)
|
||||
|
||||
settings = hardware.resolve(overrides, epochs)
|
||||
job = jobs.create(
|
||||
"train",
|
||||
params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids},
|
||||
project_id=project_id,
|
||||
message=f"{counts['train']} train / {counts['val']} val",
|
||||
)
|
||||
return job.to_dict()
|
||||
|
||||
|
||||
def listing(project_id: int) -> list:
|
||||
with db.cursor() as cur:
|
||||
cur.execute(
|
||||
"SELECT * FROM model_versions WHERE project_id = ? ORDER BY version DESC",
|
||||
(project_id,),
|
||||
)
|
||||
rows = []
|
||||
for row in cur.fetchall():
|
||||
item = dict(row)
|
||||
item["metrics"] = json.loads(item["metrics"] or "null")
|
||||
item["base_metrics"] = json.loads(item["base_metrics"] or "null")
|
||||
rows.append(item)
|
||||
return rows
|
||||
|
||||
|
||||
def get_version(model_id: int) -> Optional[dict]:
|
||||
with db.cursor() as cur:
|
||||
cur.execute("SELECT * FROM model_versions WHERE id = ?", (model_id,))
|
||||
row = cur.fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def promote(model_id: int) -> dict:
|
||||
"""Make a trained version the project's base model for the next round (REQ-064)."""
|
||||
version = get_version(model_id)
|
||||
if version is None:
|
||||
raise TrainingError("No such model version")
|
||||
project = projects.get(version["project_id"])
|
||||
base_path = os.path.join(config.project_dir(project["slug"]), "base", "model.pt")
|
||||
os.makedirs(os.path.dirname(base_path), exist_ok=True)
|
||||
shutil.copyfile(version["weights_path"], base_path)
|
||||
with db.cursor() as cur:
|
||||
cur.execute(
|
||||
"UPDATE projects SET base_model_path = ?, base_model_kind = 'trained' WHERE id = ?",
|
||||
(base_path, project["id"]),
|
||||
)
|
||||
return projects.get(project["id"])
|
||||
|
||||
|
||||
def _next_version(cur, project_id: int) -> int:
|
||||
cur.execute(
|
||||
"SELECT COALESCE(MAX(version), 0) + 1 FROM model_versions WHERE project_id = ?",
|
||||
(project_id,),
|
||||
)
|
||||
return cur.fetchone()[0]
|
||||
|
||||
|
||||
@jobs.handler("train")
|
||||
def _run_train(job) -> None:
|
||||
from ultralytics import YOLO
|
||||
|
||||
os.environ["ULTRALYTICS_OFFLINE"] = "true"
|
||||
os.environ["YOLO_OFFLINE"] = "true"
|
||||
|
||||
project = projects.get(job.params["project_id"])
|
||||
settings = job.params["settings"]
|
||||
batch_ids = job.params.get("batch_ids")
|
||||
data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids)
|
||||
|
||||
# SAM3 and a training run must not hold VRAM at the same time (REQ-065).
|
||||
from backend.sam3_engine import release_engine
|
||||
|
||||
if release_engine():
|
||||
job.log("Released SAM3 from VRAM before training")
|
||||
|
||||
start_point = project["base_model_path"] or PRETRAINED[project["label_type"]]
|
||||
if not project["base_model_path"]:
|
||||
job.log(f"No base model on this project — starting from {start_point}")
|
||||
job.log(f"Fine-tuning {os.path.basename(start_point)} for {settings['epochs']} epoch(s) "
|
||||
f"(batch={settings['batch']}, imgsz={settings['imgsz']}, "
|
||||
f"device={settings['device']})")
|
||||
|
||||
with db.cursor() as cur:
|
||||
version = _next_version(cur, project["id"])
|
||||
out_dir = os.path.join(models_dir(project["slug"]), str(version))
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
|
||||
model = YOLO(start_point)
|
||||
|
||||
def on_epoch(trainer):
|
||||
# trainer.epoch is 0-based; report a human-facing 1-based count.
|
||||
epoch = getattr(trainer, 'epoch', 0) + 1
|
||||
total = getattr(trainer, 'epochs', settings["epochs"])
|
||||
job.progress(epoch, total, f"epoch {epoch}/{total}")
|
||||
|
||||
model.add_callback("on_fit_epoch_end", on_epoch)
|
||||
job.progress(0, settings["epochs"])
|
||||
|
||||
keep_run_dir = False
|
||||
try:
|
||||
model.train(
|
||||
data=data_yaml,
|
||||
epochs=settings["epochs"],
|
||||
imgsz=settings["imgsz"],
|
||||
batch=settings["batch"],
|
||||
device=settings["device"],
|
||||
workers=settings.get("workers", 8),
|
||||
cache=False,
|
||||
project=os.path.join(out_dir, "runs"),
|
||||
name="train",
|
||||
exist_ok=True,
|
||||
amp=True,
|
||||
plots=False,
|
||||
verbose=False,
|
||||
)
|
||||
|
||||
produced = os.path.join(out_dir, "runs", "train", "weights", "best.pt")
|
||||
if not os.path.isfile(produced):
|
||||
raise TrainingError("Training finished without producing best.pt")
|
||||
weights = os.path.join(out_dir, "best.pt")
|
||||
shutil.copyfile(produced, weights)
|
||||
|
||||
job.log("Validating the base model and the new one on the same val set…")
|
||||
comparison = evaluate.compare(
|
||||
project["base_model_path"] or None, weights, data_yaml,
|
||||
[item["name"] for item in project["classes"]],
|
||||
imgsz=settings["imgsz"], device=settings["device"], batch=settings["batch"],
|
||||
)
|
||||
with open(os.path.join(out_dir, "metrics.json"), "w", encoding="utf-8") as handle:
|
||||
json.dump(comparison, handle, indent=2)
|
||||
|
||||
with db.cursor() as cur:
|
||||
cur.execute(
|
||||
"""INSERT INTO model_versions (project_id, version, weights_path,
|
||||
parent_model_path, metrics, base_metrics,
|
||||
created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)""",
|
||||
(project["id"], version, weights, project["base_model_path"],
|
||||
json.dumps(comparison["new"]), json.dumps(comparison["base"]), time.time()),
|
||||
)
|
||||
|
||||
new = comparison["new"]
|
||||
if comparison["delta"]:
|
||||
delta = comparison["delta"]
|
||||
job.log(f"v{version}: mAP50 {new['map50']:.4f} ({delta['map50']:+.4f} vs base), "
|
||||
f"mAP50-95 {new['map50_95']:.4f} ({delta['map50_95']:+.4f})")
|
||||
else:
|
||||
job.log(f"v{version}: mAP50 {new['map50']:.4f}, mAP50-95 {new['map50_95']:.4f} "
|
||||
f"— {comparison['skipped']}")
|
||||
except Exception:
|
||||
# Keep the runs directory on failure: results.csv and logs are the only record
|
||||
# of why training failed (REQ-006, REQ-064).
|
||||
keep_run_dir = True
|
||||
raise
|
||||
finally:
|
||||
# Delete run directory on success or cancellation. job.cancelled finishes training
|
||||
# without raising an exception, so it takes this delete path.
|
||||
if not keep_run_dir:
|
||||
shutil.rmtree(os.path.join(out_dir, "runs"), ignore_errors=True)
|
||||
|
||||
Reference in new issue
Block a user