File size: 3,285 Bytes
a012ee3 | 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 | 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()
|