共计 2497 个字符,预计需要花费 7 分钟才能阅读完成。
技术背景
2014 年,Ian Goodfellow 等人发表的《Generative Adversarial Networks》开创性地提出了生成对抗网络框架。其核心思想是通过生成器 (Generator) 和判别器 (Discriminator) 的对抗训练,最终让生成器能够产生与真实数据分布难以区分的样本。相比 VAE(变分自编码器)需要显式建模概率分布,或 Flow-based 模型要求可逆变换,GAN 直接通过对抗过程学习数据分布,具有更强的表达能力。

核心痛点
模式崩溃(Mode Collapse)
数学上可表示为生成器 $G$ 找到判别器 $D$ 的局部最优点:
$$ \min_G \max_D V(D,G) = \mathbb{E}{x\sim p[\log(1-D(G(z)))] $$
当生成器仅产生有限几种样本就能欺骗判别器时,就会放弃学习完整数据分布。}}[\log D(x)] + \mathbb{E}_{z\sim p_z
梯度消失问题
当判别器过于强大时,$D(G(z))$ 趋近于 0,导致生成器梯度 $\nabla_G \log(1-D(G(z)))$ 消失。实验表明,使用 $-\log D(G(z))$ 作为替代损失可缓解该问题。
PyTorch 实现原始 GAN
import torch
import torch.nn as nn
class Generator(nn.Module):
def __init__(self, latent_dim=100):
super().__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, 784), # MNIST 图像展平尺寸
nn.Tanh() # 输出归一化到[-1,1]
)
def forward(self, z):
return self.main(z)
# 判别器使用 sigmoid 输出 0 - 1 之间的概率值
class Discriminator(nn.Module):
def __init__(self):
super().__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):
x = x.view(x.size(0), -1)
return self.main(x)
关键实现细节:
-
噪声采样采用标准正态分布:
z = torch.randn(batch_size, latent_dim) -
生成器损失计算:
g_loss = torch.log(1 - D(fake_images)).mean() # 原始论文公式 # 实际常用替代形式:g_loss = -torch.log(D(fake_images)).mean()
生产级优化方案
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
)[0]
gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
return gradient_penalty
标签平滑实现
real_labels = torch.full((batch_size, 1), 0.9, device=device) # 替换原始 1.0
fake_labels = torch.full((batch_size, 1), 0.1, device=device) # 替换原始 0.0
评估与可视化
Inception Score 计算
- 使用预训练的 Inception v3 模型提取特征
- 计算生成样本的条件概率分布 $p(y|x)$
- 计算 KL 散度:
$$ \exp(\mathbb{E}_x KL(p(y|x) || p(y))) $$
结果可视化
import matplotlib.pyplot as plt
def plot_images(images, n_cols=5):
plt.figure(figsize=(10, 10))
for i in range(n_cols**2):
plt.subplot(n_cols, n_cols, i+1)
plt.imshow(images[i].detach().cpu().numpy(), cmap='gray')
plt.axis('off')
plt.tight_layout()
plt.show()
避坑实践指南
-
更新频率比:通常判别器更新次数是生成器的 2 - 5 倍,但需监控梯度幅度
-
批量归一化警告:生成器最后一层避免使用 BN,否则会导致样本间过度关联
-
显存优化技巧:
- 使用梯度累积(accumulate gradients)
- 降低
torch.float32到torch.float16 - 微批次处理:
for micro_batch in torch.split(big_batch, chunk_size): optimize(micro_batch)
开放讨论
- 如何设计更适合文本生成任务的 GAN 变体?
- 在医学图像生成中,怎样平衡生成质量与模式覆盖?
- 自监督学习能否与 GAN 框架有效结合?
正文完
发表至: 未分类
近两天内
