共计 2251 个字符,预计需要花费 6 分钟才能阅读完成。
背景介绍
生成对抗网络 (GAN) 近年来在图像生成、风格迁移、超分辨率重建等计算机视觉任务中展现出强大能力。然而对于初学者而言,GAN 的训练过程常伴随以下问题:

- 训练不稳定:判别器 (D) 和生成器 (G) 的博弈容易导致梯度震荡
- 模式崩溃:生成器倾向于产生有限种类的样本
- 梯度消失:当判别器过强时,生成器无法获得有效梯度
理论推导
GAN 的核心思想是二人极小极大博弈,其目标函数为:
$$
\min_G \max_D V(D,G) = \mathbb{E}{x\sim p[\log(1-D(G(z)))]
$$}}[\log D(x)] + \mathbb{E}_{z\sim p_z
- 判别器目标:最大化对真实样本和生成样本的区分能力
$$
L_D = -\mathbb{E}[\log D(x)] – \mathbb{E}[\log(1-D(G(z)))]
$$ - 生成器目标:最小化判别器的判断准确率
$$
L_G = \mathbb{E}[\log(1-D(G(z)))]
$$
代码实现(PyTorch)
生成器网络结构
class Generator(nn.Module):
def __init__(self, latent_dim):
super().__init__()
self.main = nn.Sequential(
# 输入: latent_dim 维噪声
nn.ConvTranspose2d(latent_dim, 256, 4, 1, 0, bias=False),
nn.BatchNorm2d(256),
nn.ReLU(True),
# 上采样至 7x7
nn.ConvTranspose2d(256, 128, 3, 2, 1, bias=False),
nn.BatchNorm2d(128),
nn.ReLU(True),
# 上采样至 14x14
nn.ConvTranspose2d(128, 64, 4, 2, 1, bias=False),
nn.BatchNorm2d(64),
nn.ReLU(True),
# 输出 28x28 的 MNIST 图像
nn.ConvTranspose2d(64, 1, 4, 2, 1, bias=False),
nn.Tanh())
判别器网络结构
class Discriminator(nn.Module):
def __init__(self):
super().__init__()
self.main = nn.Sequential(
# 输入 1x28x28 图像
nn.Conv2d(1, 64, 4, 2, 1, bias=False),
nn.LeakyReLU(0.2, inplace=True),
# 下采样至 14x14
nn.Conv2d(64, 128, 4, 2, 1, bias=False),
nn.BatchNorm2d(128),
nn.LeakyReLU(0.2, inplace=True),
# 下采样至 7x7
nn.Conv2d(128, 256, 3, 2, 1, bias=False),
nn.BatchNorm2d(256),
nn.LeakyReLU(0.2, inplace=True),
# 输出判别概率
nn.Conv2d(256, 1, 4, 1, 0, bias=False),
nn.Sigmoid())
对抗训练循环
for epoch in range(epochs):
for real_imgs, _ in dataloader:
# 训练判别器
optimizer_D.zero_grad()
z = torch.randn(batch_size, latent_dim, 1, 1)
fake_imgs = generator(z)
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()
optimizer_D.step()
# 训练生成器
optimizer_G.zero_grad()
g_loss = criterion(D(fake_imgs), real_labels)
g_loss.backward()
optimizer_G.step()
调优实践
- 学习率设置:
- 初始学习率建议 0.0002
- 使用 Adam 优化器时,beta1 设为 0.5
-
判别器和生成器可采用不同学习率
-
标签平滑:
real_labels = torch.FloatTensor(batch_size).uniform_(0.9, 1.0) fake_labels = torch.FloatTensor(batch_size).uniform_(0.0, 0.1) -
模式崩溃检测:
- 定期检查生成样本的多样性
- 计算生成样本的 FID 分数
- 使用 minibatch discrimination 技术
避坑指南
- 梯度裁剪:阈值设为 0.01~0.1
- 批量归一化:生成器最后一层和判别器第一层不要使用 BN
- 判别器强度:保持 D 和 G 的训练次数比为 1:1 或 2:1
效果验证
经过 200 轮训练后,在 MNIST 数据集上可获得清晰的数字生成效果:
Epoch [100/200] D_loss: 0.5632 G_loss: 1.8923
Epoch [150/200] D_loss: 0.5011 G_loss: 2.1034
Epoch [200/200] D_loss: 0.4876 G_loss: 2.2158
思考题
如何改进网络结构以生成更高分辨率的图像?可能的改进方向包括:
- 使用渐进式增长训练策略
- 引入注意力机制
- 采用多尺度判别器结构
- 添加谱归一化 (Spectral Norm) 约束
正文完
发表至: 未分类
近一天内
