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