Download sample_latents.py from solintellegence/SolPix: direct link, hf CLI and curl.
- Browser
- Download file 6.77 kB
-
https://huggingface.co/solintellegence/SolPix/resolve/main/sample_latents.py
- Command line
-
hf download hf://solintellegence/SolPix/sample_latents.py
-
curl -L -o sample_latents.py https://huggingface.co/solintellegence/SolPix/resolve/main/sample_latents.py
6.77 kB
| 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() | |