SolPix / generate.py
j0no12's picture
Publish SolPix step 210000 research checkpoint
8d31176 verified
Raw History Blame Contribute Delete
4.21 kB
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()