共计 3216 个字符,预计需要花费 9 分钟才能阅读完成。
GAN 的基本概念
生成对抗网络 (GAN) 是一种强大的生成模型,由生成器 (Generator) 和判别器 (Discriminator) 两部分组成。生成器负责生成假的样本,判别器则负责判断样本是真实的还是生成的。两者相互对抗,最终使得生成器能够生成逼真的样本。

GAN 在图像生成、风格迁移、超分辨率重建等领域都有广泛应用。例如,我们可以用 GAN 生成逼真的人脸图像,或者将普通照片转换成艺术风格的作品。
GAN 训练中的常见问题
- 模式崩溃(Mode Collapse)
- 生成器只学会生成部分模式的数据,而忽略了其他模式
-
表现为生成样本多样性不足
-
梯度消失(Gradient Vanishing)
- 判别器训练得太好,导致生成器无法获得有效的梯度
-
生成器的性能无法继续提升
-
训练不稳定
- 生成器和判别器的平衡难以维持
- 容易出现一方压倒另一方的情况
GAN 实现代码(Python+Pytorch)
import torch
import torch.nn as nn
import torch.optim as optim
from torchvision import datasets, transforms
from torch.utils.data import DataLoader
# 定义生成器
class Generator(nn.Module):
def __init__(self, latent_dim, img_shape):
super().__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, img_shape),
nn.Tanh())
def forward(self, z):
return self.model(z)
# 定义判别器
class Discriminator(nn.Module):
def __init__(self, img_shape):
super().__init__()
self.model = nn.Sequential(nn.Linear(img_shape, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 1),
nn.Sigmoid())
def forward(self, img):
return self.model(img)
# 训练参数
latent_dim = 100
img_shape = 28*28 # MNIST 图像大小
batch_size = 64
epochs = 200
# 初始化模型
generator = Generator(latent_dim, img_shape)
discriminator = Discriminator(img_shape)
# 优化器
g_optim = optim.Adam(generator.parameters(), lr=0.0002, betas=(0.5, 0.999))
d_optim = optim.Adam(discriminator.parameters(), lr=0.0002, betas=(0.5, 0.999))
# 损失函数
criterion = nn.BCELoss()
# 加载数据
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize([0.5], [0.5])
])
dataset = datasets.MNIST('data', train=True, download=True, transform=transform)
dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
# 训练循环
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)
# 判别器训练
d_optim.zero_grad()
# 真实图像
real_imgs = imgs.view(imgs.size(0), -1)
d_real = discriminator(real_imgs)
d_real_loss = criterion(d_real, real)
# 生成图像
z = torch.randn(imgs.size(0), latent_dim)
gen_imgs = generator(z)
d_fake = discriminator(gen_imgs.detach())
d_fake_loss = criterion(d_fake, fake)
d_loss = d_real_loss + d_fake_loss
d_loss.backward()
d_optim.step()
# 生成器训练
g_optim.zero_grad()
z = torch.randn(imgs.size(0), latent_dim)
gen_imgs = generator(z)
g_output = discriminator(gen_imgs)
g_loss = criterion(g_output, real)
g_loss.backward()
g_optim.step()
# 每个 epoch 打印损失
print(f'Epoch {epoch}, D Loss: {d_loss.item()}, G Loss: {g_loss.item()}')
GAN 架构示意图
[随机噪声] → [生成器] → [生成图像]
↓
[判别器] ← [真实图像]
- 生成器接收随机噪声作为输入,输出生成的图像
- 判别器同时接收真实图像和生成图像,输出判别结果
- 两个网络交替训练,互相促进
训练调优技巧
- 学习率设置
- 通常设置在 0.0001 到 0.0005 之间
-
可以使用学习率衰减策略
-
损失函数选择
- 基础 GAN 使用二进制交叉熵 (BCE) 损失
- WGAN 使用 Wasserstein 距离
-
LSGAN 使用最小二乘损失
-
其他技巧
- 使用批归一化 (BatchNorm) 稳定训练
- 对判别器进行多次更新后再更新生成器
- 使用标签平滑 (Label Smoothing) 防止过拟合
生产环境部署最佳实践
- 模型保存与加载
- 定期保存生成器和判别器的 checkpoint
-
使用 torch.save 和 torch.load 进行模型保存和加载
-
分布式训练
- 使用 DataParallel 或 DistributedDataParallel 进行多 GPU 训练
-
注意同步批归一化统计量
-
性能优化
- 使用混合精度训练(AMP)
- 优化数据加载管道
-
使用 ONNX 或 TorchScript 进行模型导出
-
常见问题解决方案
- 模式崩溃:尝试使用 mini-batch 判别或特征匹配
- 梯度消失:调整学习率或改用 WGAN
- 训练不稳定:使用梯度惩罚或谱归一化
可视化训练结果
经过 200 个 epoch 的训练后,我们的 GAN 已经能够生成比较清晰的 MNIST 数字图像。虽然部分数字可能还有些模糊,但已经能够辨认出 0 - 9 的各种数字形状。
改进方向
- 尝试实现 DCGAN(深度卷积 GAN),使用卷积神经网络替代全连接网络
- 实验不同的损失函数,如 WGAN-GP 或 LSGAN
- 添加条件信息,实现 cGAN(条件 GAN)
- 尝试在更高分辨率的数据集 (如 CIFAR-10) 上训练
思考题
如何评估 GAN 生成质量?可以考虑以下指标:
1. 人工评估:主观判断生成图像的逼真程度
2. Inception Score(IS):同时考虑生成图像的多样性和可识别性
3. Frechet Inception Distance(FID):比较生成图像和真实图像在特征空间的距离
4. 精确率和召回率:评估生成样本覆盖真实数据分布的程度
正文完
