File size: 13,906 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
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
#!/usr/bin/env python3
"""Oracle đo TRẦN controllability của position play (stage 2a).

Bối cảnh (14-16/07/2026): Q|pot của agent dính chặt mốc position-blind 0.53
qua MỌI config. Hai giả thuyết đầu đã bị bác:
    (1) tín hiệu yếu — POS_COEF x2 không đổi gì (14/07)
    (2) aim khoá lỗ  — aim_mode="any" fine-tune, eval 1000 cú: Q|pot 0.522
        [0.489, 0.556], vẫn = 0.53 (16/07)
Còn lại giả thuyết (3): TRẦN controllability — với skill-set 1 cú hiện tại
(aim ghost-ball + V0 + spin), Q tốt nhất CÓ THỂ đạt là bao nhiêu?

Cách đo: sample N bàn (cùng phân phối với PositionPlayEnv.reset). Mỗi bàn:
aim CỐ ĐỊNH theo ghost-ball của từng lỗ khả thi (logic _ghost_dirs_any),
grid search V0 x side x vert, simulate tất cả, lấy max Q trên các cú
(pot && !scratch). Phân phối best-Q per bàn = trần controllability.

Đọc kết quả:
    trần ~0.55-0.6  → agent (0.53) đã gần tối ưu — vấn đề là TASK, không
                      phải reward; cân nhắc nới task (bàn nhỏ, bi gần lỗ)
                      hoặc chấp nhận trần và ghi vào luận văn
    trần >= 0.75    → gap là THẬT, agent chưa học điều bi — quay lại nghĩ
                      cách dạy (curriculum, oracle-guided, reward khác)
Kèm trần NO-SPIN (a=b=0) để tách riêng: spin mua được bao nhiêu Q?

Chạy từ gốc repo (Numba JIT ~40s/worker lúc khởi động):
    python scripts/oracle_controllability.py --tables 100 --workers 4
    python scripts/oracle_controllability.py --tables 20 --workers 2  # smoke

Output: logs/oracle_<ts>/{summary.txt, tables.csv, pot_combos.csv, histogram.png}
"""

from __future__ import annotations

import argparse
import csv
import multiprocessing as mp
import sys
import time
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"))

# Mốc tham chiếu cho phần so sánh trong summary
BLIND_Q = 0.531   # model stage 1 position-blind, eval 1000 cú (14/07)
AGENT_Q = 0.522   # best_model aim-any fine-tune, eval 1000 cú (16/07)

# --- state per worker (khởi tạo 1 lần, tránh pickle env qua Pool) ---
_ENV = None
_GRIDS = None  # (v0_grid, side_grid, vert_grid, phi_jitter_deg)


def _init_worker(v0_grid, side_grid, vert_grid, phi_jitter_deg):
    """Tạo env helper + trả JIT Numba NGAY để ETA về sau chính xác."""
    global _ENV, _GRIDS
    from poolcoach_rl.envs import PositionPlayEnv

    _ENV = PositionPlayEnv()
    _GRIDS = (v0_grid, side_grid, vert_grid, phi_jitter_deg)
    _simulate_shot((0.3, 0.5), (0.6, 1.0), (0.6, 1.5), 90.0, 2.0, 0.0, 0.0)


def _simulate_shot(cue_xy, b1_xy, b2_xy, phi, v0, a, b):
    """Simulate 1 cú; trả (potted, scratch, b2_potted, q).

    Dựng System mới mỗi cú vì pt.simulate là inplace/destructive.
    q chỉ có nghĩa khi pot && !scratch (theo đúng gate của env).
    """
    import numpy as np
    import pooltool as pt
    import pooltool.constants as ptc

    balls = {
        "cue": pt.Ball.create("cue", xy=tuple(cue_xy)),
        "1": pt.Ball.create("1", xy=tuple(b1_xy)),
        "2": pt.Ball.create("2", xy=tuple(b2_xy)),
    }
    system = pt.System(table=_ENV.table, balls=balls,
                       cue=pt.Cue(cue_ball_id="cue"))
    system.cue.set_state(V0=v0, phi=phi, a=a, b=b)
    try:
        pt.simulate(system, inplace=True)
    except Exception:
        return False, False, False, 0.0

    def pocketed(bid):
        return system.balls[bid].state.s == ptc.pocketed

    potted, scratch, b2_potted = pocketed("1"), pocketed("cue"), pocketed("2")
    q = 0.0
    if potted and not scratch:
        if b2_potted:
            q = 1.0  # combo may mắn — cùng quy ước với env
        else:
            cue_f = np.asarray(system.balls["cue"].state.rvw[0][:2])
            b2_f = np.asarray(system.balls["2"].state.rvw[0][:2])
            q = _ENV._position_q(cue_f, b2_f)
    return potted, scratch, b2_potted, q


