#!/usr/bin/env python3 """BC v3 — classification head chống multi-modality (fail path §7, bước 2). Bối cảnh (20/07, dataset 10k canonical): canonical tie-break sửa được PHI (MSE 0.026 vs baseline 0.29 — chọn-lỗ hết nhảy mode) nhưng V0/side/vert vẫn dính variance baseline ở CẢ 5 config (canonical/p005/p010/p015/n2000) → spin bị regression-về-mean nghiền về ~0, Q|pot kẹt mốc blind 0.53. Chẩn đoán: label (V0, side, vert) nằm trên GRID RỜI RẠC (10×5×5 = 250 combo) và hàm obs→combo đa mode — regression MSE trung bình hoá mode; classification cross-entropy thì KHÔNG: softmax argmax trả về một mode thật. Kiến trúc v3 (net riêng, KHÔNG phải SB3): trunk MLP 2×256 ReLU ├─ phi head : regression 1D (phi đã học được bằng MSE — giữ nguyên) └─ class head : 250-way CE trên index (i_v0*25 + i_side*5 + i_vert) label = canonical tie-break (import relabel_table — đúng rule đã chạy) 3 tầng đánh giá (tách bạch để fail ở tầng nào biết tầng đó): 1. OFFLINE ORACLE-CHECK (giây, không sim): tra class dự đoán vào pot combos đã lưu của bàn val → "pot ảo" + Q — upper bound nhanh. 2. EVAL TRỰC TIẾP trong env (mặc định 1000 cú): duck-type classifier vào train_position.evaluate — KẾT QUẢ CHÍNH, đo classification có phá được multi-modality không, không dính nhiễu distill. 3. DISTILL npz → train_bc.py: clone policy classifier vào kiến trúc PPO SB3 (cần cho --init-from). Distill MSE thoát baseline hay không là câu hỏi RIÊNG (representability của 64×64), không phải câu hỏi 1-2. Chạy từ gốc repo (venv local, cần torch + pooltool): python scripts/train_bc_v3.py --run-name bc_v3_20260720 python scripts/train_bc.py --dataset data/bc_dataset_10000_124_v3distill.npz python scripts/eval_position.py models//bc_model.zip --episodes 1000 --aim-mode any Gate G1 không đổi: Q|pot > 0.65, pot >= 10% (eval 1000 cú). """ from __future__ import annotations import argparse import math import sys import time from pathlib import Path sys.stdout.reconfigure(encoding="utf-8") sys.stderr.reconfigure(encoding="utf-8") ROOT = Path(__file__).resolve().parents[1] sys.path.insert(0, str(ROOT / "src")) sys.path.insert(0, str(ROOT / "scripts")) N_V0, N_SPIN = 10, 5 N_CLS = N_V0 * N_SPIN * N_SPIN # 250 # --------------------------------------------------------------- pure helpers # (numpy thuần, không import nặng — test được trong sandbox) def _grid_idx(val: float, grid) -> int: """Index grid gần nhất (combo lưu float32 của grid float64 → nearest an toàn).""" import numpy as np return int(np.argmin(np.abs(np.asarray(grid) - val))) def _cls_of(v0: float, side: float, vert: float, grids) -> int: v0_grid, side_grid, vert_grid = grids return (_grid_idx(v0, v0_grid) * N_SPIN * N_SPIN + _grid_idx(side, side_grid) * N_SPIN + _grid_idx(vert, vert_grid)) def _cls_to_vals(cls: int, grids): v0_grid, side_grid, vert_grid = grids iv, rem = divmod(int(cls), N_SPIN * N_SPIN) isd, ivt = divmod(rem, N_SPIN) return float(v0_grid[iv]), float(side_grid[isd]), float(vert_grid[ivt]) def _norm_action(phi_deg: float, v0: float, side: float, vert: float): """Map vật lý → action [-1,1]^4 (CÙNG công thức _to_action đã verify round-trip 20/07 — nhân bản để module import được không cần pooltool).""" import numpy as np return np.clip(np.array([phi_deg / 180.0 - 1.0, (v0 - 0.5) / 3.5 * 2.0 - 1.0, side / 0.4, vert / 0.4], dtype=np.float32), -1.0, 1.0) def _circ_diff_deg(a: float, b: float) -> float: d = abs(a - b) % 360.0 return min(d, 360.0 - d) def _oracle_match(board, pred_cls: int, pred_phi_deg: float, phi_tol: float): """Tra class dự đoán vào pot combos của bàn. board = (cls_arr, phi_arr, q_arr, b2_arr). Trả (q, b2p, dphi) nếu class tồn tại và có phi (lỗ) trong tolerance của phi dự đoán; None nếu miss. """ import numpy as np cls_arr, phi_arr, q_arr, b2_arr = board hit = np.flatnonzero(cls_arr == pred_cls) if len(hit) == 0: return None dphis = np.array([_circ_diff_deg(float(p), pred_phi_deg) for p in phi_arr[hit]]) j = int(np.argmin(dphis)) if dphis[j] > phi_tol: return None k = hit[j] return float(q_arr[k]), float(b2_arr[k]), float(dphis[j]) # --------------------------------------------------------------------- main def main(): p = argparse.ArgumentParser() p.add_argument("--dataset", default="data/bc_dataset_10000_124.npz", help="npz GỐC từ gen_bc_dataset.py (cần combos + grids)") p.add_argument("--tie-margin", type=float, default=0.05) p.add_argument("--pocket-margin", type=float, default=0.0, help="giữ 0 — sweep 20/07: filter giết phi qua đường mất data") p.add_argument("--limit", type=int, default=None, help="ablation N bàn") p.add_argument("--epochs", type=int, default=300) p.add_argument("--lr", type=float, default=3e-4) p.add_argument("--batch", type=int, default=256) p.add_argument("--hidden", type=int, default=256) p.add_argument("--phi-coef", type=float, default=5.0, help="trọng số MSE phi so với CE (CE khởi điểm ~ln250≈5.5)") p.add_argument("--val-frac", type=float, default=0.1) p.add_argument("--patience", type=int, default=30) p.add_argument("--seed", type=int, default=0) p.add_argument("--phi-tol", type=float, default=3.0, help="tolerance (độ) khớp phi trong offline oracle-check") p.add_argument("--eval-episodes", type=int, default=1000, help="eval trực tiếp classifier trong env (0 = bỏ)") p.add_argument("--aim-mode", choices=["best_cut", "any"], default="any") p.add_argument("--run-name", default=None) p.add_argument("--distill-out", default=None, help="mặc định _v3distill[_p..][_n..].npz") args = p.parse_args() import numpy as np import torch from torch import nn from relabel_bc_dataset import relabel_table # đúng rule canonical đã chạy run = args.run_name or f"bc_v3_{time.strftime('%Y%m%d_%H%M%S')}" model_dir = ROOT / "models" / run log_dir = ROOT / "logs" / run model_dir.mkdir(parents=True, exist_ok=True) log_dir.mkdir(parents=True, exist_ok=True) # ------------------------------------------------------------- dataset data = np.load(args.dataset) if "combos" not in data: sys.exit("npz không có `combos` — cần bản gen 20/07") grids = (data["v0_grid"], data["side_grid"], data["vert_grid"]) obs_all = data["obs"].astype(np.float32) combos, combo_row = data["combos"], data["combo_row"] n_raw = len(obs_all) if args.limit is None else min(args.limit, len(obs_all)) starts = np.searchsorted(combo_row, np.arange(n_raw)) ends = np.searchsorted(combo_row, np.arange(n_raw) + 1) obs_l, phi_t_l, cls_l, q_l, npot_l, boards = [], [], [], [], [], [] for r in range(n_raw): c = combos[starts[r]:ends[r]] pick = relabel_table(c, args.tie_margin, args.pocket_margin, include_b2=False) if pick is None: continue phi, v0, side, vert, q, _b2p = (float(x) for x in pick) obs_l.append(obs_all[r]) phi_t_l.append(phi / 180.0 - 1.0) cls_l.append(_cls_of(v0, side, vert, grids)) q_l.append(q) npot_l.append(ends[r] - starts[r]) # lookup mọi pot combo của bàn (KỂ CẢ b2 — env tính là pot) cho # offline check + q distill cls_arr = np.array([_cls_of(float(cc[1]), float(cc[2]), float(cc[3]), grids) for cc in c], dtype=np.int32) boards.append((cls_arr, c[:, 0].copy(), c[:, 4].copy(), c[:, 5].copy())) obs = np.stack(obs_l) phi_t = np.array(phi_t_l, dtype=np.float32) cls_t = np.array(cls_l, dtype=np.int64) q_lbl = np.array(q_l, dtype=np.float32) n = len(obs) uniq, cnt = np.unique(cls_t, return_counts=True) print(f"== BC v3: {n} bàn, label canonical (tie {args.tie_margin}, " f"pocket {args.pocket_margin}) ==") print(f" class dùng: {len(uniq)}/{N_CLS}, class lớn nhất " f"{cnt.max()/n:.1%}, Q label mean {q_lbl.mean():.3f}") print(f" trunk 2x{args.hidden}, phi_coef {args.phi_coef}, lr {args.lr}, " f"batch {args.batch}, max {args.epochs} epoch, patience {args.patience}\n") rng = np.random.default_rng(args.seed) perm = rng.permutation(n) n_val = max(1, int(n * args.val_frac)) val_idx, tr_idx = perm[:n_val], perm[n_val:] # --------------------------------------------------------------- model torch.manual_seed(args.seed) device = "cuda" if torch.cuda.is_available() else "cpu" class Net(nn.Module): def __init__(self, h): super().__init__() self.trunk = nn.Sequential(nn.Linear(6, h), nn.ReLU(), nn.Linear(h, h), nn.ReLU()) self.phi_head = nn.Linear(h, 1) self.cls_head = nn.Linear(h, N_CLS) def forward(self, x): z = self.trunk(x) return self.phi_head(z).squeeze(-1), self.cls_head(z) net = Net(args.hidden).to(device) opt = torch.optim.Adam(net.parameters(), lr=args.lr) ce = nn.CrossEntropyLoss() mse = nn.functional.mse_loss obs_T = torch.as_tensor(obs, device=device) phi_T = torch.as_tensor(phi_t, device=device) cls_T = torch.as_tensor(cls_t, device=device) tr_T = torch.as_tensor(tr_idx, device=device) val_T = torch.as_tensor(val_idx, device=device) def val_metrics(): net.eval() with torch.no_grad(): ph, lg = net(obs_T[val_T]) v_phi = mse(ph, phi_T[val_T]).item() v_ce = ce(lg, cls_T[val_T]).item() top5 = lg.topk(5, dim=1).indices t1 = (top5[:, 0] == cls_T[val_T]).float().mean().item() t5 = (top5 == cls_T[val_T].unsqueeze(1)).any(1).float().mean().item() net.train() return v_phi, v_ce, t1, t5 # -------------------------------------------------------- training loop hist = {"train": [], "val": [], "top1": [], "top5": [], "phi": []} best_val, best_epoch, best_state = float("inf"), 0, None t0 = time.time() for epoch in range(1, args.epochs + 1): ep_perm = torch.randperm(len(tr_T), device=device) tl_sum, nb = 0.0, 0 for s in range(0, len(tr_T), args.batch): b = tr_T[ep_perm[s:s + args.batch]] ph, lg = net(obs_T[b]) loss = args.phi_coef * mse(ph, phi_T[b]) + ce(lg, cls_T[b]) opt.zero_grad() loss.backward() opt.step() tl_sum += loss.item() nb += 1 v_phi, v_ce, t1, t5 = val_metrics() v_total = args.phi_coef * v_phi + v_ce hist["train"].append(tl_sum / nb) hist["val"].append(v_total) hist["top1"].append(t1) hist["top5"].append(t5) hist["phi"].append(v_phi) if v_total < best_val - 1e-5: best_val, best_epoch = v_total, epoch best_state = {k: v.detach().clone() for k, v in net.state_dict().items()} if epoch == 1 or epoch % 10 == 0: print(f"epoch {epoch:3d}: train {tl_sum/nb:.4f} | val phi MSE " f"{v_phi:.4f} CE {v_ce:.4f} top1 {t1:.1%} top5 {t5:.1%}") if epoch - best_epoch >= args.patience: print(f"Early stop @ epoch {epoch} " f"(best val {best_val:.4f} tại epoch {best_epoch})") break net.load_state_dict(best_state) net.eval() v_phi, v_ce, t1, t5 = val_metrics() print(f"\nBC v3 xong trong {(time.time()-t0)/60:.1f} phút — best epoch " f"{best_epoch}: phi MSE {v_phi:.4f} (v2 canonical: 0.026), " f"top1 {t1:.1%}, top5 {t5:.1%}") print(" (top1 thấp KHÔNG tự động là xấu — near-tie nhiều mode hợp lệ; " "phán quyết ở oracle-check + eval env)") torch.save(net.state_dict(), model_dir / "classifier.pt") # ------------------------------------------- forward toàn bộ (1 lần) with torch.no_grad(): ph_all, lg_all = net(obs_T) phi_pred = ph_all.cpu().numpy() cls_pred = lg_all.argmax(1).cpu().numpy() # ------------------------------------- tầng 1: offline oracle-check (val) hits, dphis = [], [] for i in val_idx: m = _oracle_match(boards[i], int(cls_pred[i]), (float(phi_pred[i]) + 1.0) * 180.0, args.phi_tol) if m is not None: hits.append(m) dphis.append(m[2]) print(f"\n== Tầng 1 — offline oracle-check ({n_val} bàn val, phi_tol " f"{args.phi_tol}°; UPPER BOUND — phi coi như trúng) ==") if hits: qs = np.array([h[0] for h in hits]) b2s = np.array([h[1] for h in hits]) print(f" match (pot ảo) : {len(hits)/n_val:.1%}") print(f" Q | match : {qs.mean():.3f} (label ceiling: " f"{q_lbl[val_idx].mean():.3f}, mốc blind: 0.53)") print(f" b2-lucky share : {b2s.mean():.1%}") print(f" |dphi| mean : {np.mean(dphis):.2f}°") else: print(" 0 match — class dự đoán không pot được ở bàn nào (!!)") # --------------------------------- tầng 2: eval trực tiếp trong env if args.eval_episodes > 0: from train_position import evaluate, print_stats class _ClassifierPolicy: """Duck-type SB3: predict(obs) → action [-1,1]^4.""" def predict(self, o, deterministic=True): with torch.no_grad(): ph, lg = net(torch.as_tensor( np.asarray(o, dtype=np.float32), device=device).unsqueeze(0)) v0, side, vert = _cls_to_vals(int(lg.argmax(1)), grids) phi_deg = (float(np.clip(ph.item(), -1.0, 1.0)) + 1.0) * 180.0 return _norm_action(phi_deg, v0, side, vert), None print(f"\n== Tầng 2 — eval classifier TRỰC TIẾP trong env " f"({args.eval_episodes} cú, deterministic) — KẾT QUẢ CHÍNH ==") stats = evaluate(_ClassifierPolicy(), n_episodes=args.eval_episodes, aim_mode=args.aim_mode) print_stats(stats) n_pot = max(1, round(stats["pot_rate"] * args.eval_episodes)) se = 0.25 / math.sqrt(n_pot) q = stats["q_mean_on_pot"] print(f" n cú pot ≈ {n_pot} → SE(Q|pot) ≈ ±{se:.3f} " f"(Q|pot ± 2·SE = [{q-2*se:.3f}, {q+2*se:.3f}])") print(f" GATE G1: Q|pot > 0.65 và pot >= 10% " f"{'→ PASS' if q - 2*se > 0.53 and stats['pot_rate'] >= 0.10 else ''}") # ----------------------------------------- tầng 3: distill npz cho SB3 act_out = np.stack([ _norm_action((float(np.clip(phi_pred[i], -1, 1)) + 1.0) * 180.0, *_cls_to_vals(int(cls_pred[i]), grids)) for i in range(n)]).astype(np.float32) q_out = np.zeros(n, dtype=np.float32) b2_out = np.zeros(n, dtype=np.int8) for i in range(n): m = _oracle_match(boards[i], int(cls_pred[i]), (float(phi_pred[i]) + 1.0) * 180.0, args.phi_tol) if m is not None: q_out[i], b2_out[i] = m[0], int(m[1] > 0.5) d_out = (Path(args.distill_out) if args.distill_out else Path(args.dataset).with_name( Path(args.dataset).stem + "_v3distill" + (f"_p{args.pocket_margin:g}" if args.pocket_margin > 0 else "") + (f"_n{args.limit}" if args.limit else "") + ".npz")) np.savez_compressed(d_out, obs=obs, actions=act_out, q=q_out, n_pot=np.array(npot_l, dtype=np.int32), b2_lucky=b2_out) print(f"\n== Tầng 3 — distill dataset (clone policy classifier, " f"{n} sample) ==") print(f" |side|/|vert| action: {np.abs(act_out[:, 2]).mean():.2f} / " f"{np.abs(act_out[:, 3]).mean():.2f} (label canonical: 0.28/0.47 — " f"còn ~0.0 nghĩa là classifier cũng sập về mean)") print(f" q>0 (match oracle): {(q_out > 0).mean():.1%}") print(f"Dataset -> {d_out}") # ----------------------------------------------------------- loss plot import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt ep_x = np.arange(1, len(hist["val"]) + 1) fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4.5)) ax1.plot(ep_x, hist["train"], label="train total") ax1.plot(ep_x, hist["val"], label="val total") ax1.plot(ep_x, np.array(hist["phi"]) * args.phi_coef, ls="--", label=f"val phi×{args.phi_coef:g}") ax1.axvline(best_epoch, c="tab:red", ls=":", alpha=0.7, label=f"best {best_epoch}") ax1.set_yscale("log") ax1.set_xlabel("epoch") ax1.set_ylabel("loss") ax1.set_title(f"BC v3 — {n} sample, {len(uniq)} class") ax1.legend(loc="upper right") ax1.grid(alpha=0.3) ax2.plot(ep_x, hist["top1"], label="top-1") ax2.plot(ep_x, hist["top5"], label="top-5") ax2.set_xlabel("epoch") ax2.set_ylabel("val accuracy") ax2.set_ylim(0, 1) ax2.set_title("class accuracy (tham khảo — near-tie nhiều mode hợp lệ)") ax2.legend(loc="upper right") ax2.grid(alpha=0.3) fig.tight_layout() fig.savefig(log_dir / "loss_curve.png", dpi=130) print(f"Loss curve -> {log_dir / 'loss_curve.png'}") rel = d_out.relative_to(ROOT) if d_out.is_relative_to(ROOT) else d_out print(f"\nBước kế (distill → SB3 zip cho --init-from):") print(f" python scripts/train_bc.py --dataset {rel} --run-name {run}_distill") print(f" python scripts/eval_position.py models/{run}_distill/bc_model.zip " f"--episodes 1000 --aim-mode any") if __name__ == "__main__": main()