共计 2529 个字符,预计需要花费 7 分钟才能阅读完成。
背景与痛点
生成对抗网络(GAN)由 Ian Goodfellow 在 2014 年提出,迅速成为 AI 领域的热门技术。GAN 的核心思想是通过两个神经网络的对抗训练,生成高度逼真的数据。然而,尽管 GAN 潜力巨大,许多开发者在实际应用中常遇到训练不稳定、模式崩溃等问题。

- 训练不稳定 :生成器和判别器的动态平衡难以维持,常导致一方压倒另一方。
- 模式崩溃 :生成器倾向于生成有限的几种样本,缺乏多样性。
- 评估困难 :缺乏统一的量化指标来评估生成质量。
核心原理
GAN 由生成器(G)和判别器(D)组成,两者通过对抗过程共同提升。生成器试图生成逼真的假数据,而判别器则尝试区分真假数据。
数学表达
原始 GAN 的损失函数为:
[\min_G \max_D V(D, G) = \mathbb{E}{x \sim p[\log (1 – D(G(z)))] ]}(x)}[\log D(x)] + \mathbb{E}_{z \sim p_z(z)
- 生成器目标 :最小化判别器对生成样本的判别能力。
- 判别器目标 :最大化对真实和生成样本的区分能力。
变体改进
- WGAN:使用 Wasserstein 距离替代原始 GAN 的 JS 散度,缓解梯度消失问题。
- WGAN-GP:通过梯度惩罚(Gradient Penalty)进一步稳定训练。
- SNGAN:引入谱归一化(Spectral Normalization)提升判别器的鲁棒性。
代码实现
以下是一个基于 PyTorch 的 DCGAN 实现,用于生成手写数字图像(MNIST 数据集)。
数据加载
import torch
from torchvision import datasets, transforms
transform = transforms.Compose([transforms.ToTensor(),
transforms.Normalize((0.5,), (0.5,))
])
dataset = datasets.MNIST(
root='./data',
train=True,
download=True,
transform=transform
)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=64, shuffle=True)
网络结构
class Generator(nn.Module):
def __init__(self, latent_dim):
super(Generator, self).__init__()
self.main = 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):
return self.main(z)
class Discriminator(nn.Module):
def __init__(self):
super(Discriminator, self).__init__()
self.main = nn.Sequential(nn.Linear(784, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 1),
nn.Sigmoid())
def forward(self, x):
return self.main(x)
训练循环
# 初始化网络与优化器
G = Generator(latent_dim=100).to(device)
D = Discriminator().to(device)
optimizer_G = torch.optim.Adam(G.parameters(), lr=0.0002)
optimizer_D = torch.optim.Adam(D.parameters(), lr=0.0002)
for epoch in range(epochs):
for i, (real_imgs, _) in enumerate(dataloader):
real_imgs = real_imgs.view(-1, 784).to(device)
# 训练判别器
z = torch.randn(real_imgs.size(0), 100).to(device)
fake_imgs = G(z)
real_loss = torch.log(D(real_imgs)).mean()
fake_loss = torch.log(1 - D(fake_imgs.detach())).mean()
d_loss = -(real_loss + fake_loss)
optimizer_D.zero_grad()
d_loss.backward()
optimizer_D.step()
# 训练生成器
fake_loss = torch.log(1 - D(fake_imgs)).mean()
g_loss = -fake_loss
optimizer_G.zero_grad()
g_loss.backward()
optimizer_G.step()
实战技巧
- 学习率调整 :GAN 对学习率敏感,建议使用较小的学习率(如 0.0002)。
- 标签平滑 :对真实样本的标签稍低于 1(如 0.9),防止判别器过度自信。
- 梯度裁剪 :在 WGAN 中限制判别器的梯度大小,避免训练崩溃。
避坑指南
- 模式崩溃 :尝试 Mini-batch Discrimination 或 Unrolled GAN。
- 梯度消失 :改用 WGAN 或 SNGAN 架构。
- 评估困难 :结合 FID(Frechet Inception Distance)和人工观察。
总结与延伸
GAN 在图像生成、数据增强、风格迁移等领域展现出强大潜力。未来可探索:
- 文本生成 :如 SeqGAN 用于自然语言生成。
- 跨模态应用 :如文本到图像的生成(如 DALL-E)。
推荐进一步阅读:
–《Generative Adversarial Networks》by Ian Goodfellow
– WGAN 论文(arXiv:1701.07875)
– SNGAN 论文(arXiv:1802.05957)
正文完
