JerMa88's picture
Upload folder using huggingface_hub
3ce19a2 verified
Raw
History Blame Contribute Delete
3.84 kB
"""Download CIFAR-10 and extract images for training + FID evaluation.
Usage:
python prepare_cifar10.py [--data_root ./cifar10_data]
Creates:
<data_root>/cifar-10-batches-py/ (pickle files used by data.py)
<data_root>/img/ (PNG images for FID computation)
"""
import argparse
import os
import pickle
import shutil
import tarfile
import urllib.request
import numpy as np
from PIL import Image
CIFAR10_URL = "https://www.cs.toronto.edu/~kriz/cifar-10-python.tar.gz"
def is_valid_tar_gz(path: str) -> bool:
if not os.path.isfile(path):
return False
try:
with tarfile.open(path, "r:gz") as tf:
return tf.getmember("cifar-10-batches-py") is not None
except Exception:
return False
def ensure_clean_tarball(path: str):
if os.path.isfile(path) and not is_valid_tar_gz(path):
print(f"Corrupted archive detected, removing: {path}")
os.remove(path)
def has_complete_batches(batch_dir: str) -> bool:
required = [
"data_batch_1",
"data_batch_2",
"data_batch_3",
"data_batch_4",
"data_batch_5",
"test_batch",
"batches.meta",
]
return all(os.path.isfile(os.path.join(batch_dir, name)) for name in required)
def download_cifar10(data_root: str):
tar_path = os.path.join(data_root, "cifar-10-python.tar.gz")
batch_dir = os.path.join(data_root, "cifar-10-batches-py")
if os.path.isdir(batch_dir) and has_complete_batches(batch_dir):
print(f"Already exists: {batch_dir}")
return
if os.path.isdir(batch_dir):
print(f"Incomplete batch directory detected, removing: {batch_dir}")
shutil.rmtree(batch_dir, ignore_errors=True)
os.makedirs(data_root, exist_ok=True)
ensure_clean_tarball(tar_path)
if not os.path.isfile(tar_path):
print(f"Downloading CIFAR-10 to {tar_path} ...")
urllib.request.urlretrieve(CIFAR10_URL, tar_path)
print("Download complete.")
ensure_clean_tarball(tar_path)
print("Extracting ...")
try:
with tarfile.open(tar_path, "r:gz") as tf:
tf.extractall(data_root)
except Exception as e:
print(f"Extraction failed ({type(e).__name__}); re-downloading archive once ...")
if os.path.isdir(batch_dir):
shutil.rmtree(batch_dir, ignore_errors=True)
if os.path.isfile(tar_path):
os.remove(tar_path)
urllib.request.urlretrieve(CIFAR10_URL, tar_path)
with tarfile.open(tar_path, "r:gz") as tf:
tf.extractall(data_root)
print(f"Extracted to {batch_dir}")
def extract_images(data_root: str):
"""Extract all training images as PNGs into <data_root>/img/ for FID."""
img_dir = os.path.join(data_root, "img")
if os.path.isdir(img_dir) and len(os.listdir(img_dir)) >= 45000:
print(f"Image folder already populated: {img_dir}")
return
os.makedirs(img_dir, exist_ok=True)
batch_dir = os.path.join(data_root, "cifar-10-batches-py")
idx = 0
for batch_id in range(1, 6):
path = os.path.join(batch_dir, f"data_batch_{batch_id}")
with open(path, "rb") as f:
batch = pickle.load(f, encoding="bytes")
images = batch[b"data"].reshape(-1, 3, 32, 32).transpose(0, 2, 3, 1)
for img_np in images:
Image.fromarray(img_np).save(
os.path.join(img_dir, f"{idx:06d}.png"))
idx += 1
print(f"Extracted {idx} training images to {img_dir}")
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument("--data_root", type=str, default="./cifar10_data")
args = parser.parse_args()
download_cifar10(args.data_root)
extract_images(args.data_root)
print("Done. Use --data_root", args.data_root, "when launching training.")