别再死记硬背GAN公式了!用PyTorch手搓一个数字生成器,带你直观感受对抗训练全过程
用PyTorch实战GAN:从噪声到数字的魔法之旅
当第一次听说"生成对抗网络"这个词时,我脑海中浮现的是两个拳击手在擂台上互相训练的画面。但真正开始学习GAN时,却发现大多数教程都在堆砌数学公式和抽象概念,让人望而生畏。直到有一天,我决定用PyTorch亲手构建一个能生成手写数字的GAN,才真正理解了这场"对抗"的精妙之处——这不是拳击比赛,而更像是一位画家和鉴赏家之间的艺术博弈。
1. 准备你的数字画室
在开始这场创作之前,我们需要搭建好工作环境。与许多深度学习项目不同,GAN训练更像是在调教两个互相学习的学生——生成器(G)和判别器(D)。它们会互相"欺骗"和"识破",在这个过程中共同进步。
首先安装必要的库:
pip install torch torchvision matplotlib
接着导入我们将用到的模块:
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
关键工具说明 :
torch.nn:构建神经网络的核心模块torchvision:提供MNIST数据集和图像转换工具matplotlib:用于实时可视化生成结果
提示:使用Jupyter Notebook可以获得更好的交互体验,能实时观察训练过程中生成图像的演变
2. 构建我们的"艺术家"与"鉴赏家"
2.1 生成器:从噪声到艺术
生成器就像一位刚开始学画的艺术家,最初它只会胡乱涂抹,但通过判别器的反馈,它会逐渐掌握绘制逼真数字的技巧。以下是我们的生成器架构:
class Generator(nn.Module):
def __init__(self, latent_dim=100):
super(Generator, self).__init__()
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()
)
def forward(self, z):
img = self.model(z)
return img.view(-1, 1, 28, 28)
设计要点 :
- 输入:100维的随机噪声(latent space)
- 使用
LeakyReLU避免梯度消失问题 - 最终使用
Tanh将输出压缩到[-1,1]范围,与预处理后的MNIST数据匹配
2.2 判别器:火眼金睛的鉴赏家
判别器则像一位经验丰富的艺术鉴赏家,它的任务是区分真实画作和生成器的作品:
class Discriminator(nn.Module):
def __init__(self):
super(Discriminator, self).__init__()
self.model = nn.Sequential(
nn.Linear(784, 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):
flattened = img.view(-1, 784)
validity = self.model(flattened)
return validity
关键设计 :
- 使用
Dropout防止过拟合 - 最终
Sigmoid输出一个0到1的概率值 - 输入是展平后的784维向量(28x28图像)
3. 准备训练材料:MNIST数据集
任何艺术家都需要学习素材,我们的GAN将从MNIST手写数字数据集中学习:
# 数据预处理
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
# 加载数据集
train_dataset = datasets.MNIST(
root='./data',
train=True,
download=True,
transform=transform
)
# 创建数据加载器
dataloader = DataLoader(
train_dataset,
batch_size=64,
shuffle=True
)
预处理细节 :
ToTensor将图像转换为PyTorch张量并归一化到[0,1]Normalize进一步将范围调整到[-1,1],与生成器的Tanh输出匹配
4. 训练:艺术与鉴别的博弈
GAN训练的核心在于平衡生成器和判别器的学习进度。太强的判别器会让生成器学不到东西,而太弱的判别器则无法提供有用的反馈。
4.1 初始化与损失函数
# 初始化网络
generator = Generator()
discriminator = Discriminator()
# 优化器
optimizer_G = optim.Adam(generator.parameters(), lr=0.0002)
optimizer_D = optim.Adam(discriminator.parameters(), lr=0.0002)
# 损失函数
adversarial_loss = nn.BCELoss()
4.2 训练循环中的对抗过程
for epoch in range(epochs):
for i, (imgs, _) in enumerate(dataloader):
# 真实和假标签
real = torch.ones(imgs.size(0), 1)
fake = torch.zeros(imgs.size(0), 1)
# ---------------------
# 训练判别器
# ---------------------
optimizer_D.zero_grad()
# 真实图像的损失
real_loss = adversarial_loss(discriminator(imgs), real)
# 生成假图像
z = torch.randn(imgs.size(0), 100)
gen_imgs = generator(z)
# 假图像的损失
fake_loss = adversarial_loss(discriminator(gen_imgs.detach()), fake)
# 总判别器损失
d_loss = (real_loss + fake_loss) / 2
d_loss.backward()
optimizer_D.step()
# -----------------
# 训练生成器
# -----------------
optimizer_G.zero_grad()
# 生成器希望假图像被判别为真
g_loss = adversarial_loss(discriminator(gen_imgs), real)
g_loss.backward()
optimizer_G.step()
对抗过程解析 :
- 判别器训练:
- 同时看真实图像和生成图像
- 目标是最大化识别真假的能力
- 生成器训练:
- 固定判别器
- 目标是让生成的图像骗过判别器
注意:训练GAN时,保持两个网络的平衡至关重要。如果一方明显强于另一方,训练就会失效
5. 可视化:见证艺术的诞生
训练过程中最令人兴奋的部分就是观察生成器作品的演变。我们可以每隔一定epoch保存生成结果:
def sample_images(epoch):
z = torch.randn(5*5, 100)
gen_imgs = generator(z)
fig, axs = plt.subplots(5, 5)
cnt = 0
for i in range(5):
for j in range(5):
axs[i,j].imshow(gen_imgs[cnt].detach().numpy().reshape(28,28), cmap='gray')
axs[i,j].axis('off')
cnt += 1
fig.savefig(f"images/mnist_{epoch}.png")
plt.close()
典型训练过程观察 :
- 前几epoch:随机噪声
- 10-20epoch:模糊的数字形状开始出现
- 50epoch左右:可辨认的数字
- 100epoch后:清晰的数字,风格多样
6. 调优技巧:让GAN训练更稳定
GAN训练 notoriously tricky(臭名昭著地棘手)。以下是一些实际经验:
6.1 学习率调整
# 动态调整学习率
scheduler_G = optim.lr_scheduler.StepLR(optimizer_G, step_size=30, gamma=0.1)
scheduler_D = optim.lr_scheduler.StepLR(optimizer_D, step_size=30, gamma=0.1)
6.2 标签平滑
# 使用软标签
real_labels = torch.FloatTensor(batch_size, 1).uniform_(0.9, 1.0)
fake_labels = torch.FloatTensor(batch_size, 1).uniform_(0.0, 0.1)
6.3 梯度惩罚
# WGAN-GP中的梯度惩罚
def compute_gradient_penalty(D, real_samples, fake_samples):
alpha = torch.rand(real_samples.size(0), 1)
interpolates = (alpha * real_samples + (1 - alpha) * fake_samples).requires_grad_(True)
d_interpolates = D(interpolates)
gradients = torch.autograd.grad(
outputs=d_interpolates,
inputs=interpolates,
grad_outputs=torch.ones_like(d_interpolates),
create_graph=True,
retain_graph=True,
only_inputs=True,
)[0]
gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
return gradient_penalty
7. 进阶:探索GAN的创意空间
当基础GAN工作正常后,可以尝试以下扩展:
7.1 控制生成特定数字
# 在噪声z中加入类别标签
class ConditionalGenerator(nn.Module):
def __init__(self, num_classes=10):
super().__init__()
self.label_emb = nn.Embedding(num_classes, 100)
self.model = nn.Sequential(
# 其余层与之前相同
)
def forward(self, z, labels):
c = self.label_emb(labels)
x = torch.cat([z, c], 1)
return self.model(x)
7.2 使用卷积结构
class ConvGenerator(nn.Module):
def __init__(self):
super().__init__()
self.model = nn.Sequential(
nn.ConvTranspose2d(100, 512, 4, 1, 0),
nn.BatchNorm2d(512),
nn.ReLU(),
# 更多转置卷积层...
)
7.3 实现DCGAN
深度卷积GAN(DCGAN)架构建议:
- 在生成器中使用转置卷积
- 判别器中使用普通卷积
- 使用批量归一化
- 移除全连接层
- 使用ReLU(生成器)和LeakyReLU(判别器)
8. 常见问题与解决方案
在多次训练GAN的过程中,我遇到了各种"坑",以下是几个典型问题及解决方法:
模式崩溃(Mode Collapse) :
- 现象:生成器只产生几种类似的输出
- 解决方案:
- 尝试mini-batch判别
- 使用Wasserstein GAN
- 调整学习率
判别器过强 :
- 现象:生成器loss不下降
- 解决方案:
- 减少判别器的更新频率
- 降低判别器的学习率
- 添加噪声到判别器输入
生成图像模糊 :
- 原因:L2损失倾向于平均结果
- 解决方案:
- 使用L1损失
- 尝试感知损失(perceptual loss)
- 添加GAN的变体如LSGAN
训练不稳定 :
- 解决方案:
- 使用梯度惩罚
- 尝试不同的优化器
- 实现谱归一化
9. 从MNIST到更复杂的生成任务
掌握了基础GAN后,可以挑战更复杂的生成任务:
CIFAR-10 :
- 32x32彩色图像
- 需要更深的网络
- 考虑使用残差连接
人脸生成 :
- 使用CelebA数据集
- 需要更大的模型
- 考虑渐进式增长训练
风格迁移 :
- 结合GAN与风格损失
- 尝试CycleGAN架构
- 注意领域适配问题
10. 实战建议与个人心得
经过多次实验,我发现GAN训练更像是一门艺术而非严格的科学。以下是一些个人总结的经验:
- 耐心是关键 :GAN可能需要训练数百epoch才能看到好结果
- 可视化至关重要 :loss曲线可能具有欺骗性,要相信自己的眼睛
- 从小开始 :先在MNIST上验证想法,再扩展到复杂数据集
- 记录一切 :超参数、架构变化和结果要详细记录
- 尝试变体 :当标准GAN不工作时,WGAN、LSGAN等可能表现更好
第一次看到自己训练的GAN生成出可辨认的数字时,那种成就感是难以言表的。虽然最初的几次尝试可能以失败告终,但每次调整后看到生成质量的提升,都是对耐心和努力的最好回报。
更多推荐




所有评论(0)