from __future__ import annotations import argparse import json import os import numpy as np import torch from PIL import Image from diffusers import AutoencoderKL from transformers import CLIPModel, CLIPTextModel, CLIPTokenizer, T5EncoderModel, T5TokenizerFast from dit_v6 import MMDiT def patch_clip_score(): def get_image_features(self, pixel_values, **kw): pooled = self.vision_model(pixel_values=pixel_values).pooler_output return self.visual_projection(pooled) def get_text_features(self, input_ids=None, attention_mask=None, **kw): pooled = self.text_model(input_ids=input_ids, attention_mask=attention_mask).pooler_output return self.text_projection(pooled) CLIPModel.get_image_features = get_image_features CLIPModel.get_text_features = get_text_features @torch.no_grad() def sample(model, seq, mask, pool, null_seq, null_mask, null_pool, steps, cfg, dev): B = seq.shape[0] x = torch.randn(B, 4, 32, 32, device=dev) ns, nm, npo = null_seq.expand(B, -1, -1), null_mask.expand(B, -1), 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, mask, pool) vu = model(x, t, ns, nm, npo) x = x + (vu + cfg * (vc - vu)).float() * dt return x @torch.no_grad() def main(): ap = argparse.ArgumentParser() ap.add_argument("--work", default="/root/v6cache") ap.add_argument("--ckpt", default="/root/runs/pm6/best.pt") ap.add_argument("--vae", default="madebyollin/sdxl-vae-fp16-fix") ap.add_argument("--clip", default="openai/clip-vit-base-patch32") ap.add_argument("--t5", default="google/flan-t5-base") ap.add_argument("--n", type=int, default=5000) ap.add_argument("--batch", type=int, default=50) ap.add_argument("--steps", type=int, default=50) ap.add_argument("--cfg", type=float, nargs="+", default=[5.0]) ap.add_argument("--t5-len", type=int, default=32) ap.add_argument("--clip-len", type=int, default=40) ap.add_argument("--out", default="/root/runs/pm6/eval_results.jsonl") ap.add_argument("--preview", default="") args = ap.parse_args() dev = "cuda" ck = torch.load(args.ckpt, map_location=dev) c = ck["cfg"] model = MMDiT(dim=c["dim"], depth=c["depth"], heads=c["heads"], mlp_hidden=c["mlp_hidden"], t5_len=c["t5_len"]).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() vae_scale = vae.config.scaling_factor t5_tok = T5TokenizerFast.from_pretrained(args.t5) t5 = T5EncoderModel.from_pretrained(args.t5).to(dev).eval() clip_tok = CLIPTokenizer.from_pretrained(args.clip) clip_txt = CLIPTextModel.from_pretrained(args.clip).to(dev).eval() def enc(strings): te = t5_tok(strings, padding="max_length", max_length=args.t5_len, truncation=True, return_tensors="pt").to(dev) seq = t5(input_ids=te["input_ids"], attention_mask=te["attention_mask"]).last_hidden_state.float() ce = clip_tok(strings, padding="max_length", max_length=args.clip_len, truncation=True, return_tensors="pt").to(dev) pool = clip_txt(input_ids=ce["input_ids"]).pooler_output.float() return seq, te["attention_mask"].float(), pool null_seq, null_mask, null_pool = enc([""]) 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) patch_clip_score() from torchmetrics.image.fid import FrechetInceptionDistance from torchmetrics.multimodal.clip_score import CLIPScore results = [] for cfg_val in args.cfg: fid = FrechetInceptionDistance(feature=2048, normalize=True).to(dev) clip_metric = 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] seq, mask, pool = enc(cb) z = sample(model, seq, mask, pool, null_seq, null_mask, null_pool, args.steps, cfg_val, dev) img = vae.decode((z / vae_scale).half()).sample.float() img = (img.clamp(-1, 1) + 1) / 2 fid.update(img, real=False) clip_metric.update((img * 255).to(torch.uint8), cb) if args.preview and cfg_val == args.cfg[0] 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] cfg={cfg_val} generated {i}/{n}", flush=True) fid_v = float(fid.compute().item()) clip_v = float(clip_metric.compute().item()) res = {"n": n, "fid": round(fid_v, 2), "clip_score": round(clip_v, 2), "steps": args.steps, "cfg": cfg_val, "render_res": 256, "fid_size": 256, "clip_model": args.clip, "step": ck["step"]} results.append(res) print(f"[eval] cfg={cfg_val} FID={fid_v:.2f} CLIP={clip_v:.2f} (n={n}, steps={args.steps})", flush=True) with open(args.out, "a") as f: f.write(json.dumps(res) + "\n") if args.preview and cfg_val == args.cfg[0] 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) print(json.dumps(results, indent=2)) if __name__ == "__main__": main()