用PyTorch实战LSGAN:从零生成MNIST数字的避坑指南

当你第一次尝试用GAN生成手写数字时,可能会遇到这样的场景:训练了几个小时后,生成的图片依然是一团模糊的噪点,或者判别器的准确率直接飙升至100%——这意味着你的生成器彻底失败了。这正是原始GAN训练不稳定的典型表现。而LSGAN(最小二乘生成对抗网络)通过改变损失函数,让这个对抗游戏变得更加可控。

1. 环境准备与数据加载

在开始之前,确保你的Python环境已经安装了PyTorch和相关的数据处理库。如果你使用conda管理环境,可以这样设置:

conda create -n lsgan python=3.8
conda activate lsgan
pip install torch torchvision matplotlib

MNIST数据集是学习生成模型的绝佳起点——它足够简单,但又包含了真实世界数据的复杂性。PyTorch的torchvision已经内置了这个数据集,我们可以直接加载:

from torchvision import datasets, transforms

transform = transforms.Compose([
    transforms.ToTensor(),
    transforms.Normalize((0.5,), (0.5,))  # 将像素值从[0,1]归一化到[-1,1]
])

train_dataset = datasets.MNIST(root='./data', train=True, download=True, transform=transform)
train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=64, shuffle=True)

注意:归一化到[-1,1]范围是为了匹配生成器输出层的tanh激活函数。

2. 网络架构设计

2.1 生成器构建

生成器的任务是将随机噪声转换为逼真的MNIST数字图像。我们采用全连接网络作为基础架构:

class Generator(nn.Module):
    def __init__(self, latent_dim=100, img_shape=(28,28)):
        super().__init__()
        self.img_shape = img_shape
        self.img_size = img_shape[0] * img_shape[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, self.img_size),
            nn.Tanh()  # 输出在[-1,1]范围
        )
    
    def forward(self, z):
        img = self.model(z)
        return img.view(img.size(0), *self.img_shape)

关键设计选择:

  • LeakyReLU:比普通ReLU更适合GAN,可以缓解梯度消失问题
  • 逐步扩展维度:从100维噪声逐步扩展到784维(28x28)图像
  • Tanh输出:匹配归一化后的输入数据范围

2.2 判别器设计

判别器需要区分真实图像和生成图像,我们同样使用全连接网络:

class Discriminator(nn.Module):
    def __init__(self, img_shape=(28,28)):
        super().__init__()
        self.img_size = img_shape[0] * img_shape[1]
        
        self.model = nn.Sequential(
            nn.Linear(self.img_size, 512),
            nn.LeakyReLU(0.2),
            nn.Linear(512, 256),
            nn.LeakyReLU(0.2),
            nn.Linear(256, 1),
            # 注意:LSGAN不需要Sigmoid!
        )
    
    def forward(self, img):
        img_flat = img.view(img.size(0), -1)
        validity = self.model(img_flat)
        return validity

重要区别:与原始GAN不同,LSGAN的判别器最后一层不需要Sigmoid激活,因为我们直接使用MSE损失。

3. LSGAN的核心:损失函数实现

LSGAN的关键创新在于用最小二乘损失替代了原始GAN的二元交叉熵。这带来了两个主要优势:

  1. 即使判别器非常确信样本是假的,仍然会提供有意义的梯度
  2. 惩罚远离决策边界的样本,促使生成器产生更接近真实数据的样本

在PyTorch中实现LSGAN损失非常简单:

# 定义目标值
real_label = 1.0
fake_label = 0.0

# 判别器损失
real_loss = F.mse_loss(discriminator(real_images), torch.full((batch_size,1), real_label, device=device))
fake_loss = F.mse_loss(discriminator(fake_images.detach()), torch.full((batch_size,1), fake_label, device=device))
d_loss = (real_loss + fake_loss) / 2

# 生成器损失
g_loss = F.mse_loss(discriminator(fake_images), torch.full((batch_size,1), real_label, device=device))

为什么使用MSE损失?

  • 原始GAN使用Sigmoid+BCE,当判别器过于自信时梯度会消失
  • MSE对远离目标的预测给予更大惩罚,促使生成器产生更接近真实分布的样本

4. 训练循环与调试技巧

4.1 基础训练流程

完整的训练循环包含以下几个关键步骤:

for epoch in range(epochs):
    for i, (real_images, _) in enumerate(train_loader):
        
        # 准备真实数据和噪声
        real_images = real_images.to(device)
        z = torch.randn(batch_size, latent_dim, device=device)
        
        # 生成假图像
        fake_images = generator(z)
        
        # 训练判别器
        optimizer_D.zero_grad()
        d_loss = compute_discriminator_loss(real_images, fake_images)
        d_loss.backward()
        optimizer_D.step()
        
        # 训练生成器
        optimizer_G.zero_grad()
        g_loss = compute_generator_loss(fake_images)
        g_loss.backward()
        optimizer_G.step()

4.2 关键调试技巧

当LSGAN训练出现问题时,可以尝试以下调试方法:

观察损失曲线

  • 理想情况:D_loss和G_loss应该震荡下降,最终达到平衡
  • 如果D_loss快速趋近0:判别器太强,尝试降低学习率或减少判别器层数
  • 如果G_loss持续上升:生成器学习失败,检查梯度是否正常传播

学习率调整

