"""Generate images from a trained checkpoint. Run: python diffusion/sample.py [--n 64] [--out samples.png] [--seed 0] """ import argparse import os import torch from torchvision.utils import save_image from train_diffusion import UNet, Diffusion, T, DDIM_STEPS ap = argparse.ArgumentParser() ap.add_argument("--ckpt", default=os.path.join(os.path.dirname(__file__), "out_big", "checkpoint.pt")) ap.add_argument("--base", type=int, default=128, help="must match the trained model") ap.add_argument("--n", type=int, default=64, help="number of images") ap.add_argument("--steps", type=int, default=DDIM_STEPS) ap.add_argument("--seed", type=int, default=None) ap.add_argument("--out", default=os.path.join(os.path.dirname(__file__), "generated.png")) args = ap.parse_args() if args.seed is not None: torch.manual_seed(args.seed) model = UNet(base=args.base).cuda().eval() ck = torch.load(args.ckpt, map_location="cuda", weights_only=True) model.load_state_dict(ck["ema"]) imgs = Diffusion(T, "cuda").ddim_sample(model, args.n, args.steps, "cuda") save_image(imgs * 0.5 + 0.5, args.out, nrow=int(args.n ** 0.5)) print(f"saved {args.n} images to {args.out} (checkpoint epoch {ck['epoch']})")