| """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']})") |
|
|