#!/usr/bin/env python3 """Relabel BC dataset từ pot combos đã lưu (fail path §7 — chống multi-modality). Bối cảnh (BC v1, 20/07, 2k sample): FAIL G1 với chẩn đoán rõ — val MSE của V0/side/vert = ĐÚNG variance baseline (net chỉ đoán mean label), spin usage sập về ~0, pot 5-8%. Nguyên nhân: ~88 pot combo/bàn, label argmax nhảy mode giữa các bàn gần giống nhau (lỗ khác nhau + spin/V0 near-tie trong cùng lỗ) → MLP MSE regression trung bình hoá các mode thành action tồi. Hai phép xử (design doc §7), đều là POST-PROCESSING trên `combos` trong npz (không re-simulate — gen 45 phút chỉ chạy 1 lần): 1. CANONICAL TIE-BREAK (mặc định bật): trong lỗ thắng, xét các combo có Q >= best_lỗ - tie_margin, chọn combo ÍT SPIN nhất (|side|+|vert| min), rồi V0 nhỏ nhất, rồi Q lớn nhất. Near-tie hết nhảy mode → label thành hàm mượt của obs; giá phải trả <= tie_margin Q. 2. POCKET-MARGIN FILTER (--pocket-margin > 0 để bật): bỏ bàn mà lỗ tốt nhất không thắng rõ lỗ nhì (best - second < margin) — loại nguồn nhảy mode lớn nhất (chọn lỗ). In % bàn giữ lại để cân data/sạch. Chạy từ gốc repo: python scripts/relabel_bc_dataset.py data/bc_dataset_10000_124.npz python scripts/relabel_bc_dataset.py data/... --pocket-margin 0.15 python scripts/relabel_bc_dataset.py data/... --limit 2000 # ablation N python scripts/relabel_bc_dataset.py data/... --include-b2 # giữ b2-lucky Output: npz keys chuẩn (obs/actions/q/n_pot/b2_lucky) — train_bc.py đọc thẳng (label mặc định `actions`): python scripts/train_bc.py --dataset """ from __future__ import annotations import argparse import sys 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")) from gen_bc_dataset import _to_action # nghịch đảo action map (đã verify) def relabel_table(c, tie_margin: float, pocket_margin: float, include_b2: bool): """Chọn 1 label canonical từ pot combos (n,6) của 1 bàn. Trả (combo_row_6,) hoặc None nếu bàn bị loại (hết combo sau lọc b2, hoặc lỗ thắng không rõ theo pocket_margin). Cột combo: [phi, v0, side, vert, q, b2p]. """ import numpy as np if not include_b2: c = c[c[:, 5] < 0.5] if len(c) == 0: return None # Nhóm theo lỗ = theo phi (mỗi lỗ khả thi đúng 1 phi ghost-ball, không jitter) phis = np.unique(c[:, 0]) pocket_best = np.array([c[c[:, 0] == p, 4].max() for p in phis]) order = np.argsort(pocket_best)[::-1] if pocket_margin > 0 and len(phis) > 1: if pocket_best[order[0]] - pocket_best[order[1]] < pocket_margin: return None # lỗ thắng không rõ → nguồn nhảy mode → bỏ bàn win = c[c[:, 0] == phis[order[0]]] cand = win[win[:, 4] >= win[:, 4].max() - tie_margin] # canonical: primary |spin| min → secondary V0 min → tertiary Q max # (lexsort: key CUỐI là primary) pick = cand[np.lexsort((-cand[:, 4], cand[:, 1], np.abs(cand[:, 2]) + np.abs(cand[:, 3])))[0]] return pick def main(): p = argparse.ArgumentParser() p.add_argument("dataset", help="npz từ gen_bc_dataset.py (bản có combos)") p.add_argument("--tie-margin", type=float, default=0.05, help="near-tie trong lỗ thắng: Q >= best - margin đều là " "ứng viên canonical (0 = argmax thuần)") p.add_argument("--pocket-margin", type=float, default=0.0, help="> 0: bỏ bàn có best_lỗ1 - best_lỗ2 < margin " "(0 = tắt, giữ mọi bàn)") p.add_argument("--include-b2", action="store_true", help="giữ combo b2-lucky (mặc định LOẠI — knife-edge, " "46%% label v1 là fluke)") p.add_argument("--limit", type=int, default=None, help="chỉ lấy N bàn đầu (ablation kích thước dataset)") p.add_argument("--out", default=None, help="mặc định _canon[_p][_n].npz") args = p.parse_args() import numpy as np data = np.load(args.dataset) if "combos" not in data: sys.exit("npz không có `combos` — sinh lại bằng gen_bc_dataset.py " "bản 20/07 (lưu pot combos)") obs = data["obs"] combos = data["combos"] combo_row = data["combo_row"] n = len(obs) if args.limit is None else min(args.limit, len(obs)) # combo_row tăng dần theo bàn (gen ghi tuần tự) → cắt bằng searchsorted starts = np.searchsorted(combo_row, np.arange(n)) ends = np.searchsorted(combo_row, np.arange(n) + 1) keep, actions, qs, npots, b2s = [], [], [], [], [] drop_b2 = drop_pocket = 0 for r in range(n): c = combos[starts[r]:ends[r]] pick = relabel_table(c, args.tie_margin, args.pocket_margin, args.include_b2) if pick is None: # phân loại lý do bỏ (cho stats) c2 = c if args.include_b2 else c[c[:, 5] < 0.5] if len(c2) == 0: drop_b2 += 1 else: drop_pocket += 1 continue keep.append(r) phi, v0, side, vert, q, b2p = (float(x) for x in pick) actions.append(_to_action(phi, v0, side, vert)) qs.append(q) npots.append(ends[r] - starts[r]) b2s.append(int(b2p > 0.5)) if not keep: sys.exit("Không còn bàn nào sau khi lọc — nới --pocket-margin!") keep = np.array(keep) act_arr = np.stack(actions).astype(np.float32) q_arr = np.array(qs, dtype=np.float32) out = (Path(args.out) if args.out else Path(args.dataset).with_name( Path(args.dataset).stem + "_canon" + (f"_p{args.pocket_margin:g}" if args.pocket_margin > 0 else "") + (f"_n{args.limit}" if args.limit else "") + ".npz")) np.savez_compressed(out, obs=obs[keep].astype(np.float32), actions=act_arr, q=q_arr, n_pot=np.array(npots, dtype=np.int32), b2_lucky=np.array(b2s, dtype=np.int8), tie_margin=np.float64(args.tie_margin), pocket_margin=np.float64(args.pocket_margin)) # ------------------------------------------------------------- stats q_orig = data["q_xb2"][:n] if not args.include_b2 else data["q"][:n] v0_lbl = 0.5 + (act_arr[:, 1] + 1.0) / 2.0 * 3.5 print(f"== Relabel: {len(keep)}/{n} bàn giữ lại " f"(bỏ {drop_b2} hết-combo-sau-lọc-b2, {drop_pocket} lỗ-không-rõ) ==") print(f" tie_margin {args.tie_margin} | pocket_margin " f"{args.pocket_margin} | include_b2 {args.include_b2}") print(f" Q label : mean {q_arr.mean():.3f} " f"(argmax gốc trên bàn giữ lại: {q_orig[keep].mean():.3f} — " f"giá canonical: {q_orig[keep].mean() - q_arr.mean():+.3f})") print(f" V0 label : mean {v0_lbl.mean():.2f} m/s " f"(min-V0 tie-break → kỳ vọng THẤP hơn argmax)") print(f" |side|/|vert|: {np.abs(act_arr[:, 2]).mean():.2f} / " f"{np.abs(act_arr[:, 3]).mean():.2f} " f"(min-spin tie-break → kỳ vọng THẤP hơn 0.37/0.58 của v1)") print(f" spin != 0 : " f"{float(np.mean((act_arr[:, 2] != 0) | (act_arr[:, 3] != 0))):.1%}") print(f"\nDataset -> {out}") print(f"Bước kế: python scripts/train_bc.py --dataset " f"{out.relative_to(ROOT) if out.is_relative_to(ROOT) else out}") if __name__ == "__main__": main()