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()
|