GAN(Generative Adversarial Network,生成对抗网络)是深度学习里一个看起来简单、实际上门道极多的模型家族。它在 2014 年由 Ian Goodfellow 等人提出,核心思想是让两个网络互相博弈:一个负责伪造数据,一个负责辨别真伪,最终逼着伪造者生成以假乱真的样本。这篇文章对应课程笔记中第 9.1.0 节“什么是 GAN”,我会把它展开成一份可以照着读、照着写、照着排查的学习笔记,重点讲清楚 GAN 的原理、最小实现、训练观察方法和常见坑。
如果你是刚学完神经网络基础、正准备接触生成模型的读者,这篇文章可以直接收藏。我会先给出一张概念速览表,再用 PyTorch 写一个能在 MNIST 上跑通的最小 GAN,然后解释训练过程中最容易出现的模式崩塌、判别器过强、损失不收敛等问题怎么排查。硬件门槛不高,CPU 也能跑通 MNIST 实验,有 NVIDIA GPU 训练速度会更快,显存需求以实际模型和 batch size 为准,不需要也不会在这里编造具体数字。
1. GAN 核心概念速览
在深入代码之前,先把 GAN 涉及的几个关键概念说清楚。这部分内容是整个 GAN 学习路径的地基,后面所有章节都会反复用到这些术语。
| 概念 | 一句话解释 | 在 GAN 中的作用 |
|---|---|---|
| 生成器 Generator | 输入随机噪声,输出伪造样本 | 负责学习真实数据分布,生成假图 |
| 判别器 Discriminator | 输入一张图,输出真/假概率 | 负责判断输入图片是真实数据还是生成数据 |
| 对抗训练 | 生成器和判别器交替更新 | 两者互相提升,最终达到平衡 |
| 隐空间向量 z | 一个随机向量,通常是高斯分布 | 生成器的输入,决定生成样本的特征 |
| 真实数据分布 | 训练集中图片服从的分布 | 生成器最终要逼近的目标分布 |
| 损失函数 | 二分类交叉熵的变体 | 衡量判别器判断能力、生成器欺骗能力 |
| 模式崩塌 Mode Collapse | 生成器只生成少数几种样本 | GAN 训练中常见的不稳定问题 |
| 纳什均衡 | 双方都无法继续单方面改进的状态 | 理论上 GAN 训练的最终目标 |
需要特别强调的是,GAN 不是单一模型,而是“一个生成器 + 一个判别器”的整体框架。生成器想骗过判别器,判别器想识破生成器,这个博弈过程让两个网络同步进化。最终理想状态下,判别器无法区分真实样本和生成样本,生成器输出的分布就逼近了真实数据分布。
从更广义的深度学习视角看,GAN 属于生成模型的一种,它的工作方式和自编码器、变分自编码器、扩散模型都不同。GAN 不直接优化数据似然,而是通过一个辅助网络隐式地学习数据分布。这个思路在 2014 年刚提出时被认为非常反直觉,因为之前的主流生成模型都依赖显式的概率建模。理解了这一点,你就抓住了 GAN 的本质:用对抗游戏替代显式密度估计。
2. GAN 的适用场景与使用边界
很多初学者学 GAN 时会陷入一个误区:只看公式和理论,不关注它实际能解决什么问题。这里先把适用场景和使用边界划清楚,避免你学完之后不知道这东西能用来干嘛。
GAN 的主要适用场景包括四类。第一是图像生成,这是最经典的应用方向,比如生成人脸、动漫头像、风景图片、超分辨率重建,不少图像修复任务早期用的也是 GAN 思路;第二是数据增强,当某个类别训练样本不足时,可以用 GAN 生成额外样本补充分类器或目标检测器的训练集;第三是风格迁移与图像编辑,CycleGAN 可以把照片变成油画风格,StarGAN 可以改变人脸属性;第四是异常检测和半监督学习,利用判别器对真实分布和异常样本的区分能力,可以在工业质检、金融风控场景中筛查小概率异常。从热词里的“gan图像修复”也能看出,图像修复是 GAN 落地比较多的领域。
同样重要的是使用边界。GAN 不太适合对稳定性要求极高的生产环境,因为它训练不稳定,收敛状态难以精确控制;小数据集也不适合直接上 GAN,模型容量大了很容易过拟合,模型容量小了又学不到细节;如果只是做图像生成,并且对生成质量要求非常高,当前阶段扩散模型往往是更优选择。换句话说,GAN 适合做研究学习、数据增强、风格迁移和快速原型验证,但不适合零成本白嫖高质量生成效果。
这里有两条合规红线必须说清楚。训练数据如果来自开源数据集或网络图片,要确认数据集的使用许可,商业项目尤其要注意版权授权;生成人脸、声音、证件类内容时必须取得被生成对象授权,不得用于伪造、诈骗、诽谤等非法用途。GAN 只是工具,生成内容的使用责任在开发者自己。
3. GAN 原理拆解
3.1 生成器与判别器的对抗关系
用一个容易理解的类比:生成器像造假币的团伙,判别器像验钞机。造假团伙的目标是造出验钞机认不出来的假币,验钞机的目标是找出所有假币。双方不断升级手段,最终假币无限接近真币,验钞机的识别能力也无限接近完美。把“假币”替换成“假图片”,把“验钞机”替换成“二分类神经网络”,就是 GAN 的基本结构。
生成器 G 接收一个随机向量 z,输出一张图片 G(z)。随机向量 z 通常从标准正态分布或均匀分布中采样,维度一般是 64 到 512 不等。这个低维向量可以看作是生成过程的“压缩指令”,不同位置的维度对应不同视觉特征,比如形状、朝向、背景、颜色。训练完成后,任意采样一个 z 都能生成一张新图,z 空间还被发现具有线性语义:在 z 空间沿某个方向移动,生成图像的某些属性会连续变化。
判别器 D 是一个二分类网络,输入是一张图片,输出一个 0 到 1 之间的概率。输出接近 1 表示判断为真实图片,接近 0 表示判断为生成图片。训练开始时,生成器输出的图片基本是纯噪声,判别器很容易识别出来。随着对抗训练推进,生成器逐步学会产生有结构、有内容的图片,判别器也不得不提取更精细的特征来区分真假。
3.2 对抗训练的目标函数
GAN 的目标函数是一个极小极大博弈,写成标准形式是:
min_G max_D V(D, G) = E_x[log D(x)] + E_z[log(1 - D(G(z)))]
这个公式不需要死记,理解每一部分就行。第一项 E_x[log D(x)] 表示对真实图片,判别器输出要尽量接近 1,所以 log D(x) 要尽量大;第二项 E_z[log(1 - D(G(z)))] 表示对生成图片,判别器输出要尽量接近 0,也就是 1 - D(G(z)) 尽量接近 1。判别器 D 把两项一起最大化,整体接近 0 时说明判别能力最强;生成器 G 把第二项最小化,也就是想办法让 D(G(z)) 接近 1,让判别器认为假图也是真图。
实际训练时不会同时对两个网络做梯度下降,而是交替更新。每个 batch 先冻结生成器训练判别器,再冻结判别器训练生成器。这样做的原因是对抗优化不稳定,同时更新很容易震荡。梯度交替更新虽然增加了训练时间,但能显著提升稳定性。
具体训练步骤可以拆成五步:
- 从训练集采样一批真实图片 x。
- 从正态分布采样一批随机向量 z,让生成器生成一批假图 G(z)。
- 用真实图片和假图片训练判别器:真实图片目标为 1,假图片目标为 0。
- 再采样一批随机向量 z,生成新假图,训练生成器:试图让判别器对假图输出 1。
- 重复步骤 1 到 4,直到生成图片质量满足要求。
3.3 训练何时收敛
理论上,当判别器无法区分真实图片和生成图片,即对任何输入都输出 0.5 时,GAN 达到纳什均衡。此时生成器已经学习到了真实数据分布。实际中很少能达到完美的 0.5,更常见的情况是损失曲线在某个区间小幅波动,生成图片质量稳定在可接受范围。
有一点容易误解:不能单看生成器 loss 来评估效果。生成器 loss 下降说明它在一段时间内骗过了当前判别器,但如果判别器也同步变强,生成器 loss 可能长期不降。因此判断 GAN 是否训练成功,最可靠的方法是直接看生成图片的视觉效果,其次才是看损失曲线。
4. GAN 环境准备与最小实现
4.1 环境准备
这里给一个通用检查清单,具体版本按自己的环境调整:
- 操作系统:Windows、Linux、macOS 均可,Linux 下训练最稳定。
- Python:3 即可,推荐 3.8 以上版本。
- 深度学习框架:PyTorch,安装方式参考官方命令,CPU 版也可以完成 MNIST 实验。
- NVIDIA GPU:可选,有 CUDA 环境训练会快很多。
- 磁盘空间:MNIST 数据集下载下来约几十 MB,模型文件更小,总量控制在 2GB 以内足够。
确认 PyTorch 安装成功的命令:
python -c "import torch; print(torch.__version__); print(torch.cuda.is_available())"如果输出中torch.cuda.is_available()为True,说明 GPU 可用;为False也没关系,MNIST 手写数字生成用 CPU 也能跑,只是慢一些。
4.2 项目结构与完整训练代码
下面给出一个最小 GAN 训练脚本,使用 PyTorch 在 MNIST 上训练。生成器用三层全连接网络,判别器也用三层全连接网络。这个结构不是最优的,但是最容易理解、最不容易出 bug 的版本。建议先跑通这份代码,再逐步改成 DCGAN、WGAN 或 StyleGAN。
import os import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms latent_dim = 100 batch_size = 64 epochs = 50 lr = 0.0002 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") os.makedirs("gan_outputs", exist_ok=True) transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) dataset = datasets.MNIST(root="./data", train=True, transform=transform, download=True) dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) class Generator(nn.Module): def __init__(self): super().__init__() self.model = nn.Sequential( nn.Linear(latent_dim, 256), nn.ReLU(True), nn.Linear(256, 512), nn.ReLU(True), nn.Linear(512, 1024), nn.ReLU(True), nn.Linear(1024, 28 * 28), nn.Tanh() ) def forward(self, z): img = self.model(z) return img.view(-1, 1, 28, 28) class Discriminator(nn.Module): def __init__(self): super().__init__() self.model = nn.Sequential( nn.Linear(28 * 28, 1024), nn.LeakyReLU(0.2, inplace=True), nn.Linear(1024, 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): x = img.view(img.size(0), -1) return self.model(x) generator = Generator().to(device) discriminator = Discriminator().to(device) g_optim = optim.Adam(generator.parameters(), lr=lr, betas=(0.5, 0.999)) d_optim = optim.Adam(discriminator.parameters(), lr=lr, betas=(0.5, 0.999)) criterion = nn.BCELoss() for epoch in range(epochs): for i, (imgs, _) in enumerate(dataloader): real_imgs = imgs.to(device) cur_batch = real_imgs.size(0) real_label = torch.ones(cur_batch, 1, device=device) fake_label = torch.zeros(cur_batch, 1, device=device) # 训练判别器 discriminator.zero_grad() real_pred = discriminator(real_imgs) d_real_loss = criterion(real_pred, real_label) z = torch.randn(cur_batch, latent_dim, device=device) fake_imgs = generator(z) fake_pred = discriminator(fake_imgs.detach()) d_fake_loss = criterion(fake_pred, fake_label) d_loss = d_real_loss + d_fake_loss d_loss.backward() d_optim.step() # 训练生成器 generator.zero_grad() fake_pred = discriminator(fake_imgs) g_loss = criterion(fake_pred, real_label) g_loss.backward() g_optim.step() if epoch % 10 == 0: print(f"Epoch {epoch}: D_loss={d_loss.item():.4f}, G_loss={g_loss.item():.4f}") torch.save(generator.state_dict(), f"gan_outputs/generator_{epoch}.pth")这份代码可以直接保存为gan_mnist.py运行。需要说明的是,训练判别器时生成的fake_imgs在判别器反向传播之后要继续用于生成器训练,代码里用fake_imgs.detach()切断了判别器反向传播中指向生成器的梯度,避免在更新判别器时把梯度传回生成器;生成器训练时重新使用fake_imgs更新生成器参数,这个写法是 GAN 训练的标准做法,能减少一些初学者容易搞混的重复生成问题。
4.3 运行方式与预期输出
在项目目录下执行:
python gan_mnist.py如果没有报错,意味着数据下载完成、模型初始化成功、训练循环正常推进。每 10 个 epoch 会输出一行损失值,并保存一个生成器权重文件。50 个 epoch 跑完后,gan_outputs目录下会有generator_0.pth、generator_10.pth、generator_20.pth、generator_30.pth、generator_40.pth五个权重文件。如果想更快看到图片效果,可以每 5 个 epoch 保存一次,或者训练结束后额外写一个可视化脚本,把固定噪声向量批量送入生成器,用torchvision.utils.save_image保存成一张图片。
一个值得养成的习惯是固定一个测试噪声向量集合,训练过程中用同一个 z 集合反复生成图片,这样可以直接观察生成效果随 epoch 的演变。如果没有固定向量,每次保存的图片对比起来没有对应关系,很难判断模型有没有真的变好。
5. GAN 功能验证与效果判断
5.1 判断训练效果的核心指标
GAN 没有单一的验证指标,最可靠的是“人眼观察 + 损失曲线 + 多样性检查”三者结合。损失曲线方面,判别器 loss 应该在一个合理区间波动,既不能一直趋近 0,也不能长期不下降;生成器 loss 会呈现阶梯式下降的特征,也就是每过几个 epoch 突然进步一次,这是正常的。
图片效果方面,训练前期生成的图片应该是带有笔触感的噪声团块。大约训练到 20 到 30 个 epoch 时,图中会出现隐约的数字轮廓。到 50 个 epoch 时,应该能看到清晰的 0 到 9 手写数字,但边缘可能不够干净。如果你运行 50 个 epoch 后还在输出类似雪花噪点的图片,优先检查 batch size 和学习率,并确认数据是否经过了transforms.Normalize([0.5], [0.5])处理,因为生成器最后用了Tanh,输入数据需要在 [-1, 1] 区间。
多样性检查同样重要。让生成器一次生成 64 张图,如果 64 张图长得几乎一样,只有 1 到 2 种数字,就说明出现了模式崩塌。这种情况下需要降低判别器能力或调整学习率。
5.2 输出可视化脚本
训练结束后运行下面的可视化脚本,加载最新生成器权重,用固定噪声生成一张对比图:
import torch import torchvision.utils as vutils from torchvision import transforms from PIL import Image device = torch.device("cuda" if torch.cuda.is_available() else "cpu") fixed_noise = torch.randn(64, 100, device=device) # 注意:这里的 Generator 类需要从训练脚本中复制进来 loaded_gen = Generator().to(device) loaded_gen.load_state_dict(torch.load("./gan_outputs/generator_40.pth", map_location=device)) loaded_gen.eval() with torch.no_grad(): fake = loaded_gen(fixed_noise) grid = vutils.make_grid(fake, nrow=8, normalize=True, value_range=(-1, 1)) grid = transforms.ToPILImage()(grid.cpu()) grid.save("gan_result.png") print("Saved gan_result.png")扩展一下:如果想生成单张大图保存,可以用循环把不同 z 向量输入生成器并拼接;如果想对比不同 epoch 的效果,可以对每个权重文件重复上面的生成过程,并给输出文件名加上 epoch 标记。这种方式不需要额外部署 TensorBoard,适合快速验证。
5.3 CPU 与 GPU 训练观察
MNIST 这种低分辨率数据,CPU 完全能跑,但每轮速度取决于 CPU 核心数和 PyTorch 的线程调度。有 GPU 时训练速度通常会明显提升,但要求不高。训练 REC 更应该关注的不是训练速度,而是 batch size 对显存的影响:增大 batch size 会让判别器和生成器一次处理更多图片,显存占用上升;显存有限时优先降低 batch size,或者把图片缩放到更小的分辨率。实际显存占用多少受模型结构、batch size、图片分辨率共同影响,需要以你的运行环境为准,不要死记网上别人的数字。
6. GAN 训练常见问题与排查方法
GAN 训练不稳定是出了名的,实际跑起来会遇到的问题很多,这里整理成一个排查表。这张表不仅适用于上面这份最小实现,也适用于后续扩展的 DCGAN、WGAN 等结构。
| 问题现象 | 可能原因 | 排查方式 | 解决方案 |
|---|---|---|---|
| 生成图片始终是纯噪声 | 训练轮数不够、判别器过强、学习率过大 | 观察损失曲线和中间保存的图片 | 增加 epoch、降低学习率、增加判别器正则化 |
| 生成图片只有少数几种数字 | 模式崩塌 | 让生成器生成多种 z 向量比较输出 | 引入标签平滑、使用 WGAN 损失、调整网络容量 |
| 判别器 loss 接近 0 | 判别器太强,生成器完全骗不过 | 打印梯度统计,观察生成器梯度是否消失 | 降低生成器学习率或调小判别器容量 |
| 判别器 loss 长期不变 | 学习率太低或两个网络能力不匹配 | 检查两层网络的 loss 数值 | 分别调整 G 和 D 的学习率 |
| 生成图片模糊 | 数据归一化不匹配、网络容量不足 | 检查生成器最后一层是否用了 Tanh | 修正归一化范围、增加网络层数 |
| 损失曲线剧烈震荡 | 对抗训练不稳定 | 记录 loss 的滑动平均 | 使用 Adam 的 betas=(0.5, 0.999)、降低学习率 |
| CPU 训练特别慢 | 单 batch 生成 64 张图计算量大 | 查看 CPU 占用和线程设置 | 降低 batch size、减少中间全连接层宽度 |
| 下载 MNIST 失败 | 网络访问问题 | 检查网络连通性 | 手动下载数据集并放入指定目录 |
| 显存不足 CUDA OOM | batch size 过大或模型过大 | 查看显存占用 | 降低 batch size、使用混合精度或分块训练 |
| 实验结果不稳定、每次训练效果不同 | 初始化随机性 | 固定随机种子 | 设置 PyTorch 和 Python 的随机种子再训练 |
排查时的总原则是:先确认代码能跑通,再观察损失和生成图,最后调网络结构和超参数。不要一上来就乱改网络结构,否则问题定位会变得很困难。
7. GAN 最佳实践与使用建议
7.1 超参数设置与训练策略
第一次跑 GAN 实验建议固定这样一套最小配置:Adam 优化器,学习率 0.0002,batch size 64,隐向量维度 100,不使用标签平滑。这套配置在 MNIST 上容易得到可观察的结果。跑通之后再调整任意一个变量,观察它对损失和生成质量的影响。
稳定训练的三个实用技巧值得记住。第一是标签平滑,把真实样本的标签从 1 改成 0.9,可以防止判别器过于自信从而传递过于极端的梯度。第二是给判别器加梯度惩罚或谱归一化,这是 WGAN-GP 的核心思路,能从理论上缓解训练不稳定。第三是使用较大的 batch size,有研究表明较大的批量在对抗训练中能提供更稳定的梯度估计。这些技巧可以分别实验,不要全部叠加,否则无法判断是哪个改动带来的效果。
7.2 数据版权与生成内容合规
无论是用公开数据集还是自采数据训练 GAN,都要在项目文档中记录数据来源和许可协议。如果数据来自互联网,确认平台条款是否允许用于模型训练;如果数据涉及人物肖像,必须获得本人授权。生成内容发布到公开平台时,最好在元信息中标注“AI 生成”,尤其是人脸和语音相关内容。商业部署前还要考虑内容审核机制,避免生成内容被用于欺诈或造假。合规不是附加项,而是模型落地的必要条件。
7.3 工程化扩展方向
学习阶段跑通最小 GAN 之后,可以按难度顺序往三个方向扩展。第一个方向是结构改进,从全连接网络换成卷积结构,变成 DCGAN,图片质量会明显提升。第二个方向是损失改进,把二分类交叉熵换成 WGAN 的 Wasserstein 距离或 WGAN-GP,训练稳定性会显著提高。第三个方向是条件生成,给生成器和判别器都输入标签信息,做成 Conditional GAN,这样你可以控制生成数字的种类。
再往后,如果希望把训练好的生成器部署成服务,需要把模型导出为 ONNX 或 TensorRT 格式,再封装成 HTTP 接口。部署环节有几个通用做法:验证 ONNX 导出的输出是否和 PyTorch 原模型一致;接口层需要限制请求频率和输入大小,防止被恶意调用;批量生成时建议在服务层做队列管理,逐批处理而不是一次性生成大量图片。具体接口路径和请求格式需要按你的实际项目设计和验证,这里不给虚构的示例。
8. 总结与下一步
如果只想记住 GAN 最核心的三点,那么第一点是 GAN 由生成器和判别器组成,两者通过对抗博弈共同进步;第二点是它的训练目标是一个极小极大问题,实际训练采用交替更新的方式;第三点是 GAN 的难点不在理解概念,而在稳定训练,模式崩塌和不收敛是最常见的两个问题。
这篇文章给出了一个最小可运行的 MNIST GAN 项目,从环境准备、代码实现到效果验证和问题排查全部覆盖。建议你实际操作时先把 50 个 epoch 完整跑完,保存中间权重,然后逐个观察生成效果的变化。最容易踩的坑是拿生成器 loss 或判别器 loss 单独判断训练好坏,正确做法是结合生成图片的多样性、清晰度和损失曲线整体判断。
下一步可以沿三条线继续深入:一是学习 DCGAN 和 WGAN 的改进细节,理解为什么卷积结构和梯度惩罚能让训练更稳定;二是尝试把 GAN 用在风格迁移或图像修复任务中,体验真实场景下的数据需求;三是对比 GAN 和扩散模型的生成效果差异,这会帮助你理解不同生成模型的适用边界。GAN 是深度学习中值得花时间攻克的经典方向,先把这一节基础吃透,后面的路会顺很多。