Files
reTraining/backend/training.py
T
Andrew-AAAA d170cff0e4 feat: add descriptive model naming and inline rename
- Auto-generate model names: {arch}-{labelType}-{epochs}ep-{classNames}-{YYYYMMDD}
- Add PATCH /api/models/{id}/rename endpoint
- Inline rename UI on Models & Training page
- Download filename uses model name instead of v{N}
- DB migration: add name column to model_versions
- Update all docs to reflect new naming convention
2026-09-10 09:14:55 +07:00

317 lines
13 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
def _generate_model_name(project: dict, start_point: str, epochs: int,
class_ids: Optional[list] = None) -> str:
"""Build a descriptive model name: arch-labelType-epochs-classNames-YYYYMMDD."""
arch = os.path.splitext(os.path.basename(start_point))[0]
# Filter classes if specific IDs were selected.
classes = project["classes"]
if class_ids:
classes = [c for c in classes if c["class_id"] in class_ids]
class_tag = "+".join(c["name"] for c in classes) or "unknown"
date_str = time.strftime("%Y%m%d", time.localtime())
return f"{arch}-{project['label_type']}-{epochs}ep-{class_tag}-{date_str}"
from backend import (augment, base_dataset, config, dataset, datasets, 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")
#: Rough bytes one decoded 640px training image occupies in the RAM cache,
#: measured against the run that OOMed: 10.2 GB across 15,774 images.
_BYTES_PER_CACHED_IMAGE = 700_000
def _cache_mode(train_images: int, job) -> object:
"""Pick Ultralytics' `cache` argument for the memory this host actually has.
RAM caching is a large speedup and worth taking when it fits. It is only
taken with three times the headroom the raw estimate asks for: the run that
died had ~10 GB of cache on a 30 GB host and still lost, because the
dataloader workers fork after the cache is built and their copy-on-write
pages are what turn "just fits" into a kill. Two-times headroom would have
green-lit exactly the run that failed.
"""
needed = train_images * _BYTES_PER_CACHED_IMAGE
try:
import psutil
available = psutil.virtual_memory().available
except Exception:
available = 0
if available == 0:
job.log(f"Image cache: disk (cannot read free memory; {train_images} images)")
return "disk"
if needed * 3 <= available:
job.log(f"Image cache: RAM (~{needed / 1e9:.1f} GB of "
f"{available / 1e9:.1f} GB free)")
return "ram"
job.log(f"Image cache: disk (RAM cache would need ~{needed / 1e9:.1f} GB, "
f"only {available / 1e9:.1f} GB free)")
return "disk"
def start(project_id: int, epochs: int = 50, overrides: Optional[dict] = None,
batch_ids: Optional[list] = None, class_ids: Optional[list] = None,
dataset_ids: Optional[list] = None,
base_dataset_ids: Optional[list] = None) -> dict:
project = projects.get(project_id)
if project is None:
raise TrainingError("No such project")
bases = list(base_dataset_ids or [])
# No dataset picked means "everything this project has", which is what the
# single-dataset app always did. Base datasets are opt-in, so an empty pick
# never silently drags them in.
chosen = list(dataset_ids or [])
if not chosen and not bases:
chosen = [item["id"] for item in datasets.listing(project_id)]
if not chosen and not bases:
raise TrainingError(
"This project has no dataset yet — approve and merge a batch before training"
)
items = datasets.combined_items(project_id, chosen) if chosen else []
base_train = sum(base_dataset.get(bid)["image_count"] for bid in bases
if base_dataset.get(bid) is not None)
counts = {"train": sum(1 for i in items if i["split"] == "train") + base_train,
"val": sum(1 for i in items if i["split"] == "val")}
if counts["train"] == 0:
raise TrainingError(
"The chosen dataset(s) hold no training images — merge a batch before training"
)
# A base dataset is train-only, so it can never supply the val split that
# REQ-063's base-vs-new comparison is measured on.
if counts["val"] == 0:
raise TrainingError(
"Nothing to validate on — a base dataset only contributes training images, "
"so pick at least one of this project's own datasets too"
)
settings = hardware.resolve(overrides, epochs)
job = jobs.create(
"train",
params={"project_id": project_id, "settings": settings, "batch_ids": batch_ids,
"class_ids": class_ids, "dataset_ids": chosen,
"base_dataset_ids": bases},
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 rename(model_id: int, name: str) -> dict:
"""Update a model version's human-readable name."""
version = get_version(model_id)
if version is None:
raise TrainingError("No such model version")
with db.cursor() as cur:
cur.execute("UPDATE model_versions SET name = ? WHERE id = ?", (name.strip(), model_id))
return get_version(model_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")
dataset_ids = job.params.get("dataset_ids") or []
data_yaml = dataset.write_data_yaml(project, dataset_ids, batch_ids=batch_ids,
selected_class_ids=class_ids, require_val=True,
base_dataset_ids=job.params.get("base_dataset_ids") or [])
# 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}")
if job.cancelled:
trainer.stop_training = True
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
# REQ-110: explicit rather than inherited. An untouched project gets MEDIUM,
# which is Ultralytics' own default set, so this changes nothing by itself.
augmentation = augment.get(project["id"])
job.log(f"Augmentation: {augmentation['preset']} — "
+ ", ".join(f"{k}={v:g}" for k, v in sorted(augmentation["settings"].items())))
# `cache="ram"` used to be hardcoded. It holds the whole training set in
# memory, which was invisible at a few hundred images and fatal at fifteen
# thousand: the run below died mid-epoch with no traceback, killed by the
# host OOM killer, because 10 GB of cache plus per-worker copies did not fit
# in 30 GB. RAM caching is now earned, not assumed (CLAUDE.md §9).
train_list = os.path.join(os.path.dirname(data_yaml), "selected_train.txt")
with open(train_list, encoding="utf-8") as handle:
train_images = sum(1 for line in handle if line.strip())
cache_mode = _cache_mode(train_images, job)
keep_run_dir = False
try:
model.train(
data=data_yaml,
epochs=settings["epochs"],
imgsz=settings["imgsz"],
batch=settings["batch"],
**augmentation["settings"],
device=settings["device"],
workers=settings.get("workers", 8),
cache=cache_mode,
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, augment, name)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)""",
(project["id"], version, weights, project["base_model_path"],
json.dumps(comparison["new"]), json.dumps(comparison["base"]), time.time(),
json.dumps(augmentation["settings"]),
_generate_model_name(project, start_point, settings["epochs"], class_ids)),
)
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)