cifar10-ddpm / sample.py
PeterRabbit's picture
Initial upload: DDPM checkpoint, ONNX export, training code
96a0bf1 verified
Raw
History Blame Contribute Delete
1.2 kB
"""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']})")