Skip to main content
Create your own

Building and Debugging Basic GANs

Hello! Let's dive into a new and exciting paradigm of generative modeling.

Introduction

In our last two lessons, we explored Variational Autoencoders (VAEs), a class of generative models built on the principles of probabilistic inference. We first implemented a VAE and then derived its objective function, the Evidence Lower Bound (ELBO), from first principles. This approach aims to explicitly model the data's probability distribution by finding a tractable lower bound.

Today, we pivot to a fundamentally different and highly influential framework: Generative Adversarial Networks (GANs). Instead of explicit probability density estimation, GANs learn to generate data through a competitive, game-theoretic process.

Your learning outcome for this lesson is to build a basic Generative Adversarial Network (GAN) and diagnose common training issues like mode collapse. We will cover:

  • The core concept of adversarial training.
  • The two key components: a Generator and a Discriminator.
  • The mathematical formulation of the GAN as a minimax game.
  • A from-scratch implementation of a GAN in PyTorch to generate handwritten digits.
  • How to identify and understand notorious training problems like mode collapse.

This lesson will blend theory with hands-on practice, leveraging your software engineering background to bring these concepts to life.

The GAN Framework: An Adversarial Game

The core idea of a GAN, introduced by Ian Goodfellow et al. in 2014, is brilliantly simple. It pits two neural networks against each other in a zero-sum game.

  1. The Generator (G): Its goal is to create fake data that is indistinguishable from real data. It takes a random noise vector (from a latent space, similar to VAEs) as input and outputs a synthetic data sample, like an image. Think of it as a master art forger.
  2. The Discriminator (D): Its goal is to correctly identify whether a given data sample is real (from the training dataset) or fake (from the generator). It's a binary classifier, acting as an expert art detective.

The training process unfolds as follows:

  • The Generator produces a batch of fake images.
  • The Discriminator is trained on a mix of real images and these fake images. It learns to get better at telling them apart.
  • The Generator is then trained based on the Discriminator's feedback. Its goal is to produce images that the Discriminator misclassifies as real. It gets better at fooling the detective.

This cycle repeats. The Generator's forgeries become more and more convincing, forcing the Discriminator to learn ever finer details to spot the fakes. The ideal outcome, or equilibrium, is when the Generator produces such perfect fakes that the Discriminator can only guess with 50% accuracy. At this point, the Generator has learned to capture the underlying distribution of the real data.

The following video provides an excellent introduction to this concept, using the classic counterfeiter-police analogy.

Understanding GANs (Generative Adversarial Networks)

This video from DeepBean introduces the fundamental concept of GANs and how the generator and discriminator interact in an adversarial framework.

Watch the video from the beginning up to 07:11. This will cover: The counterfeiter-police analogy (00:00 - 01:17). How GANs sidestep the problem of intractable probability distributions that VAEs tackle with the ELBO (01:17 - 05:50). The architecture and roles of the Generator and Discriminator networks (05:50 - 07:11).

The Mathematics of the Minimax Game

This adversarial process is formally described as a minimax game. The Generator () and Discriminator () are competing to optimize a single value function, . The Discriminator tries to maximize it, while the Generator tries to minimize it.

From the original GAN paper, the objective function is:

Let's break this down:

  • : The Discriminator's performance on real data. is the probability that is real. The Discriminator wants to maximize this term (push towards 1 for real data).
  • : The performance on fake data. is a fake image. is the Discriminator's prediction that the fake image is real.
    • The Discriminator wants to make small (close to 0), which maximizes .
    • The Generator wants to make large (close to 1), which minimizes this term.

This minimax formulation elegantly captures the competitive dynamic. To understand how this game leads to the generator learning the true data distribution, the next part of the video provides a clear mathematical derivation.

Understanding GANs (Generative Adversarial Networks)

Continuing with the same video, let's explore the mathematical underpinnings of the GAN objective.

Watch from 07:11 to 17:37. This section explains: The derivation of the loss functions from a binary cross-entropy perspective (07:11 - 13:07). The proof that at the optimal equilibrium, the generator's distribution matches the data distribution, and the objective minimizes the Jensen-Shannon Divergence between them (13:07 - 17:37).

The key theoretical result is that training a GAN is equivalent to minimizing the Jensen-Shannon (JS) divergence between the real data distribution and the generator's distribution . The JS divergence is a way of measuring the similarity between two probability distributions. When it is zero, the distributions are identical.

Building a Basic GAN from Scratch

Now for the practical part. We'll build a GAN to generate images of handwritten digits from the MNIST dataset. The following video provides a complete walkthrough using PyTorch and PyTorch Lightning, which simplifies much of the training boilerplate code.

Since you're fluent in Python, you can focus on the core architectural components and the unique training logic of GANs.

Building a GAN From Scratch With PyTorch | Theory + Implementation

This tutorial by AssemblyAI will guide us through implementing a GAN from scratch. We will focus on the key implementation details.

