共计 2257 个字符,预计需要花费 6 分钟才能阅读完成。
1. 背景与痛点分析
生成对抗网络(GAN)近年来在图像生成领域大放异彩,但实际训练过程却充满挑战。许多开发者第一次尝试训练 GAN 时,往往会遇到模型根本不收敛的情况。这主要是因为 GAN 的训练过程本质上是一个动态博弈过程,生成器(Generator)和判别器(Discriminator)需要保持微妙的平衡。

- 常见问题:
- 梯度消失:当判别器太强时,生成器无法获得有效梯度
- 模式崩溃(Mode Collapse):生成器只学会生成有限的几种样本
-
训练不稳定:损失函数震荡剧烈,难以监控训练进度
-
与传统生成模型对比:
- VAE 生成的图像往往比较模糊
- GAN 可以产生更锐利、细节更丰富的图像
- 但 GAN 的训练难度远高于 VAE
2. DCGAN 实战实现
深度卷积 GAN(DCGAN)是 GAN 的一种改进架构,它使用卷积神经网络(CNN)作为生成器和判别器的基础组件。
网络架构设计
- 生成器结构:
- 输入:100 维的随机噪声向量
- 通过转置卷积层(Transposed Convolution)逐步上采样
-
最终输出 64×64 的 RGB 图像
-
判别器结构:
- 输入:64×64 的 RGB 图像
- 通过卷积层逐步下采样
- 最终输出一个标量,表示图像为真的概率
# 生成器核心代码示例
class Generator(nn.Module):
def __init__(self):
super().__init__()
self.main = nn.Sequential(
# 输入是 Z, 进入全连接
nn.ConvTranspose2d(100, 512, 4, 1, 0, bias=False),
nn.BatchNorm2d(512), # 批归一化加速训练
nn.ReLU(True),
# 逐步上采样...
)
3. 关键优化技巧
Wasserstein Loss 改进
传统 GAN 使用 JS 散度作为损失函数,容易导致梯度消失。Wasserstein Loss 通过以下改进解决了这个问题:
- 判别器不再输出概率,而是输出一个无约束的分数
- 要求判别器是 1 -Lipschitz 函数
- 损失函数形式更简单:L = D(x) – D(G(z))
# WGAN-GP 损失函数实现
def compute_gradient_penalty(D, real_samples, fake_samples):
# 随机插值
alpha = torch.rand(real_samples.size(0), 1, 1, 1)
interpolates = (alpha * real_samples + ((1 - alpha) * fake_samples)).requires_grad_(True)
d_interpolates = D(interpolates)
# 计算梯度范数
gradients = torch.autograd.grad(
outputs=d_interpolates,
inputs=interpolates,
grad_outputs=torch.ones_like(d_interpolates),
create_graph=True,
retain_graph=True,
)[0]
gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
return gradient_penalty
4. 训练避坑指南
- 判别器太强怎么办:
- 降低判别器的学习率
-
减少判别器的更新频率(比如每 5 次生成器更新才更新 1 次判别器)
-
生成器中的批归一化:
- 避免在输出层使用批归一化
-
可以使用层归一化(LayerNorm)作为替代
-
数据预处理:
- 将图像像素值归一化到 [-1, 1] 范围
- 对数据做适当的增强(如随机水平翻转)
5. 模型评估方法
Inception Score 计算
Inception Score (IS) 是衡量生成图像质量和多样性的常用指标:
$$
\text{IS} = \exp(\mathbb{E}_x \text{KL}(p(y|x) || p(y)))
$$
实现步骤:
- 使用预训练的 Inception v3 模型提取特征
- 计算每个生成图像的类别分布 p(y|x)
- 计算整体类别分布 p(y)
- 计算 KL 散度并取指数
6. 完整项目代码
我们提供了一个完整的 Colab Notebook 实现,包含:
- 数据加载和预处理
- DCGAN 模型定义
- 训练循环
- 可视化工具
7. 扩展应用
部署为 API 服务
使用 Flask 可以轻松将训练好的模型部署为 Web 服务:
from flask import Flask, request, jsonify
app = Flask(__name__)
@app.route('/generate', methods=['POST'])
def generate():
z = torch.randn(1, 100, 1, 1).to(device)
with torch.no_grad():
img = generator(z)
return jsonify({"image": img.tolist()})
Conditional GAN 改进
通过加入条件信息(如类别标签),可以控制生成图像的属性:
- 在生成器和判别器的输入中拼接条件向量
- 使用投影判别器(Projection Discriminator)
- 采用辅助分类器(AC-GAN)
总结
通过本文的实战指南,你应该已经掌握了稳定训练 GAN 的核心技巧。记住 GAN 训练更像是一门艺术而非科学,需要不断尝试和调整超参数。建议从小型数据集开始(如 MNIST),等模型可以稳定训练后再扩展到更复杂的数据集。
训练 GAN 需要耐心,有时即使所有设置都正确,模型也可能需要很长时间才开始生成有意义的样本。保持实验日志,记录每次调整的效果,这是提升 GAN 训练技能的最佳方式。
