How to Implement GAN Training with the GANPyTorch Notebook: A Complete Guide

You can implement GAN training by executing the GANPyTorch notebook in the microsoft/AI-For-Beginners repository, which provides a production-ready DCGAN implementation using separate Generator and Discriminator networks with PyTorch.

The microsoft/AI-For-Beginners repository contains a comprehensive educational resource for learning generative adversarial networks. The GANPyTorch.ipynb notebook located at lessons/4-ComputerVision/10-GANs/GANPyTorch.ipynb demonstrates a complete implementation of the Deep Convolutional GAN (DCGAN) architecture, offering step-by-step guidance for training GANs on image datasets.

Prerequisites and Environment Setup

Before running the notebook, install PyTorch and torchvision. The repository's environment.yml lists these dependencies, or you can install them manually:

conda install pytorch torchvision -c pytorch

# Alternative with pip

pip install torch torchvision

Ensure you have a CUDA-capable GPU for faster training, though the code will run on CPU if necessary.

Loading the Dataset (MNIST or CIFAR-10)

The notebook supports both MNIST (grayscale handwritten digits) and CIFAR-10 (RGB color images). Configure the data pipeline with appropriate transforms:

import torchvision
import torchvision.transforms as transforms

transform = transforms.Compose([
    transforms.Resize(64),
    transforms.CenterCrop(64),
    transforms.ToTensor(),
    transforms.Normalize([0.5], [0.5])  # Use [0.5, 0.5, 0.5] for RGB datasets

])

# Load MNIST example

train_set = torchvision.datasets.MNIST(
    root='./data',
    download=True,
    train=True,
    transform=transform
)
dataloader = torch.utils.data.DataLoader(
    train_set,
    batch_size=128,
    shuffle=True
)

The transform normalizes images to the range [-1, 1], matching the Tanh activation in the generator output.

Defining the Generator and Discriminator Architecture

The notebook implements the DCGAN architecture from Radford et al. (2015) with transposed convolutions for generation and strided convolutions for discrimination.

Generator Network

The Generator class takes a 100-dimensional noise vector and outputs a 64×64 image:

import torch.nn as nn

class Generator(nn.Module):
    def __init__(self, nz=100, ngf=64, nc=1):
        super(Generator, self).__init__()
        self.main = nn.Sequential(
            nn.ConvTranspose2d(nz, ngf * 8, 4, 1, 0, bias=False),
            nn.BatchNorm2d(ngf * 8),
            nn.ReLU(True),

            nn.ConvTranspose2d(ngf * 8, ngf * 4, 4, 2, 1, bias=False),
            nn.BatchNorm2d(ngf * 4),
            nn.ReLU(True),

            nn.ConvTranspose2d(ngf * 4, ngf * 2, 4, 2, 1, bias=False),
            nn.BatchNorm2d(ngf * 2),
            nn.ReLU(True),

            nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 1, bias=False),
            nn.BatchNorm2d(ngf),
            nn.ReLU(True),

            nn.ConvTranspose2d(ngf, nc, 4, 2, 1, bias=False),
            nn.Tanh()
        )

    def forward(self, input):
        return self.main(input)

Set nc=3 when training on CIFAR-10 instead of MNIST.

Discriminator Network

The Discriminator class uses LeakyReLU activations and predicts whether an image is real or fake:

class Discriminator(nn.Module):
    def __init__(self, ndf=64, nc=1):
        super(Discriminator, self).__init__()
        self.main = nn.Sequential(
            nn.Conv2d(nc, ndf, 4, 2, 1, bias=False),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(ndf, ndf * 2, 4, 2, 1, bias=False),
            nn.BatchNorm2d(ndf * 2),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(ndf * 2, ndf * 4, 4, 2, 1, bias=False),
            nn.BatchNorm2d(ndf * 4),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(ndf * 4, ndf * 8, 4, 2, 1, bias=False),
            nn.BatchNorm2d(ndf * 8),
            nn.LeakyReLU(0.2, inplace=True),

            nn.Conv2d(ndf * 8, 1, 4, 1, 0, bias=False),
            nn.Sigmoid()
        )

    def forward(self, input):
        return self.main(input).view(-1)

