"""Run one small CPU training epoch for every task supported by the WebUI.""" from __future__ import annotations import argparse import json import time import traceback from pathlib import Path from yolo_webui import TrainingConfig, TrainingRunner from yolo_webui.config import AugmentationConfig, MlflowConfig TASKS = { "detect": ("datasets/yolo26_smoke/detect/detect.yaml", "yolo26n.pt"), "segment": ("datasets/yolo26_smoke/segment/segment.yaml", "yolo26n-seg.pt"), "classify": ("datasets/yolo26_smoke/classify", "yolo26n-cls.pt"), "pose": ("datasets/yolo26_smoke/pose/pose.yaml", "yolo26n-pose.pt"), "obb": ("datasets/yolo26_smoke/obb/obb.yaml", "yolo26n-obb.pt"), } def main() -> None: parser = argparse.ArgumentParser() parser.add_argument( "--mlflow", action="store_true", help="Enable MLflow logging for every smoke-training run.", ) parser.add_argument( "--tracking-uri", default="sqlite:///runs/yolo26_mlflow_smoke/mlflow.db", ) parser.add_argument("--experiment", default="yolo26-mlflow-smoke") parser.add_argument("--project", type=Path) args = parser.parse_args() results: dict[str, dict[str, object]] = {} default_project = "runs/yolo26_mlflow_smoke/train" if args.mlflow else "runs/yolo26_smoke" project_dir = (args.project or Path(default_project)).resolve() for task, (dataset, model) in TASKS.items(): print(f"\n=== {task}: {model} ===", flush=True) config = TrainingConfig( dataset=dataset, model=model, task=task, # type: ignore[arg-type] epochs=1, image_size=64, batch_size=2, device="cpu", workers=0, patience=0, project=str(project_dir), run_name=task, augmentation=AugmentationConfig(enabled=False), mlflow=MlflowConfig( enabled=args.mlflow, tracking_uri=args.tracking_uri, experiment_name=args.experiment, run_name=f"{task}-smoke" if args.mlflow else "", ), ) started = time.monotonic() try: output = TrainingRunner().train( config, lambda event: print( f"[{event.kind}] {event.message}", flush=True, ), ) results[task] = { "status": "succeeded", "seconds": round(time.monotonic() - started, 2), "output": str(output) if output else None, } except Exception as exc: traceback.print_exc() results[task] = { "status": "failed", "seconds": round(time.monotonic() - started, 2), "error": f"{type(exc).__name__}: {exc}", } summary_path = project_dir.parent / "smoke_summary.json" if args.mlflow else project_dir / "smoke_summary.json" summary_path.parent.mkdir(parents=True, exist_ok=True) summary_path.write_text( json.dumps(results, indent=2, ensure_ascii=False) + "\n", encoding="utf-8", ) print(f"\nSummary: {summary_path.resolve()}") print(json.dumps(results, indent=2, ensure_ascii=False)) if any(result["status"] != "succeeded" for result in results.values()): raise SystemExit(1) if __name__ == "__main__": main()