从“造假者”到“艺术家”:用PyTorch复现GAN原始论文,理解对抗训练的本质

想象一下,你正在教一个孩子画画。最初他只能画出歪歪扭扭的线条,但每次画完后,你都会告诉他哪些地方画得像真实的物体,哪些不像。经过反复练习,孩子的画技越来越好,直到有一天,他的画作几乎能以假乱真——这正是生成对抗网络(GAN)的核心思想。本文将带你用PyTorch亲手实现2014年那篇开创性的GAN论文,通过代码和可视化,深入理解这个"造假者"与"鉴定专家"之间的精彩博弈。

1. GAN基础:对抗训练的双人舞

生成对抗网络由两个相互对抗的神经网络组成:生成器(Generator)和判别器(Discriminator)。这对"冤家"的较量过程可以用一个简单的比喻理解:

  • 生成器(G):就像造假币的罪犯,不断尝试制作更逼真的假币
  • 判别器(D):如同经验老道的鉴钞专家,努力分辨真币和假币

在PyTorch中,我们可以这样定义它们的基本结构:

import torch
import torch.nn as nn

class Generator(nn.Module):
    def __init__(self, input_dim, output_dim):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(input_dim, 128),
            nn.LeakyReLU(0.2),
            nn.Linear(128, 256),
            nn.LeakyReLU(0.2),
            nn.Linear(256, output_dim),
            nn.Tanh()
        )
    
    def forward(self, z):
        return self.net(z)

class Discriminator(nn.Module):
    def __init__(self, input_dim):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(input_dim, 256),
            nn.LeakyReLU(0.2),
            nn.Linear(256, 128),
            nn.LeakyReLU(0.2),
            nn.Linear(128, 1),
            nn.Sigmoid()
        )
    
    def forward(self, x):
        return self.net(x)

注意:原始论文使用MLP(多层感知机)结构,这与现代GAN常用的卷积结构不同,但更适合理解基本原理。

GAN的训练过程本质上是一个极小极大博弈(minimax game),可以用以下目标函数表示:

$$ \min_G \max_D V(D,G) = \mathbb{E}{x\sim p{data}}[\log D(x)] + \mathbb{E}_{z\sim p_z}[\log(1-D(G(z)))] $$

这个公式包含两个关键部分:

  1. 判别器D试图最大化它正确分类真实数据和生成数据的能力
  2. 生成器G试图最小化判别器的分类准确率

2. 实战:用PyTorch实现原始GAN

2.1 数据准备与模型初始化

我们将使用MNIST手写数字数据集作为训练数据。首先设置基本参数:

# 超参数设置
input_dim = 100  # 噪声向量维度
output_dim = 784  # MNIST图像展平后的维度(28x28)
batch_size = 64
epochs = 200
lr = 0.0002

# 初始化模型
generator = Generator(input_dim, output_dim)
discriminator = Discriminator(output_dim)

# 优化器
g_optimizer = torch.optim.Adam(generator.parameters(), lr=lr)
d_optimizer = torch.optim.Adam(discriminator.parameters(), lr=lr)

# 损失函数
criterion = nn.BCELoss()

2.2 训练循环的实现

GAN的训练需要交替更新生成器和判别器。以下是核心训练代码:

for epoch in range(epochs):
    for i, (real_images, _) in enumerate(train_loader):
        # 准备真实数据
        real_images = real_images.view(-1, 784)
        real_labels = torch.ones(batch_size, 1)
        fake_labels = torch.zeros(batch_size, 1)
        
        # 训练判别器
        d_optimizer.zero_grad()
        
        # 真实数据的判别损失
        real_outputs = discriminator(real_images)
        d_loss_real = criterion(real_outputs, real_labels)
        
        # 生成假数据
        z = torch.randn(batch_size, input_dim)
        fake_images = generator(z)
        
        # 假数据的判别损失
        fake_outputs = discriminator(fake_images.detach())
        d_loss_fake = criterion(fake_outputs, fake_labels)
        
        # 总判别损失
        d_loss = d_loss_real + d_loss_fake
        d_loss.backward()
        d_optimizer.step()
        
        # 训练生成器
        g_optimizer.zero_grad()
        
        # 生成器希望判别器将假数据判为真
        outputs = discriminator(fake_images)
        g_loss = criterion(outputs, real_labels)
        
        g_loss.backward()
        g_optimizer.step()

