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