Spaces:
Sleeping
Sleeping
| """Train YOLO baseline detect bi (F2 Bước 3, BRIEF 05/08/2026). | |
| Variant chọn: **YOLO11n** — thế hệ hiện tại của ultralytics, ít tham số hơn | |
| v8n một chút với mAP COCO nhỉnh hơn, cùng API; paper pix2pockets dùng YOLOv5 | |
| nhưng không có ràng buộc tương thích nào buộc theo. Fine-tune từ pretrained | |
| COCO (yolo11n.pt, tải tự động từ release chính thức ultralytics). | |
| Máy 05/08: có RTX 3070 nhưng venv dùng chung cài torch CPU-build (đổi sang | |
| CUDA sẽ lật device auto của SB3 cho cả nhánh RL → quyết định của Cowork, | |
| không tự đổi). Train CPU, setting rút gọn cho timebox; bản đầy đủ chạy đêm | |
| có launcher `run_cv_train_full.bat` ở thư mục cha. | |
| In/ghi ASCII-only: console Windows mặc định cp1252, in tiếng Việt có dấu | |
| là UnicodeEncodeError (bẫy đã dính 05/08 ngay trong fetch_dataset.py). | |
| python scripts/cv/train_baseline.py --epochs 30 --imgsz 640 --batch 8 | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import json | |
| import shutil | |
| import sys | |
| import time | |
| from pathlib import Path | |
| ROOT = Path(__file__).resolve().parents[2] # poolcoach-rl/ | |
| sys.path.insert(0, str(ROOT / "src")) | |
| DATA_YAML = ROOT / "datasets" / "pix2pockets" / "yolo" / "data.yaml" | |
| DEFAULT_ARTIFACT_DIR = ROOT.parent / "cv_baseline_20260805" | |
| RUNS_DIR = ROOT / "runs" / "cv" # runs/ đã gitignore | |
| SEED = 20260805 | |
| def _env_report() -> dict: | |
| import torch | |
| info = { | |
| "torch": torch.__version__, | |
| "cuda_available": torch.cuda.is_available(), | |
| "cpu_threads": torch.get_num_threads(), | |
| } | |
| print(f"[env] torch={info['torch']} cuda_available={info['cuda_available']} " | |
| f"cpu_threads={info['cpu_threads']}") | |
| return info | |
| def main() -> None: | |
| ap = argparse.ArgumentParser(description="Train YOLO11n ball-detection baseline") | |
| ap.add_argument("--model", default="yolo11n.pt", help="pretrained checkpoint") | |
| ap.add_argument("--epochs", type=int, default=30) | |
| ap.add_argument("--imgsz", type=int, default=640) | |
| ap.add_argument("--batch", type=int, default=8) | |
| ap.add_argument("--device", default="cpu", help="'cpu' or CUDA index") | |
| ap.add_argument("--workers", type=int, default=2) | |
| ap.add_argument("--name", default="cv_baseline_20260805", help="run name under runs/cv/") | |
| ap.add_argument("--artifact-dir", type=Path, default=DEFAULT_ARTIFACT_DIR) | |
| args = ap.parse_args() | |
| if not DATA_YAML.exists(): | |
| sys.exit(f"[ERROR] {DATA_YAML} missing - run scripts/cv/fetch_dataset.py first.") | |
| from ultralytics import YOLO | |
| env = _env_report() | |
| t0 = time.time() | |
| model = YOLO(args.model) | |
| results = model.train( | |
| data=str(DATA_YAML), | |
| epochs=args.epochs, | |
| imgsz=args.imgsz, | |
| batch=args.batch, | |
| device=args.device, | |
| workers=args.workers, | |
| seed=SEED, | |
| deterministic=True, | |
| project=str(RUNS_DIR), | |
| name=args.name, | |
| exist_ok=True, | |
| plots=True, | |
| verbose=True, | |
| ) | |
| wall_s = time.time() - t0 | |
| run_dir = Path(results.save_dir) | |
| print(f"[done] train wall time: {wall_s / 60:.1f} min, run dir: {run_dir}") | |
| # Val cuoi cung tren best.pt de lay so bao cao (train da val moi epoch, | |
| # nhung so chinh thuc lay tu best weights cho ro rang). | |
| best = run_dir / "weights" / "best.pt" | |
| metrics = YOLO(str(best)).val(data=str(DATA_YAML), device=args.device, | |
| project=str(RUNS_DIR), name=args.name + "_val", | |
| exist_ok=True) | |
| names = metrics.names | |
| per_class = {} | |
| for k, ci in enumerate(metrics.box.ap_class_index.tolist()): | |
| p, r, ap50, ap = metrics.box.class_result(k) | |
| per_class[names[ci]] = {"precision": round(p, 4), "recall": round(r, 4), | |
| "ap50": round(ap50, 4), "ap50_95": round(ap, 4)} | |
| summary = { | |
| "date": "2026-08-05", | |
| "model": args.model, | |
| "epochs": args.epochs, | |
| "imgsz": args.imgsz, | |
| "batch": args.batch, | |
| "device": args.device, | |
| "seed": SEED, | |
| "env": env, | |
| "data": str(DATA_YAML), | |
| "train_wall_min": round(wall_s / 60, 1), | |
| "map50": round(metrics.box.map50, 4), | |
| "map50_95": round(metrics.box.map, 4), | |
| "per_class": per_class, | |
| "best_weights": str(best), | |
| "run_dir": str(run_dir), | |
| } | |
| print("[metrics] " + json.dumps( | |
| {k: summary[k] for k in ("map50", "map50_95", "train_wall_min")})) | |
| for cls, m in per_class.items(): | |
| print(f"[metrics] {cls:>8}: AP50={m['ap50']:.3f} AP50-95={m['ap50_95']:.3f}") | |
| # Artifact cho Cowork archive (BRIEF buoc 3.4) | |
| art = args.artifact_dir | |
| art.mkdir(parents=True, exist_ok=True) | |
| (art / "metrics.json").write_text(json.dumps(summary, indent=2), encoding="utf-8") | |
| for f in ("results.csv", "args.yaml", "results.png", "confusion_matrix.png", | |
| "labels.jpg"): | |
| src = run_dir / f | |
| if src.exists(): | |
| shutil.copy2(src, art / f) | |
| for f in sorted(run_dir.glob("val_batch*_pred.jpg"))[:3]: | |
| shutil.copy2(f, art / f.name) | |
| for f in sorted((run_dir.parent / (args.name + "_val")).glob("val_batch*_pred.jpg"))[:3]: | |
| shutil.copy2(f, art / ("final_" + f.name)) | |
| shutil.copy2(best, art / "best.pt") | |
| print(f"[artifact] -> {art}") | |
| if __name__ == "__main__": | |
| main() | |