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