刚接触深度学习时我被 GAN 卡了很长时间。看完原理介绍感觉已经懂了一个生成器一个判别器两个网络互相博弈。但真正开始训练却发现 loss 忽高忽低生成图像要么全是噪点要么几十张输出长得一模一样。后来把课程里 9.1.1 这一节反复梳理才意识到问题不在网络结构有多深而在没有把 GAN 最核心的“训练逻辑”想透。这篇文章把 GAN 基础里的训练逻辑拆开讲配合一份可运行的 PyTorch 实战代码从概念到排错一次说清楚。如果你正在跟着深度学习系列学习恰好读到 GAN 章节或者你想从头动手实现一个最简单的 GAN这篇笔记可以直接参照。读完你会掌握 GAN 两大网络的职责、交替训练的顺序、损失函数的选择以及训练不稳定时的排查思路。1. GAN 是什么从对抗思想说起GAN 全称是 Generative Adversarial Network中文叫生成对抗网络由 Ian Goodfellow 等人于 2014 年提出。它的核心思想不是让一个网络单独完成生成任务而是让两个网络互相博弈、共同进步。1.1 生成器与判别器的角色GAN 包含两个网络生成器Generator简称 G输入一个随机噪声向量输出一张尽可能接近真实数据的图像。它的目标是“制造出以假乱真的样本”。判别器Discriminator简称 D输入一张图像输出一个 0 到 1 之间的概率值用来判断输入是真实样本还是生成样本。它的目标是“准确识别出伪造样本”。可以用一个更生活化的例子来理解生成器像一个画师判别器像一个鉴定师。画师不断练习画出的作品越来越像真迹鉴定师不断提高眼力努力分辨哪张是真画、哪张是赝品。两者在对抗过程中都会变强最终生成器画出的作品足以让鉴定师难以判断。在 GAN 基础部分我们关注的就是这种动态博弈过程生成器负责“骗过”判别器判别器负责“揭穿”生成器两者交替优化最终达到一种均衡状态。1.2 GAN 能做什么GAN 最典型的应用是图像生成但它的能力远远不止于此。常见方向包括图像生成从随机噪声生成人脸、风景、手写数字等图像。图像修复补全破损照片、去除水印区域并生成合理内容。图像超分辨率把低分辨率图像通过生成模型恢复成高分辨率图像。风格迁移把照片转成油画风格、漫画风格等。数据增广在样本量不足时生成新样本辅助其他模型训练。语音与音乐生成生成自然语音、旋律等。近两年扩散模型Diffusion Model在生成领域越来越热门但 GAN 作为生成模型的重要分支仍然是深度学习入门阶段必须掌握的核心内容。理解 GAN 的训练逻辑对后续学习各类生成模型非常有帮助。1.3 GAN 的数学目标GAN 的训练目标可以用一个极小极大博弈公式来表示min_G max_D V(D, G) E_{x~p_data}[log D(x)] E_{z~p_z}[log(1 - D(G(z)))]这里不用被公式吓到拆开来看第一项表示判别器对真实样本 x 的输出越接近 1log D(x) 越大。第二项表示判别器对生成样本 G(z) 的输出越接近 0log(1 - D(G(z))) 越大。判别器 D 希望最大化整个式子也就是让自己能正确区分真假生成器 G 希望最小化整个式子也就是让第二项变小让判别器把生成样本误判为真的概率变大。这个公式虽然简洁却包含了训练逻辑的完整答案整个训练过程不是一次前向传播就能完成的而是需要交替优化两个网络。下一步我们把环境准备好然后用代码一步步验证这个逻辑。2. 环境准备与实验设置在动手写代码之前先把实验环境准备齐全。本文以 PyTorch 作为实现框架因为它语法清晰社区资料丰富非常适合学习 GAN 基础。2.1 开发环境与依赖本文代码在以下环境中可以正常运行操作系统Windows、Linux、macOS 均可。Python3.8 及以上版本推荐 3.10。PyTorch1.8 及以上版本2.x 同样兼容。torchvision跟 PyTorch 配套安装。matplotlib用于绘制损失曲线和查看结果。CUDA可选有 GPU 会明显加速CPU 也能跑通只是速度慢一些。安装基础依赖的命令如下pip install torch torchvision matplotlib如果你使用 GPU可以到 PyTorch 官网选择对应 CUDA 版本的安装命令。这里不做具体版本指定因为 PyTorch 版本更新频繁按官网提示安装即可。CPU 版直接执行上面的默认安装命令也能运行本文全部代码。2.2 数据集选择本文选择 MNIST 手写数字数据集。MNIST 是 28×28 的灰度图像包含 0 到 9 十类数字训练集有 6 万张图像。选择它有三个原因图像尺寸小网络可以设计得比较浅训练速度快。数据下载方便torchvision 内置了 MNIST 接口。灰度图可视化直观生成结果好不好一眼就能看出来。代码中会把图像像素值归一化到 [-1, 1]这样配合生成器最后的 Tanh 激活函数输出范围一致。2.3 项目结构建议创建一个干净的目录方便后续扩展gan-demo/ ├── generator.py # 生成器网络定义 ├── discriminator.py # 判别器网络定义 ├── train.py # 训练逻辑入口 ├── output/ # 保存训练过程中的生成图像 └── data/ # 数据集缓存目录下面我们就从生成器和判别器开始写起然后重点分析训练逻辑。3. GAN 的训练逻辑拆解GAN 最精华的部分不是网络结构而是训练循环的写法。很多人把生成器和判别器写好后训练时不知道谁先更新、梯度要往哪里传播结果模型始终不收敛。本节把训练逻辑拆成清晰的几个步骤。3.1 一次迭代的整体流程每一轮迭代里需要交替做两件事更新判别器 D让它学会区分当前生成器生成的假样本和真实样本。更新生成器 G让它根据判别器的反馈调整生成方向让假样本更接近真实分布。如果只训练 D不训练 G那么 D 会变得非常强G 永远无法骗过它如果只训练 G不训练 D那么 G 失去评判标准也没有明确的优化方向。所以两者必须交替更新。3.2 第一步更新判别器训练判别器时我们把两类样本喂给它真实样本标签设为 1表示判别器应该判断为“真”。生成器当前输出的假样本标签设为 0表示判别器应该判断为“假”。对应损失函数可以用二分类交叉熵表示D_loss BCE(D(real_img), 1) BCE(D(G(z)), 0)这里有一个非常重要的细节计算假样本损失时生成器输出 G(z) 要调用 detach() 方法切断梯度。因为我们这一阶段只想更新判别器的参数不想让梯度流回生成器。如果忘记 detach()梯度会同时经过两个网络更新逻辑就会混乱。训练判别器时真实样本和假样本共享同一个判别器。在 PyTorch 中我们可以先分别计算这两部分的损失再相加后反向传播。下面给出判别器更新阶段的核心代码片段# 训练判别器 d_optimizer.zero_grad() # 真实样本部分 real_output D(real_imgs) d_real_loss criterion(real_output, real_label) # 生成样本部分 z torch.randn(batch_size, latent_dim).to(device) fake_imgs G(z) fake_output D(fake_imgs.detach()) d_fake_loss criterion(fake_output, fake_label) # 两者相加得到总损失 d_loss d_real_loss d_fake_loss d_loss.backward() d_optimizer.step()这段代码的关键点是 fake_imgs.detach()。它保证生成器网络本身不会收到来自判别器的梯度从而实现“只更新 D不更新 G”的目标。3.3 第二步更新生成器训练生成器时思路刚好反过来先生成一批假样本 G(z)。输入判别器得到输出 D(G(z))。把输出和“真实标签 1”计算损失。也就是说生成器希望判别器把假样本误判为真。损失函数可以写成G_loss BCE(D(G(z)), 1)初学者容易在这里产生误解。有人会想生成器不是要生成假样本吗那为什么不用标签 0 去训练原因在于如果生成器以“被判为假”为目标它会学习如何生成更明显的假样本这正好背离了我们的需求。生成器的目标是让判别器分不清真假所以它要学会的是让 D(G(z)) 的输出尽可能接近 1。在更新生成器时fake_imgs 不能 detach()因为我们需要梯度从判别器流回生成器让生成器知道往哪个方向调整。代码片段如下# 训练生成器 g_optimizer.zero_grad() z torch.randn(batch_size, latent_dim).to(device) fake_imgs G(z) fake_output D(fake_imgs) g_loss criterion(fake_output, real_label) g_loss.backward() g_optimizer.step()注意这里 real_label 仍然是 1只是被用来监督生成器的输出。3.4 交替训练与平衡状态在一次迭代中先更新 D再更新 G这就完成了一次交替训练。为什么不先把 D 训练到完全收敛再训练 G因为那样 D 会变得过于强大所有生成样本都会被轻松识破G 的梯度非常小几乎学不到东西。GAN 的训练必须保持两边的“动态对抗”。从博弈论角度看理想的终点是纳什均衡状态判别器无法区分真实样本和生成样本任何输入都输出 0.5。此时 D 的损失会稳定在某个合理区间不会继续下降而生成图像已经足够真实。真实训练中很难完美到达这个均衡点因此我们会看到损失曲线不是平稳下降而是带着波动变化。这也是 GAN 训练和普通分类模型最大的区别之一。3.5 损失函数为什么使用 BCELossBCELoss 全称是 Binary Cross Entropy Loss即二分类交叉熵损失。它的公式是L -[y * log(p) (1 - y) * log(1 - p)]其中 y 是真实标签p 是模型预测概率。当预测越接近真实标签时损失越小。GAN 中的判别器相当于一个二分类器输出经过 Sigmoid 后在 0 到 1 之间因此 BCELoss 是天然的匹配选择。PyTorch 中可以直接使用criterion nn.BCELoss()在后面实战代码中我们会给真样本标签设置为 1假样本标签设置为 0然后分别计算损失。4. PyTorch 实现 GAN 完整实战现在把前面讲的训练逻辑落成一份完整代码。为了保持思路清晰我们把生成器、判别器和训练循环分别放在不同文件中。4.1 定义生成器网络生成器输入一个 100 维的随机噪声向量输出一张 28×28 的灰度图像。结构上使用三层全连接中间使用 ReLU最后一层使用 Tanh。Tanh 的输出范围是 [-1, 1]与数据预处理保持一致。# 文件路径generator.py import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, latent_dim100): super(Generator, self).__init__() self.model nn.Sequential( nn.Linear(latent_dim, 256), nn.ReLU(inplaceTrue), nn.Linear(256, 512), nn.ReLU(inplaceTrue), nn.Linear(512, 28 * 28), nn.Tanh() ) def forward(self, z): img self.model(z) img img.view(-1, 1, 28, 28) return img这里 latent_dim 表示噪声向量的维度也是生成器输入向量的长度。第一层 Linear 把 100 维映射到 256 维再放大到 512 维最后输出 784 个数对应一张 28×28 图像。view 操作将一维向量重塑为图像张量方便送入判别器和保存图像。4.2 定义判别器网络判别器接收 28×28 图像输出一个概率值。内部使用 LeakyReLU 激活函数避免梯度在负区间完全消失这也是 GAN 网络设计中的常见选择。# 文件路径discriminator.py import torch import torch.nn as nn class Discriminator(nn.Module): def __init__(self): super(Discriminator, self).__init__() self.model nn.Sequential( nn.Linear(28 * 28, 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(-1, 28 * 28) return self.model(x)LeakyReLU 中的参数 0.2 表示负斜率允许负值区间也有一点梯度这能让判别器在训练过程中更稳定。4.3 编写训练循环训练循环是整个项目的核心集中体现了 GAN 的训练逻辑。下面这段代码可以直接保存为 train.py 并运行。# 文件路径train.py import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import datasets, transforms from torchvision.utils import save_image from generator import Generator from discriminator import Discriminator # 超参数 batch_size 128 # 每个批次的样本数 latent_dim 100 # 噪声向量维度 epochs 50 # 训练轮数 lr 2e-4 # 学习率 device torch.device(cuda if torch.cuda.is_available() else cpu) # 数据预处理归一化到 [-1, 1] transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) # 加载 MNIST 数据集 dataset datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) dataloader DataLoader(dataset, batch_sizebatch_size, shuffleTrue) # 初始化网络 G Generator(latent_dim).to(device) D Discriminator().to(device) # 损失函数与优化器 criterion nn.BCELoss() g_optimizer optim.Adam(G.parameters(), lrlr, betas(0.5, 0.999)) d_optimizer optim.Adam(D.parameters(), lrlr, betas(0.5, 0.999)) # 循环训练 for epoch in range(epochs): for i, (real_imgs, _) in enumerate(dataloader): real_imgs real_imgs.to(device) current_batch_size real_imgs.size(0) # 构造标签 real_label torch.ones(current_batch_size, 1).to(device) fake_label torch.zeros(current_batch_size, 1).to(device) # 1. 更新判别器 d_optimizer.zero_grad() real_output D(real_imgs) d_real_loss criterion(real_output, real_label) z torch.randn(current_batch_size, latent_dim).to(device) fake_imgs G(z) fake_output D(fake_imgs.detach()) d_fake_loss criterion(fake_output, fake_label) d_loss d_real_loss d_fake_loss d_loss.backward() d_optimizer.step() # 2. 更新生成器 g_optimizer.zero_grad() z torch.randn(current_batch_size, latent_dim).to(device) fake_imgs G(z) fake_output D(fake_imgs) g_loss criterion(fake_output, real_label) g_loss.backward() g_optimizer.step() # 每 200 个 batch 打印一次当前损失 if i % 200 0: print( fEpoch [{epoch 1}/{epochs}] Batch [{i}/{len(dataloader)}] fD_loss: {d_loss.item():.4f} G_loss: {g_loss.item():.4f} ) # 每个 epoch 结束后保存一次生成图像 with torch.no_grad(): z torch.randn(64, latent_dim).to(device) sample_imgs G(z) save_image( sample_imgs, foutput/epoch_{epoch 1:03d}.png, nrow8, normalizeTrue )训练前确保 output 目录存在否则保存图像时会报错。可以在终端先执行mkdir output4.4 运行与预期结果进入 gan-demo 目录执行python train.py如果一切正常你会看到训练日志不断输出同时在 output 目录下生成一张张 epoch 图片。对于全连接层的简单 GAN预期结果大致如下第 1 个 epoch生成图像基本是随机噪声看不出数字形状。第 5 到第 10 个 epoch图像中出现模糊的轮廓隐约能看到笔画结构。第 20 到第 50 个 epoch图像接近手写数字的样子但清晰度仍然有限。D 的损失通常在 0.5 到 1.5 之间波动G 的损失也不会有非常标准的下降曲线。这是因为两个网络在相互对抗损失数字不能像普通分类模型那样直接判断优劣。4.5 结果说明用全连接层实现的简单 GAN 虽然能学习到数字的大致结构但生成质量无法与 DCGAN、StyleGAN 等现代模型相比。这个项目的意义在于验证训练逻辑判别器学会了区分真假生成器学会了根据噪声生成有意义的图像。如果你希望观察损失变化可以在训练过程中把每一轮的 d_loss 和 g_loss 记录下来训练结束后用 matplotlib 绘制曲线。这里给一个简单的记录代码示例# 在 train.py 中增加两个列表 d_loss_history [] g_loss_history [] # 每个 epoch 结束时记录平均损失 d_loss_history.append(d_loss.item()) g_loss_history.append(g_loss.item()) # 训练结束后绘制曲线 import matplotlib.pyplot as plt plt.plot(d_loss_history, labelD Loss) plt.plot(g_loss_history, labelG Loss) plt.xlabel(Epoch) plt.ylabel(Loss) plt.legend() plt.savefig(output/loss_curve.png)通过损失曲线和生成图像的对比你能更直观地理解 GAN 训练逻辑中“对抗”的意义。5. 常见问题与排查思路GAN 训练经常被形容为玄学因为同样的代码改一个随机种子或学习率效果可能完全不同。下面总结几个最常见的问题和排查思路。5.1 训练不收敛损失剧烈震荡现象d_loss 和 g_loss 都上下跳动生成图像长时间没有明显变化。可能原因学习率过高导致参数更新幅度过大。判别器和生成器能力不平衡一方迅速压过另一方。数据集过小或 batch_size 太小时每个 batch 的分布不稳定。解决思路把学习率调整到 1e-4 到 2e-4 区间增大 batch_size固定随机种子。还有一种做法是减小判别器的学习率让生成器有足够时间追赶。5.2 判别器损失迅速归零现象训练没过多久d_loss 降到接近 0生成器 g_loss 却飙升生成图像没有细节。可能原因判别器太强能够轻易区分真实样本和生成样本。导致生成器收到的有效梯度非常小失去学习能力。解决思路降低判别器的学习率。减少判别器网络容量或者增加 Dropout。让生成器每轮更新多次平衡训练速度。5.3 模式坍塌现象生成的图像非常单一几十张图看起来像同一张或者只包含某几类数字。可能原因生成器发现某一种输出最容易骗过判别器于是停止探索多样性钻进“安全区”。模式坍塌是 GAN 最经典的难题之一在简单 GAN 上也很容易出现。解决思路尝试 WGAN 或 WGAN-GP 等改进模型。引入 Mini-batch Discrimination让判别器同时观察一个 batch 里的样本防止生成样本过于雷同。增大噪声向量维度给生成器更多表达空间。降低生成器学习率减缓它的“偷懒”速度。5.4 生成图像模糊现象图像能看出大致结构但边缘模糊缺少细节。可能原因网络容量不足全连接结构对图像空间信息利用不够或者训练轮数不足生成器还没有完全收敛。解决思路使用卷积结构例如 DCGAN增加训练轮数尝试在生成器中使用 BatchNorm 层每一项改进后都要重新观察图像效果。5.5 排查清单问题现象常见原因排查顺序解决思路训练不收敛学习率太高、两个网络能力失衡先看学习率再看网络容量调低学习率固定随机种子判别器 loss 迅速归零判别器过强检查 loss 下降速度降低 D 学习率增加 G 更新频率生成图像单一模式坍塌观察多个 epoch 图片是否雷同尝试 WGAN-GP、Mini-batch 技巧生成图像模糊网络容量或训练不足检查生成图像细节改 DCGAN 结构增加训练轮数图像全黑或全白激活函数与数据范围不匹配检查生成器输出是否在 [-1,1]使用 Tanh 输出检查 Normalize 设置排查时不要只看 loss 数字要多保存生成图像以视觉效果为准。生成模型的质量最终要通过人眼观察来判断。6. GAN 调参与工程最佳实践6.1 网络设计建议GAN 网络设计与普通分类网络有一些明显区别。生成器最后一层推荐使用 Tanh这样输出与归一化后的数据范围一致。判别器内部推荐使用 LeakyReLU尽量避免使用 ReLU防止负区间梯度全为零。判别器最后一层使用 Sigmoid保证输出是 0 到 1 的概率值。优化器推荐 Adam并且常见的 betas 参数设置为 (0.5, 0.999)。这个设置比 PyTorch 默认的 (0.9, 0.999) 更适合 GAN 训练。这些建议来自大量 GAN 实战经验虽然不是绝对真理但对于入门项目非常有效。6.2 标签平滑技巧标准 GAN 中真实标签为 1假标签为 0。但实践中过度自信的判别器会导致训练不稳定。一个常用技巧是标签平滑real_label torch.ones(current_batch_size, 1, devicedevice) * 0.9 fake_label torch.zeros(current_batch_size, 1, devicedevice) 0.1也就是说真实样本不要求判别器输出严格为 1只要接近 0.9假样本不要求输出严格为 0只要接近 0.1。这能降低判别器的置信度改善梯度信号。6.3 训练节奏与监控训 GAN 不能开盲盒边训练边监控非常关键。每个 epoch 保存一次生成图像按时间顺序排成组观察生成质量变化。记录 D 和 G 的 loss但不是只看 loss 数值而要结合图像判断。固定随机种子让实验可复现。例如torch.manual_seed(42)定期保存模型 checkpoint避免训练中断后从头再来。6.4 工程化注意事项如果 GAN 项目要走向生产环境还需要考虑更多问题训练数据的版权和隐私必须合规不能用未经授权的图片训练商用模型。生成内容要可追溯尤其在生成人脸等敏感场景时必须增加审核和过滤机制。模型部署要计算成本和响应速度GAN 推理通常比普通分类模型更重需要合理设计服务架构。生成结果存在随机性线上使用时要做好异常兜底。安全边界在生成模型项目中尤其重要。生成能力越强带来的合规风险也越高开发者在追求效果的同时也要把数据来源、内容审核和部署监控放在同等重要的位置。7. 总结与下一步学习路线通过这篇文章你应该掌握了 GAN 训练逻辑的四个关键点GAN 由生成器和判别器两个网络组成。每轮迭代先更新判别器再更新生成器。判别器学会区分真假生成器学会骗过判别器。训练过程中不能单独把某一方训练到收敛必须保持动态对抗。代码层面我们已经实现了一个可运行的 MNIST 手写数字 GAN并知道了如何通过保存图像和 loss 曲线来观察训练状态。下一步建议按照以下路线继续深入学习 DCGAN把全连接结构替换成卷积结构生成质量会大幅提升。学习 CGAN在生成器中加入类别条件让模型可以控制生成指定类别的图像。学习 WGAN 和 WGAN-GP理解它们为什么能从理论上缓解训练不稳定和模式坍塌。学习 StyleGAN了解现代 GAN 如何控制生成图像的风格和细节。如果条件允许可以自己动手改实验调整噪声维度、增加网络层数、更换数据集。每一次修改都会加深你对训练逻辑的理解。尤其是把 output 目录下不同 epoch 的图片连续播放你会直观地看到生成器如何从随机噪声慢慢学会生成数字这个过程比任何理论分析都更能帮助你掌握 GAN。