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()