近几年生成模型发展很快,你可能常听到“AI 绘画”“人脸生成”“图像修复”这些词。背后有一个绕不开的基础模型——GAN,中文通常叫生成对抗网络(Generative Adversarial Network)。如果你刚进入深度学习的学习路线,看到“9 GAN 9.1 GAN 基础 9.1.0 什么是 GAN”这一节,第一反应可能是一堆疑问:它和普通的卷积神经网络有什么区别?为什么叫“对抗”?两个网络怎么训练?
这篇文章就把 GAN 基础部分展开讲清楚。我们会从“什么是生成模型”开始,拆解生成器和判别器的关系,推导最核心的目标函数,然后基于 PyTorch 从零实现一个能在 MNIST 手写数字数据集上训练的 GAN。代码可以直接复制运行,你也可以把结果打印出来观察每一轮生成效果的变化。无论你是初学者还是正在补深度学习基础的开发者,这篇文章都适合。
1. 背景与核心概念
1.1 什么是生成模型
深度学习模型大体可以分成两类:判别模型和生成模型。
判别模型解决的是分类、回归问题。给一张图,判断里面是猫还是狗;给一段文本,判断情感是正向还是负向。这类模型学习的是条件概率分布 (p(y|x)),也就是“输入 (x) 时输出 (y) 的概率”。
生成模型解决的是“生成数据”的问题。给模型一批真实图片,让它学会这些图片的分布规律,然后从分布中采样,得到新的、之前没见过的样本。比如让它学习一万张猫的图片,之后它能自己画出新的猫。这类模型学习的是数据本身的分布 (p(x))。
听起来很简单,难点在于:真实图片的分布极其复杂。一张 (28\times28) 的 MNIST 灰度图,每个像素取值 0 到 255,如果直接建模,就是一个几万维空间中的分布,人工写出概率密度函数根本不可能。
传统方法里,我们可以用最大似然估计去拟合一个假设的分布,比如高斯混合模型,但表达能力有限。直到 GAN 出现,才真正提供了一种“用神经网络去逼近复杂分布”的可行思路。
1.2 GAN 的核心思想:让两个网络互相对抗
GAN 是 Ian Goodfellow 等人在 2014 年提出的模型。它的核心思路不是直接计算分布,而是构造两个角色:
- 生成器(Generator,简称 G):负责从随机噪声中生成伪造样本,目标是骗过判别器。
- 判别器(Discriminator,简称 D):负责区分输入样本是来自真实数据,还是来自生成器伪造的数据。
你可以把生成器理解为一个“仿造者”,把判别器理解为一个“鉴定师”。仿造者不断改进造假技术,鉴定师不断升级识别能力。两者互相竞争、互相促进。最终状态是:仿造者造出来的东西已经足以以假乱真,鉴定师无法区分真假。
这个过程在数学上,就是一个极小极大博弈(minimax game)。GAN 的训练目标可以写成:
min_G max_D V(D, G)其中目标函数 (V(D,G)) 的含义是:最大化判别器的判别能力,同时最小化生成器的生成误差。
1.3 GAN 的典型应用场景
GAN 虽然最初是为了生成图像而生,但已经扩展到非常多领域。常见应用包括:
- 图像生成:生成人脸、风景、动漫角色等。
- 图像修复:对破损图片进行补全、去噪、去雨。
- 图像超分辨率:把低分辨率图片放大成高分辨率,并补充细节。
- 风格迁移:把照片转换成油画风格、把白天的图片转换成夜晚风格。
- 数据增强:为训练集生成额外样本,缓解样本不足的问题。
- 异常检测:利用生成器的重建误差定位图像中的异常区域。
这些应用大多是在 GAN 基础理论上发展出来的。所以先把 9.1 的“什么是 GAN”学扎实,后面的进阶版本才会更容易理解。
2. GAN 核心原理拆解
2.1 生成器 G:从噪声到样本
生成器的输入是一个随机噪声向量 (z),常见维度是 100 或 128,服从标准正态分布。经过若干层神经网络的映射,输出一个和真实样本同尺寸的向量。对于 MNIST,输出就是 784 维的向量,再 reshape 成 (1 \times 28 \times 28) 的图像。
生成器的本质是:
[ G(z; \theta_g) ]
其中 (\theta_g) 是生成器的参数。训练初期,生成器输出的是纯噪声;随着训练进行,参数不断调整,输出越来越接近真实图片。它没有见过任何真实图片的“像素级标签”,只是通过判别器的反馈来改进。
这里有一个初学者容易困惑的点:生成器并没有直接计算“这张图像不像 MNIST”,它只是根据判别器给它的梯度信号来调整自己的参数。这个信号从哪来?来自判别器对它的评价。
2.2 判别器 D:真假鉴定器
判别器是一个二分类网络。输入一张图片 (x),输出一个标量 (D(x)),表示这张图片是真实图片的概率。
- 如果输入来自真实数据集,(D(x)) 的期望目标是 1。
- 如果输入来自生成器 (G(z)),(D(G(z))) 的期望目标是 0。
判别器的本质是:
[ D(x; \theta_d) ]
它的结构和普通分类网络类似,可以是一个卷积网络,也可以是一个全连接网络。在简单实验中,全连接网络往往就够用。
训练过程中,生成器和判别器是交替更新的。不能同时更新,也不能只更新其中一个,否则博弈会失衡。
2.3 目标函数与训练流程
原始 GAN 的目标函数是:
[ \min_G \max_D V(D,G) = \mathbb{E}{x \sim p{data}(x)}[\log D(x)] + \mathbb{E}_{z \sim p_z(z)}[\log(1 - D(G(z)))] ]
拆开看:
- 第一个期望项:真实样本 (x) 输入判别器,我们希望 (\log D(x)) 尽量大,也就是判别器对真实样本输出接近 1。
- 第二个期望项:固定真实样本,对随机噪声 (z),我们希望 (\log(1 - D(G(z)))) 尽量大,也就是判别器对伪造样本输出接近 0。
这是判别器的视角。但从生成器的角度看,它希望 (D(G(z))) 输出接近 1,也就是让判别器认不出伪造样本,所以它希望第二项中 (\log(1 - D(G(z)))) 尽量小。
那么训练流程就是:
- 从真实数据集中采样一批真实样本。
- 从随机噪声分布中采样一批噪声,用生成器生成一批伪造样本。
- 把真实样本和伪造样本一起输入判别器,计算判别损失,更新判别器参数。
- 固定当前判别器参数,把新的噪声输入生成器,计算生成器损失,更新生成器参数。
- 重复上述过程,直到两者达到纳什均衡。
在实现时,初学者最容易犯的错误是:把训练判别器和训练生成器合并成一次反向传播。正确的做法是交替更新,并且训练生成器时要避免让判别器这一路的梯度干扰生成器参数更新。
另外,还有一个常见的“负梯度消失”问题。原始损失函数 (\log(1 - D(G(z)))) 在判别器过于强大时,梯度会非常小,导致生成器几乎学不动。所以在实际实现中,通常会把生成器的目标改为最大化 (\log(D(G(z)))),也就是让生成器尝试“让判别器认为伪造样本是真的”。这个改进在代码里很常见,但理解了原始公式后,你会知道这本质上只是在优化目标上做了一个等价变换,目的是提供更好的梯度信号。
3. 环境准备与实验约定
3.1 软件环境
本文代码基于 PyTorch 编写。版本需要根据你的项目实际情况调整,本文示例以常见环境为例,重点演示配置思路。
- 操作系统:Windows / Linux / macOS 均可。
- Python:建议 3.8 及以上版本。
- PyTorch:建议 1.13 或 2.x 系列,CPU 版本即可运行,有 GPU 会更快。
- torchvision:与 PyTorch 版本匹配,用于加载 MNIST 数据集。
- matplotlib:用于可视化生成效果。
对于初学者,尤其注意不要只安装 torch 而忘了 torchvision。MNIST 数据集下载和预处理依赖 torchvision 的接口。
3.2 项目结构
为了便于管理,建议按下面的结构组织代码:
gan-tutorial/ ├── train.py # 主训练脚本 ├── generator.py # 生成器定义 ├── discriminator.py # 判别器定义 └── result/ # 保存生成的图片当然,为了减少文件依赖,本文会给出一个可以直接保存为train.py的完整脚本。你也可以在 Jupyter Notebook 中按小节逐步运行。
4. 从零实现一个 GAN
下面我们基于 PyTorch 实现一个最简单的 GAN,用于生成 MNIST 手写数字。这个版本重点突出基础流程,网络结构保持简单,方便你观察每一步在做什么。
4.1 数据加载与预处理
MNIST 是 28×28 的灰度图数据集。常见做法是把像素值从 [0,1] 归一化到 [-1,1]。这样做的原因是生成器最后使用 Tanh 激活函数,输出范围也是 [-1,1],两者匹配,训练更加稳定。
import torch import torch.nn as nn import torch.optim as optim import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader import os # 超参数 latent_dim = 100 # 噪声向量维度 batch_size = 128 epochs = 50 lr = 0.0002 # 学习率 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 数据预处理:转换为 Tensor,并归一化到 [-1, 1] transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset = torchvision.datasets.MNIST( root="./data", train=True, transform=transform, download=True ) train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)代码说明:
transforms.Normalize((0.5,), (0.5,))中的均值 0.5、标准差 0.5,对应 MNIST 灰度图归一化到 [-1,1] 的常见参数。download=True表示本地没有数据时会自动下载。如果你所在环境无法访问外网,可以提前把 MNIST 数据文件放到./data目录下。
4.2 定义生成器
生成器接收 100 维随机噪声,经过三层全连接层,最终输出 784 维向量,再 reshape 成 1×28×28 图像。
class Generator(nn.Module): def __init__(self, latent_dim=100): super(Generator, self).__init__() self.model = nn.Sequential( nn.Linear(latent_dim, 128), nn.ReLU(inplace=True), nn.Linear(128, 256), nn.ReLU(inplace=True), nn.Linear(256, 512), nn.ReLU(inplace=True), nn.Linear(512, 784), nn.Tanh() ) def forward(self, z): # z 形状: (batch_size, latent_dim) img = self.model(z) # 输出形状: (batch_size, 1, 28, 28) img = img.view(img.size(0), 1, 28, 28) return img几个要点:
- 隐藏层激活函数用 ReLU。ReLU 计算简单且梯度不容易饱和。
- 输出层用 Tanh,输出范围 [-1,1],与数据预处理一致。
view操作把一维向量恢复成图像形状,方便后续输入判别器和保存图片。
4.3 定义判别器
判别器接收 1×28×28 图像,先展平成 784 维向量,再经过全连接层,最终输出一个概率值。
class Discriminator(nn.Module): def __init__(self): super(Discriminator, self).__init__() self.model = nn.Sequential( nn.Linear(784, 512), nn.LeakyReLU(0.2, inplace=True), nn.Linear(512, 256), nn.LeakyReLU(0.2, inplace=True), nn.Linear(256, 1), nn.Sigmoid() ) def forward(self, img): # img 形状: (batch_size, 1, 28, 28) x = img.view(img.size(0), -1) return self.model(x)这里使用 LeakyReLU 而不是 ReLU,原因是判别器面对的是真假样本分布差异很大,LeakyReLU 在负半轴保留了一个很小的梯度,不容易让神经元“死掉”。0.2是负半轴斜率,是常见配置。
最终的 Sigmoid 输出是一个 0 到 1 的概率值,表示输入为真实样本的置信度。
4.4 定义损失函数与优化器
二分类任务使用 BCEWithLogitsLoss 或 BCELoss。由于判别器最后已经有 Sigmoid,这里使用 BCELoss 比较直观。
# 损失函数 criterion = nn.BCELoss() # 优化器 optimizer_G = optim.Adam(generator.parameters(), lr=lr, betas=(0.5, 0.999)) optimizer_D = optim.Adam(discriminator.parameters(), lr=lr, betas=(0.5, 0.999)) # 初始化网络 generator = Generator(latent_dim).to(device) discriminator = Discriminator().to(device)Adam 是 GAN 训练中最常见的优化器。值得注意的是betas=(0.5, 0.999),这和默认值(0.9, 0.999)略有区别。设置第一动量系数为 0.5,可以减小训练过程中的惯性,让模型对参数变化更敏感,这在 GAN 场景中被广泛采用。
如果你用的是 PyTorch 2.x,Lazy模块会触发形状推理,但这里没有使用,避免兼容性问题。
4.5 训练循环核心代码
训练过程中,每一轮迭代包含两部分:先更新判别器,再更新生成器。二者交替进行。
generator.train() discriminator.train() for epoch in range(epochs): for batch_idx, (real_imgs, _) in enumerate(train_loader): current_batch = real_imgs.size(0) real_imgs = real_imgs.to(device) # 真实样本标签为 1,伪造样本标签为 0 real_labels = torch.ones(current_batch, 1, device=device) fake_labels = torch.zeros(current_batch, 1, device=device) # ---------- 训练判别器 ---------- # 从正态分布采样噪声 z = torch.randn(current_batch, latent_dim, device=device) fake_imgs = generator(z) # 判别器对真实样本的损失 real_loss = criterion(discriminator(real_imgs), real_labels) # 判别器对伪造样本的损失 fake_loss = criterion(discriminator(fake_imgs.detach()), fake_labels) # 总损失是两者平均 d_loss = (real_loss + fake_loss) / 2 optimizer_D.zero_grad() d_loss.backward() optimizer_D.step() # ---------- 训练生成器 ---------- # 重新采样噪声,生成新伪造样本 z = torch.randn(current_batch, latent_dim, device=device) fake_imgs = generator(z) # 生成器希望判别器把伪造样本识别为 1 g_loss = criterion(discriminator(fake_imgs), real_labels) optimizer_G.zero_grad() g_loss.backward() optimizer_G.step()解释一下几个容易出错的地方:
fake_imgs.detach():训练判别器时,我们不需要更新生成器。如果不 detach,梯度会同时回传到生成器,导致一次迭代更新了两个网络,逻辑上混成一锅粥。- 训练生成器时,重新采样了
z,也可以直接复用上一份fake_imgs。但确保不再对判别器做 detach,因为生成器需要从判别器输出的梯度中学习。 real_loss和fake_loss求平均,是原始论文中的做法,可以避免 batch 内真假样本数量不一致带来的偏差。
4.6 保存与可视化生成效果
每个 epoch 结束后,把当前生成器产出的图片保存到result/目录,方便观察训练效果。
def save_samples(generator, epoch, save_dir="./result", num_samples=16): os.makedirs(save_dir, exist_ok=True) generator.eval() with torch.no_grad(): z = torch.randn(num_samples, latent_dim, device=device) samples = generator(z).cpu() samples = (samples + 1) / 2 # 从 [-1,1] 映射回 [0,1] grid = torchvision.utils.make_grid(samples, nrow=4) torchvision.utils.save_image(grid, os.path.join(save_dir, f"epoch_{epoch:03d}.png")) generator.train()说明:
- 推理时使用
torch.no_grad(),关闭梯度计算,节省内存并避免意外把噪声数据的梯度传入网络。 make_grid将多张图片拼接成一张网格图,便于直接查看。- 从 [-1,1] 映射回 [0,1] 是为了保存成正常可见的灰度图片。
在训练循环内部,每个 epoch 结束后调用一次:
print(f"Epoch [{epoch+1}/{epochs}] D loss: {d_loss.item():.4f}, G loss: {g_loss.item():.4f}") save_samples(generator, epoch + 1)4.7 运行与预期结果
将以上代码完整保存为train.py,然后运行:
python train.py如果没有 GPU,训练速度会比较慢。你可以把 epochs 调小,例如 10 或 20,先验证整个流程能跑通。
预期你会看到类似输出:
Epoch [1/50] D loss: 0.6912, G loss: 0.6823 Epoch [2/50] D loss: 0.5801, G loss: 1.2301 ...第一次运行时,如果目录中还没有 MNIST 数据,PyTorch 会自动下载。输出 loss 时,前几轮 D loss 和 G loss 都会在 0.6 到 1.0 附近徘徊。随着训练推进,你再打开result/epoch_010.png,会发现图片从一团噪声慢慢变成模糊的数字轮廓。
这里需要特别说明:loss 数值本身并不能直接说明生成质量。判别器 loss 降到很低,可能是判别器“太强”了,生成器跟不上;生成器 loss 很低,也不代表它生成的图片清晰、多样。GS 网络的评估最终要看视觉效果,这也是 GAN 训练中比较复杂的地方。
5. 训练过程中的常见问题与排查思路
初学者第一次训练 GAN,大概率会踩到下面几个坑。我们把常见问题整理成一张排查表。
| 问题现象 | 常见原因 | 解决思路 |
|---|---|---|
| 训练刚开始 loss 很快降到接近 0 | 判别器过于强大,生成器完全被骗不过去 | 调小判别器学习率、减少判别器更新次数,或增加生成器网络容量 |
| 生成器 loss 持续很高,图片仍然是噪声 | 梯度消失,生成器训练信号太弱 | 将生成器损失从最小化log(1-D(G(z)))改为最大化log(D(G(z))),并检查是否存在梯度为 0 的情况 |
| 生成图片只有少数几种数字 | 模式坍塌(Mode Collapse) | 尝试更小的学习率、使用标签平滑、增大噪声维度、换用更复杂的架构 |
| 生成图片有棋盘格纹路 | 上采样操作不合理 | 在卷积网络中使用ConvTranspose2d后跟上卷积层消除重叠,或改用Upsample配合卷积 |
| D loss 和 G loss 剧烈震荡 | 模型不稳定 | 降低学习率、增大 batch size、加入批归一化层 |
| 图片整体偏亮或偏暗 | 数据归一化与生成器激活函数不匹配 | 确认真实图片是否归一化到 [-1,1],生成器是否使用 Tanh 输出 |
其中模式坍塌是 GAN 训练中最经典也最棘手的问题。它表现为生成器只找到了一小部分能被判别器接受的输出,于是反复生成同一类图片。比如 MNIST 只生成“1”或“7”,缺少多样性。解决模式坍塌没有万能药,通常需要综合调节网络结构、损失函数和训练策略。
另外,初学者还容易遇到一个现象:前几轮图片噪声明显,之后突然变清晰然后又变模糊。这通常是训练过程不稳定的表现。建议你记录每个 epoch 的图片,不要只看最后一个 epoch 的结果。保存历史输出这一点,在你排查问题时能起到很大作用。
6. 最佳实践与工程建议
6.1 超参数选择与训练稳定性
根据大量 GAN 工程的实践经验,以下超参数是一个比较稳妥的起点:
- 学习率:0.0002。
- Adam 动量:betas = (0.5, 0.999)。
- batch size:64 或 128。
- 噪声维度:100 到 128。
- 隐藏层激活:生成器用 ReLU,判别器用 LeakyReLU。
- 输出层激活:生成器用 Tanh,判别器用 Sigmoid。
在训练时,可以适当“让判别器慢一点”。比如每更新一次生成器,只更新一次判别器;如果判别器 loss 下降太快,就调整为每两次生成器更新后再更新一次判别器。
6.2 标签平滑技巧
把真实样本的目标标签从 1 改成 0.9,这是一个非常简单但有效的技巧,称为标签平滑(Label Smoothing)。它能防止判别器对真实样本过于自信,从而给生成器留出更多改进空间。代码改动只需一行:
real_labels = torch.ones(current_batch, 1, device=device) * 0.96.3 从全连接网络升级到卷积网络
上面的示例使用全连接层,便于理解,但效果有限。真正应用时,推荐使用 DCGAN(Deep Convolutional GAN)架构:
- 生成器用
ConvTranspose2d逐步放大特征图。 - 判别器用
Conv2d逐步缩小特征图。 - 大量使用
BatchNorm2d和 LeakyReLU。 - 不在判别器中使用池化层,而是用带步长的卷积替代。
DCGAN 训练更稳定,生成图片也更清晰。理解基础 GAN 之后,可以把网络结构替换成 DCGAN,超参数基本不用大改,效果立刻会有明显提升。
6.4 从原始 GAN 到 WGAN
原始 GAN 的损失函数基于 JS 散度,当真实分布和生成分布重叠很少时,JS 散度会失去梯度方向。WGAN 改用 Wasserstein 距离,从理论上缓解了训练不稳定和模式坍塌问题。随后出现的 WGAN-GP 又加了梯度惩罚项,是目前应用中最常被使用的 GAN 变体之一。
所以,当你在基础 GAN 上遇到难调的稳定性问题时,不用死磕原始损失函数,可以直接跳到 WGAN 系列学习更先进的思路。
6.5 项目落地时的建议
如果你要把 GAN 用在真实项目中,下面几条建议值得记住:
- 数据隐私:训练数据可能包含人脸、身份信息,使用前需要获得合规授权。
- 生成内容审核:GAN 生成图片可能被滥用,务必在应用层加入内容过滤与来源标注。
- 模型评估:不能只看 loss,要建立评价指标,比如 FID(Fréchet Inception Distance)用于衡量真实图片和生成图片的分布距离。
- 环境一致性:使用
torch.save保存模型时,同时保存生成器和判别器的state_dict,并记录当时的超参数,方便复盘。
如果你用自己的数据集训练,注意先做数据清洗。GAN 对脏数据非常敏感,样本中混入错误标签或大量重复图片,都会让生成结果异常。
7. 进阶学习路线
如果你已经能跑通上面的代码,GAN 基础算是入门了。接下来建议按这条路线深入:
- 先复现 DCGAN,观察卷积层如何提升生成质量。
- 再学习条件 GAN(Conditional GAN,cGAN),在生成时加入标签信息,控制生成的数字类别。
- 学习 WGAN 和 WGAN-GP,理解损失函数改进背后的数学直觉。
- 尝试 CycleGAN,实现未配对图像风格迁移。
- 阿里感兴趣的方向,可以继续看 StyleGAN,这是人脸生成领域非常稳定的架构。
每一步都可以配合小实验来验证。比如在 MNIST 上先加一个条件标签,让生成器输出指定数字;再换到真实照片数据集,观察模型对复杂分布的适应能力。
GAN 的入门曲线其实不算特别高,真正的门槛在于训练稳定性和调参经验。把上面的示例代码完整跑一遍,再手动改几次超参数,你很快就能体会到“博弈”这两个字在代码里到底是怎么发生的。等你能稳定训练出一个像样的 GAN,深度学习生成模型这条支线就算真正走通了。