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