Spaces:
Sleeping
Sleeping
| # -*- 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()) | |