File size: 9,335 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
"""Eval BallID trên ảnh thật (Bước 4 BRIEF 06/08) — báo cáo, KHÔNG gate cứng.

Chạy pipeline detect + BallID thật (venv `poolcoach-cv-env`) trên các ảnh
smoke bàn giao 18 (3 ảnh val có góc đo tay trong `demo_scan.DEMO_CORNERS`),
xuất:

1. **Contact-sheet crop đánh INDEX** (`ballid_sheet_<ảnh>.png`) — mỗi bi một
   ô, chỉ ghi index, KHÔNG ghi số model đoán (Danh gán GT "mù", không bị
   mồi bởi prediction);
2. **GT mẫu** (`ballid_gt.json`, số để trống) — Danh điền tay số 1–9 theo
   contact-sheet (0 = không phải bi / không nhìn ra); file ĐÃ tồn tại thì
   không bao giờ ghi đè;
3. **Prediction** (`ballid_pred.json`) — số model gán + conf + commit hash
   lúc chạy, để chấm được cả khi checkout khác.

Có GT (≥1 ô điền) → chấm luôn: accuracy theo bi / theo ảnh + confusion các
cặp lẫn. Cách đọc số chốt TRƯỚC ở design §4: ≥7/9 đúng/ảnh = dùng được sau
van an toàn; lẫn {4↔7↔3, 1↔9} là kỳ vọng; <5/9 = tầng màu thất bại. KHÔNG
chỉnh palette/ngưỡng sau khi thấy số — đó là vòng 2, Cowork quyết.

Không cần homography: BallID chỉ đọc MÀU trong crop, không cần toạ độ bàn
(chiếu bàn đã đo riêng ở demo_scan bàn giao 18). Chọn bi như cv_worker:
lọc class bi, dedupe tâm trùng, giữ 1 cue conf cao nhất (cue thừa hạ xuống
bi thường), rồi identify_balls với WB theo cue.

    D:\\Khoa luan\\poolcoach-cv-env\\Scripts\\python.exe scripts\\cv\\eval_ballid.py
Launcher: D:\\Khoa luan\\run_eval_ballid.bat (ngoài repo, như mọi launcher).

In console ASCII thuần (bẫy cp1252 — BRIEF #4); tiếng Việt chỉ nằm trong
file JSON (utf-8).
"""

from __future__ import annotations

import argparse
import json
import subprocess
import sys
from collections import Counter
from pathlib import Path

ROOT = Path(__file__).resolve().parents[2]  # poolcoach-rl/
sys.path.insert(0, str(ROOT / "src"))
sys.path.insert(0, str(Path(__file__).resolve().parent))  # demo_scan cùng chỗ

import cv2  # noqa: E402
import numpy as np  # noqa: E402

from demo_scan import DEMO_CORNERS, VAL_IMAGES  # noqa: E402 — ảnh smoke 18
from poolcoach_cv.ballid import identify_balls  # noqa: E402

DEFAULT_WEIGHTS = Path(r"D:\Khoa luan\cv_full_20260805\best.pt")  # = cv_worker
OP_CONF = 0.2993          # đồng bộ cv_worker.OP_CONF (pick_conf 05/08)
DEFAULT_OUT = ROOT.parent / "cv_full_20260805" / "ballid"
BALL_CLASSES = {"Black", "Cue", "Solid", "Striped"}   # như cv_worker, Dot loại
TILE = 96                 # cạnh ô contact-sheet (px)
DEDUP_FRAC = 0.6          # 2 tâm gần hơn 0.6×cỡ bbox = double-detect, giữ conf cao


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:  # noqa: BLE001 — thiếu git không được chặn eval
        return "unknown"


def detect_balls(model, im: np.ndarray, conf: float) -> tuple[list[dict], dict | None]:
    """YOLO → (bi thường đã dedupe, bi cue) — nhại đúng nếp chọn của cv_worker:
    lọc class bi, dedupe tâm trùng theo conf giảm dần, cue lấy ĐÚNG 1 con
    conf cao nhất, cue thừa hạ xuống bi thường."""
    res = model.predict(im, conf=conf, imgsz=640, verbose=False)[0]
    names = res.names
    cands = []
    for b in res.boxes:
        cls = names[int(b.cls)]
        if cls not in BALL_CLASSES:
            continue
        x1, y1, x2, y2 = (float(v) for v in b.xyxy[0])
        cands.append({"cls": cls, "conf": float(b.conf), "box": (x1, y1, x2, y2),
                      "cx": (x1 + x2) / 2, "cy": (y1 + y2) / 2,
                      "size": ((x2 - x1) + (y2 - y1)) / 2})
    kept: list[dict] = []
    for c in sorted(cands, key=lambda c: -c["conf"]):
        thr = DEDUP_FRAC * c["size"]
        if any(np.hypot(c["cx"] - k["cx"], c["cy"] - k["cy"]) < thr for k in kept):
            continue
        kept.append(c)
    cue = None
    others = []
    for c in kept:
        if c["cls"] == "Cue" and cue is None:
            cue = c
        else:
            others.append(c)          # cue thừa rơi vào đây như worker
    return others, cue


