203 lines
7.5 KiB
Python
203 lines
7.5 KiB
Python
"""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, class_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, "class_ids": class_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")
|
|
class_ids = job.params.get("class_ids")
|
|
data_yaml = dataset.write_data_yaml(project, batch_ids=batch_ids, selected_class_ids=class_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"])
|
|
|
|
import torch
|
|
if torch.cuda.is_available():
|
|
torch.backends.cudnn.benchmark = True
|
|
|
|
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="ram",
|
|
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)
|
|
|