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

Function train_fn

ProGAN/demo.py:45–126  ·  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

43
44
45def 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)

Callers 1

mainFunction · 0.70

Calls 2

gradient_penaltyFunction · 0.70
plot_to_tensorboardFunction · 0.70

Tested by

no test coverage detected