# Training Loop # Lists to keep track of progress img_list = [] G_losses = [] D_losses = [] iters = 0 print("Starting Training Loop...") # For each epoch for epoch in range(num_epochs): # For each batch in the dataloader for i, data in enumerate([dataloader](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader "torch.utils.data.DataLoader"), 0): ############################ # (1) Update D network: maximize log(D(x)) + log(1 - D(G(z))) ########################### ## Train with all-real batch [netD.zero_grad](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.zero_grad "torch.nn.Module.zero_grad")() # Format batch [real_cpu](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = data[0].to([device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")) b_size = [real_cpu](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").size(0) [label](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [torch.full](https://docs.pytorch.org/docs/stable/generated/torch.full.html#torch.full "torch.full")((b_size,), real_label, dtype=[torch.float](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.dtype "torch.dtype"), [device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")=[device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")) # Forward pass real batch through D [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = netD([real_cpu](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")).view(-1) # Calculate loss on all-real batch [errD_real](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [criterion](https://docs.pytorch.org/docs/stable/generated/torch.nn.BCELoss.html#torch.nn.BCELoss "torch.nn.BCELoss")([output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), [label](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # Calculate gradients for D in backward pass [errD_real.backward](https://docs.pytorch.org/docs/stable/generated/torch.Tensor.backward.html#torch.Tensor.backward "torch.Tensor.backward")() D_x = [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").mean().item() ## Train with all-fake batch # Generate batch of latent vectors [noise](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [torch.randn](https://docs.pytorch.org/docs/stable/generated/torch.randn.html#torch.randn "torch.randn")(b_size, nz, 1, 1, [device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")=[device](https://docs.pytorch.org/docs/stable/tensor_attributes.html#torch.device "torch.device")) # Generate fake image batch with G [fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = netG([noise](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) [label](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").fill_(fake_label) # Classify all fake batch with D [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = netD([fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").detach()).view(-1) # Calculate D's loss on the all-fake batch [errD_fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [criterion](https://docs.pytorch.org/docs/stable/generated/torch.nn.BCELoss.html#torch.nn.BCELoss "torch.nn.BCELoss")([output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), [label](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # Calculate the gradients for this batch, accumulated (summed) with previous gradients [errD_fake.backward](https://docs.pytorch.org/docs/stable/generated/torch.Tensor.backward.html#torch.Tensor.backward "torch.Tensor.backward")() D_G_z1 = [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").mean().item() # Compute error of D as sum over the fake and the real batches [errD](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [errD_real](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") + [errD_fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") # Update D [optimizerD.step](https://docs.pytorch.org/docs/stable/generated/torch.optim.Adam.html#torch.optim.Adam.step "torch.optim.Adam.step")() ############################ # (2) Update G network: maximize log(D(G(z))) ########################### [netG.zero_grad](https://docs.pytorch.org/docs/stable/generated/torch.nn.Module.html#torch.nn.Module.zero_grad "torch.nn.Module.zero_grad")() [label](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").fill_(real_label) # fake labels are real for generator cost # Since we just updated D, perform another forward pass of all-fake batch through D [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = netD([fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")).view(-1) # Calculate G's loss based on this output [errG](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = [criterion](https://docs.pytorch.org/docs/stable/generated/torch.nn.BCELoss.html#torch.nn.BCELoss "torch.nn.BCELoss")([output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), [label](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")) # Calculate gradients for G [errG.backward](https://docs.pytorch.org/docs/stable/generated/torch.Tensor.backward.html#torch.Tensor.backward "torch.Tensor.backward")() D_G_z2 = [output](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").mean().item() # Update G [optimizerG.step](https://docs.pytorch.org/docs/stable/generated/torch.optim.Adam.html#torch.optim.Adam.step "torch.optim.Adam.step")() # Output training stats if i % 50 == 0: print('[%d/%d][%d/%d]\tLoss_D: %.4f\tLoss_G: %.4f\tD(x): %.4f\tD(G(z)): %.4f / %.4f' % (epoch, num_epochs, i, len([dataloader](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader "torch.utils.data.DataLoader")), [errD](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").item(), [errG](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").item(), D_x, D_G_z1, D_G_z2)) # Save Losses for plotting later G_losses.append([errG](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").item()) D_losses.append([errD](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor").item()) # Check how the generator is doing by saving G's output on fixed_noise if (iters % 500 == 0) or ((epoch == num_epochs-1) and (i == len([dataloader](https://docs.pytorch.org/docs/stable/data.html#torch.utils.data.DataLoader "torch.utils.data.DataLoader"))-1)): with [torch.no_grad](https://docs.pytorch.org/docs/stable/generated/torch.no_grad.html#torch.no_grad "torch.no_grad")(): [fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor") = netG([fixed_noise](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor")).detach().cpu() img_list.append([vutils.make_grid](https://docs.pytorch.org/vision/stable/generated/torchvision.utils.make_grid.html#torchvision.utils.make_grid "torchvision.utils.make_grid")([fake](https://docs.pytorch.org/docs/stable/tensors.html#torch.Tensor "torch.Tensor"), padding=2, normalize=True)) iters += 1