def _eval_table(args):
    """Grid search 1 bàn. Trả (idx, stats dict, list pot-combo rows)."""
    import numpy as np

    idx, cue_xy, b1_xy, b2_xy = args
    v0_grid, side_grid, vert_grid, jitter = _GRIDS
    cue_xy, b1_xy, b2_xy = map(np.asarray, (cue_xy, b1_xy, b2_xy))

    # phi ứng viên: ghost-ball của mọi lỗ khả thi (+ jitter tuỳ chọn)
    phis = []
    for d in _ENV._ghost_dirs_any(cue_xy, b1_xy):
        phi0 = float(np.degrees(np.arctan2(d[1], d[0])) % 360.0)
        offsets = [0.0] if jitter <= 0 else [-jitter, 0.0, +jitter]
        phis.extend((phi0 + o) % 360.0 for o in offsets)

    n_sims = n_pot = 0
    best = {"q": -1.0, "phi": np.nan, "v0": np.nan, "a": np.nan, "b": np.nan}
    best_nospin = -1.0
    best_xb2 = -1.0  # trần LOẠI combo b2 rớt lỗ (Q=1 may mắn thổi phồng trần)
    pot_rows = []
    for phi in phis:
        for v0 in v0_grid:
            for a in side_grid:
                for b in vert_grid:
                    n_sims += 1
                    potted, scratch, b2p, q = _simulate_shot(
                        cue_xy, b1_xy, b2_xy, phi, float(v0), float(a), float(b))
                    if not (potted and not scratch):
                        continue
                    n_pot += 1
                    pot_rows.append([idx, round(phi, 2), float(v0),
                                     float(a), float(b), round(q, 4), int(b2p)])
                    if q > best["q"]:
                        best = {"q": q, "phi": phi, "v0": float(v0),
                                "a": float(a), "b": float(b)}
                    if not b2p and q > best_xb2:
                        best_xb2 = q
                    if a == 0.0 and b == 0.0 and q > best_nospin:
                        best_nospin = q

    stats = {
        "idx": idx,
        "cue_x": cue_xy[0], "cue_y": cue_xy[1],
        "b1_x": b1_xy[0], "b1_y": b1_xy[1],
        "b2_x": b2_xy[0], "b2_y": b2_xy[1],
        "n_phis": len(phis), "n_sims": n_sims, "n_pot": n_pot,
        "best_q": best["q"] if n_pot else float("nan"),
        "best_q_excl_b2": best_xb2 if best_xb2 >= 0 else float("nan"),
        "best_q_nospin": best_nospin if best_nospin >= 0 else float("nan"),
        "best_phi": best["phi"], "best_v0": best["v0"],
        "best_side": best["a"], "best_vert": best["b"],
    }
    return idx, stats, pot_rows


def sample_tables(n: int, seed: int):
    """Sample vị trí 3 bi — cùng phân phối với PositionPlayEnv.reset."""
    import numpy as np

    from poolcoach_rl.envs import PositionPlayEnv
    from poolcoach_rl.envs.position_env import BALL_R

    env = PositionPlayEnv()
    rng = np.random.default_rng(seed)
    margin = 4 * BALL_R
    tables = []
    for i in range(n):
        placed = []
        while len(placed) < 3:
            xy = np.array([rng.uniform(margin, env.w - margin),
                           rng.uniform(margin, env.l - margin)])
            if all(np.linalg.norm(xy - q) > 4 * BALL_R for q in placed):
                placed.append(xy)
        tables.append((i, tuple(placed[0]), tuple(placed[1]), tuple(placed[2])))
    return tables


