File size: 4,894 Bytes
a012ee3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import argparse, json, os, sys, numpy as np, torch
from diffusers import AutoencoderKL
from transformers import CLIPTextModel, CLIPTokenizer, CLIPModel
sys.path.insert(0, '/root')
from dit import DiT

SCALE = 0.18215

def _img_feats(self, pixel_values=None, **kw):
    return self.visual_projection(self.vision_model(pixel_values=pixel_values).pooler_output)

def _txt_feats(self, input_ids=None, attention_mask=None, **kw):
    return self.text_projection(self.text_model(input_ids=input_ids,
                                                attention_mask=attention_mask).pooler_output)

CLIPModel.get_image_features = _img_feats
CLIPModel.get_text_features = _txt_feats

@torch.no_grad()
def sample(model, seq, pool, ns, npool, steps, cfg, dev):
    B = seq.shape[0]
    x = torch.randn(B, 4, 32, 32, device=dev)
    a, b = ns.expand(B, -1, -1), npool.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, a, b)
        x = x + (vu + cfg * (vc - vu)).float() * dt
    return x

@torch.no_grad()
def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--ckpt', default='/root/runs/pm5/best.pt')
    ap.add_argument('--data', default='/root/pm5eval/eval_256.npz')
    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('--cfgs', default='2,3,4,5,6')
    ap.add_argument('--out', default='/root/pm5_eval.json')
    args = ap.parse_args()
    dev = 'cuda'

    ck = torch.load(args.ckpt, map_location=dev)
    model = DiT(dim=384, depth=12, heads=6).to(dev).eval()
    model.load_state_dict(ck['ema'])
    print(f"[eval] ckpt step {ck.get('step')} val {ck.get('val')}", flush=True)

    vae = AutoencoderKL.from_pretrained('stabilityai/sd-vae-ft-mse').to(dev).half().eval()
    tok = CLIPTokenizer.from_pretrained('openai/clip-vit-base-patch32')
    txt = CLIPTextModel.from_pretrained('openai/clip-vit-base-patch32').to(dev).half().eval()
    e = tok([''], padding='max_length', max_length=40, truncation=True, return_tensors='pt').to(dev)
    o = txt(**e)
    ns, npool = o.last_hidden_state.float(), o.pooler_output.float()

    d = np.load(args.data, allow_pickle=True)
    real = d['images'][:args.n]
    caps = [str(x) for x in d['captions'][:args.n]]
    n = len(caps)
    print(f'[eval] {n} real images and captions', flush=True)

    from torchmetrics.image.fid import FrechetInceptionDistance
    from torchmetrics.multimodal.clip_score import CLIPScore

    ref = CLIPScore(model_name_or_path='openai/clip-vit-base-patch32').to(dev)
    for i in range(0, n, args.batch):
        rb = torch.from_numpy(real[i:i + args.batch]).permute(0, 3, 1, 2).to(dev)
        ref.update(rb, caps[i:i + args.batch])
    real_clip = float(ref.compute().item())
    print(f'[eval] REAL images CLIP score {real_clip:.2f} (metric sanity check / ceiling)', flush=True)
    del ref
    torch.cuda.empty_cache()

    results = []
    for cfg in [float(x) for x in args.cfgs.split(',')]:
        fid = FrechetInceptionDistance(feature=2048, normalize=True).to(dev)
        clip = CLIPScore(model_name_or_path='openai/clip-vit-base-patch32').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)
        for i in range(0, n, args.batch):
            cb = caps[i:i + args.batch]
            t = tok(cb, padding='max_length', max_length=40, truncation=True, return_tensors='pt').to(dev)
            oo = txt(**t)
            z = sample(model, oo.last_hidden_state.float(), oo.pooler_output.float(),
                       ns, npool, args.steps, 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)
        f = float(fid.compute().item()); c = float(clip.compute().item())
        results.append({'cfg': cfg, 'fid': round(f, 2), 'clip_score': round(c, 2),
                        'n': n, 'steps': args.steps})
        print(f'[eval] cfg {cfg}: FID {f:.2f}  CLIP {c:.2f}', flush=True)
        del fid, clip
        torch.cuda.empty_cache()

    best = min(results, key=lambda r: r['fid'])
    json.dump({'ckpt_step': ck.get('step'), 'val_loss': ck.get('val'),
               'real_image_clip_score': round(real_clip, 2),
               'results': results, 'best_fid': best}, open(args.out, 'w'), indent=2)
    print(f"[eval] BEST FID {best['fid']} at cfg {best['cfg']} (CLIP {best['clip_score']})", flush=True)
    print('EVALDONE', flush=True)

if __name__ == '__main__':
    main()