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()