共计 2811 个字符,预计需要花费 8 分钟才能阅读完成。
生成对抗网络(Generative Adversarial Networks,GAN)在图像生成领域展现了惊人的能力,从风格迁移到超分辨率重建,再到如今火热的 AI 绘画,背后都有 GAN 的身影。对于初学者来说,理解 GAN 的核心思想并不复杂——它通过让两个神经网络相互博弈来学习数据分布。今天我们就用 PyTorch 来实现一个最简单的 GAN 模型,生成 MNIST 手写数字。

一、GAN 的核心原理
GAN 由生成器(Generator)和判别器(Discriminator)组成。生成器负责生成假数据,判别器则要判断输入是真实数据还是生成器造的假数据。两者就像警察和小偷,在对抗中不断进步。
- 对抗训练机制 :
- 生成器 G 试图生成足以乱真的假数据
- 判别器 D 努力区分真假数据
-
两者通过反向传播交替优化
-
数学表达 :训练过程可以表示为最小最大博弈:
$$\min_G \max_D V(D,G) = \mathbb{E}{x\sim p[\log(1-D(G(z)))]$$}(x)}[\log D(x)] + \mathbb{E}_{z\sim p_z(z) -
损失函数 :实际使用时会对生成器采用改进的交叉熵损失:
- 原始公式可能导致梯度消失
- 常见变体是将生成器目标改为最大化 D(G(z)) 而非最小化 1 -D(G(z))
二、PyTorch 代码实现
以下是一个完整的 MNIST 生成示例(Python 3.8+,PyTorch 1.10+):
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# 1. 数据准备
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,)) # MNIST 单通道
])
train_set = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = DataLoader(train_set, batch_size=64, shuffle=True)
# 2. 网络定义
class Generator(nn.Module):
def __init__(self):
super().__init__()
self.main = nn.Sequential(nn.Linear(100, 256), # 输入噪声维度 100
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 784), # MNIST 28x28=784
nn.Tanh() # 输出归一化到 [-1,1]
)
def forward(self, x):
return self.main(x)
class Discriminator(nn.Module):
def __init__(self):
super().__init__()
self.main = nn.Sequential(nn.Linear(784, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 1),
nn.Sigmoid() # 输出概率值)
def forward(self, x):
x = x.view(x.size(0), -1) # 展平
return self.main(x)
# 3. 训练循环
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
G = Generator().to(device)
D = Discriminator().to(device)
criterion = nn.BCELoss()
optim_G = optim.Adam(G.parameters(), lr=0.0002, betas=(0.5, 0.999))
optim_D = optim.Adam(D.parameters(), lr=0.0002, betas=(0.5, 0.999))
for epoch in range(50):
for i, (real_imgs, _) in enumerate(train_loader):
# 真实数据
real_imgs = real_imgs.to(device)
real_labels = torch.ones(real_imgs.size(0), 1).to(device)
# 生成假数据
noise = torch.randn(real_imgs.size(0), 100).to(device)
fake_imgs = G(noise)
fake_labels = torch.zeros(real_imgs.size(0), 1).to(device)
# 训练判别器
optim_D.zero_grad()
real_loss = criterion(D(real_imgs), real_labels)
fake_loss = criterion(D(fake_imgs.detach()), fake_labels)
d_loss = real_loss + fake_loss
d_loss.backward()
optim_D.step()
# 训练生成器
optim_G.zero_grad()
g_loss = criterion(D(fake_imgs), real_labels) # 让生成的图片被判别为真
g_loss.backward()
optim_G.step()
关键点说明:
– 第 18、32 行:使用 LeakyReLU 防止梯度消失
– 第 45 行:detach() 切断生成器梯度传播
– 第 53 行:生成器目标是让 D(G(z)) 接近 1
三、实战避坑指南
- 模式崩溃(Mode Collapse):
- 现象:生成器只产出有限几种样本
-
解决方案:
- 增加噪声维度(如从 100 调整到 256)
- 尝试 Wasserstein GAN 等改进结构
-
参数调整 :
- 学习率通常在 0.0001-0.0005 之间
-
批大小建议 64/128,过大可能导致训练不稳定
-
可视化监控 :
- 每 100 次迭代保存生成样本
- 使用 TensorBoard 记录损失曲线
- 观察判别器准确率应保持在 50-60% 左右
四、延伸思考
- 判别器过强的表现是训练早期准确率就接近 100%,这时可以:
- 降低判别器学习率
- 减少判别器层数
-
添加 Dropout 层
-
DCGAN 通过以下改进提升生成质量:
- 使用转置卷积替代全连接
- 引入 Batch Normalization
- 移除池化层改用步长卷积
通过这个基础实现,你应该已经感受到了 GAN 的奇妙之处。建议尝试修改网络结构或调整超参数,观察对生成结果的影响。当你能稳定生成清晰的手写数字时,就迈出了掌握生成模型的第一步!
