Spaces:
Sleeping
Sleeping
| """Bước 3 P0 (BRIEF 11/08/2026): detector zero-shot + track cue ball thô. | |
| Chạy trên ``poolcoach-cv-env`` (CUDA). Với mỗi ``shot_XX`` (đã có frames + | |
| homography.json): | |
| 1. Detector 5-class hiện có (``cv_full_20260805/best.pt``) chạy zero-shot | |
| từng frame, imgsz 640 (đúng cấu hình vận hành cv_worker — KHÔNG chỉnh | |
| để làm đẹp số; muốn báo thêm ở conf/imgsz khác thì ghi rõ nhãn). | |
| Detection lưu từ conf floor 0.10 (chẩn đoán), cột ``at_op`` đánh dấu | |
| det đạt conf vận hành 0.2993 — MỌI số gate chỉ dùng det đạt op. | |
| 2. Track cue ball thô: nearest-neighbor qua frame, gating vận tốc | |
| VMAX = 11 m/s (design §4.4) theo dt TIMESTAMP THẬT + slack nhiễu; | |
| mất det thì coast (không tính coverage), tự bắt lại trong bán kính | |
| gating nở dần theo thời gian mất dấu. | |
| 3. Xuất per-shot: ``detections.csv`` (mọi det), ``track.csv`` (cue ball, | |
| toạ độ bàn mm + timestamp), ``overlay.mp4`` + ``contact_sheet.jpg`` | |
| (bắt buộc — mắt người kiểm track, số coverage suông không tin được). | |
| D:\\Khoa luan\\poolcoach-cv-env\\Scripts\\python.exe \ | |
| scripts/broadcast/track_p0.py | |
| """ | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import sys | |
| from pathlib import Path | |
| ROOT = Path(__file__).resolve().parents[2] # poolcoach-rl/ | |
| sys.path.insert(0, str(ROOT / "src")) | |
| import cv2 # noqa: E402 | |
| import numpy as np # noqa: E402 | |
| from poolcoach_cv.homography import TableHomography # noqa: E402 | |
| PILOT_ROOT = ROOT / "datasets" / "bb9_pilot" | |
| DEFAULT_WEIGHTS = Path(r"D:\Khoa luan\cv_full_20260805\best.pt") # = cv_worker | |
| OP_CONF = 0.2993 # conf vận hành (pick_conf 05/08) — đồng bộ cv_worker | |
| CONF_FLOOR = 0.10 # lưu det từ đây cho chẩn đoán; KHÔNG dùng cho số gate | |
| IMGSZ = 640 # cùng imgsz train/val/worker | |
| BATCH = 16 | |
| VMAX = 11.0 # m/s — gating vận tốc (design §4.4) | |
| SLACK_M = 0.06 # nhiễu vị trí + tâm bbox lệch tâm bi (~2R) | |
| IN_TOL_M = 0.05 # tâm chiếu lệch ra ngoài mặt bàn quá mức này → loại | |
| BALL_CLASSES = {"Black", "Cue", "Solid", "Striped"} # Dot = nút thành gỗ | |
| TRAIL_N = 90 # số điểm track vẽ đuôi trên overlay | |
| def load_frames_meta(shot_dir: Path) -> list[tuple[str, float]]: | |
| rows = [] | |
| with open(shot_dir / "frames_meta.csv", encoding="utf-8") as f: | |
| for r in csv.DictReader(f): | |
| rows.append((r["frame_file"], float(r["t_video_s"]))) | |
| return rows | |
| def frame_diffs(shot_dir: Path, meta: list[tuple[str, float]], | |
| th: TableHomography) -> dict[str, float]: | |
| """Mean abs diff (gray) vùng bàn giữa frame liên tiếp — lộ frame nhân đôi. | |
| Stream YouTube 30fps upconvert từ nguồn 25fps sẽ có ~1/6 frame gần trùng | |
| (đo 11/08 trên trận Florida Open); measure_p0 dùng cột này để đề xuất | |
| ngưỡng trên chuỗi đã loại frame trùng, có dán nhãn. | |
| """ | |
| corners = th.corners_px.astype(int) | |
| x0, x1 = corners[:, 0].min(), corners[:, 0].max() | |
| y0, y1 = corners[:, 1].min(), corners[:, 1].max() | |
| diffs: dict[str, float] = {} | |
| prev = None | |
| for name, _t in meta: | |
| im = cv2.imread(str(shot_dir / "frames" / name), cv2.IMREAD_GRAYSCALE) | |
| crop = im[max(0, y0):y1, max(0, x0):x1].astype(np.int16) | |
| diffs[name] = (float(np.abs(crop - prev).mean()) | |
| if prev is not None else -1.0) | |
| prev = crop | |
| return diffs | |
| def detect_shot(model, shot_dir: Path, th: TableHomography, conf_floor: float, | |
| op_conf: float, imgsz: int, device) -> list[dict]: | |
| """YOLO trên mọi frame của cú → list det dict (đã chiếu toạ độ bàn).""" | |
| meta = load_frames_meta(shot_dir) | |
| dets: list[dict] = [] | |
| for i in range(0, len(meta), BATCH): | |
| chunk = meta[i:i + BATCH] | |
| imgs = [] | |
| for name, _t in chunk: | |
| im = cv2.imread(str(shot_dir / "frames" / name)) | |
| if im is None: | |
| sys.exit(f"[ERROR] Khong doc duoc {shot_dir.name}/frames/{name}") | |
| imgs.append(im) | |
| results = model.predict(imgs, conf=conf_floor, imgsz=imgsz, | |
| verbose=False, device=device) | |
| for (name, t), res in zip(chunk, results): | |
| names = res.names | |
| for b in res.boxes: | |
| cls = names[int(b.cls)] | |
| x1, y1, x2, y2 = (float(v) for v in b.xyxy[0]) | |
| cx, cy = (x1 + x2) / 2.0, (y1 + y2) / 2.0 | |
| tx, ty = th.px_to_table((cx, cy)) | |
| w, l = th.table_w, th.table_l | |
| inside = (-IN_TOL_M <= tx <= w + IN_TOL_M | |
| and -IN_TOL_M <= ty <= l + IN_TOL_M) | |
| dets.append({ | |
| "frame_file": name, "t_s": t, "cls": cls, | |
| "conf": float(b.conf), | |
| "px_cx": cx, "px_cy": cy, | |
| "px_w": x2 - x1, "px_h": y2 - y1, | |
| "table_x_m": float(tx), "table_y_m": float(ty), | |
| "in_table": int(inside), | |
| "at_op": int(float(b.conf) >= op_conf), | |
| }) | |
| return dets | |
| def track_cue(meta: list[tuple[str, float]], dets: list[dict], | |
| op_conf: float, diffs: dict[str, float]) -> list[dict]: | |
| """NN + gating vận tốc trên det class Cue đạt op-conf, trong mặt bàn.""" | |
| by_frame: dict[str, list[dict]] = {} | |
| for d in dets: | |
| if d["cls"] == "Cue" and d["at_op"] and d["in_table"]: | |
| by_frame.setdefault(d["frame_file"], []).append(d) | |
| rows = [] | |
| last_pos = None # (x, y) bàn, m | |
| last_t = None | |
| for name, t in meta: | |
| cands = by_frame.get(name, []) | |
| chosen = None | |
| if cands: | |
| if last_pos is None: | |
| chosen = max(cands, key=lambda d: d["conf"]) | |
| else: | |
| dt = max(t - last_t, 1e-3) | |
| gate = VMAX * dt + SLACK_M | |
| best = None | |
| for d in cands: | |
| dist = float(np.hypot(d["table_x_m"] - last_pos[0], | |
| d["table_y_m"] - last_pos[1])) | |
| if dist <= gate and (best is None or dist < best[0]): | |
| best = (dist, d) | |
| chosen = best[1] if best else None | |
| if chosen is not None: | |
| last_pos = (chosen["table_x_m"], chosen["table_y_m"]) | |
| last_t = t | |
| rows.append({ | |
| "frame_file": name, "t_s": t, | |
| "covered": int(chosen is not None), | |
| "px_x": chosen["px_cx"] if chosen else "", | |
| "px_y": chosen["px_cy"] if chosen else "", | |
| "table_x_m": chosen["table_x_m"] if chosen else "", | |
| "table_y_m": chosen["table_y_m"] if chosen else "", | |
| "conf": chosen["conf"] if chosen else "", | |
| "n_cands": len(cands), | |
| "gap_s": 0.0 if chosen is not None or last_t is None | |
| else round(t - last_t, 4), | |
| "img_diff": round(diffs.get(name, -1.0), 4), | |
| }) | |
| return rows | |
| def render_overlay(shot_dir: Path, meta, dets, track_rows, fps: float) -> None: | |
| by_frame_dets: dict[str, list[dict]] = {} | |
| for d in dets: | |
| if d["at_op"]: | |
| by_frame_dets.setdefault(d["frame_file"], []).append(d) | |
| track_by_frame = {r["frame_file"]: r for r in track_rows} | |
| first = cv2.imread(str(shot_dir / "frames" / meta[0][0])) | |
| h, w = first.shape[:2] | |
| vw = cv2.VideoWriter(str(shot_dir / "overlay.mp4"), | |
| cv2.VideoWriter_fourcc(*"mp4v"), fps, (w, h)) | |
| trail: list[tuple[int, int]] = [] | |
| n_cov = 0 | |
| sheet_idx = set(np.linspace(0, len(meta) - 1, 12, dtype=int).tolist()) | |
| thumbs = [] | |
| for i, (name, t) in enumerate(meta): | |
| im = cv2.imread(str(shot_dir / "frames" / name)) | |
| for d in by_frame_dets.get(name, []): | |
| x1 = int(d["px_cx"] - d["px_w"] / 2); y1 = int(d["px_cy"] - d["px_h"] / 2) | |
| x2 = int(d["px_cx"] + d["px_w"] / 2); y2 = int(d["px_cy"] + d["px_h"] / 2) | |
| col = (180, 180, 180) if d["cls"] != "Cue" else (0, 220, 220) | |
| cv2.rectangle(im, (x1, y1), (x2, y2), col, 1) | |
| r = track_by_frame[name] | |
| if r["covered"]: | |
| n_cov += 1 | |
| p = (int(r["px_x"]), int(r["px_y"])) | |
| trail.append(p) | |
| cv2.circle(im, p, 14, (0, 220, 0), 2) | |
| if len(trail) >= 2: | |
| cv2.polylines(im, [np.asarray(trail[-TRAIL_N:], np.int32)], | |
| False, (0, 220, 0), 2) | |
| status = "TRACK" if r["covered"] else "NO TRACK" | |
| col = (0, 220, 0) if r["covered"] else (0, 0, 255) | |
| cv2.rectangle(im, (0, 0), (w, 30), (0, 0, 0), -1) | |
| cv2.putText(im, f"{shot_dir.name} t={t:9.3f}s frame {i + 1}/{len(meta)} " | |
| f"cov {100.0 * n_cov / (i + 1):5.1f}% [{status}]", | |
| (8, 22), cv2.FONT_HERSHEY_SIMPLEX, 0.6, col, 2) | |
| vw.write(im) | |
| if i in sheet_idx: | |
| thumbs.append(cv2.resize(im, (480, 270))) | |
| vw.release() | |
| rows = [np.hstack(thumbs[r * 4:(r + 1) * 4]) for r in range(3) | |
| if len(thumbs[r * 4:(r + 1) * 4]) == 4] | |
| if rows: | |
| cv2.imwrite(str(shot_dir / "contact_sheet.jpg"), np.vstack(rows), | |
| [cv2.IMWRITE_JPEG_QUALITY, 90]) | |
| def write_csv(path: Path, rows: list[dict]) -> None: | |
| with open(path, "w", newline="", encoding="utf-8") as f: | |
| w = csv.DictWriter(f, fieldnames=list(rows[0].keys())) | |
| w.writeheader() | |
| w.writerows(rows) | |
| def main() -> None: | |
| ap = argparse.ArgumentParser(description="Zero-shot detect + track cue ball (P0)") | |
| ap.add_argument("--root", type=Path, default=PILOT_ROOT) | |
| ap.add_argument("--weights", type=Path, default=DEFAULT_WEIGHTS) | |
| ap.add_argument("--conf", type=float, default=OP_CONF, | |
| help="conf van hanh cho so gate (mac dinh 0.2993 — dung cham)") | |
| ap.add_argument("--conf-floor", type=float, default=CONF_FLOOR) | |
| ap.add_argument("--imgsz", type=int, default=IMGSZ, | |
| help="doi khac 640 thi so KHONG phai so gate — ghi ro nhan") | |
| ap.add_argument("--device", default=0) | |
| ap.add_argument("--only", type=int, nargs="*", default=None) | |
| args = ap.parse_args() | |
| if not args.weights.exists(): | |
| sys.exit(f"[ERROR] Khong thay weights: {args.weights} — DUNG, hoi Danh.") | |
| from ultralytics import YOLO | |
| model = YOLO(str(args.weights)) | |
| shots = sorted(d for d in args.root.glob("shot_*") if d.is_dir()) | |
| if args.only: | |
| want = {f"shot_{i:02d}" for i in args.only} | |
| shots = [d for d in shots if d.name in want] | |
| if not shots: | |
| sys.exit("[ERROR] Khong co shot_XX nao trong root.") | |
| print(f"weights={args.weights} op_conf={args.conf} imgsz={args.imgsz} " | |
| f"floor={args.conf_floor}") | |
| if args.conf != OP_CONF or args.imgsz != IMGSZ: | |
| print("[NOTE] Cau hinh KHAC van hanh — so sinh ra KHONG phai so gate G0.") | |
| summary = [] | |
| for d in shots: | |
| hpath = d / "homography.json" | |
| if not hpath.exists(): | |
| print(f"[WARN] {d.name}: chua co homography.json — bo qua.") | |
| continue | |
| hdata = json.loads(hpath.read_text(encoding="utf-8")) | |
| th = TableHomography(np.asarray(hdata["oriented_px"], dtype=np.float64), | |
| table_w=hdata["table_w_m"], | |
| table_l=hdata["table_l_m"]) | |
| meta = load_frames_meta(d) | |
| fps = json.loads((d / "meta.json").read_text(encoding="utf-8"))[ | |
| "fps_nominal"] | |
| dets = detect_shot(model, d, th, args.conf_floor, args.conf, | |
| args.imgsz, args.device) | |
| diffs = frame_diffs(d, meta, th) | |
| track_rows = track_cue(meta, dets, args.conf, diffs) | |
| write_csv(d / "detections.csv", dets) if dets else None | |
| write_csv(d / "track.csv", track_rows) | |
| render_overlay(d, meta, dets, track_rows, fps) | |
| n = len(meta) | |
| n_cue_frames = len({r["frame_file"] for r in track_rows | |
| if r["n_cands"] > 0}) | |
| n_cov = sum(r["covered"] for r in track_rows) | |
| gaps = [r["gap_s"] for r in track_rows if r["gap_s"]] | |
| s = {"shot": d.name, "frames": n, | |
| "cue_det_rate": round(n_cue_frames / n, 4), | |
| "coverage_frames": round(n_cov / n, 4), | |
| "max_gap_s": max(gaps) if gaps else 0.0} | |
| summary.append(s) | |
| print(f"{d.name}: frames={n} cue_det={s['cue_det_rate']:.1%} " | |
| f"cov={s['coverage_frames']:.1%} max_gap={s['max_gap_s']:.2f}s") | |
| if summary: | |
| write_csv(args.root / "track_summary.csv", summary) | |
| n80 = sum(1 for s in summary if s["coverage_frames"] >= 0.80) | |
| print(f"\n[G0 so bo] {n80}/{len(summary)} cu coverage >= 80% " | |
| f"(so chot o measure_p0 — coverage theo THOI LUONG, day la theo frame)") | |
| if __name__ == "__main__": | |
| main() | |