Spaces:
Sleeping
Sleeping
File size: 8,087 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 | #!/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()
|