MiniAI-TITS2 / prepare_data2.py
Codeminute's picture
T.I.T.S.2 — 93M DiT, rectified flow, 256px, 24 epochs on cleaned CC12M
2ca14fe verified
Raw History Blame Contribute Delete
7.5 kB
"""
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()