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