Files
feedmill-auto-label/backend/training.py
T
2026-08-05 15:56:11 +07:00

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)