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