def main():
    p = argparse.ArgumentParser()
    p.add_argument("--tables", type=int, default=100)
    p.add_argument("--v0-steps", type=int, default=10,
                   help="số mức V0 trong [0.5, 4.0] (khớp action map của env)")
    p.add_argument("--spin-steps", type=int, default=5,
                   help="số mức side/vert trong [-0.4, 0.4]; nên LẺ để có 0")
    p.add_argument("--phi-jitter", type=float, default=0.0,
                   help="thêm ±X độ quanh ghost aim (x3 chi phí; mặc định tắt)")
    p.add_argument("--workers", type=int, default=4)
    p.add_argument("--seed", type=int, default=42)
    p.add_argument("--run-name", default=None)
    args = p.parse_args()

    import numpy as np

    run = args.run_name or f"oracle_{time.strftime('%Y%m%d_%H%M%S')}"
    out_dir = ROOT / "logs" / run
    out_dir.mkdir(parents=True, exist_ok=True)

    v0_grid = np.linspace(0.5, 4.0, args.v0_steps)
    side_grid = np.linspace(-0.4, 0.4, args.spin_steps)
    vert_grid = np.linspace(-0.4, 0.4, args.spin_steps)

    tables = sample_tables(args.tables, args.seed)
    per_pocket = args.v0_steps * args.spin_steps ** 2
    per_pocket *= 3 if args.phi_jitter > 0 else 1
    print(f"== Oracle controllability: {args.tables} bàn, "
          f"~{per_pocket} sim/lỗ khả thi (TB ~2.2 lỗ/bàn) ==")
    print(f"   grid: V0 {args.v0_steps} mức x side/vert {args.spin_steps} mức"
          f"{f' x phi ±{args.phi_jitter}°' if args.phi_jitter > 0 else ''}")
    print(f"   {args.workers} worker — JIT Numba ~40s lúc khởi động...\n")

    t0 = time.time()
    all_stats, all_combos = [], []
    with mp.Pool(args.workers, initializer=_init_worker,
                 initargs=(v0_grid, side_grid, vert_grid, args.phi_jitter)) as pool:
        for k, (idx, stats, rows) in enumerate(
                pool.imap_unordered(_eval_table, tables), 1):
            all_stats.append(stats)
            all_combos.extend(rows)
            el = time.time() - t0
            eta = el / k * (len(tables) - k)
            bq = stats["best_q"]
            print(f"  bàn {idx:3d} ({k}/{len(tables)}): "
                  f"pot {stats['n_pot']}/{stats['n_sims']}, "
                  f"best Q = {'—' if np.isnan(bq) else f'{bq:.3f}'}"
                  f"   [{el/60:.1f} phút, còn ~{eta/60:.1f}]")

    all_stats.sort(key=lambda s: s["idx"])
    total_sims = sum(s["n_sims"] for s in all_stats)
    el = time.time() - t0
    print(f"\nXong {total_sims} sim trong {el/60:.1f} phút "
          f"({total_sims/el:.0f} sim/s)\n")

    # ---------------------------------------------------------------- CSV
    with open(out_dir / "tables.csv", "w", newline="") as f:
        wr = csv.DictWriter(f, fieldnames=list(all_stats[0].keys()))
        wr.writeheader()
        wr.writerows(all_stats)
    with open(out_dir / "pot_combos.csv", "w", newline="") as f:
        wr = csv.writer(f)
        wr.writerow(["table_idx", "phi", "v0", "side", "vert", "q", "b2_potted"])
        wr.writerows(all_combos)

    # ------------------------------------------------------------- summary
    best_q = np.array([s["best_q"] for s in all_stats])
    best_x = np.array([s["best_q_excl_b2"] for s in all_stats])
    best_ns = np.array([s["best_q_nospin"] for s in all_stats])
    potable = ~np.isnan(best_q)
    bq, bns = best_q[potable], best_ns[~np.isnan(best_ns)]
    bx = best_x[~np.isnan(best_x)]

    lines = [
        f"== Oracle controllability — {args.tables} bàn, {total_sims} sim ==",
        f"grid: V0 {args.v0_steps} mức [0.5,4.0] x side/vert "
        f"{args.spin_steps} mức [-0.4,0.4]"
        + (f" x phi ±{args.phi_jitter}°" if args.phi_jitter > 0 else ""),
        "",
        f"Bàn pot được (>=1 combo pot && !scratch): "
        f"{potable.sum()}/{args.tables} ({100*potable.mean():.0f}%)",
        "",
        "TRẦN Q (best-Q per bàn, chỉ trên bàn pot được):",
        f"  mean   : {bq.mean():.3f}   (gồm cả combo b2 rớt lỗ, Q=1 may mắn)",
        f"  median : {np.median(bq):.3f}",
        f"  p25/p75: {np.percentile(bq, 25):.3f} / {np.percentile(bq, 75):.3f}",
        f"  p10/p90: {np.percentile(bq, 10):.3f} / {np.percentile(bq, 90):.3f}",
        "",
        f"TRẦN LOẠI b2-potted: mean {bx.mean():.3f}"
        f"   <-- TRẦN controllability THẬT (điều bi, không tính golf-in)"
        if len(bx) else "TRẦN LOẠI b2-potted: (không có)",
        "",
        f"TRẦN NO-SPIN (a=b=0): mean {bns.mean():.3f}"
        f"  -> spin mua thêm ~{bq.mean()-bns.mean():+.3f} Q" if len(bns) else
        "TRẦN NO-SPIN: không có combo no-spin nào pot được",
        "",
        "So sánh:",
        f"  position-blind baseline (14/07): Q|pot = {BLIND_Q:.3f}",
        f"  agent aim-any best     (16/07): Q|pot = {AGENT_Q:.3f}",
        f"  -> gap agent vs trần thật: {bx.mean()-AGENT_Q:+.3f}" if len(bx)
        else "  -> gap: n/a",
        "",
        f"% bàn có trần thật > 0.53 (mốc blind): {100*(bx > BLIND_Q).mean():.0f}%",
        f"% bàn có trần thật > 0.70            : {100*(bx > 0.70).mean():.0f}%",
        f"% bàn có trần thật > 0.80            : {100*(bx > 0.80).mean():.0f}%",
        "",
        "Đọc kết quả:",
        "  trần ~0.55-0.6 -> agent đã gần tối ưu, vấn đề là TASK",
        "  trần >= 0.75   -> gap THẬT, agent chưa học điều bi",
    ]
    summary = "\n".join(lines)
    print(summary)
    (out_dir / "summary.txt").write_text(summary, encoding="utf-8")

    # ---------------------------------------------------------------- plot
    import matplotlib

    matplotlib.use("Agg")
    import matplotlib.pyplot as plt

    fig, ax = plt.subplots(figsize=(9, 5.5))
    bins = np.linspace(0, 1, 21)
    ax.hist(bx if len(bx) else bq, bins=bins, alpha=0.65,
            label="best Q (loại b2-potted)")
    if len(bns):
        ax.hist(bns, bins=bins, alpha=0.5, label="best Q (no-spin)")
    ax.axvline(AGENT_Q, color="tab:red", ls="--",
               label=f"agent Q|pot ({AGENT_Q:.2f})")
    ref = bx.mean() if len(bx) else bq.mean()
    ax.axvline(ref, color="tab:green", ls="-",
               label=f"trần thật mean ({ref:.2f})")
    ax.set_xlabel("best-Q per bàn (grid oracle)")
    ax.set_ylabel("số bàn")
    ax.set_title(f"Trần controllability — {potable.sum()} bàn pot được / "
                 f"{args.tables}")
    ax.legend(loc="upper right")
    fig.tight_layout()
    fig.savefig(out_dir / "histogram.png", dpi=120)
    print(f"\nOutput -> {out_dir}")


if __name__ == "__main__":
    main()