poolcoach / scripts /train_bc.py
masterdanh's picture
deploy: snapshot for HF Space
78738de
Raw
History Blame Contribute Delete
12.4 kB
#!/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()