File size: 13,057 Bytes
78738de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
"""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()