You can follow along with this video to build the GAN. Pay special attention to the following segments: Overview (00:00 - 03:32): A quick recap of the GAN concept. Network Definitions (06:42 - 10:07): Understand the structure of the Discriminator (a standard CNN for classification) and the Generator (using ConvTranspose2d for upsampling). The training_step (17:46 - 28:27): This is the most critical part. Observe how the training is split into two parts using optimizer_idx: Generator training (optimizer_idx == 0): The generator creates fake images, which are passed to the discriminator. The generator's loss is calculated based on how well it fools the discriminator (using real labels as the target). Discriminator training (optimizer_idx == 1): The discriminator's loss is the sum of its performance on real images (should be close to 1) and fake images (should be close to 0). Note the use of .detach() on the fake images. This is crucial as it prevents gradients from flowing back into the generator during the discriminator's update step.

Test your understanding!

In the training_step for the discriminator, the fake images generated by self(z) are passed as fake.detach(). Why is the .detach() call necessary here? What would happen if it were omitted?

Show answer

The .detach() method creates a new tensor that shares the same data but is detached from the current computation graph. This means that no gradients will be backpropagated through it.

This is necessary because when we train the discriminator, we only want to update the discriminator's weights. The fake images are treated as fixed inputs for this step. If we didn't use .detach(), the gradients from the discriminator's loss calculation on fake images would flow all the way back through the generator network. This would effectively be training the generator with the discriminator's objective, which is the opposite of what we want and would destabilize the entire training process.

Diagnosing GAN Training Issues

As the PyTorch DCGAN tutorial (LINK) notes, "training GANs is somewhat of an art form." Unlike standard deep learning models with a clear, monotonically decreasing loss, GANs involve a delicate equilibrium that can easily be disrupted. This leads to several common failure modes.

The learning outcome specifically asks you to be able to diagnose these issues, with a focus on the most famous one: mode collapse.

Common Failure Modes

  • Vanishing Gradients: If the discriminator becomes too powerful too quickly, its loss on fake images will be near-perfect, and the gradients passed back to the generator can become vanishingly small. The generator stops learning because it gets no useful feedback.
  • Non-Convergence: The generator and discriminator losses may oscillate wildly and never reach a stable point (the Nash equilibrium). The models are essentially undoing each other's progress without making any net improvement.
  • Mode Collapse: The generator discovers one or a few "modes" (types of output) that can reliably fool the discriminator. It then stops exploring and produces only these limited outputs, failing to capture the full diversity of the training data. For MNIST, this might mean the GAN only ever generates the digit '1', because it found that particular '1' to be a very safe bet.
Mode Collapse vs. Stable GAN Training
This diagram illustrates the difference between a healthy GAN and one suffering from mode collapse. On the left, the generator's distribution (green) learns to match the multi-modal data distribution (blue). On the right, the generator has collapsed to a single mode (red), ignoring the others.
GAN Mode Collapse Demonstration
A practical example of mode collapse on MNIST. The generator initially produces noise, then learns to produce the digit '1', and later shifts to only producing '9's, never learning to generate the full set of digits.

How to Identify Failure Modes

The article GANs Failure Modes: How to Identify and Monitor Them provides an excellent, practical guide.

GANs Failure Modes: How to Identify and Monitor Them

This article from Neptune.ai is a practical guide to spotting when your GAN training is going wrong. It uses the same MNIST example we just built.

Read the sections 'GAN failure modes' and 'Evaluating failure modes'. Focus on the two key methods for diagnosis: Looking at intermediate images: Are the generated images improving in quality and diversity over time? Or are they becoming repetitive? Observing loss graphs: The article shows examples of healthy vs. unhealthy loss graphs. A healthy graph shows fluctuation, but the losses don't systematically go to zero or explode. If D_loss drops to zero, the discriminator is winning too easily (vanishing gradients for G). If G_loss drops to zero while D_loss skyrockets, the generator is winning too easily (the discriminator is useless).

A few key rules of thumb for stable training, often formalized in architectures like DCGANs which we'll see next, include:

  • Using a lower learning rate (e.g., 0.0002).
  • Using the Adam optimizer with a lower momentum term for beta1 (e.g., 0.5).
  • Normalizing inputs to the range [-1, 1] and using Tanh as the final activation in the generator.
  • Using LeakyReLU in the discriminator to prevent sparse gradients.
  • Applying Batch Normalization in both the generator and discriminator.

Conclusion

In this lesson, we transitioned from the probabilistic world of VAEs to the game-theoretic framework of GANs. You've seen how a competitive dynamic between two networks can be a powerful mechanism for learning to generate realistic data.

Key Takeaways:

  • GANs consist of a Generator that creates fakes and a Discriminator that spots them.
  • Training is a minimax game where the discriminator tries to maximize a value function, and the generator tries to minimize it.
  • The ideal equilibrium is a Nash Equilibrium, where the generator has learned the true data distribution and the discriminator is forced to guess randomly.
  • The training_step in a GAN requires careful, alternating optimization of the two networks, using .detach() to prevent unwanted gradient flow.
  • GANs are notoriously unstable. Key failure modes include vanishing gradients, non-convergence, and mode collapse.
  • You can diagnose these issues by observing the diversity of generated images over time and analyzing the behavior of the generator and discriminator loss curves.

Preview of the next lesson:
The basic GAN we built today is powerful but can be unstable. In our next lesson, we will implement a Deep Convolutional GAN (DCGAN). The DCGAN paper introduced a set of key architectural guidelines (like using strided convolutions, batch normalization, and specific activations) that dramatically improved the stability and quality of GANs, paving the way for many of the advanced models that followed.

Can't find a good explanation? Sign up and we'll make it for you

Sign up