Download prepare_data2.py from M1n1A1/MiniAI-TITS2: direct link, hf CLI and curl.
- Browser
- Download file 7.5 kB
-
https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/prepare_data2.py
- Command line
-
hf download hf://M1n1A1/MiniAI-TITS2/prepare_data2.py
-
curl -L -o prepare_data2.py https://huggingface.co/M1n1A1/MiniAI-TITS2/resolve/main/prepare_data2.py
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() | |