告别GAN训练不稳定:用PyTorch手把手实现LSGAN生成MNIST数字(附完整代码)
用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的二元交叉熵。这带来了两个主要优势:
- 即使判别器非常确信样本是假的,仍然会提供有意义的梯度
- 惩罚远离决策边界的样本,促使生成器产生更接近真实数据的样本
在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次生成器的策略也值得尝试,这可以防止判别器过强导致生成器无法学习。
更多推荐
所有评论(0)