"""Download CIFAR-10 and extract images for training + FID evaluation. Usage: python prepare_cifar10.py [--data_root ./cifar10_data] Creates: /cifar-10-batches-py/ (pickle files used by data.py) /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 /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.")