""" T.I.T.S.2 data prep — streams CC12M, keeps only the rows OpenDiffusionAI's cleaned list vouches for (watermark/junk filtered, LLaVA-written captions), crops to 256px, encodes with the SD VAE, and writes latent shards. Why the join: `opendiffusionai/cc12m-cleaned` has the good captions but is URL-only (~2 img/s once you account for dead links). `pixparse/cc12m-wds` has the actual image bytes and streams fast. Same source, same order, so we walk the image stream and keep a rolling window of the cleaned list to look up each url. Resume works at tar-file granularity: CC12M is 2176 tars, so we record which tar we're on rather than "row N", because skipping N rows means re-downloading them. The cleaned list is parquet (cheap to skip), so its position is tracked as a row offset. Nothing full-size ever hits disk: images are encoded to 4x32x32 fp16 latents (8KB each) and written in shards of 10k alongside their captions. Usage: python prepare_data2.py --out_dir data2 --num 600000 """ import argparse import collections import glob import json import os import re import time import numpy as np import torch from datasets import load_dataset from huggingface_hub import HfApi from PIL import Image SHARD_SIZE = 10000 WDS_REPO = "pixparse/cc12m-wds" CLEAN_REPO = "opendiffusionai/cc12m-cleaned" def parse_args(): p = argparse.ArgumentParser() p.add_argument("--out_dir", type=str, default="data2") p.add_argument("--num", type=int, default=600000, help="how many usable images to collect") p.add_argument("--image_size", type=int, default=256) p.add_argument("--min_side", type=int, default=256, help="skip images smaller than this") p.add_argument("--max_aspect", type=float, default=1.6, help="skip images more oblong than this") p.add_argument("--window", type=int, default=60000, help="rolling lookup window into the cleaned list") p.add_argument("--batch_size", type=int, default=32, help="VAE encode batch") p.add_argument("--vae", type=str, default="stabilityai/sd-vae-ft-mse") p.add_argument("--device", type=str, default="cuda" if torch.cuda.is_available() else "cpu") return p.parse_args() def center_crop_resize(img, size): w, h = img.size side = min(w, h) img = img.crop(((w - side) // 2, (h - side) // 2, (w + side) // 2, (h + side) // 2)) return img.resize((size, size), Image.BICUBIC) def tar_urls(): files = sorted(f for f in HfApi().list_repo_files(WDS_REPO, repo_type="dataset") if f.endswith(".tar")) return [f"https://huggingface.co/datasets/{WDS_REPO}/resolve/main/{f}" for f in files] def main(): args = parse_args() os.makedirs(args.out_dir, exist_ok=True) state_path = os.path.join(args.out_dir, "state.json") from diffusers import AutoencoderKL vae = AutoencoderKL.from_pretrained(args.vae, torch_dtype=torch.float16).to(args.device).eval() vae.requires_grad_(False) existing = sorted(glob.glob(os.path.join(args.out_dir, "shard_*.pt"))) next_shard = max((int(re.search(r"(\d+)", os.path.basename(p)).group(1)) for p in existing), default=-1) + 1 state = {"tar_index": 0, "clean_rows": 0, "images": 0} if os.path.exists(state_path) and existing: state.update(json.load(open(state_path))) print(f"resuming: {state['images']} images, tar {state['tar_index']}, " f"cleaned-list row {state['clean_rows']}", flush=True) urls = tar_urls() print(f"{len(urls)} source tars", flush=True) clean = load_dataset(CLEAN_REPO, split="train", streaming=True) if state["clean_rows"]: clean = clean.skip(state["clean_rows"]) # parquet: cheap to skip clean = iter(clean) clean_pulled = state["clean_rows"] window = collections.OrderedDict() def top_up(): nonlocal clean_pulled while len(window) < args.window: try: r = next(clean) except StopIteration: return False window[r["url"]] = r["caption_llava_short"] or r["caption_llava"] clean_pulled += 1 return True top_up() batch_imgs, batch_caps, shard_lat, shard_cap = [], [], [], [] n = state["images"] t0, t_imgs, miss = time.time(), 0, 0 def flush_batch(): if not batch_imgs: return x = torch.stack(batch_imgs).to(args.device, torch.float16) with torch.no_grad(): lat = vae.encode(x).latent_dist.sample() * vae.config.scaling_factor shard_lat.append(lat.cpu().to(torch.float16)) shard_cap.extend(batch_caps) batch_imgs.clear() batch_caps.clear() def save_state(tar_index): # Re-read the still-pending window entries on resume rather than lose them. json.dump({"tar_index": tar_index, "clean_rows": max(0, clean_pulled - len(window)), "images": n}, open(state_path, "w")) def flush_shard(tar_index): nonlocal next_shard, shard_lat, shard_cap if not shard_cap: return path = os.path.join(args.out_dir, f"shard_{next_shard:05d}.pt") torch.save({"latents": torch.cat(shard_lat), "captions": list(shard_cap)}, path) next_shard += 1 shard_lat, shard_cap = [], [] save_state(tar_index) rate = t_imgs / (time.time() - t0 + 1e-9) print(f"saved {path} ({n} images total, tar {tar_index}, {rate:.0f} img/s)", flush=True) for ti in range(state["tar_index"], len(urls)): if n >= args.num: break try: ds = load_dataset("webdataset", data_files={"train": [urls[ti]]}, split="train", streaming=True) for ex in ds: caption = window.pop(ex["json"].get("url"), None) if caption is None: # Not in the cleaned list -> watermarked/junk. But a long miss streak # means the window has drifted behind the image stream (e.g. after a # resume), so walk it forward until matches resume. miss += 1 if miss > 3000: for _ in range(min(20000, len(window))): window.popitem(last=False) top_up() miss = 0 continue miss = 0 if len(window) < args.window // 2: top_up() img = ex["jpg"] w, h = img.size if min(w, h) < args.min_side or max(w, h) / min(w, h) > args.max_aspect: continue img = center_crop_resize(img.convert("RGB"), args.image_size) batch_imgs.append(torch.from_numpy(np.asarray(img).copy()).permute(2, 0, 1).float() / 127.5 - 1.0) batch_caps.append(caption.strip()) n += 1 t_imgs += 1 if len(batch_imgs) >= args.batch_size: flush_batch() if len(shard_cap) >= SHARD_SIZE: flush_shard(ti) if n >= args.num: break except Exception as e: # Truncated tar / dropped connection: lose this file, not the run. print(f"tar {ti} failed ({type(e).__name__}: {e}) — skipping", flush=True) save_state(ti + 1) flush_batch() flush_shard(min(state["tar_index"] + 1, len(urls))) print(f"done: {n} images", flush=True) if __name__ == "__main__": main()