File size: 2,648 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
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):
    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

@torch.no_grad()
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()