optimizer_G = torch.optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))
optimizer_D = torch.optim.Adam(discriminator.parameters(), lr=0.0001, betas=(0.5, 0.999))
  • 判别器学习率通常设为生成器的一半
  • 使用较小的beta1值(如0.5)有助于稳定训练

梯度裁剪

torch.nn.utils.clip_grad_norm_(discriminator.parameters(), max_norm=1.0)
torch.nn.utils.clip_grad_norm_(generator.parameters(), max_norm=1.0)
  • 防止梯度爆炸,特别是训练初期

5. 结果可视化与评估

训练过程中定期查看生成结果非常重要。我们可以每10个epoch保存一次生成样本:

def save_sample_images(epoch):
    with torch.no_grad():
        z = torch.randn(25, latent_dim, device=device)
        gen_imgs = generator(z).cpu()
        
        fig, axs = plt.subplots(5, 5, figsize=(5,5))
        cnt = 0
        for i in range(5):
            for j in range(5):
                axs[i,j].imshow(gen_imgs[cnt,0,:,:], cmap='gray')
                axs[i,j].axis('off')
                cnt += 1
        fig.savefig(f"images/mnist_{epoch}.png")
        plt.close()

评估指标建议:

  • Inception Score (IS):衡量生成图像的多样性和质量
  • Fréchet Inception Distance (FID):比较生成图像与真实图像的分布距离
  • 人工评估:定期查看生成的数字是否清晰可辨

6. 进阶优化策略

当基础模型能够生成可辨认的数字后,可以尝试以下进阶技巧:

添加Dropout层

self.model = nn.Sequential(
    nn.Linear(latent_dim, 256),
    nn.LeakyReLU(0.2),
    nn.Dropout(0.3),
    ...
)
  • 防止过拟合,特别是当判别器太强时

使用谱归一化

from torch.nn.utils import spectral_norm

self.model = nn.Sequential(
    spectral_norm(nn.Linear(latent_dim, 256)),
    nn.LeakyReLU(0.2),
    ...
)
  • 稳定判别器的Lipschitz常数
  • 特别适合深层网络

渐进式训练

  • 先从低分辨率(7x7)开始训练
  • 逐步增加分辨率到14x14,最后到28x28
  • 显著提高生成质量,特别是对于更复杂的数据集

7. 常见问题解决方案

模式崩溃(Mode Collapse)

  • 现象:生成器只产生少量几种样本,缺乏多样性
  • 解决方案:
    • 增加噪声向量的维度
    • 尝试mini-batch判别
    • 调整学习率

梯度消失

  • 现象:生成器停止学习,损失不再变化
  • 解决方案:
    • 检查是否使用了LeakyReLU
    • 确保生成器和判别器的能力平衡
    • 尝试Wasserstein GAN的梯度惩罚

生成图像模糊

  • 现象:数字可辨认但缺乏清晰边缘
  • 解决方案:
    • 在网络中添加卷积层
    • 尝试使用PixelShuffle上采样
    • 增加判别器的感受野

8. 完整代码实现

以下是整合了所有优化技巧的完整LSGAN实现:

import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
import matplotlib.pyplot as plt
import numpy as np

# 参数设置
latent_dim = 100
img_shape = (28, 28)
batch_size = 64
epochs = 200
lr_D = 0.0001
lr_G = 0.0002

# 初始化网络
generator = Generator(latent_dim, img_shape).to(device)
discriminator = Discriminator(img_shape).to(device)

# 优化器
optimizer_G = optim.Adam(generator.parameters(), lr=lr_G, betas=(0.5, 0.999))
optimizer_D = optim.Adam(discriminator.parameters(), lr=lr_D, betas=(0.5, 0.999))

# 训练循环
for epoch in range(epochs):
    for i, (real_images, _) in enumerate(train_loader):
        real_images = real_images.to(device)
        
        # 训练判别器
        optimizer_D.zero_grad()
        
        # 真实图像损失
        real_loss = F.mse_loss(discriminator(real_images), 
                              torch.full((real_images.size(0), 1), 1.0, device=device))
        
        # 生成图像
        z = torch.randn(real_images.size(0), latent_dim, device=device)
        fake_images = generator(z)
        
        # 假图像损失
        fake_loss = F.mse_loss(discriminator(fake_images.detach()),
                              torch.full((real_images.size(0), 1), 0.0, device=device))
        
        d_loss = (real_loss + fake_loss) / 2
        d_loss.backward()
        optimizer_D.step()
        
        # 训练生成器
        optimizer_G.zero_grad()
        g_loss = F.mse_loss(discriminator(fake_images),
                           torch.full((real_images.size(0), 1), 1.0, device=device))
        g_loss.backward()
        optimizer_G.step()
        
    # 每个epoch保存样本
    if epoch % 10 == 0:
        save_sample_images(epoch)
        print(f"[Epoch {epoch}/{epochs}] D_loss: {d_loss.item():.4f}, G_loss: {g_loss.item():.4f}")

# 保存最终模型
torch.save(generator.state_dict(), 'lsgan_generator.pth')
torch.save(discriminator.state_dict(), 'lsgan_discriminator.pth')

在实际项目中,我发现将判别器的学习率设为生成器的一半,同时使用较小的beta1值(0.5)能显著提升训练稳定性。另外,每训练5次判别器再训练1次生成器的策略也值得尝试,这可以防止判别器过强导致生成器无法学习。

Logo

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

更多推荐