File size: 11,789 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
"""Metric any-ball (F2, BRIEF 05/08/2026 bàn giao 17, bước 4).

Gộp **Black + Cue + Solid + Striped thành 1 class "ball"**, LOẠI **Dot**
(nút kim cương trên thành gỗ — không phải bi, không được tính là bi) khỏi
cả prediction lẫn ground-truth, rồi tính AP50 / AP50-95 trên split val
cho một file weights cho trước. Trả lời câu hỏi "tìm được bi bất kể loại
hay không" — tách khỏi lỗi phân loại solid/striped.

Cách tính (BRIEF cho tự quyết giữa "ultralytics val merge class" và "tự
tính AP" — chọn phương án THỨ NHẤT, ghi rõ ở đây):

- Chạy đúng pipeline ``model.val()`` của ultralytics (dataloader rect,
  NMS ``multi_label=True``, conf=0.001 · iou=0.7 · max_det=300 ·
  imgsz=640) qua một subclass ``DetectionValidator`` chỉ chen MỘT khâu:
  remap class id của prediction + GT ngay trước khâu match/AP
  (``_prepare_pred`` / ``_prepare_batch``). Matching và AP là code
  ultralytics nguyên bản, không tự chế.
- KHÔNG chạy lại NMS sau khi gộp: hai detection cùng một bi ở hai class
  bi khác nhau (NMS class-aware giữ cả hai) thành duplicate cùng class
  "ball" — cái thứ hai tính là false positive. Trung thực với detector
  đang có, không che lỗi double-detection.
- **Self-check bắt buộc trước khi tin số merged**: chạy cùng validator
  với remap ĐỒNG NHẤT (giữ nguyên 5 class), per-class AP50/AP50-95 phải
  khớp metrics.json của lần train trong ``--tolerance`` (mặc định 0.01);
  lệch hơn là DỪNG, không ghi anyball.json — số sai tệ hơn không có số.
  Device lấy theo trường ``device`` trong metrics.json (baseline đo CPU,
  full đo GPU) để so cùng numerics với reference.
- Bài học buộc phải đi đường này: bản đầu dùng ``model.predict`` + port
  matcher, self-check lệch tới 0.20 ở Cue — vì ``val`` NMS với
  ``multi_label=True`` còn ``predict`` là ``multi_label=False`` (một box
  chỉ giữ class argmax), class dễ lẫn như Cue/Dot trắng lệch nặng nhất.
  Hai "AP" đó không cùng một thước — không được trộn.

Console in ASCII-only (bẫy cp1252 Windows đã trả giá 05/08 sáng).
Kết quả ghi ``anyball.json`` vào ``--artifact-dir``.

    python scripts/cv/eval_anyball.py --weights "D:/Khoa luan/cv_baseline_20260805/best.pt" \
        --artifact-dir "D:/Khoa luan/cv_baseline_20260805"
"""

from __future__ import annotations

import argparse
import json
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"))

DATA_YAML = ROOT / "datasets" / "pix2pockets" / "yolo" / "data.yaml"
RUNS_DIR = ROOT / "runs" / "cv"  # runs/ đã gitignore

BALL_NAMES = ("Black", "Cue", "Solid", "Striped")
EXCLUDE_NAME = "Dot"


def build_remap(names: dict[int, str], merged: bool) -> dict[int, int | None]:
    """Bảng đổi class id cho validator.

    merged=False: đồng nhất (self-check). merged=True: 4 class bi -> 0,
    Dot -> None (None = LOẠI box khỏi cả pred lẫn GT).
    """
    name_to_id = {n: i for i, n in names.items()}
    missing = [n for n in (*BALL_NAMES, EXCLUDE_NAME) if n not in name_to_id]
    if missing:
        raise ValueError(f"data.yaml thieu class {missing}; co: {sorted(name_to_id)}")
    if not merged:
        return {i: i for i in names}
    remap: dict[int, int | None] = {name_to_id[n]: 0 for n in BALL_NAMES}
    remap[name_to_id[EXCLUDE_NAME]] = None
    return remap


def count_gt_labels(lbl_dir: Path) -> Counter:
    """Đếm instance từng class id từ file nhãn YOLO txt — độc lập validator,
    dùng đối chứng số Dot bị loại."""
    counts: Counter = Counter()
    for f in lbl_dir.glob("*.txt"):
        for line in f.read_text(encoding="utf-8").splitlines():
            parts = line.split()
            if parts:
                counts[int(parts[0])] += 1
    return counts


