File size: 22,383 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
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
#!/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()