"""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)