Spaces:
Sleeping
Sleeping
| # -*- coding: utf-8 -*- | |
| """Harness eval trên held-out synthetic + analytic baseline (BG26 bước 3). | |
| VAI TRÒ KÉP — thiết kế cho BG27 cắm net vào (BRIEF bối cảnh 1): | |
| - ``run_eval(shots, predict_fn)`` là harness TRUNG LẬP: nhận iterator cú | |
| (loader chung ``gen_synth_shots.iter_shots``) + một predictor bất kỳ, | |
| trả bảng số. BG27 chỉ việc đưa ``predict_fn`` thứ hai (net) — CÙNG | |
| held-out, CÙNG 3 chỉ số gate: median |Δφ|, sai số tương đối V0, | |
| accuracy dấu spin theo trục trên cú identifiable. | |
| - ``baseline_predict`` là predictor giải tích: dựng ``rows``/``others`` | |
| đúng shape pipeline thật rồi gọi ``poolcoach_cv.broadcast.analyze_track`` | |
| — nó CHỈ thấy quỹ đạo bẩn, không thấy label. | |
| Predictor contract: predict_fn(shot) → dict với | |
| ``v0`` — ước lượng V0 GẬY (m/s) | None (label là tốc độ gậy!) | |
| ``phi`` — hướng đánh độ [0,360) | None | |
| ``vert`` — "follow"/"draw"/"stun" | None | |
| ``side`` — "side-L"/"side-R" | None | |
| (baseline trả kèm chẩn đoán ``v0_raw``/``n_collisions``/``confidence``.) | |
| Bàn: synth sinh trên bàn pooltool default (spec.json ``table``) — harness | |
| PATCH ``broadcast.TABLE_W_M/TABLE_L_M`` theo spec trước khi chạy (script-level | |
| adapter, KHÔNG sửa broadcast.py; app thật vẫn bàn giải 1.27×2.54). | |
| Quy đổi V0 gậy↔bi (bảng hằng, nguồn probe BG26b trong spec): | |
| v_bi = V0_gậy · 2/(1.3 + 2.5(a²+b²)). Baseline không biết (a, b) nên dùng | |
| hằng kỳ vọng K_V = 2/(1.3 + 2.5·E[a²+b²]) với E[a²+b²] = 2·(0.4²/3) — | |
| giải tích thuần, không fit label. Báo cáo CẢ HAI: ``raw`` (so thẳng chữ | |
| BRIEF "V0 trực tiếp") và ``phys`` (chia K_V — bar công bằng cho BG27). | |
| GT spin từ label (dấu trục, ngưỡng "nhỏ" tự quyết vào bảng hằng): | |
| dọc: b > +B_STUN_MAX → follow · b < −B_STUN_MAX → draw · còn lại stun | |
| ngang: a > +A_SIDE_MIN → side-L · a < −A_SIDE_MIN → side-R · còn lại | |
| neutral (không chấm accuracy, chỉ đếm false-side riêng). Chiều a↔side | |
| nghiệm thu bằng probe settle-chord han2005 (spec ``side_sign``): a>0 | |
| (đánh mép TRÁI, ωz<0) → side-L. | |
| null tính RIÊNG mọi chỗ, không đổ vào sai (BRIEF bước 3.2). | |
| Chạy (venv app): | |
| python scripts/broadcast/eval_baseline_synth.py ^ | |
| --data "D:/Khoa luan/datasets/bb9_synth" --out-dir "D:/Khoa luan" | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import math | |
| import subprocess | |
| import sys | |
| import time | |
| from pathlib import Path | |
| import numpy as np | |
| ROOT = Path(__file__).resolve().parents[2] | |
| sys.path.insert(0, str(ROOT / "src")) | |
| sys.path.insert(0, str(ROOT / "scripts" / "broadcast")) | |
| from gen_synth_shots import AB_MAX, iter_shots # noqa: E402 | |
| from poolcoach_cv import broadcast as bc # noqa: E402 | |
| # ------------------------------------------------------------ bảng hằng | |
| B_STUN_MAX = 0.10 # |b| ≤ mức này = stun GT — ~25% dải [−0.4,0.4]; dưới đó | |
| # đường cong follow/draw per physics chỉ vài độ (tự quyết | |
| # BRIEF "Được tự quyết", một chỗ đổi) | |
| A_SIDE_MIN = 0.10 # |a| ≤ mức này = neutral GT (side không chấm) — probe: | |
| # |a|=0.3 mới lệch gương ~21° (sát ngưỡng đọc 20°) | |
| K_V = 2.0 / (1.3 + 2.5 * (2 * AB_MAX ** 2 / 3.0)) | |
| # = 1.2766 — quy đổi bi→gậy kỳ vọng (docstring module) | |
| V0_SLICE_MPS = 1.5 # cắt lớp vùng chậm/mù đã biết (BRIEF bước 3.3) | |
| PREFIX = "bb9_synth_baseline_" | |
| PREFIX_NET = "bb9_synth_shotnet_" | |
| # Lát target-matched (BRIEF 28 bước 1): chạm đầu của cue (min t_first_bb/ | |
| # t_first_cush, GT sim trong npz) ≥ 0.3s sau strike — quần thể khớp clip | |
| # thật P0 (chạm ≥0.3s, vùng thiết kế broadcast.py; slice BG27 đã lộ 74.8% | |
| # held-out chạm <0.2s là thủ phạm chính). Cú KHÔNG chạm (cả hai NaN) không | |
| # thuộc lát — nó không có "chạm đầu". | |
| T_FIRST_TARGET_S = 0.3 | |
| PREFIX_TARGET = "bb9_synth_target_" # prefix MỚI — không đè CSV 26–27 | |
| PREFIX_P2PRIME = "bb9_synth_p2prime_" | |
| # BG29: bộ số của vòng P2′ (net c4 trên CÙNG lát ≥0.3s). | |
| # Mỗi vòng một prefix, không bao giờ đè — bảng so c3↔c4 | |
| # dựng được là nhờ hai bộ CSV còn nguyên cạnh nhau. | |
| # Truyền qua ``--prefix`` (logic chấm KHÔNG đổi một dòng). | |
| # Gate G-27.3 (BRIEF 27 — chốt trước, không đổi): bar tuyệt đối + vế thắng | |
| # baseline; tập con side/vert lấy đúng từ per-shot CSV của baseline. | |
| GATE_BARS = {"dphi_med": 2.0, "v0_relerr_med_raw": 0.10, | |
| "side_acc": 0.80, "vert_acc": 0.80} | |
| N_SIDE_BASELINE = 630 # số cú baseline đọc được side (HANDOFF 26b) | |
| N_VERT_BASELINE = 186 # số cú baseline đọc được vert | |
| # ------------------------------------------------------------- predictor | |
| def shot_to_rows_others(shot): | |
| """Dựng đúng input analyze_track từ một cú synthetic: track cue (mọi | |
| frame, cờ covered) + detection KHÔNG-cue rời rạc (bi 1..9, frame nào | |
| thấy thì có mặt) — cùng shape dữ liệu pipeline thật.""" | |
| t, xy, cov, diffs = (shot["t"], shot["xy"], shot["covered"], | |
| shot["img_diff"]) | |
| rows = [] | |
| for i in range(len(t)): | |
| on = bool(cov[i, 0]) | |
| rows.append({"frame_file": f"{i:05d}", "t_s": float(t[i]), | |
| "covered": int(on), | |
| "table_x_m": float(xy[i, 0, 0]) if on else "", | |
| "table_y_m": float(xy[i, 0, 1]) if on else "", | |
| "img_diff": float(diffs[i])}) | |
| others = [{"t_s": float(t[i]), "x_m": float(xy[i, b, 0]), | |
| "y_m": float(xy[i, b, 1])} | |
| for b in range(1, xy.shape[1]) | |
| for i in np.flatnonzero(cov[:, b])] | |
| return rows, others | |
| def _split_spin(spin_class): | |
| vert = side = None | |
| for part in (spin_class or "").split("+"): | |
| if part in ("follow", "draw", "stun"): | |
| vert = part | |
| elif part.startswith("side"): | |
| side = part | |
| return vert, side | |
| def baseline_predict(shot): | |
| """Analytic baseline = broadcast.analyze_track trên quỹ đạo bẩn.""" | |
| rows, others = shot_to_rows_others(shot) | |
| try: | |
| out = bc.analyze_track(rows, others=others) | |
| except ValueError: | |
| return {"v0": None, "phi": None, "vert": None, "side": None, | |
| "v0_raw": None, "n_collisions": 0, "confidence": None, | |
| "error": "track_rong"} | |
| v0_raw = out["metrics"]["v0_mps"] | |
| vert, side = _split_spin(out["spin_class"]) | |
| return {"v0": (v0_raw / K_V) if v0_raw is not None else None, | |
| "phi": out["metrics"]["phi_deg"], | |
| "vert": vert, "side": side, | |
| "v0_raw": v0_raw, | |
| "n_collisions": out["metrics"]["n_collisions"], | |
| "confidence": out["spin_confidence"]} | |
| def make_shotnet_predict(ckpt_path: Path, device: str = "cuda"): | |
| """Predictor thứ hai (BG27): ShotNet đã train — cùng contract với | |
| ``baseline_predict``, cùng harness, cùng thước. | |
| Map (a, b) dự đoán → lớp spin theo thước đã chốt (BRIEF 27 bối cảnh 1, | |
| net KHÔNG abstain): | |
| - ``side``: dấu của a thuần (side-L nếu a>0, side-R nếu a<0 — không | |
| dead-zone; cú GT neutral không được chấm side_acc trong harness nên | |
| vùng chết chỉ làm thước "sign acc" mất nghĩa); | |
| - ``vert``: 3 lớp follow/stun/draw quantize theo B_STUN_MAX — cùng | |
| ngưỡng GT (harness chấm exact-match 3 lớp, baseline cũng vậy). | |
| ``v0``/``v0_raw`` trả CÙNG một số: net dự đoán thẳng V0 GẬY (thước | |
| ``raw`` — quyết định Cowork từ HANDOFF 26b câu 1), không có quy đổi | |
| K_V nào ở đây. | |
| """ | |
| import torch # lazy — chỉ nhánh shotnet cần (venv CV) | |
| from poolcoach_cv.shotnet import (ShotNet, ShotNetConfig, featurize_shot, | |
| spin_classes) | |
| ck = torch.load(ckpt_path, map_location=device) | |
| net_cfg = ShotNetConfig(**ck["model_config"]) | |
| model = ShotNet(net_cfg).to(device) | |
| model.load_state_dict(ck["state_dict"]) | |
| model.eval() | |
| def predict(shot): | |
| feats, t = featurize_shot(shot, deltas=net_cfg.use_deltas) | |
| x = torch.from_numpy(feats)[None].to(device) | |
| tt = torch.from_numpy(t)[None].to(device) | |
| mask = torch.ones(1, len(t), dtype=torch.bool, device=device) | |
| p = model.predict(x, tt, mask) | |
| a, b = float(p["a"][0]), float(p["b"][0]) | |
| vert, _side = spin_classes(np.array([a]), np.array([b])) | |
| v0 = float(p["v0"][0]) | |
| return {"v0": v0, "phi": float(p["phi_deg"][0]), | |
| "vert": str(vert[0]), | |
| "side": "side-L" if a > 0 else "side-R", | |
| "v0_raw": v0, "n_collisions": 0, | |
| "confidence": round(float(p["p_ident"][0]), 3)} | |
| return predict, ck | |
| # ------------------------------------------------------- lát target-matched | |
| def t_first_contact_s(shot) -> float: | |
| """Chạm đầu của cue kể từ strike: min(t_first_bb, t_first_cush) — NaN | |
| nếu cú không có va chạm nào (non-identifiable).""" | |
| vals = [float(v) for v in (shot["t_first_bb"], shot["t_first_cush"]) | |
| if not math.isnan(float(v))] | |
| return min(vals) if vals else math.nan | |
| def is_target_matched(shot) -> bool: | |
| """Cú thuộc lát target-matched (BRIEF 28): CÓ chạm đầu và chạm đầu | |
| ≥ T_FIRST_TARGET_S sau strike.""" | |
| tf = t_first_contact_s(shot) | |
| return (not math.isnan(tf)) and tf >= T_FIRST_TARGET_S | |
| # ------------------------------------------------------------ GT + metric | |
| def gt_axes(shot): | |
| b, a = shot["label_b"], shot["label_a"] | |
| vert = ("follow" if b > B_STUN_MAX else | |
| "draw" if b < -B_STUN_MAX else "stun") | |
| side = ("side-L" if a > A_SIDE_MIN else | |
| "side-R" if a < -A_SIDE_MIN else None) # None = neutral | |
| return vert, side | |
| def circ_diff_deg(x, y): | |
| return abs(((x - y + 180.0) % 360.0) - 180.0) | |
| def run_eval(shots, predict_fn): | |
| """Chạy predictor trên iterator cú → list record per-shot (số thô, | |
| aggregate tách riêng ở ``summarize`` — BG27 dùng lại cả hai).""" | |
| recs = [] | |
| for shot in shots: | |
| p = predict_fn(shot) | |
| gt_vert, gt_side = gt_axes(shot) | |
| r = {"shot_idx": shot["shot_idx"], | |
| "label_v0": shot["label_v0"], "label_phi": shot["label_phi"], | |
| "label_a": shot["label_a"], "label_b": shot["label_b"], | |
| "v0_ball": shot["v0_ball"], "phi_ball": shot["phi_ball"], | |
| "identifiable": int(shot["identifiable"]), | |
| "fps": shot["fps"], "upconvert": int(shot["upconvert"]), | |
| "scratch": int(shot["scratch"]), | |
| "gt_vert": gt_vert, "gt_side": gt_side or "", | |
| "pred_v0": p["v0"], "pred_v0_raw": p.get("v0_raw"), | |
| "pred_phi": p["phi"], "pred_vert": p["vert"] or "", | |
| "pred_side": p["side"] or "", | |
| "pred_conf": p.get("confidence") or "", | |
| "n_collisions": p.get("n_collisions", 0)} | |
| r["v0_relerr"] = (abs(p["v0"] - r["label_v0"]) / r["label_v0"] | |
| if p["v0"] is not None else None) | |
| r["v0_relerr_raw"] = (abs(p["v0_raw"] - r["label_v0"]) | |
| / r["label_v0"] | |
| if p.get("v0_raw") is not None else None) | |
| r["dphi"] = (circ_diff_deg(p["phi"], r["label_phi"]) | |
| if p["phi"] is not None else None) | |
| recs.append(r) | |
| return recs | |
| def _q(vals, q): | |
| return float(np.percentile(np.asarray(vals, dtype=float), q)) \ | |
| if vals else float("nan") | |
| def summarize(recs): | |
| """Một hàng số cho một tập record (slice).""" | |
| n = len(recs) | |
| v0 = [r["v0_relerr"] for r in recs if r["v0_relerr"] is not None] | |
| v0r = [r["v0_relerr_raw"] for r in recs | |
| if r["v0_relerr_raw"] is not None] | |
| ph = [r["dphi"] for r in recs if r["dphi"] is not None] | |
| ident = [r for r in recs if r["identifiable"]] | |
| vert_read = [r for r in ident if r["pred_vert"]] | |
| side_gt = [r for r in ident if r["gt_side"]] | |
| side_read = [r for r in side_gt if r["pred_side"]] | |
| side_neutral = [r for r in ident if not r["gt_side"]] | |
| return { | |
| "n": n, "n_identifiable": len(ident), | |
| "v0_null_rate": round(1 - len(v0) / n, 4) if n else None, | |
| "v0_relerr_med_phys": round(_q(v0, 50), 4), | |
| "v0_relerr_p90_phys": round(_q(v0, 90), 4), | |
| "v0_relerr_med_raw": round(_q(v0r, 50), 4), | |
| "v0_relerr_p90_raw": round(_q(v0r, 90), 4), | |
| "phi_null_rate": round(1 - len(ph) / n, 4) if n else None, | |
| "dphi_med_deg": round(_q(ph, 50), 2), | |
| "dphi_p90_deg": round(_q(ph, 90), 2), | |
| "vert_null_rate": (round(1 - len(vert_read) / len(ident), 4) | |
| if ident else None), | |
| "vert_acc": (round(np.mean([r["pred_vert"] == r["gt_vert"] | |
| for r in vert_read]), 4) | |
| if vert_read else None), | |
| "n_vert_scored": len(vert_read), | |
| "side_null_rate": (round(1 - len(side_read) / len(side_gt), 4) | |
| if side_gt else None), | |
| "side_acc": (round(np.mean([r["pred_side"] == r["gt_side"] | |
| for r in side_read]), 4) | |
| if side_read else None), | |
| "n_side_scored": len(side_read), | |
| "side_false_on_neutral": (round(np.mean( | |
| [bool(r["pred_side"]) for r in side_neutral]), 4) | |
| if side_neutral else None), | |
| } | |
| def slice_recs(recs): | |
| """Các lát cắt BRIEF bước 3.3: tổng, V0 <1.5/≥1.5, fps (30 tách dup).""" | |
| out = [("all", recs), | |
| ("v0_lt_1.5", [r for r in recs if r["label_v0"] < V0_SLICE_MPS]), | |
| ("v0_ge_1.5", [r for r in recs | |
| if r["label_v0"] >= V0_SLICE_MPS])] | |
| for fps in (25, 30, 50, 60): | |
| sub = [r for r in recs if r["fps"] == fps] | |
| if fps == 30: | |
| out.append(("fps30_dup", [r for r in sub if r["upconvert"]])) | |
| out.append(("fps30_sach", [r for r in sub | |
| if not r["upconvert"]])) | |
| else: | |
| out.append((f"fps{fps}", sub)) | |
| return out | |
| def write_csvs(recs, out_dir: Path, data_dir: Path, elapsed_s: float, | |
| prefix: str = PREFIX, | |
| predictor_name: str = "analytic_baseline_broadcast", | |
| run_extra: list | None = None): | |
| rows = [] | |
| for name, sub in slice_recs(recs): | |
| row = {"slice": name} | |
| row.update(summarize(sub)) | |
| rows.append(row) | |
| with open(out_dir / f"{prefix}summary.csv", "w", newline="", | |
| encoding="utf-8") as f: | |
| w = csv.DictWriter(f, fieldnames=list(rows[0].keys())) | |
| w.writeheader() | |
| w.writerows(rows) | |
| per_shot_fields = list(recs[0].keys()) | |
| with open(out_dir / f"{prefix}per_shot.csv", "w", newline="", | |
| encoding="utf-8") as f: | |
| w = csv.DictWriter(f, fieldnames=per_shot_fields) | |
| w.writeheader() | |
| w.writerows(recs) | |
| try: | |
| commit = subprocess.run(["git", "-C", str(ROOT), "rev-parse", "HEAD"], | |
| capture_output=True, text=True, | |
| check=True).stdout.strip() | |
| except Exception: | |
| commit = "unknown" | |
| spec = json.loads((data_dir / "spec.json").read_text(encoding="utf-8")) | |
| with open(out_dir / f"{prefix}run.csv", "w", newline="", | |
| encoding="utf-8") as f: | |
| w = csv.writer(f) | |
| w.writerow(["key", "value"]) | |
| for k, v in ([("commit_repo_luc_chay", commit), | |
| ("data_commit", spec.get("commit")), | |
| ("n_heldout", len(recs)), ("K_V", round(K_V, 4)), | |
| ("B_STUN_MAX", B_STUN_MAX), ("A_SIDE_MIN", A_SIDE_MIN), | |
| ("table_w_m", bc.TABLE_W_M), | |
| ("table_l_m", bc.TABLE_L_M), | |
| ("elapsed_s", round(elapsed_s, 1)), | |
| ("predictor", predictor_name)] | |
| + list(run_extra or [])): | |
| w.writerow([k, v]) | |
| return rows | |
| # ------------------------------------------------ bảng gate net↔baseline | |
| def _acc(rows, pred_key, gt_key): | |
| return (round(float(np.mean([r[pred_key] == r[gt_key] for r in rows])), | |
| 4) if rows else float("nan")) | |
| def write_gate_csv(recs, baseline_per_shot: Path, out_path: Path, | |
| n_side: int = N_SIDE_BASELINE, | |
| n_vert: int = N_VERT_BASELINE): | |
| """Bảng G-27.3: 4 chỉ số × (bar tuyệt đối trên TOÀN identifiable/5k + | |
| vế thắng baseline trên đúng tập con baseline đọc được). Tập con lấy | |
| theo ``shot_idx`` từ per-shot CSV của baseline — số baseline TÍNH LẠI | |
| từ chính CSV đó (không chạy lại predictor), phải khớp số mốc HANDOFF | |
| 26b (``n_side``/``n_vert``); lệch là DỪNG.""" | |
| base = list(csv.DictReader(open(baseline_per_shot, encoding="utf-8"))) | |
| side_rows = [r for r in base if r["identifiable"] == "1" | |
| and r["gt_side"] and r["pred_side"]] | |
| vert_rows = [r for r in base if r["identifiable"] == "1" | |
| and r["pred_vert"]] | |
| if len(side_rows) != n_side or len(vert_rows) != n_vert: | |
| sys.exit(f"tap con baseline lech so moc: side {len(side_rows)} != " | |
| f"{n_side} / vert {len(vert_rows)} != {n_vert} — DUNG, " | |
| f"kiem {baseline_per_shot}") | |
| side_ids = {int(r["shot_idx"]) for r in side_rows} | |
| vert_ids = {int(r["shot_idx"]) for r in vert_rows} | |
| base_dphi = [float(r["dphi"]) for r in base if r["dphi"] != ""] | |
| base_v0r = [float(r["v0_relerr_raw"]) for r in base | |
| if r["v0_relerr_raw"] != ""] | |
| ident = [r for r in recs if r["identifiable"]] | |
| net_side_all = _acc([r for r in ident if r["gt_side"] and r["pred_side"]], | |
| "pred_side", "gt_side") | |
| net_vert_all = _acc([r for r in ident if r["pred_vert"]], | |
| "pred_vert", "gt_vert") | |
| net_dphi = [r["dphi"] for r in recs if r["dphi"] is not None] | |
| net_v0r = [r["v0_relerr_raw"] for r in recs | |
| if r["v0_relerr_raw"] is not None] | |
| def med(v): | |
| return round(float(np.median(v)), 4) if v else float("nan") | |
| rows = [ | |
| {"metric": "dphi_med_deg", "scope_abs": "all_5k", | |
| "bar_abs": GATE_BARS["dphi_med"], "net_abs": med(net_dphi), | |
| "baseline": med(base_dphi), | |
| "net_vs_baseline_scope": "all_5k (null baseline tinh rieng)", | |
| "net_on_subset": med(net_dphi), | |
| "dat_abs": int(med(net_dphi) <= GATE_BARS["dphi_med"]), | |
| "dat_beat": int(med(net_dphi) < med(base_dphi))}, | |
| {"metric": "v0_relerr_med_raw", "scope_abs": "all_5k", | |
| "bar_abs": GATE_BARS["v0_relerr_med_raw"], "net_abs": med(net_v0r), | |
| "baseline": med(base_v0r), | |
| "net_vs_baseline_scope": "all_5k (null baseline tinh rieng)", | |
| "net_on_subset": med(net_v0r), | |
| "dat_abs": int(med(net_v0r) <= GATE_BARS["v0_relerr_med_raw"]), | |
| "dat_beat": int(med(net_v0r) < med(base_v0r))}, | |
| {"metric": "side_acc", "scope_abs": "toan identifiable co GT side", | |
| "bar_abs": GATE_BARS["side_acc"], "net_abs": net_side_all, | |
| "baseline": _acc(side_rows, "pred_side", "gt_side"), | |
| "net_vs_baseline_scope": f"{len(side_ids)} cu baseline doc", | |
| "net_on_subset": _acc([r for r in ident | |
| if r["shot_idx"] in side_ids], | |
| "pred_side", "gt_side"), | |
| "dat_abs": int(net_side_all >= GATE_BARS["side_acc"])}, | |
| {"metric": "vert_acc", "scope_abs": "toan identifiable", | |
| "bar_abs": GATE_BARS["vert_acc"], "net_abs": net_vert_all, | |
| "baseline": _acc(vert_rows, "pred_vert", "gt_vert"), | |
| "net_vs_baseline_scope": f"{len(vert_ids)} cu baseline doc", | |
| "net_on_subset": _acc([r for r in ident | |
| if r["shot_idx"] in vert_ids], | |
| "pred_vert", "gt_vert"), | |
| "dat_abs": int(net_vert_all >= GATE_BARS["vert_acc"])}, | |
| ] | |
| for r in rows[2:]: | |
| r["dat_beat"] = int(r["net_on_subset"] > r["baseline"]) | |
| # p90 tham khảo (BRIEF bước 3.2 "thêm p90") — không phải ô gate | |
| for name, net_v, base_v in [("dphi_p90_deg", net_dphi, base_dphi), | |
| ("v0_relerr_p90_raw", net_v0r, base_v0r)]: | |
| rows.append({"metric": name + "_info", "scope_abs": "all_5k", | |
| "bar_abs": "", "net_abs": round(_q(net_v, 90), 4), | |
| "baseline": round(_q(base_v, 90), 4), | |
| "net_vs_baseline_scope": "all_5k", | |
| "net_on_subset": "", "dat_abs": "", "dat_beat": ""}) | |
| fields = ["metric", "scope_abs", "bar_abs", "net_abs", "dat_abs", | |
| "baseline", "net_vs_baseline_scope", "net_on_subset", | |
| "dat_beat"] | |
| with open(out_path, "w", newline="", encoding="utf-8") as f: | |
| w = csv.DictWriter(f, fieldnames=fields) | |
| w.writeheader() | |
| w.writerows(rows) | |
| return rows | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("--data", default=r"D:\Khoa luan\datasets\bb9_synth") | |
| ap.add_argument("--out-dir", default=r"D:\Khoa luan") | |
| ap.add_argument("--limit", type=int, default=0, | |
| help="chỉ chạy N cú đầu (smoke)") | |
| ap.add_argument("--predictor", default="baseline", | |
| choices=["baseline", "shotnet"]) | |
| ap.add_argument("--ckpt", default="", | |
| help="best.pt của run shotnet (bắt buộc với --predictor " | |
| "shotnet; chọn theo VAL, không theo held-out)") | |
| ap.add_argument("--device", default="cuda") | |
| ap.add_argument("--baseline-per-shot", | |
| default=r"D:\Khoa luan\bb9_synth_baseline_per_shot.csv", | |
| help="per-shot CSV baseline đã chạy — nguồn id tập con " | |
| "630/186 cho bảng gate") | |
| ap.add_argument("--target-slice", action="store_true", | |
| help="chỉ chấm lát target-matched: chạm đầu >= " | |
| f"{T_FIRST_TARGET_S:g}s sau strike (BRIEF 28 " | |
| "bước 1; prefix CSV thành bb9_synth_target_*)") | |
| ap.add_argument("--prefix", default="", | |
| help="ghi đè tiền tố CSV (vd " | |
| f"{PREFIX_P2PRIME}shotnet_ cho vòng P2' BG29) — " | |
| "mỗi vòng một prefix để KHÔNG đè bộ số vòng trước") | |
| args = ap.parse_args() | |
| data_dir = Path(args.data) | |
| spec = json.loads((data_dir / "spec.json").read_text(encoding="utf-8")) | |
| # bàn synth = bàn pooltool default — patch hằng bàn của broadcast theo | |
| # spec (adapter mức script; xem docstring) | |
| bc.TABLE_W_M = float(spec["table"]["w_m"]) | |
| bc.TABLE_L_M = float(spec["table"]["l_m"]) | |
| shards = sorted(data_dir.glob("heldout_*.npz")) | |
| if not shards: | |
| sys.exit("khong thay shard heldout_*.npz — chay gen_synth_shots truoc") | |
| def shots(): | |
| k = 0 | |
| for p in shards: | |
| for s in iter_shots(p): | |
| if args.target_slice and not is_target_matched(s): | |
| continue | |
| if args.limit and k >= args.limit: | |
| return | |
| k += 1 | |
| yield s | |
| if args.predictor == "shotnet": | |
| if not args.ckpt: | |
| sys.exit("--predictor shotnet can --ckpt <best.pt>") | |
| predict_fn, ck = make_shotnet_predict(Path(args.ckpt), args.device) | |
| prefix = PREFIX_NET | |
| run_extra = [("ckpt", args.ckpt), | |
| ("ckpt_train_commit", ck.get("commit", "?")), | |
| ("ckpt_best_epoch", ck.get("epoch", "?")), | |
| ("device", args.device)] | |
| predictor_name = "shotnet" | |
| else: | |
| predict_fn, prefix, run_extra = baseline_predict, PREFIX, [] | |
| predictor_name = "analytic_baseline_broadcast" | |
| if args.target_slice: | |
| # prefix MỚI cho lát — không đè bộ số 5k của BG26b/27 | |
| prefix = PREFIX_TARGET + ("shotnet_" if args.predictor == "shotnet" | |
| else "baseline_") | |
| run_extra = list(run_extra) + [ | |
| ("slice", f"target-matched: t_first >= {T_FIRST_TARGET_S:g}s " | |
| "(min t_first_bb/t_first_cush, GT sim)")] | |
| if args.prefix: | |
| prefix = args.prefix | |
| t0 = time.time() | |
| recs = run_eval(shots(), predict_fn) | |
| elapsed = time.time() - t0 | |
| rows = write_csvs(recs, Path(args.out_dir), data_dir, elapsed, | |
| prefix=prefix, predictor_name=predictor_name, | |
| run_extra=run_extra) | |
| print(f"eval {len(recs)} cu / {elapsed / 60:.1f} phut -> " | |
| f"{Path(args.out_dir) / (prefix + '*.csv')}") | |
| for row in rows: | |
| print(" " + json.dumps(row)) | |
| if args.predictor == "shotnet": | |
| if args.limit or args.target_slice: | |
| print("smoke --limit / --target-slice: bo qua bang gate " | |
| "(tap con 630/186 can du 5k cu)") | |
| else: | |
| gate = write_gate_csv(recs, Path(args.baseline_per_shot), | |
| Path(args.out_dir) / f"{prefix}gate.csv") | |
| print("GATE G-27.3:") | |
| for g in gate: | |
| print(" " + json.dumps(g)) | |
| if __name__ == "__main__": | |
| main() | |