poolcoach / scripts /broadcast /diag_height_compensation.py
masterdanh's picture
deploy: snapshot for HF Space
78738de
Raw
History Blame Contribute Delete
11.5 kB
# -*- coding: utf-8 -*-
"""BG30 Bước 3 — thí nghiệm bù độ cao tâm bi OFFLINE (điều kiện: G-30.2 KHỚP).
KHÔNG đụng cv_worker.py / broadcast.py / shotnet.py / app — script chẩn
đoán đứng riêng, chỉ IMPORT (read-only) đúng hàm pipeline để hai phương
pháp chạy trên cùng code như production.
Mô hình đã xác nhận ở Bước 2 (bb9_diag30_offset_fit.csv, commit f763db7):
track là ảnh của tâm bi (cao R) chiếu qua mặt phẳng chuẩn homography —
vị trí track X' = P + k·(P − C_xy) với P = chân tâm bi thật, C = chân
camera, k = (R − z_c)/(h − R) theo phiên click. BÙ = giải ngược:
P = (X' + k·C_xy) / (1 + k)
Camera + k lấy từ Bước 2 (hằng số bên dưới, kèm nguồn). Hai biến thể bù:
k phiên click pilot (mặt phẳng z_c≈21mm — ĐÚNG homography đã dựng track
này) và k mô hình mặt vải z_c=0 (độ nhạy).
Hai pha (hai venv, đúng bẫy "GPU/CV ở poolcoach-cv-env"):
--extract (venv CV): dựng lại track cú 11/12 từ cu11/cu12.mp4 (clip gốc
trên đĩa) + corners ĐỒNG THUẬN = median 4 góc của 10 lần click
pilot (camera tĩnh — đã kiểm 10 bộ click lệch nhau ≤ 11px,
KHÔNG cần Danh click lại). Chạy bc.analyze_clip + YOLO, bắt
rows/others như worker (adapter mức module quanh MỘT lần gọi
analyze_track — nếp cv_worker BG28). Ghi
bb9_diag30_cu1112_tracks.json.
(mặc định) (venv app, CPU): trên các cú PILOT có sự kiện băng + cú 11/12
nếu đã extract: chạy (a) analyze_track và (b) c4b inference
(ĐÚNG glue worker: cv_worker.load_shotnet + shotnet_infer,
device cpu) TRƯỚC và SAU bù → bb9_diag30_compensation.csv.
Artifact (mới, không đè): bb9_diag30_cu1112_tracks.json,
bb9_diag30_compensation.csv.
"""
from __future__ import annotations
import csv
import json
import subprocess
import sys
from pathlib import Path
import numpy as np
ROOT = Path(__file__).resolve().parents[2]
for _p in (ROOT / "src", ROOT, ROOT / "scripts"):
if str(_p) not in sys.path:
sys.path.insert(0, str(_p))
from poolcoach_cv import broadcast as bc # noqa: E402
WORKSPACE = ROOT.parent
PILOT_DIR = ROOT / "datasets" / "bb9_pilot"
OUT_TRACKS = WORKSPACE / "bb9_diag30_cu1112_tracks.json"
OUT_CSV = WORKSPACE / "bb9_diag30_compensation.csv"
CLIPS = {"cu11": WORKSPACE / "cu11.mp4", "cu12": WORKSPACE / "cu12.mp4"}
R = bc.BALL_R_M
W, L = bc.TABLE_W_M, bc.TABLE_L_M
# ---- số Bước 2 (bb9_diag30_offset_fit.csv, sinh tại commit f763db7) ----
# camera hệ CHUẨN (phía y > L), decompose median 10 homography pilot:
CAM_CANON = (0.673, 5.046) # chân camera (m); h = 2.410 m
K_PILOT = 0.00328 # k phiên click pilot (z_c ≈ +20.8mm)
K_Z0 = 0.01200 # k mặt vải z_c = 0 (R/(h−R))
# Mọi cú trong bài đều flip=1 (camera phía y < 0 hệ cú) → camera HỆ CÚ:
CAM_SHOT = (W - CAM_CANON[0], L - CAM_CANON[1]) # = (0.597, −2.506)
# cú pilot có sự kiện băng (bb9_diag30_cushion_events.csv, cùng commit)
PILOT_SHOTS = ("shot_01", "shot_02", "shot_04", "shot_06", "shot_07",
"shot_08", "shot_09")
def git_head() -> str:
try:
return subprocess.run(["git", "-C", str(ROOT), "rev-parse",
"--short", "HEAD"], capture_output=True,
text=True, check=True).stdout.strip()
except Exception:
return "unknown"
# ------------------------------------------------------------------ bù
def compensate(rows: list[dict], others: list[dict], k: float
) -> tuple[list[dict], list[dict]]:
"""Bù độ cao: X' → P = (X' + k·C)/(1+k) cho MỌI toạ độ bàn (track cue
covered + detections không-cue). Trả bản sao — không sửa input."""
cx, cy = CAM_SHOT
out_rows = []
for r in rows:
r2 = dict(r)
if str(r.get("covered", "0")) in ("1", "True", "true"):
x, y = float(r["table_x_m"]), float(r["table_y_m"])
r2["table_x_m"] = (x + k * cx) / (1 + k)
r2["table_y_m"] = (y + k * cy) / (1 + k)
out_rows.append(r2)
out_others = []
for o in others:
o2 = dict(o)
o2["x_m"] = (float(o["x_m"]) + k * cx) / (1 + k)
o2["y_m"] = (float(o["y_m"]) + k * cy) / (1 + k)
out_others.append(o2)
return out_rows, out_others
# ------------------------------------------------------- pha extract (CV)
def extract_cu1112() -> int:
"""Venv CV: YOLO + analyze_clip trên cu11/cu12.mp4 với corners đồng
thuận pilot; bắt rows/others; ghi JSON."""
if OUT_TRACKS.exists():
print(f"BO QUA extract: {OUT_TRACKS} da ton tai")
return 0
corners = np.median(np.stack(
[np.asarray(json.loads(p.read_text(encoding="utf-8"))["oriented_px"])
for p in sorted(PILOT_DIR.glob("shot_*/homography.json"))]), axis=0)
print(f"corners dong thuan (median 10 click pilot):\n{corners.round(1)}")
from ultralytics import YOLO
from cv_worker import DEFAULT_WEIGHTS, OP_CONF # noqa: E402 — hằng số
model = YOLO(str(DEFAULT_WEIGHTS))
data = {"corners_px": corners.tolist(),
"corners_source": "median oriented_px 10 shot pilot "
"(camera tinh, click lech <= 11px)",
"weights": str(DEFAULT_WEIGHTS), "conf": OP_CONF,
"commit_luc_chay": git_head(), "clips": {}}
orig = bc.analyze_track
for name, clip in CLIPS.items():
if not clip.exists():
print(f"{name}: khong thay {clip} -- bo qua")
continue
captured = {}
def _cap(rows, others=None):
captured["rows"], captured["others"] = rows, others
return orig(rows, others=others)
bc.analyze_track = _cap
try:
result = bc.analyze_clip(clip, corners, model, op_conf=OP_CONF)
finally:
bc.analyze_track = orig
data["clips"][name] = {
"rows": captured["rows"], "others": captured["others"],
"metrics": result["metrics"], "collisions": result["collisions"],
"spin_class": result["spin_class"],
"spin_confidence": result["spin_confidence"]}
print(f"{name}: {result['metrics']['n_frames']} frame, "
f"V0={result['metrics']['v0_mps']} "
f"phi={result['metrics']['phi_deg']} "
f"colls={[(c['t_s'], c.get('contact')) for c in result['collisions']]}")
OUT_TRACKS.write_text(json.dumps(data), encoding="utf-8")
print(f"da ghi {OUT_TRACKS}")
return 0
# ------------------------------------------------- pha so sánh (venv app)
def spin_axes(out: dict) -> dict:
"""Rút (vert, lat) từ spin_class dạng 'follow+side-L' | None."""
vert = lat = ""
for part in (out.get("spin_class") or "").split("+"):
if part in ("follow", "draw", "stun"):
vert = part
elif part.startswith("side"):
lat = part
return {"vert": vert, "lat": lat, "conf": out.get("spin_confidence") or ""}
def run_both(rows: list[dict], others: list[dict], sn) -> dict:
"""(a) analytic + (b) c4b trên MỘT bộ rows/others."""
from cv_worker import shotnet_infer
out = bc.analyze_track(rows, others=others)
contacts = [(c["t_s"], c.get("contact") or "?") for c in
out["collisions"]]
res = {"v0_ana": out["metrics"]["v0_mps"],
"phi_ana": out["metrics"]["phi_deg"],
"contacts": ";".join(f"{t}:{c}" for t, c in contacts),
"n_unknown": sum(1 for _t, c in contacts if c == "unknown"),
"n_cushion": sum(1 for _t, c in contacts if c == "cushion"),
**{f"spin_{k}": v for k, v in spin_axes(out).items()}}
if sn is not None:
blk = shotnet_infer(sn, rows, others)
res.update({"v0_net": blk["v0_cue_mps"], "phi_net": blk["phi_deg"],
"a_net": blk["a"], "b_net": blk["b"],
"vert_net": blk["spin_vert"],
"side_net": blk["spin_side"]})
if res["phi_ana"] is not None:
d = abs((blk["phi_deg"] - res["phi_ana"] + 180.0) % 360.0 - 180.0)
res["dphi_net_ana"] = round(d, 1)
return res
def load_pilot(shot: str):
rows = list(csv.DictReader((PILOT_DIR / shot / "track.csv").open(
encoding="utf-8")))
others = []
with (PILOT_DIR / shot / "detections.csv").open(encoding="utf-8") as f:
for d in csv.DictReader(f):
if d["cls"] != "Cue" and d.get("in_table", "1") == "1":
others.append({"t_s": float(d["t_s"]),
"x_m": float(d["table_x_m"]),
"y_m": float(d["table_y_m"])})
return rows, others
def compare() -> int:
commit = git_head()
from cv_worker import load_shotnet
sn = load_shotnet("cpu")
if sn is None:
print("CANH BAO: khong load duoc ShotNet c4b -- chi chay analytic")
cases = [(s, *load_pilot(s)) for s in PILOT_SHOTS]
if OUT_TRACKS.exists():
data = json.loads(OUT_TRACKS.read_text(encoding="utf-8"))
for name, blob in data["clips"].items():
cases.append((name, blob["rows"], blob["others"]))
print(f"co track cu11/12 dung lai tu clip goc ({OUT_TRACKS.name})")
else:
print("chua co bb9_diag30_cu1112_tracks.json (chay --extract o venv "
"CV truoc) -- chi chay pilot")
out_rows = []
for name, rows, others in cases:
base = run_both(rows, others, sn)
for tag, k in (("k_pilot", K_PILOT), ("k_z0", K_Z0)):
rows2, others2 = compensate(rows, others, k)
after = run_both(rows2, others2, sn)
rec = {"cu": name, "bien_the": tag, "k": k,
**{f"{kk}_truoc": v for kk, v in base.items()},
**{f"{kk}_sau": v for kk, v in after.items()},
"commit_luc_chay": commit}
rec["flip_unknown_sang_cushion"] = (
base["n_unknown"] - after["n_unknown"]
if base["n_unknown"] >= after["n_unknown"] else
-(after["n_unknown"] - base["n_unknown"]))
out_rows.append(rec)
print(f"{name:<8} {tag:<8} contacts: {base['contacts']} -> "
f"{after['contacts']}")
if "phi_net" in base:
print(f"{'':<8} {'':<8} phi_net {base['phi_net']} -> "
f"{after['phi_net']} | dphi(net,ana) "
f"{base.get('dphi_net_ana')} -> "
f"{after.get('dphi_net_ana')} | spin ana "
f"[{base['spin_vert']}/{base['spin_lat']}] -> "
f"[{after['spin_vert']}/{after['spin_lat']}]")
keys = sorted({k for r in out_rows for k in r}, key=lambda s: (
s != "cu", s != "bien_the", s))
if OUT_CSV.exists():
print(f"BO QUA ghi {OUT_CSV}: da ton tai")
else:
with OUT_CSV.open("w", newline="", encoding="utf-8") as f:
wr = csv.DictWriter(f, fieldnames=keys)
wr.writeheader()
wr.writerows(out_rows)
print(f"da ghi {OUT_CSV} ({len(out_rows)} dong)")
return 0
if __name__ == "__main__":
sys.exit(extract_cu1112() if "--extract" in sys.argv else compare())