
GANGenerative 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 下训练最稳定。Python3 即可推荐 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_okTrue) transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) dataset datasets.MNIST(root./data, trainTrue, transformtransform, downloadTrue) dataloader DataLoader(dataset, batch_sizebatch_size, shuffleTrue) 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, inplaceTrue), nn.Linear(1024, 512), nn.LeakyReLU(0.2, inplaceTrue), nn.Linear(512, 256), nn.LeakyReLU(0.2, inplaceTrue), 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(), lrlr, betas(0.5, 0.999)) d_optim optim.Adam(discriminator.parameters(), lrlr, 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, devicedevice) fake_label torch.zeros(cur_batch, 1, devicedevice) # 训练判别器 discriminator.zero_grad() real_pred discriminator(real_imgs) d_real_loss criterion(real_pred, real_label) z torch.randn(cur_batch, latent_dim, devicedevice) 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(fEpoch {epoch}: D_loss{d_loss.item():.4f}, G_loss{g_loss.item():.4f}) torch.save(generator.state_dict(), fgan_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, devicedevice) # 注意这里的 Generator 类需要从训练脚本中复制进来 loaded_gen Generator().to(device) loaded_gen.load_state_dict(torch.load(./gan_outputs/generator_40.pth, map_locationdevice)) loaded_gen.eval() with torch.no_grad(): fake loaded_gen(fixed_noise) grid vutils.make_grid(fake, nrow8, normalizeTrue, 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 OOMbatch size 过大或模型过大查看显存占用降低 batch size、使用混合精度或分块训练实验结果不稳定、每次训练效果不同初始化随机性固定随机种子设置 PyTorch 和 Python 的随机种子再训练排查时的总原则是先确认代码能跑通再观察损失和生成图最后调网络结构和超参数。不要一上来就乱改网络结构否则问题定位会变得很困难。7. GAN 最佳实践与使用建议7.1 超参数设置与训练策略第一次跑 GAN 实验建议固定这样一套最小配置Adam 优化器学习率 0.0002batch 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 是深度学习中值得花时间攻克的经典方向先把这一节基础吃透后面的路会顺很多。