从零理解GAN:2014年Goodfellow提出的生成对抗网络核心原理与实战指南

1次阅读
没有评论

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

image.webp

背景解析:GAN 的核心博弈原理

2014 年 Ian Goodfellow 提出的生成对抗网络 (GAN),本质上是一个让两个神经网络相互对抗的游戏。就像古董鉴定师和造假者之间的博弈:

从零理解 GAN:2014 年 Goodfellow 提出的生成对抗网络核心原理与实战指南

  • 生成器 (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)

这个公式的意思是:

  1. 判别器试图最大化识别真实数据的能力(第一个期望项)和识别假数据的能力(第二个期望项)
  2. 生成器试图最小化判别器识别假数据的能力

训练过程中两者的 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)

进阶思考

  1. 生成质量评估
  2. 如何定量评估生成图像的真实性?
  3. Inception Score 和 FID 指标各有什么优缺点?

  4. 条件生成

  5. 如何让 GAN 生成指定类别的数字(如仅生成数字 ”7″)
  6. 对比 AC-GAN 和 CGAN 两种条件生成方案的差异

经过这次实践,我深刻体会到 GAN 就像在训练两个相互竞争的运动员——需要精心调节他们的训练强度,才能让双方都不断进步。建议初学者多尝试调整网络结构和超参数,观察训练过程中的 loss 变化和生成效果,这是理解 GAN 行为的最佳方式。

正文完
 0
评论(没有评论)