PixelModel v4: tiny latent diffusion (DiT + rectified flow), FID 39.54 / CLIP 28.04 (part 2)
6c9c825 verified | from __future__ import annotations | |
| import argparse | |
| import json | |
| import os | |
| import numpy as np | |
| import torch | |
| from PIL import Image | |
| from safetensors.torch import load_file | |
| from diffusers import AutoencoderKL | |
| from transformers import CLIPTextModel, CLIPTokenizer | |
| from dit import DiT | |
| SCALE = 0.18215 | |
| 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, npool = null_seq.expand(B, -1, -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, pool) | |
| vu = model(x, t, ns, npool) | |
| x = x + (vu + cfg * (vc - vu)).float() * dt | |
| return x | |
| def main(): | |
| ap = argparse.ArgumentParser() | |
| ap.add_argument("prompt") | |
| ap.add_argument("--out", default="out.png") | |
| ap.add_argument("--cfg", type=float, default=6.0) | |
| ap.add_argument("--steps", type=int, default=50) | |
| ap.add_argument("--device", default="cuda") | |
| ap.add_argument("--safetensors", default="model.safetensors") | |
| ap.add_argument("--config", default="config.json") | |
| ap.add_argument("--vae", default="stabilityai/sd-vae-ft-mse") | |
| ap.add_argument("--clip", default="openai/clip-vit-base-patch32") | |
| ap.add_argument("--max-tokens", type=int, default=40) | |
| args = ap.parse_args() | |
| dev = args.device | |
| d = json.load(open(args.config))["dit"] if os.path.exists(args.config) else {"dim": 384, "depth": 12, "heads": 6} | |
| model = DiT(dim=d["dim"], depth=d["depth"], heads=d["heads"]).to(dev).eval() | |
| model.load_state_dict(load_file(args.safetensors)) | |
| 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() | |
| def enc(strings): | |
| t = tok(strings, padding="max_length", max_length=args.max_tokens, truncation=True, return_tensors="pt").to(dev) | |
| o = txt(**t) | |
| return o.last_hidden_state.float(), o.pooler_output.float() | |
| seq, pool = enc([args.prompt]) | |
| null_seq, null_pool = enc([""]) | |
| 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)[0].permute(1, 2, 0).cpu().numpy() | |
| Image.fromarray((img * 255).round().astype(np.uint8)).save(args.out) | |
| print(f'[main] "{args.prompt}" -> {args.out} (cfg {args.cfg}, {args.steps} steps)') | |
| if __name__ == "__main__": | |
| main() | |