#!/usr/bin/env python3 """Behavior cloning từ oracle dataset → PPO .zip (stage 2a warm start). Kiến trúc = Y HỆT PPO của train_position.py (MlpPolicy mặc định SB3: MLP 64x64, Gaussian + state-independent log_std, n_steps=128, batch=256) → save ra SB3 .zip, `train_position.py --init-from` load lại được không đổi một dòng code. Loss (design doc §4): actor : MSE(mean_action(obs), label) — đường forward vi phân: extract_features → mlp_extractor.forward_actor → action_net value : MSE(V(obs), POT + POS_COEF*Q + AIM_COEF*0.95) — value ấm để advantage đầu fine-tune đỡ loạn (bật mặc định, rẻ) 2 bẫy đã xử theo design doc: - log_std: SB3 mặc định std=1.0 ≈ random trong [-1,1] — rollout đầu của PPO sẽ xoá sạch hành vi BC. Set log_std = log(0.15) trước save. - value lạnh → advantage loạn → PPO phá policy ngay batch đầu. Chạy từ gốc repo (venv, cần torch — chạy trên máy local): python scripts/train_bc.py --dataset data/bc_dataset_2000_123.npz python scripts/train_bc.py --dataset ... --epochs 300 --lr 1e-4 # tune Output: models/bc_/bc_model.zip — SB3 zip (dùng cho eval_position + --init-from) logs/bc_/loss_curve.png — train/val loss theo epoch + eval nhanh 200 cú cuối script (tham khảo — gate G1 CHÍNH THỨC 1000 cú) Gate G1 (design doc §5): Q|pot > 0.65 (pot% ≥ 10% chấp nhận được): python scripts/eval_position.py models/bc_/bc_model.zip --episodes 1000 --aim-mode any Pass → sinh dataset 10-20k qua đêm; fail → §7 (multi-modality fail path). """ 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")) def forward_actor_value(policy, obs_t): """Mean action + value — đường forward VI PHÂN được của ActorCriticPolicy. (policy.predict là no_grad + numpy → không dùng để train được.) Xử cả 2 nhánh share_features_extractor True/False cho chắc. """ feats = policy.extract_features(obs_t) if isinstance(feats, tuple): # share_features_extractor=False pi_f, vf_f = feats else: pi_f = vf_f = feats mean_actions = policy.action_net(policy.mlp_extractor.forward_actor(pi_f)) values = policy.value_net(policy.mlp_extractor.forward_critic(vf_f)) return mean_actions, values def main(): p = argparse.ArgumentParser() p.add_argument("--dataset", required=True, help="npz từ gen_bc_dataset.py") p.add_argument("--labels", choices=["b2", "xb2"], default="b2", help="b2 = argmax Q kể cả b2-lucky (khớp reward env, " "mặc định theo design doc); xb2 = loại b2-lucky " "(điều bi 'thật' — smoke 20/07: b2-lucky ~50% label, " "double-pot knife-edge khó imitate)") p.add_argument("--epochs", type=int, default=200) p.add_argument("--lr", type=float, default=3e-4) p.add_argument("--batch", type=int, default=256) p.add_argument("--log-std", type=float, default=0.15, help="STD Gaussian sau BC — set log_std = log(giá trị này); " "0.15 theo design doc (std=1.0 mặc định sẽ xoá BC)") p.add_argument("--value-pretrain", action=argparse.BooleanOptionalAction, default=True, help="pre-train value head từ Q label") p.add_argument("--val-frac", type=float, default=0.1) p.add_argument("--patience", type=int, default=20, help="early stop: dừng nếu val actor loss không cải thiện N epoch") p.add_argument("--seed", type=int, default=0) p.add_argument("--aim-mode", choices=["best_cut", "any"], default="any", help="cho eval cuối (chỉ ảnh hưởng aim_cos/reward mean)") p.add_argument("--eval-episodes", type=int, default=200, help="0 = bỏ eval") p.add_argument("--run-name", default=None) args = p.parse_args() import numpy as np import torch from stable_baselines3 import PPO from stable_baselines3.common.vec_env import DummyVecEnv from poolcoach_rl.envs import PositionPlayEnv from poolcoach_rl.envs.position_env import AIM_COEF, POS_COEF, POT_REWARD run = args.run_name or f"bc_{args.labels}_{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) obs = data["obs"].astype(np.float32) if args.labels == "xb2": if "actions_xb2" not in data: sys.exit("npz cũ không có label xb2 — sinh lại bằng " "gen_bc_dataset.py bản 20/07 (2 bộ label)") actions = data["actions_xb2"].astype(np.float32) q = data["q_xb2"].astype(np.float32) else: actions = data["actions"].astype(np.float32) q = data["q"].astype(np.float32) n = len(obs) # V(s) ≈ reward kỳ vọng của cú best: pot + position term + aim ~0.95 v_target = (POT_REWARD + POS_COEF * q + AIM_COEF * 0.95).astype(np.float32) 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:] print(f"== BC train: {n} sample ({len(tr_idx)} train / {n_val} val), " f"label={args.labels}, Q label mean {q.mean():.3f} ==") print(f" lr {args.lr}, batch {args.batch}, max {args.epochs} epoch, " f"patience {args.patience}, value_pretrain={args.value_pretrain}\n") # --------------------------- model (kiến trúc KHỚP train_position.py) torch.manual_seed(args.seed) env = DummyVecEnv([lambda: PositionPlayEnv(aim_mode=args.aim_mode)]) model = PPO("MlpPolicy", env, n_steps=128, batch_size=256, seed=args.seed, verbose=0) policy = model.policy device = policy.device obs_t = torch.as_tensor(obs, device=device) act_t = torch.as_tensor(actions, device=device) v_t = torch.as_tensor(v_target, device=device).unsqueeze(1) tr_idx_t = torch.as_tensor(tr_idx, device=device) val_idx_t = torch.as_tensor(val_idx, device=device) opt = torch.optim.Adam(policy.parameters(), lr=args.lr) mse = torch.nn.functional.mse_loss # ------------------------------------------------------ training loop hist = {"train_actor": [], "val_actor": [], "val_value": []} best_val, best_epoch, best_state = float("inf"), 0, None t0 = time.time() policy.set_training_mode(True) for epoch in range(1, args.epochs + 1): ep_perm = torch.randperm(len(tr_idx_t), device=device) ta_sum, nb = 0.0, 0 for s in range(0, len(tr_idx_t), args.batch): b = tr_idx_t[ep_perm[s:s + args.batch]] mean_a, values = forward_actor_value(policy, obs_t[b]) loss_a = mse(mean_a, act_t[b]) loss = loss_a if args.value_pretrain: loss = loss + 0.5 * mse(values, v_t[b]) opt.zero_grad() loss.backward() opt.step() ta_sum += loss_a.item() nb += 1 with torch.no_grad(): mean_a, values = forward_actor_value(policy, obs_t[val_idx_t]) va = mse(mean_a, act_t[val_idx_t]).item() vv = mse(values, v_t[val_idx_t]).item() hist["train_actor"].append(ta_sum / nb) hist["val_actor"].append(va) hist["val_value"].append(vv) if va < best_val - 1e-5: best_val, best_epoch = va, epoch best_state = {k: v.detach().clone() for k, v in policy.state_dict().items()} if epoch == 1 or epoch % 10 == 0: msg = (f"epoch {epoch:3d}: actor train {ta_sum/nb:.5f} / val {va:.5f}") if args.value_pretrain: msg += f" | value val {vv:.4f}" print(msg) if epoch - best_epoch >= args.patience: print(f"Early stop @ epoch {epoch} " f"(best val {best_val:.5f} tại epoch {best_epoch})") break policy.load_state_dict(best_state) policy.set_training_mode(False) print(f"\nBC xong trong {(time.time()-t0)/60:.1f} phút — " f"restore best epoch {best_epoch} (val actor {best_val:.5f})") # --------------------------------------- chẩn đoán per-component (§7) with torch.no_grad(): mean_a, _ = forward_actor_value(policy, obs_t[val_idx_t]) per_comp = ((mean_a - act_t[val_idx_t]) ** 2).mean(dim=0).cpu().numpy() baseline = actions[val_idx].var(axis=0) # MSE nếu chỉ đoán mean label print("\nVal MSE per component (baseline = var label = đoán mean):") n_flat = 0 for name, m, b in zip(["a0 phi ", "a1 V0 ", "a2 side", "a3 vert"], per_comp, baseline): note = "" if b > 0 and m > 0.85 * b: note = " <-- ~= baseline: regression về MEAN (multi-modality §7)" n_flat += 1 print(f" {name}: {m:.4f} (baseline {b:.4f}){note}") if n_flat >= 2: print(" => label đa mode đang bị trung bình hoá — dùng " "relabel_bc_dataset.py (canonical + pocket-margin) rồi BC lại") # -------------------------- log_std: bẫy số 1 của design doc — PHẢI set with torch.no_grad(): policy.log_std.fill_(math.log(args.log_std)) print(f"\nlog_std := log({args.log_std}) = {math.log(args.log_std):.3f} " f"(mặc định std=1.0 sẽ xoá BC ngay rollout đầu)") model.save(model_dir / "bc_model") print(f"Model -> {model_dir / 'bc_model.zip'}") # ----------------------------------------------------------- loss plot import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt epochs_x = np.arange(1, len(hist["val_actor"]) + 1) fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(12, 4.5)) ax1.plot(epochs_x, hist["train_actor"], label="train") ax1.plot(epochs_x, hist["val_actor"], label="val") ax1.axhline(float(baseline.mean()), c="gray", ls="--", label=f"đoán-mean baseline ({baseline.mean():.3f})") ax1.axvline(best_epoch, c="tab:red", ls=":", alpha=0.7, label=f"best epoch {best_epoch}") ax1.set_yscale("log") ax1.set_xlabel("epoch") ax1.set_ylabel("actor MSE") ax1.set_title(f"BC loss — {n} sample") ax1.legend(loc="upper right") ax1.grid(alpha=0.3) ax2.plot(epochs_x, hist["val_value"], c="tab:green") ax2.set_yscale("log") ax2.set_xlabel("epoch") ax2.set_ylabel("value MSE (val)") ax2.set_title("Value head pretrain" if args.value_pretrain else "Value (KHÔNG pretrain — chỉ theo dõi)") 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'}") # ------------------------------------------------- eval nhanh + gate G1 if args.eval_episodes > 0: from train_position import evaluate, print_stats print(f"\n== Eval nhanh BC thuần (deterministic, {args.eval_episodes} cú" f" — THAM KHẢO, SE lớn) ==") stats = evaluate(model, n_episodes=args.eval_episodes, aim_mode=args.aim_mode) print_stats(stats) env.close() rel = (model_dir / "bc_model.zip").relative_to(ROOT) print(f"\nGate G1 CHÍNH THỨC (eval 1000 cú — bài học 14+16/07):") print(f" python scripts/eval_position.py {rel} --episodes 1000 --aim-mode any") print(" PASS (Q|pot > 0.65, pot >= 10%) -> sinh 10-20k qua đêm, BC lại") print(" FAIL -> design doc §7: lọc mode thắng rõ / thêm pocket-id vào obs") print(f"\nPPO fine-tune (sau G1):") print(f" python scripts/train_position.py --aim-mode any --init-from {rel} " f"--total-steps 300000") if __name__ == "__main__": main()