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