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