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