poolcoach / scripts /cv /fetch_dataset.py
masterdanh's picture
deploy: snapshot for HF Space
78738de
Raw
History Blame Contribute Delete
6.29 kB
"""Tải dataset pix2pockets (F2 Bước 2, BRIEF 05/08/2026) — idempotent.
Nguồn CHÍNH THỨC: repo của paper pix2pockets (arXiv 2504.12045, SCIA 2025)
https://github.com/viktorseba/pix2pockets — tác giả ship sẵn bản export
Roboflow `8-Ball-Pool-3.zip` (~41 MB) ngay trong root repo, nên KHÔNG cần
Roboflow API key (project gốc: universe.roboflow.com/bachelorthesis/
8-ball-pool-l530o, **License: CC BY 4.0** — ghi trong README.dataset.txt
và data.yaml của bản export).
Zip v3 dồn CẢ 247 ảnh vào train/ (val/test trong data.yaml gốc trỏ thư mục
không tồn tại) → script dựng thêm split 80/20 TẤT ĐỊNH (sort tên + shuffle
seed 20260805) ở `yolo/`, copy ảnh chứ không symlink (Windows), bản gốc
`8-Ball-Pool-3/` giữ nguyên không sửa.
Chạy lại thoải mái: mỗi bước tự skip nếu sản phẩm đã có. Xoá thư mục
`datasets/pix2pockets/` nếu muốn ép làm mới toàn bộ.
python scripts/cv/fetch_dataset.py # tải + giải + split + tóm tắt
python scripts/cv/fetch_dataset.py --summary # chỉ in tóm tắt
"""
from __future__ import annotations
import argparse
import random
import shutil
import sys
import urllib.request
import zipfile
from collections import Counter
from pathlib import Path
ROOT = Path(__file__).resolve().parents[2] # poolcoach-rl/
ZIP_URL = "https://github.com/viktorseba/pix2pockets/raw/main/8-Ball-Pool-3.zip"
DATASET_DIR = ROOT / "datasets" / "pix2pockets"
ZIP_PATH = DATASET_DIR / "8-Ball-Pool-3.zip"
RAW_DIR = DATASET_DIR / "8-Ball-Pool-3" # bản gốc từ zip — KHÔNG sửa
YOLO_DIR = DATASET_DIR / "yolo" # split 80/20 dựng lại cho ultralytics
SPLIT_SEED = 20260805
VAL_FRACTION = 0.2
CLASS_NAMES = ["Black", "Cue", "Dot", "Solid", "Striped"]
IMG_EXTS = {".jpg", ".jpeg", ".png", ".bmp"}
def _download() -> None:
if ZIP_PATH.exists() and zipfile.is_zipfile(ZIP_PATH):
print(f"[skip] zip already present and valid: {ZIP_PATH}")
return
DATASET_DIR.mkdir(parents=True, exist_ok=True)
print(f"[down] {ZIP_URL}")
tmp = ZIP_PATH.with_suffix(".zip.part")
urllib.request.urlretrieve(ZIP_URL, tmp) # noqa: S310 — URL cố định https github
if not zipfile.is_zipfile(tmp):
tmp.unlink(missing_ok=True)
sys.exit("[ERROR] Downloaded file is not a zip - check URL/network and retry.")
tmp.replace(ZIP_PATH)
print(f"[ok] {ZIP_PATH} ({ZIP_PATH.stat().st_size / 1e6:.1f} MB)")
def _extract() -> None:
if (RAW_DIR / "data.yaml").exists():
print(f"[skip] already extracted: {RAW_DIR}")
return
print(f"[unzip] -> {DATASET_DIR}")
with zipfile.ZipFile(ZIP_PATH) as zf:
zf.extractall(DATASET_DIR)
def _make_split() -> None:
"""Split 80/20 tất định từ 247 ảnh train gốc — copy, không đụng bản gốc."""
if (YOLO_DIR / "data.yaml").exists():
print(f"[skip] split already built: {YOLO_DIR}")
return
src_img = RAW_DIR / "train" / "images"
src_lbl = RAW_DIR / "train" / "labels"
stems = sorted(p.stem for p in src_img.iterdir() if p.suffix.lower() in IMG_EXTS)
if not stems:
sys.exit(f"[ERROR] No images found under {src_img}")
rng = random.Random(SPLIT_SEED)
rng.shuffle(stems)
n_val = round(len(stems) * VAL_FRACTION)
splits = {"valid": set(stems[:n_val]), "train": set(stems[n_val:])}
print(f"[split] seed={SPLIT_SEED}: train={len(splits['train'])} valid={n_val}")
for split, chosen in splits.items():
(YOLO_DIR / split / "images").mkdir(parents=True, exist_ok=True)
(YOLO_DIR / split / "labels").mkdir(parents=True, exist_ok=True)
for img in src_img.iterdir():
if img.suffix.lower() not in IMG_EXTS:
continue
split = "valid" if img.stem in splits["valid"] else "train"
shutil.copy2(img, YOLO_DIR / split / "images" / img.name)
lbl = src_lbl / (img.stem + ".txt")
if lbl.exists():
shutil.copy2(lbl, YOLO_DIR / split / "labels" / lbl.name)
names_yaml = "\n".join(f" - {n}" for n in CLASS_NAMES)
(YOLO_DIR / "data.yaml").write_text(
"# Split 80/20 tat dinh tu pix2pockets 8-Ball-Pool-3 (CC BY 4.0)\n"
f"# sinh boi scripts/cv/fetch_dataset.py, seed={SPLIT_SEED}\n"
f"path: {YOLO_DIR.as_posix()}\n"
"train: train/images\n"
"val: valid/images\n"
f"nc: {len(CLASS_NAMES)}\n"
f"names:\n{names_yaml}\n",
encoding="utf-8",
)
print(f"[ok] {YOLO_DIR / 'data.yaml'}")
def _count_split(base: Path) -> None:
for split in ("train", "valid", "test"):
img_dir = base / split / "images"
lbl_dir = base / split / "labels"
if not img_dir.is_dir():
continue
imgs = [p for p in img_dir.iterdir() if p.suffix.lower() in IMG_EXTS]
n_boxes = 0
cls_counter: Counter[int] = Counter()
if lbl_dir.is_dir():
for lbl in lbl_dir.glob("*.txt"):
for line in lbl.read_text().splitlines():
parts = line.split()
if parts:
n_boxes += 1
cls_counter[int(parts[0])] += 1
per_class = {CLASS_NAMES[k]: v for k, v in sorted(cls_counter.items())}
print(f" {split:>5}: {len(imgs):4d} imgs, {n_boxes:5d} box, per-class {per_class}")
def _summary() -> None:
print("\n===== DATASET SUMMARY =====")
readme = RAW_DIR / "README.dataset.txt"
if readme.exists():
print(f"--- {readme.name} ---")
print(readme.read_text(encoding="utf-8", errors="replace").strip())
else:
print("[!] README.dataset.txt NOT found - license unconfirmed from zip.")
print("\n--- raw export (as shipped) ---")
_count_split(RAW_DIR)
print(f"--- rebuilt split (seed={SPLIT_SEED}) ---")
_count_split(YOLO_DIR)
def main() -> None:
ap = argparse.ArgumentParser(description="Fetch pix2pockets dataset (idempotent)")
ap.add_argument("--summary", action="store_true", help="only print summary, no download")
args = ap.parse_args()
if not args.summary:
_download()
_extract()
_make_split()
_summary()
if __name__ == "__main__":
main()