File size: 1,200 Bytes
96a0bf1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
"""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']})")