| 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 |
|
|
| @torch.no_grad() |
| def sample(model, seq, pool, null_seq, null_pool, steps, cfg, dev, seed=None): |
| B = seq.shape[0] |
| g = None |
| if seed is not None: |
| g = torch.Generator(device=dev).manual_seed(seed) |
| x = torch.randn(B, 4, 32, 32, device=dev, generator=g) |
| 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 |
|
|
| @torch.no_grad() |
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("prompts", nargs="+") |
| ap.add_argument("--out", default="out.png") |
| ap.add_argument("--cfg", type=float, default=5.0) |
| ap.add_argument("--steps", type=int, default=50) |
| ap.add_argument("--seed", type=int, default=None) |
| ap.add_argument("--device", default="cuda") |
| ap.add_argument("--weights", 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() |
| sd = load_file(args.weights) |
| model.load_state_dict({k[len("dit."):]: v for k, v in sd.items() if k.startswith("dit.")}) |
|
|
| vae = AutoencoderKL.from_pretrained(args.vae).to(dev).half().eval() |
| tok = CLIPTokenizer.from_pretrained(args.clip) |
| txt = CLIPTextModel.from_pretrained(args.clip).to(dev).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.prompts) |
| null_seq, null_pool = enc([""]) |
| z = sample(model, seq, pool, null_seq, null_pool, args.steps, args.cfg, dev, args.seed) |
| img = vae.decode((z / SCALE).half()).sample.float() |
| img = ((img.clamp(-1, 1) + 1) / 2).permute(0, 2, 3, 1).cpu().numpy() |
|
|
| n = len(args.prompts) |
| if n == 1: |
| Image.fromarray((img[0] * 255).round().astype(np.uint8)).save(args.out) |
| paths = [args.out] |
| else: |
| root, ext = os.path.splitext(args.out) |
| paths = [] |
| for i in range(n): |
| p = f"{root}_{i}{ext}" |
| Image.fromarray((img[i] * 255).round().astype(np.uint8)).save(p) |
| paths.append(p) |
| for p, s in zip(paths, args.prompts): |
| print(f'[pixelmodel] "{s}" -> {p} (cfg {args.cfg}, {args.steps} steps)') |
|
|
| if __name__ == "__main__": |
| main() |
|
|