poolcoach / scripts /train_qfield.py
masterdanh's picture
deploy: snapshot for HF Space
78738de
Raw
History Blame Contribute Delete
22.4 kB
#!/usr/bin/env python3
"""BC v4 — Q-landscape distillation: học TRƯỜNG kết cục, không học action.
Bối cảnh (20-21/07/2026): chuỗi v1 (argmax) → v2 (canonical) → v3
(classification) đóng nhánh imitation đơn-action: label (V0, side, vert)
gần như không mang thông tin điều kiện theo obs (val CE ≈ marginal entropy).
Vấn đề là TARGET, không phải loss.
v4 đổi target (Documents/PoolCoach_Stage2_QField_Design.md): dùng 870k pot
combo đã lưu + negatives khôi phục từ obs (grid × lỗ khả thi; combo vắng mặt
= không pot sạch) → ~560 điểm supervision/bàn thay vì 1:
Net([obs(6), cos φ, sin φ]) → 250 cell × (p̂ pot sạch, Q̂ position)
loss = BCE(pot) + Q_COEF × MSE(Q̂) masked theo (pot && !b2-lucky)
Multi-modality biến mất theo định nghĩa (mỗi cell có đúng 1 nhãn oracle).
Phi KHÔNG cần học: aim = ghost-ball analytic của lỗ chọn ở inference.
Inference: argmax score trên (mọi lỗ khả thi × 250 cell) → action.
3 tầng đánh giá (tầng 1 giờ CHÍNH XÁC tuyệt đối — cell chọn nằm trên grid
đã simulate, kết cục tra từ dataset, không sim):
1. Offline selection eval (val): virtual pot% + virtual Q|pot; track
theo epoch, vẽ vào learning curve. KILL-SWITCH (§7 design doc):
sau run đầu + tối đa 2 chỉnh mà Q|pot(xb2) < 0.60 hoặc pot < 25%
→ DỪNG nhánh BC, pivot app oracle-at-inference.
2. Eval trực tiếp trong env (--eval-episodes, mặc định 1000 cú).
Gate G1: Q|pot > 0.65 (ngoài 2·SE) và pot >= 10%.
3. Distill npz → train_bc.py → SB3 zip cho --init-from (pipeline v3).
Chạy từ gốc repo (venv local, cần torch + pooltool):
python scripts/train_qfield.py --run-name qfield_20260721
python scripts/train_qfield.py --select qgate # nếu EV bảo thủ về Q
"""
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
POS_COEF = 0.5 # khớp env — dùng trong score EV = p̂(1 + POS_COEF·Q̂)
PHI_MATCH_TOL = 0.5 # (độ) match combo phi ↔ phi ứng viên recompute từ obs
# --------------------------------------------------------------- pure helpers
# (numpy thuần, không import nặng — test được trong sandbox)
def _grid_idx(val: float, grid) -> int:
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:
"""Quy ước index y hệt v3: i_v0*25 + i_side*5 + i_vert."""
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 đã verify round-trip)."""
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 build_board_targets(phis, combos_board, grids):
"""Dense targets 1 bàn từ pot combos + K phi ứng viên.
phis: list K phi (độ, recompute từ obs — đúng logic _ghost_dirs_any).
combos_board: (n,6) [phi, v0, side, vert, q, b2p] — MỌI pot combo.
Trả (pot (K,250) f32, q (K,250) f32, qmask (K,250) bool, n_unmatched):
pot[k,c]=1 nếu có combo (phi_k, cell c); combo vắng = không pot sạch
q[k,c] = Q combo (b2-lucky: q=1.0, GIỮ để đọc kết cục env-consistent)
qmask[k,c] = pot && !b2 — chỉ cell này vào Q-loss (b2 thổi phồng Q)
n_unmatched: combo không match phi nào trong tol (kỳ vọng 0 — bug nếu >0).
"""
import numpy as np
K = len(phis)
pot = np.zeros((K, N_CLS), dtype=np.float32)
q = np.zeros((K, N_CLS), dtype=np.float32)
qm = np.zeros((K, N_CLS), dtype=bool)
unmatched = 0
for row in combos_board:
phi_c, v0, side, vert, qq, b2p = (float(x) for x in row)
dphis = [_circ_diff_deg(phi_c, p) for p in phis]
k = int(np.argmin(dphis))
if dphis[k] > PHI_MATCH_TOL:
unmatched += 1
continue
c = _cls_of(v0, side, vert, grids)
pot[k, c] = 1.0
q[k, c] = qq
if b2p < 0.5:
qm[k, c] = True
return pot, q, qm, unmatched
def select_action(p_hat, q_hat, mode: str = "ev", pot_thresh: float = 0.5):
"""Chọn (example k, cell c) từ p̂/q̂ shape (K, 250). Trả (k, c).
ev : argmax p̂ × (1 + POS_COEF·Q̂) — kỳ vọng reward env, mặc định.
qgate : trong cell p̂ >= thresh argmax Q̂; không cell nào đạt → max p̂.
"""
import numpy as np
if mode == "ev":
score = p_hat * (1.0 + POS_COEF * q_hat)
else: # qgate
ok = p_hat >= pot_thresh
if not ok.any():
k, c = np.unravel_index(int(np.argmax(p_hat)), p_hat.shape)
return int(k), int(c)
score = np.where(ok, q_hat, -1.0)
k, c = np.unravel_index(int(np.argmax(score)), score.shape)
return int(k), int(c)
# --------------------------------------------------------------------- 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("--limit", type=int, default=None, help="ablation N bàn")
p.add_argument("--q-coef", type=float, default=5.0,
help="trọng số MSE Q so với BCE pot")
p.add_argument("--epochs", type=int, default=200)
p.add_argument("--lr", type=float, default=3e-4)
p.add_argument("--batch", type=int, default=512)
p.add_argument("--hidden", type=int, default=256,
help="giữ 256 = v3 để so công bằng capacity")
p.add_argument("--val-frac", type=float, default=0.1)
p.add_argument("--patience", type=int, default=20)
p.add_argument("--seed", type=int, default=0)
p.add_argument("--select", choices=["ev", "qgate"], default="ev")
p.add_argument("--pot-thresh", type=float, default=0.5,
help="ngưỡng p̂ cho --select qgate")
p.add_argument("--eval-episodes", type=int, default=1000,
help="tầng 2: eval trực tiếp trong env (0 = bỏ)")
p.add_argument("--aim-mode", choices=["best_cut", "any"], default="any")
p.add_argument("--no-distill", action="store_true",
help="bỏ tầng 3 (distill npz cho train_bc.py)")
p.add_argument("--run-name", default=None)
args = p.parse_args()
import numpy as np
import torch
from torch import nn
from torch.nn import functional as F
from poolcoach_rl.envs import PositionPlayEnv
run = args.run_name or f"qfield_{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)
# env helper: CHỈ lấy geometry (_ghost_dirs_any, w, l) — không simulate
env_h = PositionPlayEnv()
w, l = env_h.w, env_h.l
# ------------------------------------------------------------- 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"])
assert len(grids[0]) == N_V0 and len(grids[1]) == N_SPIN \
and len(grids[2]) == N_SPIN, "grid npz không khớp N_V0/N_SPIN"
obs_all = data["obs"].astype(np.float32)
combos, combo_row = data["combos"], data["combo_row"]
n = len(obs_all) if args.limit is None else min(args.limit, len(obs_all))
starts = np.searchsorted(combo_row, np.arange(n))
ends = np.searchsorted(combo_row, np.arange(n) + 1)
print(f"== BC v4 Q-field: {n} bàn — build dense targets "
f"(negatives khôi phục từ obs, không re-sim) ==")
t0 = time.time()
x_l, pot_l, q_l, qm_l = [], [], [], []
board_phis = [] # list[list[phi]] per bàn — dùng lại ở distill
ex_start = np.zeros(n + 1, dtype=np.int64) # slice example của từng bàn
total_unmatched = 0
for r in range(n):
cue = np.array([obs_all[r][0] * w, obs_all[r][1] * l])
b1 = np.array([obs_all[r][2] * w, obs_all[r][3] * l])
phis = [float(np.degrees(np.arctan2(d[1], d[0])) % 360.0)
for d in env_h._ghost_dirs_any(cue, b1)]
pot_b, q_b, qm_b, um = build_board_targets(
phis, combos[starts[r]:ends[r]], grids)
total_unmatched += um
for k, phi in enumerate(phis):
rad = math.radians(phi)
x_l.append(np.concatenate([
obs_all[r],
np.array([math.cos(rad), math.sin(rad)], dtype=np.float32)]))
pot_l.append(pot_b)
q_l.append(q_b)
qm_l.append(qm_b)
board_phis.append(phis)
ex_start[r + 1] = ex_start[r] + len(phis)
X = np.stack(x_l).astype(np.float32)
POT = np.concatenate(pot_l)
Q = np.concatenate(q_l)
QM = np.concatenate(qm_l)
m = len(X)
n_pos = int(POT.sum())
print(f" {m} example ({m/n:.2f} lỗ/bàn), {m * N_CLS} điểm supervision, "
f"positives {n_pos} ({n_pos/(m*N_CLS):.1%}) "
f"[build {time.time()-t0:.0f}s]")
if total_unmatched:
frac = total_unmatched / max(1, int(POT.sum()) + total_unmatched)
print(f" !! {total_unmatched} combo không match phi ứng viên "
f"({frac:.2%}) — biên feasibility float32")
if frac > 0.01:
sys.exit("Unmatched > 1% — bug recompute phi, DỪNG (design §3)")
# split THEO BÀN (không theo example — tránh leak cùng bàn qua 2 phía)
rng = np.random.default_rng(args.seed)
perm = rng.permutation(n)
n_val = max(1, int(n * args.val_frac))
val_boards, tr_boards = perm[:n_val], perm[n_val:]
tr_ex = np.concatenate([np.arange(ex_start[b], ex_start[b + 1])
for b in tr_boards])
val_ex = np.concatenate([np.arange(ex_start[b], ex_start[b + 1])
for b in val_boards])
print(f" split theo bàn: {len(tr_boards)} train / {n_val} val "
f"({len(tr_ex)}/{len(val_ex)} example)\n")
# --------------------------------------------------------------- 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(8, h), nn.ReLU(),
nn.Linear(h, h), nn.ReLU())
self.pot_head = nn.Linear(h, N_CLS)
self.q_head = nn.Linear(h, N_CLS)
def forward(self, x):
z = self.trunk(x)
return self.pot_head(z), torch.sigmoid(self.q_head(z))
net = Net(args.hidden).to(device)
opt = torch.optim.Adam(net.parameters(), lr=args.lr)
X_T = torch.as_tensor(X, device=device)
POT_T = torch.as_tensor(POT, device=device)
Q_T = torch.as_tensor(Q, device=device)
QM_T = torch.as_tensor(QM, device=device)
tr_T = torch.as_tensor(tr_ex, device=device)
def _losses(idx):
lg, qs = net(X_T[idx])
bce = F.binary_cross_entropy_with_logits(lg, POT_T[idx])
mask = QM_T[idx]
qmse = (((qs - Q_T[idx])[mask]) ** 2).mean() if mask.any() \
else torch.zeros((), device=device)
return bce, qmse
@torch.no_grad()
def _forward_np(idx):
lg, qs = net(X_T[torch.as_tensor(idx, device=device)])
return torch.sigmoid(lg).cpu().numpy(), qs.cpu().numpy()
def offline_select(board_ids):
"""Tầng 1 — selection trên grid, kết cục tra từ dataset (CHÍNH XÁC).
Trả dict: pot% / Q|pot loại b2 / Q|pot kể b2 (=1.0) / b2 share
+ (board→(k, c)) để tái dùng ở distill.
"""
picks, pots, q_x, q_all, b2s = {}, [], [], [], []
for b in board_ids:
idx = np.arange(ex_start[b], ex_start[b + 1])
p_hat, q_hat = _forward_np(idx)
k, c = select_action(p_hat, q_hat, args.select, args.pot_thresh)
picks[int(b)] = (k, c)
gi = ex_start[b] + k
hit = POT[gi, c] > 0.5
pots.append(hit)
if hit:
b2 = not QM[gi, c]
b2s.append(b2)
q_all.append(Q[gi, c])
if not b2:
q_x.append(Q[gi, c])
return {
"pot": float(np.mean(pots)),
"q_xb2": float(np.mean(q_x)) if q_x else float("nan"),
"q_all": float(np.mean(q_all)) if q_all else float("nan"),
"b2_share": float(np.mean(b2s)) if b2s else 0.0,
"n_pot_x": len(q_x),
"picks": picks,
}
def val_metrics():
net.eval()
with torch.no_grad():
bce, qmse = _losses(torch.as_tensor(val_ex, device=device))
net.train()
return bce.item(), qmse.item()
# -------------------------------------------------------- training loop
pi = float(POT[val_ex].mean())
bce_base = -(pi * math.log(pi) + (1.0 - pi) * math.log(1.0 - pi))
print(f" trunk 2x{args.hidden}, q_coef {args.q_coef}, lr {args.lr}, "
f"batch {args.batch}, select {args.select}\n"
f" mốc thoát-baseline: val BCE < {bce_base:.4f} "
f"(entropy positives {pi:.1%}) — dính mốc này = pot head "
f"không điều kiện hoá được theo obs (như v3)\n")
hist = {"train": [], "vbce": [], "vqmse": [], "vpot": [], "vq": []}
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]]
bce, qmse = _losses(b)
loss = bce + args.q_coef * qmse
opt.zero_grad()
loss.backward()
opt.step()
tl_sum += loss.item()
nb += 1
v_bce, v_qmse = val_metrics()
v_total = v_bce + args.q_coef * v_qmse
net.eval()
sel = offline_select(val_boards)
net.train()
hist["train"].append(tl_sum / nb)
hist["vbce"].append(v_bce)
hist["vqmse"].append(v_qmse)
hist["vpot"].append(sel["pot"])
hist["vq"].append(sel["q_xb2"])
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 BCE "
f"{v_bce:.4f} qMSE {v_qmse:.4f} | virtual pot "
f"{sel['pot']:.1%} Q|pot(xb2) {sel['q_xb2']:.3f}")
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()
torch.save(net.state_dict(), model_dir / "qfield.pt")
print(f"\nTrain xong {(time.time()-t0)/60:.1f} phút — best epoch "
f"{best_epoch}. Model -> {model_dir / 'qfield.pt'}")
# ------------------------------------------- tầng 1: offline eval (val)
sel = offline_select(val_boards)
ceiling = float(data["q_xb2"][val_boards].mean())
se = 0.25 / math.sqrt(max(1, sel["n_pot_x"]))
print(f"\n== Tầng 1 — offline selection eval ({n_val} bàn val, "
f"select={args.select}; kết cục tra từ grid, KHÔNG xấp xỉ) ==")
print(f" virtual pot% : {sel['pot']:.1%} "
f"(random cell ~{POT[val_ex].mean():.1%}; kill-switch: < 25%)")
print(f" virtual Q|pot (xb2) : {sel['q_xb2']:.3f} ± {2*se:.3f} (2·SE, "
f"n={sel['n_pot_x']})")
print(f" virtual Q|pot (env) : {sel['q_all']:.3f} | b2 share "
f"{sel['b2_share']:.1%}")
print(f" Mốc: label ceiling {ceiling:.3f} | random-pot 0.524 | "
f"blind 0.531 | gate 0.65")
killed = sel["q_xb2"] < 0.60 or sel["pot"] < 0.25
print(f" KILL-SWITCH (§7): {'FAIL — cân nhắc dừng nhánh BC' if killed else 'qua'}")
# --------------------------------- tầng 2: eval trực tiếp trong env
if args.eval_episodes > 0:
from train_position import evaluate, print_stats
class _QFieldPolicy:
"""Duck-type SB3: predict(obs) → action [-1,1]^4."""
def predict(self, o, deterministic=True):
o = np.asarray(o, dtype=np.float32)
cue = np.array([o[0] * w, o[1] * l])
b1 = np.array([o[2] * w, o[3] * l])
phis = [float(np.degrees(np.arctan2(d[1], d[0])) % 360.0)
for d in env_h._ghost_dirs_any(cue, b1)]
xs = np.stack([np.concatenate([
o, np.array([math.cos(math.radians(ph)),
math.sin(math.radians(ph))],
dtype=np.float32)]) for ph in phis])
with torch.no_grad():
lg, qs = net(torch.as_tensor(xs, device=device))
p_hat = torch.sigmoid(lg).cpu().numpy()
q_hat = qs.cpu().numpy()
k, c = select_action(p_hat, q_hat, args.select,
args.pot_thresh)
v0, side, vert = _cls_to_vals(c, grids)
return _norm_action(phis[k], v0, side, vert), None
print(f"\n== Tầng 2 — eval Q-field TRỰC TIẾP trong env "
f"({args.eval_episodes} cú, deterministic) — KẾT QUẢ CHÍNH ==")
stats = evaluate(_QFieldPolicy(), 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"]
ok = q - 2 * se > 0.65 and stats["pot_rate"] >= 0.10
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 ngoài 2·SE, pot >= 10%): "
f"{'PASS' if ok else 'chưa pass'}")
# ----------------------------------------- tầng 3: distill npz cho SB3
if not args.no_distill:
acts, q_out, b2_out = [], [], []
all_sel = offline_select(np.arange(n))
for r in range(n):
k, c = all_sel["picks"][r]
v0, side, vert = _cls_to_vals(c, grids)
acts.append(_norm_action(board_phis[r][k], v0, side, vert))
gi = ex_start[r] + k
hit = POT[gi, c] > 0.5
q_out.append(Q[gi, c] if hit else 0.0)
b2_out.append(int(hit and not QM[gi, c]))
d_out = Path(args.dataset).with_name(
Path(args.dataset).stem + "_v4distill"
+ (f"_n{args.limit}" if args.limit else "") + ".npz")
np.savez_compressed(
d_out, obs=obs_all[:n], actions=np.stack(acts).astype(np.float32),
q=np.array(q_out, dtype=np.float32),
n_pot=(ends - starts).astype(np.int32),
b2_lucky=np.array(b2_out, dtype=np.int8))
print(f"\n== Tầng 3 — distill dataset ({n} sample, virtual pot toàn "
f"bộ {all_sel['pot']:.1%}) ==")
print(f"Dataset -> {d_out}")
rel = d_out.relative_to(ROOT) if d_out.is_relative_to(ROOT) else d_out
print(f"Bước kế (chỉ khi tầng 2 pass):")
print(f" python scripts/train_bc.py --dataset {rel} "
f"--run-name {run}_distill")
print(f" python scripts/eval_position.py "
f"models/{run}_distill/bc_model.zip --episodes 1000 "
f"--aim-mode any")
# ----------------------------------------------------------- loss plot
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
ep_x = np.arange(1, len(hist["train"]) + 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["vbce"], label="val BCE")
ax1.plot(ep_x, np.array(hist["vqmse"]) * args.q_coef, ls="--",
label=f"val qMSE×{args.q_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 v4 Q-field — {m} example, {n} bàn")
ax1.legend(loc="upper right")
ax1.grid(alpha=0.3)
ax2.plot(ep_x, hist["vpot"], label="virtual pot%")
ax2.plot(ep_x, hist["vq"], label="virtual Q|pot (xb2)")
ax2.axhline(0.531, c="gray", ls="--", alpha=0.7, label="blind 0.531")
ax2.axhline(0.65, c="tab:green", ls="--", alpha=0.7, label="gate 0.65")
ax2.set_xlabel("epoch")
ax2.set_ylim(0, 1)
ax2.set_title("tầng 1 theo epoch (val, exact trên grid)")
ax2.legend(loc="lower right")
ax2.grid(alpha=0.3)
fig.tight_layout()
fig.savefig(log_dir / "loss_curve.png", dpi=130)
print(f"\nLoss curve -> {log_dir / 'loss_curve.png'}")
if __name__ == "__main__":
main()