File size: 4,967 Bytes
6c9c825 | 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 | from __future__ import annotations
import argparse, json, os
import numpy as np
import torch
from PIL import Image
from diffusers import AutoencoderKL
from transformers import CLIPTextModel, CLIPTokenizer
from dit import DiT
SCALE = 0.18215
@torch.no_grad()
def sample(model, seq, pool, null_seq, null_pool, steps, cfg, dev):
B = seq.shape[0]
x = torch.randn(B, 4, 32, 32, device=dev)
ns = null_seq.expand(B, -1, -1)
np_ = 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, pool)
vu = model(x, t, ns, np_)
v = vu + cfg * (vc - vu)
x = x + v.float() * dt
return x
@torch.no_grad()
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--work", default="/root/pm4")
ap.add_argument("--ckpt", default="/root/pm4/ckpt/final.pt")
ap.add_argument("--vae", default="stabilityai/sd-vae-ft-mse")
ap.add_argument("--clip", default="openai/clip-vit-base-patch32")
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("--cfg", type=float, default=3.0)
ap.add_argument("--max-tokens", type=int, default=40)
ap.add_argument("--out", default="/root/pm4/eval_dit.json")
ap.add_argument("--preview", default="")
args = ap.parse_args()
dev = "cuda"
ck = torch.load(args.ckpt, map_location=dev)
c = ck["cfg"]
model = DiT(dim=c["dim"], depth=c["depth"], heads=c["heads"]).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()
tok = CLIPTokenizer.from_pretrained(args.clip)
txt = CLIPTextModel.from_pretrained(args.clip).to(dev).half().eval()
null_seq = txt(**tok([""], padding="max_length", max_length=args.max_tokens,
truncation=True, return_tensors="pt").to(dev)).last_hidden_state.float()
null_pool = txt(**tok([""], padding="max_length", max_length=args.max_tokens,
truncation=True, return_tensors="pt").to(dev)).pooler_output.float()
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)
from torchmetrics.image.fid import FrechetInceptionDistance
from torchmetrics.multimodal.clip_score import CLIPScore
fid = FrechetInceptionDistance(feature=2048, normalize=True).to(dev)
clip = 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]
t = tok(cb, padding="max_length", max_length=args.max_tokens, truncation=True, return_tensors="pt").to(dev)
o = txt(**t)
seq = o.last_hidden_state.float()
pool = o.pooler_output.float()
z = sample(model, seq, pool, null_seq, null_pool, args.steps, args.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)
if args.preview 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] generated {i}/{n}", flush=True)
fid_v = float(fid.compute().item())
clip_v = float(clip.compute().item())
res = {"n": n, "fid": round(fid_v, 2), "clip_score": round(clip_v, 2),
"steps": args.steps, "cfg": args.cfg, "render_res": 256, "fid_size": 256,
"clip_model": args.clip, "step": ck["step"]}
with open(args.out, "w") as f:
json.dump(res, f, indent=2)
print(f"[eval] FID={fid_v:.2f} CLIP={clip_v:.2f} (n={n}, cfg={args.cfg}, steps={args.steps})", flush=True)
if args.preview 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)
if __name__ == "__main__":
main()
|