提示:在实际训练中,通常会先更新判别器多次(如5次),再更新生成器1次,以保持两者能力平衡。

2.3 训练动态的可视化

理解GAN训练过程的关键是观察三个方面的变化:

  1. 损失曲线:生成器和判别器的损失变化
  2. 生成样本质量:随着训练进行生成图像的演变
  3. 判别器输出分布:对真实和生成样本的判别概率

以下是绘制训练动态的代码示例:

def plot_training_process(g_losses, d_losses, samples):
    plt.figure(figsize=(12, 4))
    
    # 损失曲线
    plt.subplot(1, 3, 1)
    plt.plot(g_losses, label='Generator Loss')
    plt.plot(d_losses, label='Discriminator Loss')
    plt.legend()
    
    # 生成样本
    plt.subplot(1, 3, 2)
    plt.imshow(samples[-1][0].reshape(28, 28), cmap='gray')
    
    # 判别器输出分布
    plt.subplot(1, 3, 3)
    sns.histplot(real_outputs.detach().numpy(), color='blue', label='Real')
    sns.histplot(fake_outputs.detach().numpy(), color='orange', label='Fake')
    plt.legend()
    
    plt.show()

3. 深入理解对抗训练的动态平衡

3.1 纳什均衡与训练稳定性

GAN的训练目标是达到纳什均衡点,此时:

  • 生成器产生的数据分布 $p_g$ 完全匹配真实数据分布 $p_{data}$
  • 判别器对所有输入都输出0.5(完全无法区分真假)

数学上可以证明,当 $p_g = p_{data}$ 时,最优判别器为:

$$ D_G^*(x) = \frac{p_{data}(x)}{p_{data}(x) + p_g(x)} = \frac{1}{2} $$

然而在实际训练中,GAN常常面临以下挑战:

问题类型 表现特征 解决方案
模式崩溃 生成器只产生有限的几种样本 小批量判别、添加噪声
梯度消失 判别器过早变得太强 标签平滑、降低学习率
振荡不收敛 损失函数剧烈波动 使用Wasserstein距离替代

3.2 原始GAN的改进技巧

虽然原始GAN论文提出了基本框架,但实践中我们发现几个关键改进点:

  1. 损失函数调整

    • 原始公式 $\min_G \log(1-D(G(z)))$ 在训练早期梯度很小
    • 实际使用 $\max_G \log D(G(z))$ 能提供更强梯度
  2. 网络结构设计

    • 使用LeakyReLU代替ReLU防止梯度消失
    • 在判别器中使用Dropout增加鲁棒性
  3. 训练策略优化

    • 对真实样本标签使用0.9而非1.0(标签平滑)
    • 对生成样本标签使用0.1而非0.0
# 改进后的标签设置
real_labels = torch.full((batch_size, 1), 0.9)
fake_labels = torch.full((batch_size, 1), 0.1)

4. 从理论到实践:GAN的现代应用

虽然我们复现的是原始GAN,但理解这些基础对掌握现代GAN变种至关重要。以下是几个关键发展脉络:

  1. 架构演进

    • DCGAN:首次将卷积网络引入GAN
    • WGAN:使用Wasserstein距离改进训练稳定性
    • StyleGAN:实现前所未有的生成质量
  2. 应用场景

    • 图像生成(如艺术创作、人脸合成)
    • 数据增强(医疗影像等领域的小样本学习)
    • 跨模态生成(文本到图像,如DALL·E)
  3. 评估指标

    • Inception Score (IS)
    • Fréchet Inception Distance (FID)
    • Precision & Recall for Generative Models

在完成这个基础实现后,建议尝试以下扩展实验:

  • 将MLP结构改为卷积网络(DCGAN架构)
  • 尝试不同的损失函数(如Wasserstein损失)
  • 在更复杂的数据集(如CIFAR-10)上测试

通过PyTorch实现原始GAN论文,我们不仅理解了"对抗训练"的精妙之处,也为后续探索更复杂的生成模型打下了坚实基础。在实际项目中,GAN的训练往往需要大量调参经验——比如发现判别器loss降为0时,通常意味着生成器已经失败,这时需要调整两者的学习率比例。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