PixelModel-v6 / main.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
3.5 kB
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, T5EncoderModel, T5TokenizerFast
from dit_v6 import MMDiT
@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("prompt")
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("--device", default="cuda" if torch.cuda.is_available() else "cpu")
ap.add_argument("--safetensors", default="model.safetensors")
ap.add_argument("--config", default="config.json")
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("--t5-len", type=int, default=32)
ap.add_argument("--clip-len", 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": 512, "depth": 16, "heads": 8, "mlp_hidden": 1408, "t5_len": 32}
model = MMDiT(dim=d["dim"], depth=d["depth"], heads=d["heads"], mlp_hidden=d["mlp_hidden"],
t5_len=d["t5_len"]).to(dev).eval()
model.load_state_dict(load_file(args.safetensors))
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
seq, mask, pool = enc([args.prompt])
null_seq, null_mask, null_pool = enc([""])
z = sample(model, seq, mask, pool, null_seq, null_mask, null_pool, args.steps, args.cfg, dev)
img = vae.decode((z / vae_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()