(
critic,
gen,
loader,
dataset,
step,
alpha,
opt_critic,
opt_gen,
tensorboard_step,
writer,
scaler_gen,
scaler_critic,
)
| 43 | |
| 44 | |
| 45 | def train_fn( |
| 46 | critic, |
| 47 | gen, |
| 48 | loader, |
| 49 | dataset, |
| 50 | step, |
| 51 | alpha, |
| 52 | opt_critic, |
| 53 | opt_gen, |
| 54 | tensorboard_step, |
| 55 | writer, |
| 56 | scaler_gen, |
| 57 | scaler_critic, |
| 58 | ): |
| 59 | loop = tqdm(loader, leave=True) |
| 60 | for batch_idx, (real, _) in enumerate(loop): |
| 61 | real = real.to(DEVICE) |
| 62 | cur_batch_size = real.shape[0] |
| 63 | |
| 64 | # Train Critic: max E[critic(real)] - E[critic(fake)] <-> min -E[critic(real)] + E[critic(fake)] |
| 65 | # which is equivalent to minimizing the negative of the expression |
| 66 | noise = torch.randn(cur_batch_size, Z_DIM, 1, 1).to(DEVICE) |
| 67 | |
| 68 | # with torch.cuda.amp.autocast(): |
| 69 | fake = gen(noise, alpha, step) |
| 70 | critic_real = critic(real, alpha, step) |
| 71 | critic_fake = critic(fake.detach(), alpha, step) |
| 72 | gp = gradient_penalty(critic, real, fake, alpha, step, device=DEVICE) |
| 73 | loss_critic = ( |
| 74 | -(torch.mean(critic_real) - torch.mean(critic_fake)) |
| 75 | + LAMBDA_GP * gp |
| 76 | + (0.001 * torch.mean(critic_real ** 2)) |
| 77 | ) |
| 78 | |
| 79 | opt_critic.zero_grad() |
| 80 | loss_critic.backward() |
| 81 | opt_critic.step() |
| 82 | # opt_critic.zero_grad() |
| 83 | # scaler_critic.scale(loss_critic).backward() |
| 84 | # scaler_critic.step(opt_critic) |
| 85 | # scaler_critic.update() |
| 86 | |
| 87 | # Train Generator: max E[critic(gen_fake)] <-> min -E[critic(gen_fake)] |
| 88 | # with torch.cuda.amp.autocast(): |
| 89 | gen_fake = critic(fake, alpha, step) |
| 90 | loss_gen = -torch.mean(gen_fake) |
| 91 | |
| 92 | opt_gen.zero_grad() |
| 93 | loss_gen.backward() |
| 94 | opt_gen.step() |
| 95 | # opt_gen.zero_grad() |
| 96 | # scaler_gen.scale(loss_gen).backward() |
| 97 | # scaler_gen.step(opt_gen) |
| 98 | # scaler_gen.update() |
| 99 | |
| 100 | # Update alpha and ensure less than 1 |
| 101 | alpha += cur_batch_size / ( |
| 102 | (PROGRESSIVE_EPOCHS[step] * 0.5) * len(dataset) |
no test coverage detected