from __future__ import annotations import argparse import io import os import tarfile import time from concurrent.futures import ThreadPoolExecutor import numpy as np import torch from PIL import Image from diffusers import AutoencoderKL from transformers import CLIPTokenizer, T5TokenizerFast from huggingface_hub import hf_hub_download REPO = "undefined443/cc12m-wds-coco-recaptioned" def csr(img, size): img = img.convert("RGB") w, h = img.size s = min(w, h) l, t = (w - s) // 2, (h - s) // 2 return np.asarray(img.crop((l, t, l + s, t + s)).resize((size, size), Image.BICUBIC), dtype=np.uint8) def load_shard_items(tar_path, size): t = tarfile.open(tar_path) raw = {} for m in t.getmembers(): if not m.isfile(): continue key, ext = m.name.rsplit(".", 1) raw.setdefault(key, {})[ext] = t.extractfile(m).read() t.close() def proc(kv): _, d = kv if "jpg" not in d or "txt" not in d: return None try: arr = csr(Image.open(io.BytesIO(d["jpg"])), size) cap = d["txt"].decode("utf-8", "ignore").strip() if not cap: return None return arr, cap except Exception: return None results = [] with ThreadPoolExecutor(max_workers=32) as pool: for r in pool.map(proc, raw.items()): if r is not None: results.append(r) return results @torch.no_grad() def main(): ap = argparse.ArgumentParser() ap.add_argument("--out", default="/root/v6cache/shards") ap.add_argument("--tmp", default="/root/v6cache/tars") ap.add_argument("--size", type=int, default=256) ap.add_argument("--t5-len", type=int, default=32) ap.add_argument("--clip-len", type=int, default=40) ap.add_argument("--batch", type=int, default=128) ap.add_argument("--vae", default="madebyollin/sdxl-vae-fp16-fix") ap.add_argument("--clip", default="openai/clip-vit-base-patch32") ap.add_argument("--t5", default="google/flan-t5-base") ap.add_argument("--start", type=int, default=0) ap.add_argument("--end", type=int, default=598) ap.add_argument("--prefetch", type=int, default=2) args = ap.parse_args() os.makedirs(args.out, exist_ok=True) os.makedirs(args.tmp, exist_ok=True) dev = "cuda" vae = AutoencoderKL.from_pretrained(args.vae).to(dev).half().eval() scale = vae.config.scaling_factor print(f"[prep] vae={args.vae} scaling_factor={scale}", flush=True) clip_tok = CLIPTokenizer.from_pretrained(args.clip) t5_tok = T5TokenizerFast.from_pretrained(args.t5) shard_names = [f"cc12m-coco-{i:04d}.tar" for i in range(args.start, args.end)] def fetch(name): return hf_hub_download(REPO, name, repo_type="dataset", local_dir=args.tmp) fpool = ThreadPoolExecutor(max_workers=args.prefetch) futures = {} def ensure_fetch(idx): if idx < len(shard_names) and idx not in futures: futures[idx] = fpool.submit(fetch, shard_names[idx]) for k in range(args.prefetch): ensure_fetch(k) t0 = time.time() total = 0 for i, name in enumerate(shard_names): out_path = f"{args.out}/shard_{args.start+i:04d}.npz" if os.path.exists(out_path): total += np.load(out_path)["latents"].shape[0] futures.pop(i, None) ensure_fetch(i + args.prefetch) continue tar_path = futures.pop(i).result() ensure_fetch(i + args.prefetch) items = load_shard_items(tar_path, args.size) os.remove(tar_path) if not items: print(f"[prep] shard {args.start+i:04d} EMPTY, skipping", flush=True) continue imgs = [a for a, c in items] caps = [c for a, c in items] lat_chunks = [] for j in range(0, len(imgs), args.batch): chunk = np.stack(imgs[j:j + args.batch]).astype(np.float32) / 127.5 - 1.0 x = torch.from_numpy(chunk).permute(0, 3, 1, 2).to(dev).half() z = vae.encode(x).latent_dist.mean * scale lat_chunks.append(z.cpu().numpy().astype(np.float16)) latents = np.concatenate(lat_chunks) t5o = t5_tok(caps, padding="max_length", max_length=args.t5_len, truncation=True, return_tensors="np") clip_ids = clip_tok(caps, padding="max_length", max_length=args.clip_len, truncation=True, return_tensors="np")["input_ids"] np.savez(out_path, latents=latents, t5_ids=t5o["input_ids"].astype(np.int32), t5_mask=t5o["attention_mask"].astype(np.int8), clip_ids=clip_ids.astype(np.int64)) total += len(imgs) el = time.time() - t0 print(f"[prep] shard {args.start+i:04d} +{len(imgs)} total={total} " f"({total/el:.1f} img/s, {el/3600:.2f}h elapsed)", flush=True) print(f"[prep] DONE total={total}", flush=True) if __name__ == "__main__": main()