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.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
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
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.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
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
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
# 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.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
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 |