File size: 5,045 Bytes
6c9c825
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
from __future__ import annotations
import argparse, io, json, os, random, zipfile
from concurrent.futures import ThreadPoolExecutor
import numpy as np
import requests
import torch
from PIL import Image
from diffusers import AutoencoderKL
from transformers import CLIPTextModel, CLIPTokenizer

SCALE = 0.18215
ANN_URL = "http://images.cocodataset.org/annotations/annotations_trainval2014.zip"

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 fetch_one(item, size):
    url, cap = item
    for _ in range(3):
        try:
            r = requests.get(url, timeout=15)
            if r.status_code == 200:
                return csr(Image.open(io.BytesIO(r.content)), size), cap
        except Exception:
            pass
    return None

@torch.no_grad()
def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--work", default="/root/pm4")
    ap.add_argument("--n-train", type=int, default=85000)
    ap.add_argument("--size", type=int, default=256)
    ap.add_argument("--max-tokens", type=int, default=40)
    ap.add_argument("--batch", type=int, default=64)
    ap.add_argument("--workers", type=int, default=48)
    ap.add_argument("--vae", default="stabilityai/sd-vae-ft-mse")
    ap.add_argument("--clip", default="openai/clip-vit-base-patch32")
    args = ap.parse_args()
    os.makedirs(args.work, exist_ok=True)
    dev = "cuda"

    ann_path = os.path.join(args.work, "captions_train2014.json")
    if not os.path.exists(ann_path):
        print("[data] downloading annotations", flush=True)
        z = os.path.join(args.work, "ann.zip")
        with requests.get(ANN_URL, stream=True, timeout=120) as r:
            with open(z, "wb") as f:
                for chunk in r.iter_content(1 << 20):
                    f.write(chunk)
        with zipfile.ZipFile(z) as zf:
            with zf.open("annotations/captions_train2014.json") as src, open(ann_path, "wb") as dst:
                dst.write(src.read())
        os.remove(z)
    ann = json.load(open(ann_path))
    url_by_id = {im["id"]: im["coco_url"] for im in ann["images"]}
    cap_by_id = {}
    for a in ann["annotations"]:
        cap_by_id.setdefault(a["image_id"], a["caption"])
    items = [(url_by_id[i], cap_by_id[i]) for i in cap_by_id if i in url_by_id]
    random.Random(0).shuffle(items)
    print(f"[data] {len(items)} train2014 image/caption pairs available; target {args.n_train}", flush=True)

    vae = AutoencoderKL.from_pretrained(args.vae).to(dev).half().eval()
    tok = CLIPTokenizer.from_pretrained(args.clip)
    txt = CLIPTextModel.from_pretrained(args.clip).to(dev).half().eval()

    lat_list, seq_list, pool_list = [], [], []
    pool = ThreadPoolExecutor(max_workers=args.workers)
    got, idx, nb = 0, 0, 0
    print("[data] starting download/encode loop", flush=True)
    while got < args.n_train and idx < len(items):
        chunk = items[idx:idx + args.batch]
        idx += args.batch
        nb += 1
        try:
            results = [r for r in pool.map(lambda it: fetch_one(it, args.size), chunk) if r is not None]
            if not results:
                print(f"[data] batch {nb}: 0 ok (skipped)", flush=True)
                continue
            imgs = np.stack([r[0] for r in results]).astype(np.float32) / 127.5 - 1.0
            caps = [r[1] for r in results]
            x = torch.from_numpy(imgs).permute(0, 3, 1, 2).to(dev).half()
            lat_list.append((vae.encode(x).latent_dist.mean * SCALE).cpu().numpy().astype(np.float16))
            t = tok(caps, padding="max_length", max_length=args.max_tokens, truncation=True, return_tensors="pt").to(dev)
            o = txt(**t)
            seq_list.append(o.last_hidden_state.cpu().numpy().astype(np.float16))
            pool_list.append(o.pooler_output.cpu().numpy().astype(np.float16))
            got += len(results)
        except Exception as e:
            print(f"[data] batch {nb} ERROR {type(e).__name__}: {str(e)[:100]}", flush=True)
            continue
        if nb % 10 == 0:
            print(f"[data] cached {got}/{args.n_train} (batch {nb})", flush=True)

    lat = np.concatenate(lat_list)[:args.n_train]
    seq = np.concatenate(seq_list)[:args.n_train]
    pool_ = np.concatenate(pool_list)[:args.n_train]
    np.save(os.path.join(args.work, "latents.npy"), lat)
    np.save(os.path.join(args.work, "text_seq.npy"), seq)
    np.save(os.path.join(args.work, "text_pool.npy"), pool_)
    t = tok([""], padding="max_length", max_length=args.max_tokens, truncation=True, return_tensors="pt").to(dev)
    o = txt(**t)
    np.save(os.path.join(args.work, "null_seq.npy"), o.last_hidden_state.cpu().numpy().astype(np.float16))
    np.save(os.path.join(args.work, "null_pool.npy"), o.pooler_output.cpu().numpy().astype(np.float16))
    print(f"[data] DONE latents{lat.shape} seq{seq.shape} -> {args.work}", flush=True)

if __name__ == "__main__":
    main()