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