import os, sys, io, gc, time, tarfile, argparse, numpy as np, torch from concurrent.futures import ThreadPoolExecutor from PIL import Image from huggingface_hub import hf_hub_download from diffusers import AutoencoderKL from transformers import CLIPTokenizer REPO = 'undefined443/cc12m-wds-coco-recaptioned' OUT = '/root/v5data' TMP = '/root/v5tmp' SCALE = 0.18215 RES = 256 MAXTOK = 40 ap = argparse.ArgumentParser() ap.add_argument('--start', type=int, default=0) ap.add_argument('--end', type=int, default=598) ap.add_argument('--batch', type=int, default=64) args = ap.parse_args() os.makedirs(OUT, exist_ok=True) os.makedirs(TMP, exist_ok=True) dev = 'cuda' vae = AutoencoderKL.from_pretrained('stabilityai/sd-vae-ft-mse').to(dev).half().eval() for p in vae.parameters(): p.requires_grad_(False) tok = CLIPTokenizer.from_pretrained('openai/clip-vit-base-patch32') def prep_img(raw): try: im = Image.open(io.BytesIO(raw)).convert('RGB') except Exception: return None w, h = im.size s = RES / min(w, h) im = im.resize((max(RES, int(w * s + 0.5)), max(RES, int(h * s + 0.5))), Image.BICUBIC) w, h = im.size l, t = (w - RES) // 2, (h - RES) // 2 return np.asarray(im.crop((l, t, l + RES, t + RES)), dtype=np.uint8) pool = ThreadPoolExecutor(16) t_start = time.time() done_total = 0 def fetch(si): return hf_hub_download(REPO, f'cc12m-coco-{si:04d}.tar', repo_type='dataset', local_dir=TMP) todo = [i for i in range(args.start, args.end) if not os.path.exists(f'{OUT}/shard_{i:04d}.npz')] print(f'{len(todo)} shards to process', flush=True) dl = ThreadPoolExecutor(2) nxt = dl.submit(fetch, todo[0]) if todo else None for k, si in enumerate(todo): dst = f'{OUT}/shard_{si:04d}.npz' try: p = nxt.result() except Exception as e: print(f' shard {si} download failed: {type(e).__name__}', flush=True) nxt = dl.submit(fetch, todo[k + 1]) if k + 1 < len(todo) else None continue nxt = dl.submit(fetch, todo[k + 1]) if k + 1 < len(todo) else None pairs = {} try: with tarfile.open(p) as tf: for m in tf: if not m.isfile(): continue key, _, ext = m.name.partition('.') if ext == 'jpg': pairs.setdefault(key, {})['img'] = tf.extractfile(m).read() elif ext == 'txt': pairs.setdefault(key, {})['cap'] = tf.extractfile(m).read().decode('utf-8', 'ignore').strip() except Exception as e: print(f' shard {si} tar failed: {type(e).__name__}', flush=True) os.remove(p); continue os.remove(p) items = [(v['img'], v['cap']) for v in pairs.values() if 'img' in v and 'cap' in v and v['cap']] del pairs if not items: continue lat_all, tok_all = [], [] for i in range(0, len(items), args.batch): chunk = items[i:i + args.batch] arrs = list(pool.map(prep_img, [c[0] for c in chunk])) keep = [j for j, a in enumerate(arrs) if a is not None] if not keep: continue x = np.stack([arrs[j] for j in keep]) t = torch.from_numpy(x).to(dev).half().permute(0, 3, 1, 2).div_(127.5).sub_(1.0) with torch.no_grad(): lat = vae.encode(t).latent_dist.sample() * SCALE lat_all.append(lat.cpu().numpy().astype(np.float16)) ids = tok([chunk[j][1] for j in keep], padding='max_length', max_length=MAXTOK, truncation=True, return_tensors='np')['input_ids'] tok_all.append(ids.astype(np.int32)) if not lat_all: continue lat = np.concatenate(lat_all); ids = np.concatenate(tok_all) np.savez(dst, latents=lat, tokens=ids) done_total += len(lat) el = time.time() - t_start print(f' shard {si:04d}: {len(lat)} imgs total {done_total} {done_total/el:.0f} img/s ' f'{el/60:.1f} min elapsed', flush=True) del lat_all, tok_all, lat, ids, items gc.collect() print('PREPDONE', done_total, flush=True)