Weight Initialization Strategy

DCGAN training requires specific weight initialization to prevent mode collapse. The notebook defines a weights_init function that applies a normal distribution with mean 0 and standard deviation 0.02:

def weights_init(m):
    classname = m.__class__.__name__
    if classname.find('Conv') != -1:
        nn.init.normal_(m.weight.data, 0.0, 0.02)
    elif classname.find('BatchNorm') != -1:
        nn.init.normal_(m.weight.data, 1.0, 0.02)
        nn.init.constant_(m.bias.data, 0)

Apply this to both networks before training:

netG = Generator().to(device)
netG.apply(weights_init)

netD = Discriminator().to(device)
netD.apply(weights_init)

Loss Functions and Optimizers

The implementation uses Binary Cross-Entropy Loss (nn.BCELoss) and the Adam optimizer with specific hyperparameters crucial for stability:

criterion = nn.BCELoss()
lr = 0.0002
beta1 = 0.5

optimizerD = torch.optim.Adam(netD.parameters(), lr=lr, betas=(beta1, 0.999))
optimizerG = torch.optim.Adam(netG.parameters(), lr=lr, betas=(beta1, 0.999))

According to the DCGAN paper, using a learning rate of 0.0002 and beta1 of 0.5 (rather than the default 0.9) significantly improves training stability.

The Training Loop Implementation

The training loop follows the standard GAN update procedure with three distinct phases per iteration:

  1. Update Discriminator with real images
  2. Update Discriminator with fake images
  3. Update Generator to fool the discriminator
num_epochs = 25
real_label = 1.
fake_label = 0.

for epoch in range(num_epochs):
    for i, (data, _) in enumerate(dataloader):
        # (a) Train Discriminator on real data

        netD.zero_grad()
        real = data.to(device)
        b_size = real.size(0)
        label = torch.full((b_size,), real_label, dtype=torch.float, device=device)
        output = netD(real)
        errD_real = criterion(output, label)
        errD_real.backward()
        D_x = output.mean().item()

        # (b) Train Discriminator on fake data

        noise = torch.randn(b_size, 100, 1, 1, device=device)
        fake = netG(noise)
        label.fill_(fake_label)
        output = netD(fake.detach())
        errD_fake = criterion(output, label)
        errD_fake.backward()
        D_G_z1 = output.mean().item()
        errD = errD_real + errD_fake
        optimizerD.step()

        # (c) Train Generator

        netG.zero_grad()
        label.fill_(real_label)  # Generator wants D(G(z)) = 1

        output = netD(fake)
        errG = criterion(output, label)
        errG.backward()
        D_G_z2 = output.mean().item()
        optimizerG.step()

        # Logging

        if i % 100 == 0:
            print(f'[{epoch}/{num_epochs}][{i}/{len(dataloader)}] '
                  f'Loss_D: {errD.item():.4f} Loss_G: {errG.item():.4f} '
                  f'D(x): {D_x:.4f} D(G(z)): {D_G_z1:.4f}/{D_G_z2:.4f}')

Using .detach() when feeding fake images to the discriminator prevents gradients from flowing back to the generator during discriminator updates.

Visualizing Generated Images

Monitor training progress by generating a fixed grid of images after each epoch:

import torchvision.utils as vutils
import matplotlib.pyplot as plt

# Create fixed noise vector for consistent comparison

fixed_noise = torch.randn(64, 100, 1, 1, device=device)

# Generate images

with torch.no_grad():
    fake = netG(fixed_noise).cpu()
img_grid = vutils.make_grid(fake, padding=2, normalize=True)

plt.figure(figsize=(8,8))
plt.axis('off')
plt.title(f'Generated Images - Epoch {epoch}')
plt.imshow(img_grid.permute(1, 2, 0))
plt.show()

The vutils.make_grid function arranges the batch into a visual grid, while normalize=True scales the output from [-1, 1] to [0, 1] for display.

Complete Self-Contained Example

Here is a runnable script combining all components from the notebook:

import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
import torchvision.utils as vutils
import matplotlib.pyplot as plt

# Hyperparameters

batch_size = 128
image_size = 64
nz = 100
ngf = 64
ndf = 64
nc = 1
lr = 0.0002
beta1 = 0.5
num_epochs = 25
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