def make_sheet(im: np.ndarray, others: list[dict], out_png: Path) -> None:
    """Contact-sheet 1 hàng: crop bbox NGUYÊN (không co 60% — Danh cần ngữ
    cảnh nhận bi), đánh index, KHÔNG in số model đoán (GT mù)."""
    tiles = []
    for i, c in enumerate(others):
        x1, y1, x2, y2 = (int(round(v)) for v in c["box"])
        crop = im[max(0, y1):y2, max(0, x1):x2]
        if crop.size == 0:
            crop = np.zeros((8, 8, 3), np.uint8)
        tile = cv2.resize(crop, (TILE, TILE), interpolation=cv2.INTER_NEAREST)
        canvas = np.full((TILE + 26, TILE, 3), 32, np.uint8)
        canvas[26:] = tile
        cv2.putText(canvas, f"#{i}", (6, 19), 0, 0.6, (255, 255, 255), 2)
        tiles.append(canvas)
    sheet = cv2.hconcat(tiles) if tiles else np.zeros((TILE, TILE, 3), np.uint8)
    out_png.parent.mkdir(parents=True, exist_ok=True)
    cv2.imwrite(str(out_png), sheet)


def score(gt: dict, pred: dict) -> None:
    """Chấm khi GT có ô điền: accuracy theo bi / theo ảnh + confusion."""
    per_image = {}
    n_ok = n_all = 0
    confusion: Counter[tuple[str, str]] = Counter()
    for img, balls in pred["images"].items():
        gt_img = gt.get(img, {})
        ok = tot = 0
        for idx, p in balls.items():
            if idx.startswith("_"):           # "_wb" — metadata, không phải bi
                continue
            g = gt_img.get(idx)
            if not isinstance(g, int) or not 1 <= g <= 9:
                continue                      # trống / 0 = ngoài chấm
            tot += 1
            if p["number"] == g:
                ok += 1
            else:
                confusion[(str(g), str(p["number"]))] += 1
        if tot:
            per_image[img] = (ok, tot)
            n_ok += ok
            n_all += tot
    if not n_all:
        print("[score] GT trong — chua cham duoc (Danh dien ballid_gt.json roi chay lai).")
        return
    print(f"\n[score] accuracy theo bi: {n_ok}/{n_all} = {n_ok / n_all:.3f}")
    for img, (ok, tot) in per_image.items():
        print(f"  {img[:40]:42} {ok}/{tot}")
    if confusion:
        print("[score] confusion (GT -> pred, so lan):")
        for (g, p), n in confusion.most_common():
            print(f"  {g} -> {p}: {n}")


def main() -> None:
    ap = argparse.ArgumentParser(description="BallID eval on handover-18 smoke images")
    ap.add_argument("--weights", type=Path, default=DEFAULT_WEIGHTS)
    ap.add_argument("--conf", type=float, default=OP_CONF)
    ap.add_argument("--out", type=Path, default=DEFAULT_OUT)
    args = ap.parse_args()

    from ultralytics import YOLO

    model = YOLO(str(args.weights))
    head = _git_head()
    pred = {"meta": {"weights": str(args.weights), "conf": args.conf,
                     "commit": head}, "images": {}}
    gt_path = args.out / "ballid_gt.json"
    gt_template: dict = {
        "_huong_dan": "Dien SO 1-9 cho tung index theo contact-sheet cung ten; "
                      "0 = khong phai bi / khong nhin ra; de null = bo qua. "
                      "Nguoi dien: Danh (khong phai model).",
    }

    for name in sorted(DEMO_CORNERS):
        im = cv2.imread(str(VAL_IMAGES / name))
        if im is None:
            sys.exit(f"[ERROR] Cannot read smoke image {name}")
        others, cue = detect_balls(model, im, args.conf)
        assigns, wb = identify_balls(im, [c["box"] for c in others],
                                     cue["box"] if cue else None)
        # tên ảnh val dạng "103_png.rf.<hash>.jpg" — lấy mỗi số đầu cho gọn
        make_sheet(im, others, args.out / f"ballid_sheet_{name.split('_')[0]}.png")
        pred["images"][name] = {
            str(i): {"number": num, "number_conf": conf,
                     "det_cls": c["cls"], "det_conf": round(c["conf"], 3)}
            for i, (c, (num, conf)) in enumerate(zip(others, assigns))
        }
        pred["images"][name]["_wb"] = wb
        gt_template[name] = {str(i): None for i in range(len(others))}
        nums = [a[0] for a in assigns]
        print(f"{name[:40]:42} balls={len(others)} cue={'y' if cue else 'n'} "
              f"wb={'y' if wb else 'n'} numbers={nums}")

    args.out.mkdir(parents=True, exist_ok=True)
    (args.out / "ballid_pred.json").write_text(
        json.dumps(pred, ensure_ascii=False, indent=2), encoding="utf-8")
    if gt_path.exists():
        print(f"[gt] {gt_path} da ton tai — KHONG ghi de (giu cong suc dien tay).")
    else:
        gt_path.write_text(json.dumps(gt_template, ensure_ascii=False, indent=2),
                           encoding="utf-8")
        print(f"[gt] mau GT (so de trong) -> {gt_path}")

    # per-ball "number" trong pred so voi GT — chay duoc nhieu lan, cham khi co GT
    gt = json.loads(gt_path.read_text(encoding="utf-8")) if gt_path.exists() else {}
    score(gt, pred)
    print(f"\n[done] artifacts -> {args.out} (commit {head})")


if __name__ == "__main__":
    main()