Spaces:
Sleeping
Sleeping
| #!/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 <file_relabel.npz> | |
| """ | |
| 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 <input>_canon[_p<margin>][_n<limit>].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() | |