# Data loading

transform = transforms.Compose([
    transforms.Resize(image_size),
    transforms.CenterCrop(image_size),
    transforms.ToTensor(),
    transforms.Normalize([0.5], [0.5])
])
dataset = torchvision.datasets.MNIST(root='./data', train=True, download=True, transform=transform)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=batch_size, shuffle=True)

# Model definitions (as defined above)

# ... Generator and Discriminator classes ...

# ... weights_init function ...

# Initialize

netG = Generator().to(device)
netD = Discriminator().to(device)
netG.apply(weights_init)
netD.apply(weights_init)

criterion = nn.BCELoss()
optimizerD = optim.Adam(netD.parameters(), lr=lr, betas=(beta1, 0.999))
optimizerG = optim.Adam(netG.parameters(), lr=lr, betas=(beta1, 0.999))

# Training loop

real_label = 1.
fake_label = 0.
fixed_noise = torch.randn(64, nz, 1, 1, device=device)

for epoch in range(num_epochs):
    for i, (images, _) in enumerate(dataloader):
        # Update D

        netD.zero_grad()
        real = images.to(device)
        b_size = real.size(0)
        label = torch.full((b_size,), real_label, device=device)
        output = netD(real)
        errD_real = criterion(output, label)
        errD_real.backward()
        
        noise = torch.randn(b_size, nz, 1, 1, device=device)
        fake = netG(noise)
        label.fill_(fake_label)
        output = netD(fake.detach())
        errD_fake = criterion(output, label)
        errD_fake.backward()
        optimizerD.step()
        
        # Update G

        netG.zero_grad()
        label.fill_(real_label)
        output = netD(fake)
        errG = criterion(output, label)
        errG.backward()
        optimizerG.step()
    
    # Visualize

    with torch.no_grad():
        fake = netG(fixed_noise).cpu()
    plt.imshow(vutils.make_grid(fake, padding=2, normalize=True).permute(1,2,0))
    plt.show()

Summary

  • Location: The GANPyTorch notebook resides at lessons/4-ComputerVision/10-GANs/GANPyTorch.ipynb in the microsoft/AI-For-Beginners repository.
  • Architecture: Implements DCGAN with specific Generator (transposed convolutions) and Discriminator (strided convolutions) designs.
  • Initialization: Custom weight initialization using normal distribution (0.0, 0.02) is essential for stable training.
  • Optimization: Use Adam with learning rate 0.0002 and beta1 0.5, paired with BCELoss.
  • Training: Alternate between updating the discriminator on real and fake batches, then update the generator to maximize fooling probability.
  • Monitoring: Visualize progress using torchvision.utils.make_grid to ensure the generator is learning meaningful patterns.

Frequently Asked Questions

What is the GANPyTorch notebook?

The GANPyTorch notebook is an educational Jupyter notebook in the microsoft/AI-For-Beginners repository that demonstrates how to implement a Deep Convolutional Generative Adversarial Network (DCGAN) using PyTorch. It provides complete code for dataset loading, model architecture definition, training loops, and visualization of generated images.

Which dataset should I use when learning to implement GAN training?

Start with MNIST (grayscale, 28×28 digits) because it trains faster and requires less memory than color datasets. Once you understand the training dynamics, switch to CIFAR-10 (32×32 RGB images) to practice with more complex, colorized outputs. Both datasets are automatically downloadable through torchvision.

Why does the DCGAN implementation use beta1=0.5 instead of the default 0.9?

The DCGAN paper specifically recommends beta1=0.5 because the original Adam default (0.9) causes training instability in GANs, often leading to mode collapse or oscillating loss values. The lower beta1 value helps the optimizer maintain more momentum from recent gradients, which is crucial when the discriminator and generator are competing against each other.

How can I tell if my GAN is training correctly?

Monitor three key metrics during training: the discriminator loss on real images should stay around 0.5, the generator loss should gradually decrease, and the D(x) and D(G(z)) values should approach 0.5 (indicating the discriminator cannot distinguish real from fake). If losses diverge to 0 or explode, check your learning rate and ensure you applied the custom weight initialization function.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →