PyTorch Generative Adversarial Network (GAN)

Generative Adversarial Network (GAN)is one of the most creative model architectures in deep learning. By making two neural networks compete against and learn from each other, it can ultimately generate very realistic data. GANs are widely used in scenarios such as image generation, style transfer, and data augmentation.


1. Core Principles of GAN

The core idea of GAN comes from the "zero-sum game" in game theory. It consists of two competing networks:

  • Generator: Learns to generate fake data, with the goal of making the discriminator unable to distinguish generated data from real data
  • Discriminator: Learns to distinguish real data from generated data, with the goal of making judgments as accurate as possible

The two compete against each other during training and continuously improve, ultimately reaching a Nash equilibrium state.

1.1 Objective Function of GAN

The training objective of GAN can be expressed as the following minimax game:

\[ \min_G \max_D \mathbb{E}_{x \sim p_{data}(x)}[\log D(x)] + \mathbb{E}_{z \sim p_z(z)}[\log(1 - D(G(z)))] \]

where:

  • \(G\) denotes the generator network
  • \(D\) denotes the discriminator network
  • \(x\) denotes real data
  • \(z\) denotes the random noise vector (usually following a standard normal distribution)
  • \(G(z)\) denotes the fake data generated by the generator from noise

1.2 Understanding the Training Process

GAN training is divided into two stages:

Stage 1: Train the discriminator

Freeze the generator and improve the discriminator's discrimination ability:

\[ \max_D \mathbb{E}_{x \sim p_{data}}[\log D(x)] + \mathbb{E}_{z \sim p_z}[\log(1 - D(G(z)))] \]

Stage 2: Train the generator

Freeze the discriminator and improve the generator's deception ability:

\[ \min_G \mathbb{E}_{z \sim p_z}[\log(1 - D(G(z)))] \]

In actual training, the discriminator is usually trained for k steps first, then the generator for 1 step, to maintain balance.


2. Basic GAN Implementation

Below is a minimal GAN implementation — used to generate two-dimensional data points.

2.1 Define the Generator and Discriminator

Example

import torch
import torch.nn as nn
import torch.optim as optim
import matplotlib.pyplot as plt

# Set random seed
torch.manual_seed(42)

# ── Generator network ──────────────────────────────────────
class Generator(nn.Module):
    """
Generator: Generate data from random noise
Input: noise vector (batch_size, noise_dim)
Output: generated data (batch_size, data_dim)
    """

    def __init__(self, noise_dim, data_dim, hidden_dim=64):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(noise_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, hidden_dim),
            nn.ReLU(),
            nn.Linear(hidden_dim, data_dim),
            # Output is not activated; GAN will learn the appropriate distribution
        )

    def forward(self, x):
        return self.net(x)


# ── Discriminator network ──────────────────────────────────────
class Discriminator(nn.Module):
    """
Discriminator: Distinguish real data from generated data
Input: data points (batch_size, data_dim)
Output: probability of real data (batch_size, 1)
    """

    def __init__(self, data_dim, hidden_dim=64):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(data_dim, hidden_dim),
            nn.LeakyReLU(0.2),  # LeakyReLU prevents vanishing gradients
            nn.Linear(hidden_dim, hidden_dim),
            nn.LeakyReLU(0.2),
            nn.Linear(hidden_dim, 1),
            nn.Sigmoid()  # Output probability
        )

    def forward(self, x):
        return self.net(x)


# Hyperparameters
NOISE_DIM = 16
DATA_DIM = 2
HIDDEN_DIM = 64
BATCH_SIZE = 128

# Create networks
generator = Generator(NOISE_DIM, DATA_DIM, HIDDEN_DIM)
discriminator = Discriminator(DATA_DIM, HIDDEN_DIM)

print(f"Generator parameter count: {sum(p.numel() for p in generator.parameters()):,}")
print(f"Discriminator parameter count: {sum(p.numel() for p in discriminator.parameters()):,}")

2.2 Training Loop

Example

# ── Optimizer ──────────────────────────────────────
lr = 0.001
g_optimizer = optim.Adam(generator.parameters(), lr=lr)
d_optimizer = optim.Adam(discriminator.parameters(), lr=lr)

# Loss function: binary cross-entropy
criterion = nn.BCELoss()

# ── Training data: ring distribution ──────────────────────────
def generate_real_data(batch_size):
    """Generate real data with a ring distribution"""
    angles = torch.rand(batch_size) * 2 * torch.pi
    radius = 1.0 + torch.randn(batch_size) * 0.1  # Radius is approximately 1
    x = radius * torch.cos(angles)
    y = radius * torch.sin(angles)
    return torch.stack([x, y], dim=1)


# ── Training loop ──────────────────────────────────────
NUM_EPOCHS = 1000
d_losses = []
g_losses = []

for epoch in range(NUM_EPOCHS):
    # 1. Train the discriminator
    # Generate fake data
    noise = torch.randn(BATCH_SIZE, NOISE_DIM)
    fake_data = generator(noise).detach()  # detach to avoid computing generator gradients

    # Generate real data
    real_data = generate_real_data(BATCH_SIZE)

    # Discriminator loss
    real_pred = discriminator(real_data)
    fake_pred = discriminator(fake_data)
    d_loss = criterion(real_pred, torch.ones_like(real_pred)) + \
             criterion(fake_pred, torch.zeros_like(fake_pred))

    # Update discriminator
    d_optimizer.zero_grad()
    d_loss.backward()
    d_optimizer.step()

    # 2. Train the generator
    # Generate a new batch of fake data
    noise = torch.randn(BATCH_SIZE, NOISE_DIM)
    fake_data = generator(noise)

    # Generator loss: make the discriminator believe the generated data is real
    fake_pred = discriminator(fake_data)
    g_loss = criterion(fake_pred, torch.ones_like(fake_pred))

    # Update generator
    g_optimizer.zero_grad()
    g_loss.backward()
    g_optimizer.step()

    # Record losses
    d_losses.append(d_loss.item())
    g_losses.append(g_loss.item())

    if (epoch + 1) % 100 == 0:
        print(f"Epoch {epoch+1:4d} | D_loss: {d_loss:.4f} | G_loss: {g_loss:.4f}")

print("Training complete!")

2.3 Visualize Generated Results

Example

# Generate data and visualize
def visualize_results(generator, num_samples=1000):
    noise = torch.randn(num_samples, NOISE_DIM)
    generated_data = generator(noise).detach().numpy()

    plt.figure(figsize=(6, 6))
    plt.scatter(generated_data[:, 0], generated_data[:, 1],
                alpha=0.5, s=10, c='blue', label='Generated')
    plt.xlim(-2, 2)
    plt.ylim(-2, 2)
    plt.xlabel('x')
    plt.ylabel('y')
    plt.title('GAN Generated Data')
    plt.legend()
    plt.grid(True, alpha=0.3)
    plt.show()


# View generation results
visualize_results(generator)

3. DCGAN - Deep Convolutional GAN

DCGAN is a classic architecture that introduces convolutional neural networks into GAN, greatly improving image generation quality.

3.1 Key Points of DCGAN Architecture

  • Use transposed convolution for upsampling to generate images
  • Use strided convolution for downsampling to discriminate images
  • Use BatchNorm in both generator and discriminator (but not in the output layer or input layer)
  • Generator uses ReLU, discriminator uses LeakyReLU

3.2 DCGAN Implementation

Example

import torch
import torch.nn as nn

# ── DCGAN Generator ─────────────────────────────────
class DCGenerator(nn.Module):
    """
DCGAN Generator: upsample using transposed convolution
    """

    def __init__(self, noise_dim=100, channels=3, features_g=64):
        super().__init__()
        self.noise_dim = noise_dim

        # Input: noise_dim x 1 x 1
        self.net = nn.Sequential(
            # Transposed convolution: (batch, features_g*8, 4, 4)
            nn.ConvTranspose2d(noise_dim, features_g * 8, 4, 1, 0, bias=False),
            nn.BatchNorm2d(features_g * 8),
            nn.ReLU(True),

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

            # (batch, features_g*2, 16, 16)
            nn.ConvTranspose2d(features_g * 4, features_g * 2, 4, 2, 1, bias=False),
            nn.BatchNorm2d(features_g * 2),
            nn.ReLU(True),

            # (batch, features_g, 32, 32)
            nn.ConvTranspose2d(features_g * 2, features_g, 4, 2, 1, bias=False),
            nn.BatchNorm2d(features_g),
            nn.ReLU(True),

            # Output: (batch, channels, 64, 64)
            nn.ConvTranspose2d(features_g, channels, 4, 2, 1, bias=False),
            nn.Tanh()  # Output range [-1, 1]
        )

    def forward(self, x):
        # x: (batch, noise_dim) -> (batch, noise_dim, 1, 1)
        x = x.view(x.size(0), x.size(1), 1, 1)
        return self.net(x)


# ── DCGAN Discriminator ─────────────────────────────────
class DCDiscriminator(nn.Module):
    """
DCGAN Discriminator: downsample using convolution
    """

    def __init__(self, channels=3, features_d=64):
        super().__init__()

        # Input: (batch, channels, 64, 64)
        self.net = nn.Sequential(
            # (batch, features_d, 32, 32)
            nn.Conv2d(channels, features_d, 4, 2, 1, bias=False),
            nn.LeakyReLU(0.2, inplace=True),

            # (batch, features_d*2, 16, 16)
            nn.Conv2d(features_d, features_d * 2, 4, 2, 1, bias=False),
            nn.BatchNorm2d(features_d * 2),
            nn.LeakyReLU(0.2, inplace=True),

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

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

            # Output: (batch, 1, 1, 1)
            nn.Conv2d(features_d * 8, 1, 4, 1, 0, bias=False),
            nn.Sigmoid()
        )

    def forward(self, x):
        return self.net(x).view(x.size(0), -1)


# Test network
noise_dim = 100
generator = DCGenerator(noise_dim=noise_dim, channels=3, features_g=64)
discriminator = DCDiscriminator(channels=3, features_d=64)

# Test forward pass
noise = torch.randn(2, noise_dim)
fake_images = generator(noise)
print(f"Generated image shape: {fake_images.shape}")  # torch.Size([2, 3, 64, 64])

decision = discriminator(fake_images)
print(f"Discrimination result shape: {decision.shape}")     # torch.Size([2, 1])

3.3 Complete DCGAN Training Code

Example

# ── Training configuration ─────────────────────────────────────
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print(f"Using device: {device}")

NOISE_DIM = 100
LEARNING_RATE = 0.0002
BETA1 = 0.5  # Adam parameters

# Create networks
generator = DCGenerator(noise_dim=NOISE_DIM).to(device)
discriminator = DCDiscriminator().to(device)

# Optimizer
g_optimizer = optim.Adam(generator.parameters(), lr=LEARNING_RATE, betas=(BETA1, 0.999))
d_optimizer = optim.Adam(discriminator.parameters(), lr=LEARNING_RATE, betas=(BETA1, 0.999))

criterion = nn.BCELoss()

# ── Training loop ─────────────────────────────────────
fixed_noise = torch.randn(64, NOISE_DIM, device=device)  # For visualization

def train_dcgan(generator, discriminator, g_optimizer, d_optimizer, criterion,
                num_epochs, device, fixed_noise):
    G_losses = []
    D_losses = []

    for epoch in range(num_epochs):
        for batch_idx in range(100):  # Assume each epoch has 100 batches
            # Train the discriminator
            discriminator.zero_grad()

            # Real images (assumed to be available)
            # real_images = ...
            # Use random noise to simulate here
            real_images = torch.randn(32, 3, 64, 64).to(device)

            batch_size = real_images.size(0)
            labels_real = torch.ones(batch_size, 1).to(device)
            labels_fake = torch.zeros(batch_size, 1).to(device)

            # Real image loss
            output = discriminator(real_images)
            d_loss_real = criterion(output, labels_real)

            # Generated image loss
            noise = torch.randn(batch_size, NOISE_DIM).to(device)
            fake_images = generator(noise)
            output = discriminator(fake_images.detach())
            d_loss_fake = criterion(output, labels_fake)

            # Total loss
            d_loss = d_loss_real + d_loss_fake
            d_loss.backward()
            d_optimizer.step()

            # Train the generator
            generator.zero_grad()

            noise = torch.randn(batch_size, NOISE_DIM).to(device)
            fake_images = generator(noise)
            output = discriminator(fake_images)
            g_loss = criterion(output, labels_real)  # Hope the generated images are judged as real

            g_loss.backward()
            g_optimizer.step()

            # Record the loss
            if batch_idx % 50 == 0:
                G_losses.append(g_loss.item())
                D_losses.append(d_loss.item())
                print(f"[{epoch}/{num_epochs}][{batch_idx}/100] "
                      f"D_loss: {d_loss:.4f} | G_loss: {g_loss:.4f}")

    return G_losses, D_losses


# Start training
# G_losses, D_losses = train_dcgan(generator, discriminator, g_optimizer,
#                                   d_optimizer, criterion, 5, device, fixed_noise)
print("DCGAN architecture has been defined, training can begin!")

4. GAN Training Tips

4.1 Common Problems and Solutions

Problem Cause Solution
Mode Collapse The generator only produces a limited variety of samples Use WGAN, add minibatch discrimination, use label smoothing
Discriminator too strong Generator gradients vanish, making it unable to learn Train the generator multiple times, reduce the discriminator learning rate, use LeakyReLU
Unstable training GAN objective function is non-convex and difficult to converge Use spectral normalization, gradient penalty, learning rate warmup
Poor generation quality Insufficient network capacity or insufficient training Increase network depth, use more training data, train longer

4.2 Loss Function Improvements

The original GAN uses JS divergence, which has vanishing gradient problems. WGAN uses Wasserstein distance, which is more stable:

Example

# WGAN loss function (replaces BCE)
def wgan_d_loss(real_pred, fake_pred):
    """Discriminator loss: real samples score high, generated samples score low"""
    return -(torch.mean(real_pred) - torch.mean(fake_pred))

def wgan_g_loss(fake_pred):
    """Generator loss: make generated samples score high"""
    return -torch.mean(fake_pred)

# Gradient Penalty - WGAN-GP
def gradient_penalty(discriminator, real_images, fake_images, device):
    """WGAN-GP gradient penalty term"""
    batch_size = real_images.size(0)
    alpha = torch.rand(batch_size, 1, 1, 1).to(device)

    # Interpolate between real and generated images
    interpolated = alpha * real_images + (1 - alpha) * fake_images
    interpolated.requires_grad = True

    # Compute the discriminator output for interpolated images
    pred = discriminator(interpolated)

    # Compute the gradient
    gradients = torch.autograd.grad(
        outputs=pred,
        inputs=interpolated,
        grad_outputs=torch.ones_like(pred),
        create_graph=True,
        retain_graph=True,
        only_inputs=True
    )[0]

    # Compute the gradient norm
    gradients = gradients.view(batch_size, -1)
    gradient_norm = gradients.norm(2, dim=1)
    penalty = ((gradient_norm - 1) ** 2).mean()

    return penalty

4.3 Spectral Normalization

Spectral normalization can stabilize GAN training by controlling the Lipschitz constant of the discriminator:

Example

import torch.nn.utils.spectral_norm as spectral_norm

# Discriminator using spectral normalization
class SNDiscriminator(nn.Module):
    def __init__(self, channels=3, features_d=64):
        super().__init__()
        self.net = nn.Sequential(
            spectral_norm(nn.Conv2d(channels, features_d, 4, 2, 1)),
            nn.LeakyReLU(0.2, inplace=True),

            spectral_norm(nn.Conv2d(features_d, features_d * 2, 4, 2, 1)),
            nn.LeakyReLU(0.2, inplace=True),

            spectral_norm(nn.Conv2d(features_d * 2, features_d * 4, 4, 2, 1)),
            nn.LeakyReLU(0.2, inplace=True),

            spectral_norm(nn.Conv2d(features_d * 4, 1, 4, 1, 0)),
            nn.Sigmoid()
        )

    def forward(self, x):
        return self.net(x).view(x.size(0), -1)

5. Conditional GAN (cGAN)

Conditional GAN allows specifying class labels for generated data, enabling conditional generation.

5.1 cGAN Architecture

Example

import torch
import torch.nn as nn

class ConditionalGenerator(nn.Module):
    """Conditional generator: receives both noise and class labels"""
    def __init__(self, noise_dim, num_classes, embed_dim, img_channels, features_g=64):
        super().__init__()
        self.label_emb = nn.Embedding(num_classes, embed_dim)

        # Concatenate noise and label embeddings
        self.net = nn.Sequential(
            nn.Linear(noise_dim + embed_dim, features_g * 8 * 4 * 4),
            nn.BatchNorm1d(features_g * 8 * 4 * 4),
            nn.ReLU(),
            # Then apply transposed convolutions (similar to DCGAN)
            # ...
        )

    def forward(self, noise, labels):
        # Embed class labels to the same dimension as noise
        label_embedding = self.label_emb(labels)
        # Concatenate noise and label embeddings
        x = torch.cat([noise, label_embedding], dim=1)
        return self.net(x)


class ConditionalDiscriminator(nn.Module):
    """Conditional discriminator: receives both images and class labels"""
    def __init__(self, img_channels, num_classes, embed_dim, features_d=64):
        super().__init__()
        self.label_emb = nn.Embedding(num_classes, embed_dim)

        # Concatenate image and label embeddings
        self.net = nn.Sequential(
            nn.Conv2d(img_channels + embed_dim, features_d, 4, 2, 1),
            nn.LeakyReLU(0.2),
            # ...
        )

    def forward(self, img, labels):
        # Reshape label embeddings to the same spatial size as images
        label_embedding = self.label_emb(labels)
        # Adjust shape for concatenation
        label_embedding = label_embedding.unsqueeze(2).unsqueeze(3)
        label_embedding = label_embedding.expand(-1, -1, img.size(2), img.size(3))
        # Concatenate image and label
        x = torch.cat([img, label_embedding], dim=1)
        return self.net(x)

6. GAN Evaluation Metrics

6.1 Common Evaluation Metrics

Metric Description Advantages Disadvantages
Inception Score (IS) Use Inception v3 to evaluate the quality and diversity of generated images Simple to compute, has some correlation with human judgment Does not evaluate overfitting and cannot detect mode collapse
Fréchet Inception Distance (FID) Computes the distance between real and generated images in feature space More sensitive to noise and mode collapse Requires a large number of samples and is slow to compute
Human evaluation Humans judge the quality of generated images Most accurately reflects human perception Subjective and time-consuming

6.2 FID Calculation Implementation

Example

import numpy as np
from scipy import linalg

def calculate_fid(real_activations, fake_activations):
    """
Compute Fréchet Inception Distance
real_activations: feature vectors of real images (N, dim)
fake_activations: feature vectors of generated images (N, dim)
    """

    # Compute the mean and covariance
    mu1, sigma1 = real_activations.mean(axis=0), np.cov(real_activations, rowvar=False)
    mu2, sigma2 = fake_activations.mean(axis=0), np.cov(fake_activations, rowvar=False)

    # Compute FID
    diff = mu1 - mu2
    # Calculate the sum of the covariance matrices
    covmean, _ = linalg.sqrtm(sigma1.dot(sigma2), disp=False)

    # Avoid complex numbers
    if np.iscomplexobj(covmean):
        covmean = covmean.real

    fid = diff.dot(diff) + np.trace(sigma1 + sigma2 - 2 * covmean)
    return fid


# Simplified example: use random data
np.random.seed(42)
real_acts = np.random.randn(1000, 2048)  # Inception v3 output dimension
fake_acts = np.random.randn(1000, 2048)

fid_score = calculate_fid(real_acts, fake_acts)
print(f"FID Score: {fid_score:.2f}")
# Lower FID is better; usually below 50 indicates good generation quality

7. Common GAN Variants

Since its development, GAN has produced many variants suitable for different application scenarios:

Model Full Name Features Applicable Scenarios
DCGAN Deep Convolutional GAN Uses convolutional networks, generates high-quality images Image generation
WGAN Wasserstein GAN Uses Wasserstein distance, more stable training Stable training
WGAN-GP WGAN with Gradient Penalty Gradient penalty replaces weight clipping Stable training
CGAN Conditional GAN Adds conditional information, controllable generation Conditional generation
CycleGAN Cycle-Consistent GAN Unsupervised image-to-image translation Style transfer
StyleGAN Style-Based GAN Style control, high-quality face generation Face generation
BigGAN Big GAN Large-scale, high-quality image generation High-resolution images
ProGAN Progressive Growing GAN Progressively increasing resolution High-resolution generation
Other extensions