(args)
| 59 | |
| 60 | |
| 61 | def 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 | |
| 94 | def launch(): |
no test coverage detected