def make_validator(remap: dict[int, int | None], counters: dict):
    """DetectionValidator + remap class trước khâu match; mọi khâu khác nguyên bản."""
    import torch
    from ultralytics.models.yolo.detect import DetectionValidator

    class RemapValidator(DetectionValidator):
        def _remap_cls(self, cls: "torch.Tensor"):
            keep = torch.ones_like(cls, dtype=torch.bool)
            new = cls.clone()
            for old, tgt in remap.items():
                m = cls == old
                if tgt is None:
                    keep &= ~m
                else:
                    new[m] = tgt
            return keep, new

        def _prepare_batch(self, si, batch):
            pbatch = super()._prepare_batch(si, batch)
            keep, new = self._remap_cls(pbatch["cls"])
            counters["gt_kept"] += int(keep.sum())
            counters["gt_dropped"] += int((~keep).sum())
            pbatch["cls"] = new[keep]
            pbatch["bboxes"] = pbatch["bboxes"][keep]
            return pbatch

        def _prepare_pred(self, pred):
            predn = super()._prepare_pred(pred)
            keep, new = self._remap_cls(predn["cls"])
            counters["pred_kept"] += int(keep.sum())
            counters["pred_dropped"] += int((~keep).sum())
            out = {k: v[keep] for k, v in predn.items()}
            out["cls"] = new[keep]
            return out

    return RemapValidator


def run_val(weights: Path, merged: bool, device: str, run_name: str):
    """model.val() với validator remap; trả (DetMetrics, counters)."""
    from ultralytics import YOLO
    import yaml

    data = yaml.safe_load(DATA_YAML.read_text(encoding="utf-8"))
    names = {i: n for i, n in enumerate(data["names"])}
    counters = {"gt_kept": 0, "gt_dropped": 0, "pred_kept": 0, "pred_dropped": 0}
    model = YOLO(str(weights))
    if dict(model.names) != names:
        sys.exit(f"[ERROR] names trong checkpoint {model.names} khac data.yaml {names}")
    metrics = model.val(
        validator=make_validator(build_remap(names, merged), counters),
        data=str(DATA_YAML), device=device,
        project=str(RUNS_DIR), name=run_name, exist_ok=True,
        plots=False, verbose=False,
    )
    return metrics, counters


