File size: 6,374 Bytes
6c311ad | 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 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 | 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()
|