生成对抗网络(GAN)补疑:从理论推导到实战优化的新手入门指南

1次阅读
没有评论

共计 2251 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景介绍

生成对抗网络 (GAN) 近年来在图像生成、风格迁移、超分辨率重建等计算机视觉任务中展现出强大能力。然而对于初学者而言,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

  1. 判别器目标:最大化对真实样本和生成样本的区分能力
    $$
    L_D = -\mathbb{E}[\log D(x)] – \mathbb{E}[\log(1-D(G(z)))]
    $$
  2. 生成器目标:最小化判别器的判断准确率
    $$
    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()

调优实践

  1. 学习率设置
  2. 初始学习率建议 0.0002
  3. 使用 Adam 优化器时,beta1 设为 0.5
  4. 判别器和生成器可采用不同学习率

  5. 标签平滑

    real_labels = torch.FloatTensor(batch_size).uniform_(0.9, 1.0)
    fake_labels = torch.FloatTensor(batch_size).uniform_(0.0, 0.1)

  6. 模式崩溃检测

  7. 定期检查生成样本的多样性
  8. 计算生成样本的 FID 分数
  9. 使用 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) 约束
正文完
 0
评论(没有评论)