File size: 4,213 Bytes
8d31176
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from __future__ import annotations

import argparse
from pathlib import Path

import torch
from huggingface_hub import hf_hub_download
from PIL import Image
from transformers import AutoTokenizer, T5EncoderModel

from solpix import AutoencoderDCSol, SolPixConfig, SolPixTransformer2D


REPO_ID = "solintellegence/SolPix"
TEXT_MODEL = "google/flan-t5-base"
VAE_MODEL = "mit-han-lab/dc-ae-f32c32-sana-1.1-diffusers"
VAE_REVISION = "df0d9d634aea77793c1fb685d4b9db092c99e686"


def encode(text: str, tokenizer, encoder, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
    tokens = tokenizer(text, return_tensors="pt", truncation=True, max_length=96)
    input_ids = tokens["input_ids"].to(device)
    mask = tokens["attention_mask"].to(device).bool()
    with torch.inference_mode():
        embeddings = encoder(input_ids=input_ids, attention_mask=mask).last_hidden_state
    return embeddings.float(), mask


def main() -> None:
    parser = argparse.ArgumentParser(description="Generate one image with SolPix and the matching SANA DC-AE.")
    parser.add_argument("--prompt", required=True)
    parser.add_argument("--output", default="solpix.png")
    parser.add_argument("--checkpoint", default=None, help="Local full checkpoint; defaults to the Hub release checkpoint.")
    parser.add_argument("--seed", type=int, default=1234)
    parser.add_argument("--steps", type=int, default=40)
    parser.add_argument("--guidance-scale", type=float, default=3.5)
    args = parser.parse_args()

    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    dtype = torch.bfloat16 if device.type == "cuda" and torch.cuda.is_bf16_supported() else torch.float32
    checkpoint_path = args.checkpoint or hf_hub_download(REPO_ID, "step_00210000.pt")
    checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=False)

    model_config = SolPixConfig(**checkpoint["model_config"])
    model = SolPixTransformer2D(model_config).to(device).eval()
    model_keys = model.state_dict()
    ema = checkpoint.get("ema") or checkpoint["model"]
    compatible = {name: value for name, value in ema.items() if name in model_keys}
    missing = set(model_keys) - set(compatible)
    if missing:
        raise ValueError(f"Checkpoint is missing model tensors: {sorted(missing)[:8]}")
    model.load_state_dict(compatible, strict=True)

    tokenizer = AutoTokenizer.from_pretrained(TEXT_MODEL)
    encoder = T5EncoderModel.from_pretrained(TEXT_MODEL, torch_dtype=dtype).to(device).eval()
    encoder.requires_grad_(False)
    conditional, conditional_mask = encode(args.prompt, tokenizer, encoder, device)
    unconditional, unconditional_mask = encode("", tokenizer, encoder, device)

    generator = torch.Generator(device="cpu").manual_seed(args.seed)
    latents = torch.randn(1, 32, 16, 16, generator=generator).to(device)
    time_grid = torch.linspace(1.0, 0.0, args.steps + 1, device=device)
    with torch.inference_mode():
        for current, following in zip(time_grid[:-1], time_grid[1:]):
            time = current.expand(1)
            with torch.autocast(device_type=device.type, dtype=dtype, enabled=device.type == "cuda"):
                v_cond = model(latents, time, conditional, conditional_mask)
                v_uncond = model(latents, time, unconditional, unconditional_mask)
                velocity = v_uncond + args.guidance_scale * (v_cond - v_uncond)
            latents = latents + (following - current) * velocity.float()

    vae = AutoencoderDCSol.from_pretrained(
        VAE_MODEL, revision=VAE_REVISION, torch_dtype=dtype
    ).to(device).eval()
    scale = float(vae.config.scaling_factor or 0.41407)
    with torch.inference_mode():
        decoded = vae.decode((latents / scale).to(dtype=dtype), return_dict=False)[0]
    pixels = (decoded.float().squeeze(0).permute(1, 2, 0).cpu().clamp(-1, 1) * 127.5 + 127.5)
    image = Image.fromarray(pixels.to(torch.uint8).numpy(), mode="RGB")
    output = Path(args.output)
    output.parent.mkdir(parents=True, exist_ok=True)
    image.save(output, format="PNG", optimize=True)
    print(f"Saved {output} (checkpoint step {checkpoint.get('step')}, seed {args.seed})")


if __name__ == "__main__":
    main()