File size: 12,430 Bytes
78738de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
#!/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()