File size: 4,967 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, json, os
import numpy as np
import torch
from PIL import Image
from diffusers import AutoencoderKL
from transformers import CLIPTextModel, CLIPTokenizer
from dit import DiT

SCALE = 0.18215

@torch.no_grad()
def sample(model, seq, pool, null_seq, null_pool, steps, cfg, dev):
    B = seq.shape[0]
    x = torch.randn(B, 4, 32, 32, device=dev)
    ns = null_seq.expand(B, -1, -1)
    np_ = null_pool.expand(B, -1)
    dt = 1.0 / steps
    for i in range(steps):
        t = torch.full((B,), i * dt, device=dev)
        with torch.autocast("cuda", dtype=torch.bfloat16):
            vc = model(x, t, seq, pool)
            vu = model(x, t, ns, np_)
        v = vu + cfg * (vc - vu)
        x = x + v.float() * dt
    return x

@torch.no_grad()
def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--work", default="/root/pm4")
    ap.add_argument("--ckpt", default="/root/pm4/ckpt/final.pt")
    ap.add_argument("--vae", default="stabilityai/sd-vae-ft-mse")
    ap.add_argument("--clip", default="openai/clip-vit-base-patch32")
    ap.add_argument("--n", type=int, default=5000)
    ap.add_argument("--batch", type=int, default=100)
    ap.add_argument("--steps", type=int, default=50)
    ap.add_argument("--cfg", type=float, default=3.0)
    ap.add_argument("--max-tokens", type=int, default=40)
    ap.add_argument("--out", default="/root/pm4/eval_dit.json")
    ap.add_argument("--preview", default="")
    args = ap.parse_args()
    dev = "cuda"

    ck = torch.load(args.ckpt, map_location=dev)
    c = ck["cfg"]
    model = DiT(dim=c["dim"], depth=c["depth"], heads=c["heads"]).to(dev).eval()
    model.load_state_dict(ck["ema"])
    print(f"[eval] loaded {args.ckpt} step {ck['step']} params {model.num_params():,}", 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()
    null_seq = txt(**tok([""], padding="max_length", max_length=args.max_tokens,
                          truncation=True, return_tensors="pt").to(dev)).last_hidden_state.float()
    null_pool = txt(**tok([""], padding="max_length", max_length=args.max_tokens,
                          truncation=True, return_tensors="pt").to(dev)).pooler_output.float()

    d = np.load(os.path.join(args.work, "eval_256.npz"), allow_pickle=True)
    real = d["images"][:args.n]
    caps = [str(x) for x in d["captions"][:args.n]]
    n = len(caps)

    from torchmetrics.image.fid import FrechetInceptionDistance
    from torchmetrics.multimodal.clip_score import CLIPScore
    fid = FrechetInceptionDistance(feature=2048, normalize=True).to(dev)
    clip = CLIPScore(model_name_or_path=args.clip).to(dev)

    for i in range(0, n, args.batch):
        rb = torch.from_numpy(real[i:i + args.batch].astype(np.float32) / 255.0).permute(0, 3, 1, 2).to(dev)
        fid.update(rb, real=True)

    preview_imgs = []
    for i in range(0, n, args.batch):
        cb = caps[i:i + args.batch]
        t = tok(cb, padding="max_length", max_length=args.max_tokens, truncation=True, return_tensors="pt").to(dev)
        o = txt(**t)
        seq = o.last_hidden_state.float()
        pool = o.pooler_output.float()
        z = sample(model, seq, pool, null_seq, null_pool, args.steps, args.cfg, dev)
        img = vae.decode((z / SCALE).half()).sample.float()
        img = (img.clamp(-1, 1) + 1) / 2
        fid.update(img, real=False)
        clip.update((img * 255).to(torch.uint8), cb)
        if args.preview and len(preview_imgs) < 12:
            for j in range(min(len(cb), 12 - len(preview_imgs))):
                a = (img[j].permute(1, 2, 0).cpu().numpy() * 255).astype(np.uint8)
                preview_imgs.append((a, cb[j]))
        if i % (args.batch * 10) == 0:
            print(f"[eval] generated {i}/{n}", flush=True)

    fid_v = float(fid.compute().item())
    clip_v = float(clip.compute().item())
    res = {"n": n, "fid": round(fid_v, 2), "clip_score": round(clip_v, 2),
           "steps": args.steps, "cfg": args.cfg, "render_res": 256, "fid_size": 256,
           "clip_model": args.clip, "step": ck["step"]}
    with open(args.out, "w") as f:
        json.dump(res, f, indent=2)
    print(f"[eval] FID={fid_v:.2f} CLIP={clip_v:.2f} (n={n}, cfg={args.cfg}, steps={args.steps})", flush=True)

    if args.preview and preview_imgs:
        cell, pad = 256, 8
        cols = 4
        rows = (len(preview_imgs) + cols - 1) // cols
        sheet = Image.new("RGB", (cols * cell + (cols + 1) * pad, rows * cell + (rows + 1) * pad), (245, 246, 248))
        for k, (a, cap) in enumerate(preview_imgs):
            r, cc = divmod(k, cols)
            sheet.paste(Image.fromarray(a), (pad + cc * (cell + pad), pad + r * (cell + pad)))
        sheet.save(args.preview)
        print(f"[eval] wrote preview {args.preview}", flush=True)

if __name__ == "__main__":
    main()