Spaces:
Sleeping
Sleeping
| #!/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_<ts>/bc_model.zip — SB3 zip (dùng cho eval_position + --init-from) | |
| logs/bc_<ts>/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_<ts>/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() | |