def main() -> None:
    ap_cli = argparse.ArgumentParser(description="Any-ball AP50/AP50-95 on val split")
    ap_cli.add_argument("--weights", type=Path, required=True)
    ap_cli.add_argument("--artifact-dir", type=Path, required=True,
                        help="noi ghi anyball.json + doc metrics.json lam reference self-check")
    ap_cli.add_argument("--metrics-json", type=Path, default=None,
                        help="reference self-check (mac dinh <artifact-dir>/metrics.json)")
    ap_cli.add_argument("--device", default=None,
                        help="mac dinh: truong 'device' cua metrics.json de cung numerics voi reference")
    ap_cli.add_argument("--tolerance", type=float, default=0.01,
                        help="nguong lech AP self-check identity vs metrics.json")
    args = ap_cli.parse_args()

    import torch
    import yaml

    metrics_json = args.metrics_json or (args.artifact_dir / "metrics.json")
    ref = json.loads(metrics_json.read_text(encoding="utf-8"))
    device = args.device if args.device is not None else str(ref.get("device", "cpu"))
    print(f"[env] torch={torch.__version__} cuda_available={torch.cuda.is_available()} "
          f"device={device} (reference: {metrics_json.name})")

    data = yaml.safe_load(DATA_YAML.read_text(encoding="utf-8"))
    names = {i: n for i, n in enumerate(data["names"])}
    name_to_id = {n: i for i, n in names.items()}
    img_dir = Path(data["path"]) / data["val"]
    lbl_dir = Path(str(img_dir).replace("images", "labels"))
    gt_counts = count_gt_labels(lbl_dir)
    n_dot_gt = gt_counts[name_to_id[EXCLUDE_NAME]]
    n_ball_gt = sum(gt_counts[name_to_id[n]] for n in BALL_NAMES)
    print(f"[data] val GT tu file nhan: ball {n_ball_gt} + Dot {n_dot_gt} "
          f"= {sum(gt_counts.values())} boxes")

    tag = args.weights.resolve().parent.name  # vd cv_baseline_20260805
    # --- Self-check: remap dong nhat phai tai hien metrics.json ---
    m_id, c_id = run_val(args.weights, merged=False, device=device,
                         run_name=f"anyball_selfcheck_{tag}")
    got = {}
    for k, ci in enumerate(m_id.box.ap_class_index.tolist()):
        _p, _r, ap50_c, ap_c = m_id.box.class_result(k)
        got[names[ci]] = {"ap50": round(float(ap50_c), 4), "ap50_95": round(float(ap_c), 4)}
    diffs = {}
    print(f"[selfcheck] identity-validator vs {metrics_json}")
    print(f"[selfcheck] {'class':>8} {'ap50 here':>10} {'ap50 ref':>9} {'diff':>7}")
    for cname, m in ref["per_class"].items():
        here = got.get(cname, {"ap50": 0.0, "ap50_95": 0.0})
        d = max(abs(here["ap50"] - m["ap50"]), abs(here["ap50_95"] - m["ap50_95"]))
        diffs[cname] = round(d, 4)
        print(f"[selfcheck] {cname:>8} {here['ap50']:>10.4f} {m['ap50']:>9.4f} {d:>7.4f}")
    max_diff = max(diffs.values())
    if max_diff > args.tolerance:
        sys.exit(f"[FAIL] self-check lech {max_diff:.4f} > tolerance {args.tolerance} "
                 f"- KHONG ghi anyball.json.")
    if c_id["gt_dropped"] != 0 or c_id["pred_dropped"] != 0:
        sys.exit(f"[FAIL] identity remap ma drop box: {c_id} - logic remap sai.")
    print(f"[selfcheck] OK, max diff {max_diff:.4f} <= {args.tolerance}")

    # --- Any-ball: 4 class bi -> "ball", Dot -> loai ---
    m_ball, c_ball = run_val(args.weights, merged=True, device=device,
                             run_name=f"anyball_{tag}")
    if c_ball["gt_dropped"] != n_dot_gt:
        sys.exit(f"[FAIL] validator loai {c_ball['gt_dropped']} GT box nhung file nhan "
                 f"dem duoc {n_dot_gt} Dot - lech, khong ghi so.")
    ap50 = round(float(m_ball.box.map50), 4)
    ap50_95 = round(float(m_ball.box.map), 4)
    print(f"[anyball] ball = {'+'.join(BALL_NAMES)}; Dot EXCLUDED "
          f"({n_dot_gt} GT instances, xac nhan validator drop du)")
    print(f"[anyball] GT ball {c_ball['gt_kept']}; pred giu {c_ball['pred_kept']}, "
          f"pred Dot bo {c_ball['pred_dropped']}")
    print(f"[anyball] AP50={ap50:.4f}  AP50-95={ap50_95:.4f}  ({args.weights})")

    out = {
        "date": "2026-08-05",
        "weights": str(args.weights),
        "data": str(DATA_YAML),
        "split": "val",
        "n_images": len(list(img_dir.glob("*.jpg"))),
        "method": ("ultralytics val pipeline (rect dataloader, NMS multi_label=True, "
                   "conf=0.001 iou=0.7 max_det=300 imgsz=640) via DetectionValidator "
                   "subclass remapping classes before matching: "
                   "Black+Cue+Solid+Striped -> ball, Dot dropped from pred+GT; "
                   "no re-NMS after merge; identity self-check vs metrics.json"),
        "anyball": {"ap50": ap50, "ap50_95": ap50_95,
                    "n_gt_ball": c_ball["gt_kept"],
                    "n_gt_dot_excluded": c_ball["gt_dropped"],
                    "n_pred_kept": c_ball["pred_kept"],
                    "n_pred_dot_dropped": c_ball["pred_dropped"]},
        "selfcheck": {"reference": str(metrics_json), "max_abs_diff": max_diff,
                      "tolerance": args.tolerance, "per_class_max_diff": diffs},
        "env": {"torch": torch.__version__, "device": device},
    }
    args.artifact_dir.mkdir(parents=True, exist_ok=True)
    (args.artifact_dir / "anyball.json").write_text(json.dumps(out, indent=2),
                                                    encoding="utf-8")
    print(f"[artifact] -> {args.artifact_dir / 'anyball.json'}")


if __name__ == "__main__":
    main()