共计 2778 个字符,预计需要花费 7 分钟才能阅读完成。
背景解析:GAN 的核心博弈原理
2014 年 Ian Goodfellow 提出的生成对抗网络 (GAN),本质上是一个让两个神经网络相互对抗的游戏。就像古董鉴定师和造假者之间的博弈:

- 生成器 (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)
这个公式的意思是:
- 判别器试图最大化识别真实数据的能力(第一个期望项)和识别假数据的能力(第二个期望项)
- 生成器试图最小化判别器识别假数据的能力
训练过程中两者的 loss 变化会呈现此消彼长的趋势:
graph LR
A[判别器 loss 下降] --> B[生成器生成质量提升]
B --> C[判别器 loss 上升]
C --> D[判别器调整参数]
D --> A
GAN vs 其他生成模型
与其他生成模型相比,GAN 的特点非常鲜明:
| 特性 | GAN | VAE | 自回归模型 |
|---|---|---|---|
| 生成质量 | ★★★★★ | ★★★☆☆ | ★★★★☆ |
| 训练稳定性 | ★★☆☆☆ | ★★★★★ | ★★★★☆ |
| 多样性 | ★★★☆☆ | ★★★★☆ | ★★★★★ |
| 计算效率 | ★★★★☆ | ★★★☆☆ | ★★☆☆☆ |
| 理论可解释性 | ★★☆☆☆ | ★★★★★ | ★★★★☆ |
关键区别在于:
- VAE 通过编码 - 解码结构学习数据分布,生成结果往往较模糊
- GAN 通过对抗机制直接学习生成策略,能产生更清晰的图像
- 自回归模型(如 PixelCNN)逐个像素生成,计算成本高但可控性强
PyTorch 实战:MNIST 生成
网络结构定义
import torch.nn as nn
# 生成器:将随机噪声转为 28x28 图像
class Generator(nn.Module):
def __init__(self, latent_dim=100):
super().__init__()
self.main = nn.Sequential(nn.Linear(latent_dim, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 1024),
nn.LeakyReLU(0.2),
nn.Linear(1024, 28*28),
nn.Tanh() # 输出归一化到 [-1,1]
)
def forward(self, z):
return self.main(z).view(-1,1,28,28)
# 判别器:判断图像真伪
class Discriminator(nn.Module):
def __init__(self):
super().__init__()
self.main = nn.Sequential(nn.Linear(28*28, 1024),
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(1024, 512),
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Dropout(0.3),
nn.Linear(256, 1),
nn.Sigmoid())
def forward(self, img):
img_flat = img.view(img.size(0), -1)
return self.main(img_flat)
训练循环关键代码
# 超参数设置依据:# - 学习率:GAN 对学习率敏感,通常取较小值 (论文推荐 0.0002)
# - batch_size:较大 batch 能提供更稳定的梯度,但会占用更多显存
opt_G = torch.optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))
opt_D = torch.optim.Adam(discriminator.parameters(), lr=0.0002, betas=(0.5, 0.999))
for epoch in range(epochs):
for real_imgs, _ in dataloader:
# 训练判别器
z = torch.randn(real_imgs.size(0), latent_dim)
fake_imgs = generator(z)
real_loss = criterion(discriminator(real_imgs), real_labels)
fake_loss = criterion(discriminator(fake_imgs.detach()), fake_labels)
loss_D = (real_loss + fake_loss) / 2
opt_D.zero_grad()
loss_D.backward()
opt_D.step()
# 训练生成器
output = discriminator(fake_imgs)
loss_G = criterion(output, real_labels) # 骗过判别器
opt_G.zero_grad()
loss_G.backward()
opt_G.step()
新手避坑指南
问题 1:模式坍塌(Mode Collapse)
- 现象 :生成器只产生几种固定样本,缺乏多样性
- 原因 :生成器找到判别器的弱点后停止探索
- 解决 :
- 改用 Wasserstein GAN(WGAN)
- 在判别器中使用 Mini-batch Discrimination
- 增加噪声输入:
z += torch.randn_like(z)*0.1
问题 2:梯度消失
- 现象 :判别器 loss 降为 0,生成器停止更新
- 原因 :判别器过于强大导致生成器梯度消失
- 解决 :
- 使用 LeakyReLU 替代 ReLU
- 适度降低判别器学习率
- 采用标签平滑:
real_labels = torch.FloatTensor(batch_size).uniform_(0.9, 1.0)
问题 3:训练震荡
- 现象 :loss 剧烈波动无法收敛
- 原因 :生成器与判别器学习速度不匹配
- 解决 :
- 采用 TTUR(Two Time-scale Update Rule)
- 定期保存模型快照
- 使用梯度裁剪:
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
进阶思考
- 生成质量评估 :
- 如何定量评估生成图像的真实性?
-
Inception Score 和 FID 指标各有什么优缺点?
-
条件生成 :
- 如何让 GAN 生成指定类别的数字(如仅生成数字 ”7″)
- 对比 AC-GAN 和 CGAN 两种条件生成方案的差异
经过这次实践,我深刻体会到 GAN 就像在训练两个相互竞争的运动员——需要精心调节他们的训练强度,才能让双方都不断进步。建议初学者多尝试调整网络结构和超参数,观察训练过程中的 loss 变化和生成效果,这是理解 GAN 行为的最佳方式。
正文完
发表至: 未分类
近三天内
