MCPcopy Create free account
hub / github.com/dome272/Diffusion-Models-pytorch / train

Function train

ddpm.py:61–91  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

59
60
61def train(args):
62 setup_logging(args.run_name)
63 device = args.device
64 dataloader = get_data(args)
65 model = UNet().to(device)
66 optimizer = optim.AdamW(model.parameters(), lr=args.lr)
67 mse = nn.MSELoss()
68 diffusion = Diffusion(img_size=args.image_size, device=device)
69 logger = SummaryWriter(os.path.join("runs", args.run_name))
70 l = len(dataloader)
71
72 for epoch in range(args.epochs):
73 logging.info(f"Starting epoch {epoch}:")
74 pbar = tqdm(dataloader)
75 for i, (images, _) in enumerate(pbar):
76 images = images.to(device)
77 t = diffusion.sample_timesteps(images.shape[0]).to(device)
78 x_t, noise = diffusion.noise_images(images, t)
79 predicted_noise = model(x_t, t)
80 loss = mse(noise, predicted_noise)
81
82 optimizer.zero_grad()
83 loss.backward()
84 optimizer.step()
85
86 pbar.set_postfix(MSE=loss.item())
87 logger.add_scalar("MSE", loss.item(), global_step=epoch * l + i)
88
89 sampled_images = diffusion.sample(model, n=images.shape[0])
90 save_images(sampled_images, os.path.join("results", args.run_name, f"{epoch}.jpg"))
91 torch.save(model.state_dict(), os.path.join("models", args.run_name, f"ckpt.pt"))
92
93
94def launch():

Callers 1

launchFunction · 0.70

Calls 8

sample_timestepsMethod · 0.95
noise_imagesMethod · 0.95
sampleMethod · 0.95
UNetClass · 0.90
setup_loggingFunction · 0.85
get_dataFunction · 0.85
save_imagesFunction · 0.85
DiffusionClass · 0.70

Tested by

no test coverage detected