PixelModel-v6 / eval_v6.py
TobiasLogic's picture
Publish PixelModel v6: MMDiT + REPA, 150k steps, FID 23.62 at cfg 3.0
6c311ad verified
Raw
History Blame Contribute Delete
6.37 kB
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()