1. 项目概述:用GAN生成手写数字的实践指南

在计算机视觉领域,生成对抗网络(GAN)已经成为图像生成任务的标杆技术。我最近完成了一个用GAN生成MNIST手写数字的项目,这个看似简单的任务实际上包含了GAN训练中的所有核心挑战。MNIST数据集作为入门级的28x28灰度图像,是理解GAN工作原理的理想起点,但要让生成器产出逼真的数字,需要处理好模式坍塌、训练不稳定等典型问题。

这个项目特别适合两类开发者:一是刚接触生成模型的初学者,可以通过相对简单的架构理解GAN的核心机制;二是有一定深度学习基础的研究者,可以在此基础上升级网络结构、尝试不同的损失函数。整个项目使用PyTorch框架实现,代码量控制在200行以内,但涵盖了数据加载、网络定义、训练循环等完整流程。

2. 核心架构设计解析

2.1 GAN的双网络工作原理

GAN的核心在于生成器(Generator)和判别器(Discriminator)的对抗训练。在我的实现中:

生成器接收100维的随机噪声(服从标准正态分布),通过全连接层逐步上采样到784维(28x28),最后用Tanh激活将输出压缩到[-1,1]区间。关键设计点是:

self.model = 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, 784),
    nn.Tanh()
)

判别器则是标准的二分类网络,采用LeakyReLU防止梯度消失:

self.model = nn.Sequential(
    nn.Linear(784, 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()
)

2.2 损失函数与优化器选择

采用二进制交叉熵损失(BCELoss)作为目标函数:

criterion = nn.BCELoss()

优化器使用Adam,这是GAN训练的常见选择。关键参数设置:

  • 生成器学习率:0.0002
  • 判别器学习率:0.0002
  • Beta1:0.5(比默认值0.9更稳定)

提示:判别器和生成器使用不同的优化器实例,这是为了让它们可以独立更新参数。

3. 完整训练流程实现

3.1 数据预处理管道

MNIST数据需要做以下处理:

  1. 标准化:将像素值从[0,255]线性变换到[-1,1]
  2. 展平:将28x28图像转为784维向量
  3. 批处理:设置batch_size=64
transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,)),
    transforms.Lambda(lambda x: x.view(-1))
])

3.2 训练循环的关键步骤

每个epoch包含:

  1. 训练判别器:

    • 用真实数据计算loss_real
    • 用生成数据计算loss_fake
    • 总loss = (loss_real + loss_fake)/2
  2. 训练生成器:

    • 让判别器对生成图像的评分接近1
    • 只反向传播生成器的参数
for epoch in range(epochs):
    for i, (real_images, _) in enumerate(dataloader):
        
        # 训练判别器
        optimizer_D.zero_grad()
        
        # 真实图像损失
        real_loss = criterion(D(real_images), real_labels)
        
        # 生成图像损失
        z = torch.randn(batch_size, latent_dim)
        fake_images = G(z)
        fake_loss = criterion(D(fake_images.detach()), fake_labels)
        
        d_loss = (real_loss + fake_loss) / 2
        d_loss.backward()
        optimizer_D.step()
        
        # 训练生成器
        optimizer_G.zero_grad()
        
        g_loss = criterion(D(fake_images), real_labels)
        g_loss.backward()
        optimizer_G.step()

3.3 训练监控与可视化

每100个batch保存一次生成样本:

if i % 100 == 0:
    with torch.no_grad():
        test_z = torch.randn(16, latent_dim)
        generated = G(test_z)
        save_image(generated.view(16, 1, 28, 28), 
                 f"output/epoch{epoch}_batch{i}.png", 
                 nrow=4, normalize=True)

同时记录损失曲线,这是诊断训练状态的重要依据。典型的健康训练表现为:

  • 判别器损失在0.5附近震荡
  • 生成器损失呈缓慢下降趋势

4. 实战问题排查指南

4.1 模式坍塌的识别与解决

当生成器开始反复输出相同数字时,说明发生了模式坍塌。解决方法包括:

  1. 增加判别器的Dropout率(如从0.3提到0.5)
  2. 在生成器损失中加入特征匹配损失:
feature_loss = torch.norm(D.features(real_images) - D.features(fake_images), p=2)
g_loss += 0.1 * feature_loss
  1. 尝试不同的噪声维度(如从100增加到256)

4.2 梯度消失/爆炸的处理

如果训练早期损失就停滞不变:

  1. 检查权重初始化:使用Xavier初始化
def weights_init(m):
    if isinstance(m, nn.Linear):
        nn.init.xavier_normal_(m.weight)
        nn.init.constant_(m.bias, 0)
G.apply(weights_init)
D.apply(weights_init)
  1. 调整学习率(通常需要降低)
  2. 尝试梯度裁剪:
torch.nn.utils.clip_grad_norm_(D.parameters(), 1.0)
torch.nn.utils.clip_grad_norm_(G.parameters(), 1.0)

4.3 生成图像模糊的改进

如果生成的数字边缘不清晰:

  1. 在生成器最后层加入PixelNorm:
class PixelNorm(nn.Module):
    def forward(self, x):
        return x / torch.sqrt(torch.mean(x**2, dim=1, keepdim=True) + 1e-8)
  1. 改用Wasserstein GAN(WGAN)架构
  2. 增加生成器的卷积层(转置卷积)

5. 进阶优化方向

基础模型稳定后,可以尝试以下改进:

5.1 网络架构升级

  • 将全连接层替换为卷积层(DCGAN架构)
  • 添加自注意力机制(SAGAN)
  • 尝试渐进式增长(ProGAN)

5.2 损失函数改进

  • 改用Wasserstein距离(WGAN-GP)
  • 加入感知损失(Perceptual Loss)
  • 使用对比学习(Contrastive Learning)

5.3 评估指标引入

  • 计算FID分数(需要预训练分类器)
  • 记录IS分数(Inception Score)
  • 人工评估生成质量

我在实际训练中发现,保持判别器比生成器稍强一些(约快20%的学习速度)有助于稳定训练。另一个实用技巧是在训练初期(前10个epoch)使用较高的学习率(如0.002),然后逐步衰减到0.0001,这可以加速收敛。

Logo

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

更多推荐