GAN实战:从Goodfellow原始论文到生产环境避坑指南

1次阅读
没有评论

共计 2497 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

技术背景

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

GAN 实战:从 Goodfellow 原始论文到生产环境避坑指南

核心痛点

模式崩溃(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)

关键实现细节:

  1. 噪声采样采用标准正态分布:

    z = torch.randn(batch_size, latent_dim)

  2. 生成器损失计算:

    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 计算

  1. 使用预训练的 Inception v3 模型提取特征
  2. 计算生成样本的条件概率分布 $p(y|x)$
  3. 计算 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()

避坑实践指南

  1. 更新频率比:通常判别器更新次数是生成器的 2 - 5 倍,但需监控梯度幅度

  2. 批量归一化警告:生成器最后一层避免使用 BN,否则会导致样本间过度关联

  3. 显存优化技巧

  4. 使用梯度累积(accumulate gradients)
  5. 降低 torch.float32torch.float16
  6. 微批次处理:
    for micro_batch in torch.split(big_batch, chunk_size):
        optimize(micro_batch)

开放讨论

  1. 如何设计更适合文本生成任务的 GAN 变体?
  2. 在医学图像生成中,怎样平衡生成质量与模式覆盖?
  3. 自监督学习能否与 GAN 框架有效结合?
正文完
 0
评论(没有评论)