from __future__ import annotations import argparse from pathlib import Path import torch from solpix import SolPix, SolPixConfig def load_text(path: str, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]: payload = torch.load(path, map_location="cpu", weights_only=True) if isinstance(payload, torch.Tensor): embeddings, mask = payload, None elif isinstance(payload, dict) and isinstance(payload.get("text_embeddings"), torch.Tensor): embeddings = payload["text_embeddings"] mask = payload.get("text_mask") else: raise ValueError("text file must be a tensor or contain text_embeddings and optional text_mask") if embeddings.ndim == 2: embeddings = embeddings.unsqueeze(0) if embeddings.ndim != 3: raise ValueError("text_embeddings must have shape [L,D] or [B,L,D]") if mask is None: mask = torch.ones(embeddings.shape[:2], dtype=torch.bool) elif mask.ndim == 1: mask = mask.unsqueeze(0) if mask.shape != embeddings.shape[:2]: raise ValueError("text_mask must match the first two text_embeddings dimensions") return embeddings.to(device), mask.to(device=device, dtype=torch.bool) def main() -> None: parser = argparse.ArgumentParser(description="Euler-sample SolPix latent images from text embeddings") parser.add_argument("--checkpoint", required=True, help="Training checkpoint containing EMA weights") parser.add_argument("--text", required=True, help=".pt file with precomputed text_embeddings and optional text_mask") parser.add_argument("--empty-text", default=None, help="Required for classifier-free guidance when scale differs from 1") parser.add_argument("--output", required=True, help="Path for sampled latent .pt file") parser.add_argument("--height", type=int, default=16, help="Latent height; 16 means 512px at 32x compression") parser.add_argument("--width", type=int, default=16, help="Latent width; 16 means 512px at 32x compression") parser.add_argument("--steps", type=int, default=25) parser.add_argument("--guidance-scale", type=float, default=4.0) parser.add_argument("--seed", type=int, default=1234) parser.add_argument("--device", choices=("auto", "cuda", "mps", "cpu"), default="auto") parser.add_argument("--precision", choices=("auto", "bf16", "fp16", "fp32"), default="auto") parser.add_argument("--use-raw-weights", action="store_true", help="Use raw model weights instead of EMA weights") args = parser.parse_args() if args.steps < 1 or args.height < 1 or args.width < 1: raise ValueError("steps and latent dimensions must be positive") if args.device == "auto": if torch.cuda.is_available(): device = torch.device("cuda") elif torch.backends.mps.is_available(): device = torch.device("mps") else: device = torch.device("cpu") else: device = torch.device(args.device) checkpoint = torch.load(args.checkpoint, map_location="cpu", weights_only=False) config = SolPixConfig(**checkpoint.get("model_config", {})) model = SolPix(config).to(device).eval() weight_key = "model" if args.use_raw_weights or checkpoint.get("ema") is None else "ema" checkpoint_state = checkpoint[weight_key] # FP8-trained checkpoints retain float master weights plus TorchAO # bookkeeping buffers. Ignore only FP8-specific extra entries at sampling. model_keys = model.state_dict() compatible_state = {name: value for name, value in checkpoint_state.items() if name in model_keys} missing = set(model_keys) - set(compatible_state) if missing: raise ValueError(f"checkpoint is missing sampler weights: {sorted(missing)[:8]}") model.load_state_dict(compatible_state, strict=True) if args.precision == "auto": dtype = torch.bfloat16 if device.type == "cuda" and torch.cuda.is_bf16_supported() else ( torch.float16 if device.type == "cuda" else None ) elif args.precision == "fp32": dtype = None elif args.precision == "bf16": if device.type == "mps": raise ValueError("bf16 is not supported by this sampler setting on MPS") dtype = torch.bfloat16 else: dtype = torch.float16 conditional, conditional_mask = load_text(args.text, device) if conditional.shape[0] == 1: conditional = conditional.expand(1, -1, -1) batch = conditional.shape[0] unconditional = unconditional_mask = None if args.guidance_scale != 1.0: if args.empty_text is None: raise ValueError("pass --empty-text for classifier-free guidance, or set --guidance-scale 1") unconditional, unconditional_mask = load_text(args.empty_text, device) if unconditional.shape[0] == 1 and batch > 1: unconditional = unconditional.expand(batch, -1, -1) unconditional_mask = unconditional_mask.expand(batch, -1) if unconditional.shape[0] != batch: raise ValueError("conditional and empty-prompt batch sizes must match") generator = torch.Generator(device="cpu").manual_seed(args.seed) latents = torch.randn( batch, config.latent_channels, args.height, args.width, generator=generator, ).to(device) time_grid = torch.linspace(1.0, 0.0, args.steps + 1, device=device) autocast_enabled = dtype is not None with torch.inference_mode(): for current, following in zip(time_grid[:-1], time_grid[1:]): time = current.expand(batch) with torch.autocast(device_type=device.type, dtype=dtype, enabled=autocast_enabled): conditional_velocity = model(latents, time, conditional, conditional_mask) if unconditional is not None: unconditional_velocity = model(latents, time, unconditional, unconditional_mask) velocity = unconditional_velocity + args.guidance_scale * ( conditional_velocity - unconditional_velocity ) else: velocity = conditional_velocity latents = latents + (following - current) * velocity.float() output = Path(args.output) output.parent.mkdir(parents=True, exist_ok=True) torch.save( { "latents": latents.detach().float().cpu(), "height": args.height, "width": args.width, "steps": args.steps, "guidance_scale": args.guidance_scale, "seed": args.seed, "checkpoint_step": checkpoint.get("step"), }, output, ) print(f"Saved {batch} latent sample(s) to {output}") if __name__ == "__main__": main()