MCPcopy Create free account
hub / github.com/KeepTryingTo/Pytorch-GAN / train_fn

Function train_fn

ProGAN/train.py:54–135  ·  view source on GitHub ↗
(
    critic,
    gen,
    loader,
    dataset,
    step,
    alpha,
    opt_critic,
    opt_gen,
    tensorboard_step,
    writer,
    scaler_gen,
    scaler_critic,
)

Source from the content-addressed store, hash-verified

52
53
54def 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)

Callers 1

mainFunction · 0.70

Calls 2

gradient_penaltyFunction · 0.90
plot_to_tensorboardFunction · 0.90

Tested by

no test coverage detected