
第一次动手写GAN代码的时候我卡在了一个现在回头看特别基础的地方判别器和生成器的损失函数到底该怎么写。原论文那个 max_D min_G 的公式明明没有负号为什么代码里全是一大串 BCEWithLogitsLoss后来我花了一个周末把最小可跑的GAN从零怼出来才彻底明白这些看起来“反直觉”的地方背后全是数值计算和梯度传播的细节。这篇文章就记录一下我练GAN代码时的完整过程包括网络结构、损失函数、训练调参和几个经典变体的实现思路给同样想从代码入手啃GAN的同学一条能直接照着走的路。1. 为什么我建议从MLP版GAN开始写代码很多人一上来就复现DCGAN、StyleGAN结果被卷积、归一化、渐进式训练这些外围细节淹没反而没搞明白GAN最核心的对抗逻辑。我自己的经验是想要练透GAN的代码第一步应该写一个没有任何卷积的MLP版GAN在MNIST上跑通。这个版本大概只有一百多行代码却能把“生成器、判别器、对抗损失、反向传播”这条主线完整走一遍。1.1 从原理到代码的映射关系GAN的原始思想是让两个网络互相博弈生成器G把随机噪声映射成假样本判别器D判断输入是真实样本还是假样本。训练目标是最小化生成器损失、最大化判别器判对的能力。对应到代码上其实就三块数据加载准备好真实样本。生成器网络输入一个低维噪声向量输出一张图片。判别器网络输入一张图片输出一个0到1之间的“真实性”分数。用MNIST做例子输入噪声维度可以设置成100图片是28×28的灰度图展平后是784维。判别器输入784维输出1维的logit生成器输入100维输出784维。这里有一个初学者容易忽略的细节生成器最后一层必须加Tanh把输出约束到[-1,1]区间。因为你在数据预处理时会把真实图片像素从[0,255]归一化到[-1,1]如果生成器输出的值域和真实数据不一致判别器学到的判别标准就是错的。1.2 MNIST这个数据集为什么最适合练手MNIST的图片小、类别简单、灰度图单通道训练一张图的成本极低CPU都能跑。我用笔记本的CPU训练一个MLP版GAN一个epoch大概不到两分钟很快就能看到生成效果的变化。相比CIFAR-10或者CelebAMNIST能让你的调试迭代周期短很多。数据加载部分直接用torchvision就行from torch.utils.data import DataLoader from torchvision import datasets, transforms transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) dataset datasets.MNIST(root./data, trainTrue, transformtransform, downloadTrue) dataloader DataLoader(dataset, batch_size128, shuffleTrue, drop_lastTrue)这里Normalize((0.5,), (0.5,))不是随便写的。它把像素值从[0,1]线性变换到[-1,1]公式是(x - 0.5) / 0.5。生成器输出做Tanh后也是[-1,1]两边就对齐了。1.3 环境准备与训练成本控制PyTorch 2.x就行了不需要额外装复杂依赖。没有GPU的话把隐藏层宽度设小一点比如256再用CPU训练效果也不会差。我实际测过MNIST这个任务GPU和CPU的差异主要在迭代速度但不影响你理解代码逻辑。练手阶段不建议开TensorBoard之类的可视化工具直接在epoch结束的时候把生成图片保存成网格图扫一眼就知道训练有没有正常。越简单的工具越能逼你关注模型本身。2. 最小可跑通的代码骨架生成器、判别器与训练循环这一节直接给代码。我不会贴完整的、可以直接复制跑的几百行文件而是会把关键模块拆开讲因为写GAN代码写到最后真正重要的不是“跑通”而是“知道每一行为什么这么写”。如果只是复制粘贴跑通了下次换个任务照样不会。2.1 判别器实现细节判别器本质上是一个二分类网络。输入图片展平成784维经过三个全连接层中间用LeakyReLU激活最后一层输出一个logit。注意最后一层不要接Sigmoid因为损失函数会用BCEWithLogitsLoss这个损失函数内部自带Sigmoid直接在输出层再接Sigmoid会降低数值稳定性。import torch import torch.nn as nn class Discriminator(nn.Module): def __init__(self, img_dim784): super().__init__() self.model nn.Sequential( nn.Linear(img_dim, 256), nn.LeakyReLU(0.2, inplaceTrue), nn.Linear(256, 256), nn.LeakyReLU(0.2, inplaceTrue), nn.Linear(256, 1), ) def forward(self, x): return self.model(x)LeakyReLU的负斜率取0.2是DCGAN论文里的常见配置。ReLU会把负数全部置零在判别器这种需要输出连续分数的地方LeakyReLU能保留少量负信息训练时梯度更稳定。2.2 生成器的实现与输出约束生成器是把噪声向量“放大”成图片。输入100维经过两个隐含层最后输出784维。中间激活函数用ReLU最后一层用Tanh。class Generator(nn.Module): def __init__(self, noise_dim100, img_dim784): super().__init__() self.model nn.Sequential( nn.Linear(noise_dim, 256), nn.ReLU(inplaceTrue), nn.Linear(256, 256), nn.ReLU(inplaceTrue), nn.Linear(256, img_dim), nn.Tanh(), ) def forward(self, z): return self.model(z)为什么生成器不用LeakyReLU而用ReLU生成器输出的目标分布是图像像素像素值是有正有负的连续值实验经验表明ReLU在生成器里表现更好。这一点不需要过度纠结跟着主流做法走就行。2.3 训练循环里容易被绕晕的标签设置这是整个训练流程的核心也是很多初学者绕圈子的地方。GAN训练是交替进行的每个batch先更新判别器再更新生成器而不是像普通分类网络那样一个batch反向传播一次就行。criterion nn.BCEWithLogitsLoss() d_optim torch.optim.Adam(disc.parameters(), lr2e-4, betas(0.5, 0.999)) g_optim torch.optim.Adam(gen.parameters(), lr2e-4, betas(0.5, 0.999)) for epoch in range(epochs): for real_imgs, _ in dataloader: batch_size real_imgs.size(0) real_imgs real_imgs.view(batch_size, -1) # 真实图片标签为1生成图片标签为0 real_labels torch.ones(batch_size, 1) fake_labels torch.zeros(batch_size, 1) # 生成噪声 z torch.randn(batch_size, noise_dim) fake_imgs gen(z) # 1. 训练判别器 d_real_logits disc(real_imgs) d_fake_logits disc(fake_imgs.detach()) d_loss criterion(d_real_logits, real_labels) criterion(d_fake_logits, fake_labels) d_optim.zero_grad() d_loss.backward() d_optim.step() # 2. 训练生成器 z torch.randn(batch_size, noise_dim) fake_imgs gen(z) g_logits disc(fake_imgs) # 关键生成器的目标是让判别器把假图当真实图片所以标签是1 g_loss criterion(g_logits, real_labels) g_optim.zero_grad() g_loss.backward() g_optim.step()生成器更新时用fake_imgs.detach()是为了确保梯度只流经生成器参数不流向判别器。而第二次调用disc(fake_imgs)时不加detach()因为这一步就是要让梯度穿过判别器、再传到生成器。这两个地方一混整个训练就乱了。还有个细节更新判别器后生成器拿到的判别器已经是“更新过一次的”新判别器。训练生成器时另外采样一批新的噪声z而不是复用之前的噪声。这样做可以让生成器看到更多样的输入减少不同batch之间的相关性。2.4 训练完成后的可视化检查每隔一个epoch把生成器输出的图片保存一次。保存时不能直接把[-1,1]的预测值当作图片显示要先映射回[0,255]def save_images(gen, epoch, path./results): z torch.randn(64, noise_dim) imgs gen(z).view(-1, 1, 28, 28) imgs (imgs 1) / 2 # 从[-1,1]映射回[0,1] grid torchvision.utils.make_grid(imgs, nrow8, normalizeTrue) torchvision.utils.save_image(grid, f{path}/epoch_{epoch}.png)训练早期你会看到一堆噪点中期开始出现数字的轮廓后期有些数字已经比较清晰。如果发现所有图片都长得差不多那就是模式崩溃后面专门讲。3. 损失函数这一关交叉熵负号问题、数值稳定性与Label Smoothing很多人在“原始GAN公式的交叉熵为什么没有负号”这个问题上卡住而且卡很久。我尽量把它讲透因为它背后不是数学难题而是“目标函数的形式”和“代码里loss的定义方式”之间的映射关系。3.1 原论文目标函数中的“没有负号”到底是怎么回事原始GAN的目标函数写为min_G max_D V(D,G) E[log D(x)] E[log(1 - D(G(z)))]其中E是期望log D(x)和log(1-D(G(z)))两项都没有显式的负号。看起来和信息论里的交叉熵不太一样因为交叉熵定义是H(p,q) -Σ p(x) log q(x)前面有个负号。为什么GAN的公式不写负号因为原始GAN用的是最大化max而非最小化min。信息论中的交叉熵是在做最小化所以前面有负号。GAN里对判别器D来说目标是尽可能正确区分真假这就等价于最大化把真实样本判真、把生成样本判假的对数似然。当你把“最大化”变成“最小化loss”时负号就会自然出现。更直白地说max_D V(D,G)等价于min_D (-V(D,G))而-V(D,G)正是交叉熵的标准形式正样本交叉熵 负样本交叉熵。所以原论文公式没写负号不代表没有负号只是被max吸收了。3.2 代码里为什么用BCEWithLogitsLoss而不是手写log有网友会自己手写生成器损失-torch.log(d_fake)来模仿原论文结果训练时经常遇到NaN或者梯度爆炸。原因有两点。第一D(G(z))是经过Sigmoid后的概率值当它落到0附近时log值趋近负无穷梯度非常大。第二PyTorch在反向传播中对这类极端值处理时要经过Sigmoid的导数一旦数值下溢梯度就变成NaN。BCEWithLogitsLoss内部把Sigmoid和Binary Cross Entropy的数值计算融合在一起会在实现上做数值稳定化处理。所以你在代码里传的是判别器最后一层的logit而不是Sigmoid之后的概率这个函数会自动完成后续计算l -[ y * log(sigmoid(x)) (1 - y) * log(1 - sigmoid(x)) ]这才是完整带负号的二进制交叉熵。它对应的正是生成器和判别器各自要最小化的目标。3.3 Label Smoothing与随机标签翻转小数据集上直接训GAN判别器很容易快速“过拟合”到训练集输出很极端的置信度导致生成器拿到的梯度变得无意义。解决办法之一是用One-sided Label Smoothing真实标签不用1而用0.9假标签还是0。这样判别器不会追求在logit上输出非常大或非常小的值收敛更稳。实现方式很简单real_labels torch.full((batch_size, 1), 0.9)还有一招是随机标签翻转以极小的概率比如0.05把真实标签和生成标签互换。这是一种正则化手段防止判别器记住训练集。不过经验上label smoothing的效果更基础建议优先尝试。3.4 判别器和生成器损失组合参数速查组合方式判别器loss生成器loss特点原始minimaxBCE(D(real), 1) BCE(D(fake), 0)BCE(D(fake), 0)生成器梯度容易饱和非饱和损失同上BCE(D(fake), 1)生成器梯度更充足推荐新手使用LSGANMSE(D(real), 1) MSE(D(fake), 0)MSE(D(fake), 1)更平滑训练稳定但可能产生模糊图像WGAN-GPD(real) - D(fake) 梯度惩罚-D(fake)常用于高质量生成任务我练手期用的就是第二行“非饱和损失”因为生成器在训练早中期能持续获得有意义的梯度不会因为判别器太强而直接“躺平”。你手里的BCEWithLogitsLoss(fake_logits, real_labels)就是非饱和损失的一种具体代码形式。4. 训练不收敛时我排查的完整思路DCGAN论文里有一句著名的话训练GAN就像在玩“猫鼠游戏”经常处于不稳定的状态。我第一次练的时候确实踩了不少坑下面这些现象和排查路径是反复试出来的。4.1 模式崩溃所有生成样本长得一模一样现象epoch过半保存的图片网格里每个格子的图像几乎完全相同但每一张又能看出来是某个数字。这就是典型的mode collapse生成器发现了某个“以假乱真”的捷径于是只输出这一种样本放弃了对整个数据分布的覆盖。我当时的排查流程是这样的先看判别器的loss是不是下降得非常快。如果判别器很快就收敛到接近0说明它对真实样本和生成样本区分得太容易了生成器完全没有机会学到东西。检查生成器输出的标准差。对固定一批噪声生成的图片计算像素标准差如果标准差非常小说明所有图片几乎一样基本可以确定模式崩溃。把z从标准正态采样换成在某个固定点旁边加噪声看生成的图片是否变化。如果几乎不变化说明生成器把输入噪声“忽略”了。缓解办法按优先级排序先降低判别器的学习率比如从2e-4降到1e-4再用feature matching下一节讲最后可以考虑用Mini-batch Discrimination。前两种方法在小模型上效果非常明显。4.2 判别器太强或太弱梯度信号失衡训练中判别器loss持续下降而生成器loss持续上升这种情况通常是判别器“太强”了生成器的梯度被压制。反过来如果判别器loss一直在0.69附近不动说明判别器根本没学会区分真假生成器也就无法被引导。0.69这个数值得解释一下对二分类问题如果判别器完全无法区分它的预测概率是0.5那么BCE损失就是-log(0.5) ≈ 0.693。所以看到0.69不要慌它表示“判别器在瞎猜”。在代码层面我会这样调如果D loss太小0.1、G loss太大把判别器的学习率调低或者每训练2次判别器才训练1次生成器让生成器有更多机会学习。如果D loss一直在0.69附近波动且不下降检查真实标签和生成标签是否设反了或者生成器输出值域和真实图片不一致。如果D loss和G loss都在剧烈跳动降低学习率采用betas(0.5, 0.999)这个DCGAN标配减少动量带来的震荡。4.3 监控指标比肉眼观察更可靠肉眼观察生成图片有一定滞后性我习惯在训练过程中每个epoch打印三样东西判别器对真实样本的平均输出、判别器对生成样本的平均输出、生成器损失。如果前两个都在0.5附近说明判别器还在挣扎生成器仍有希望如果第一个输出接近1、第二个输出接近0且生成器损失很大说明判别器已经碾压生成器了。一个更进阶的做法是保存“固定噪声向量”的生成结果。每次评估都用同一个z这样生成图片的演化过程就可以横向对比方便确认模型进步还是退化。我踩过的一个坑是每次评估重新采样z导致上一轮看起来模糊、这一轮看起来清晰误以为模型进步了其实只是随机噪声不同。5. 从基础GAN进阶Feature Matching与图像修复的代码思路当你把最基础的MLP版GAN跑通、调稳就值得往两个实用方向拓展了。这两个方向都能在现有代码上小改实现不需要推倒重来。5.1 Feature Matching让生成器模仿中间层特征Feature Matching的思路来自Salimans等人在2016年发表的《Improved Techniques for Training GANs》。它不去直接对抗判别器的最终输出分数而是让生成器学习让判别器中间层的特征分布接近真实数据的特征分布。直观理解判别器可以看成是一个“特征提取器 二分类头”。它的中间层特征包含了自己学到的数据的判别性信息。生成器如果能在特征空间上和真实图片对齐就更容易生成符合数据分布的内容。实现时需要在判别器forward里返回中间层特征class Discriminator(nn.Module): def __init__(self, img_dim784): super().__init__() self.fc1 nn.Linear(img_dim, 256) self.fc2 nn.Linear(256, 256) self.fc_out nn.Linear(256, 1) def forward(self, x, return_featFalse): h1 torch.relu(self.fc1(x)) h2 torch.relu(self.fc2(h1)) logit self.fc_out(h2) if return_feat: return logit, h2 return logit生成器的feature matching loss是这样计算的_, real_feat disc(real_imgs, return_featTrue) _, fake_feat disc(fake_imgs, return_featTrue) fm_loss torch.mean((real_feat.mean(dim0) - fake_feat.mean(dim0)) ** 2)然后生成器的总损失就是g_loss fm_weight * fm_lossfm_weight一般取10左右。这样生成器既保留了让判别器判真的对抗压力又增加了一个更明确的学习目标训练起来稳很多。这个trick在特征空间维度较高时效果尤其明显。5.2 GAN图像修复用生成器和判别器做“脑补”图像修复inpainting是GAN非常成功的应用方向之一思路也很有意思与其直接生成缺失像素不如在一个训练好的GAN的潜在空间中搜索一个向量让生成结果在已知区域尽可能接近原图在缺失区域看起来“合理”。具体分成两步。第一步先训练一个GAN或者直接用预训练模型把生成器锁定住。第二步针对每一张待修复图片随机初始化一个z用梯度下降更新z让生成图在“已知区域”像素上和原图接近同时用判别器打分为“真实性”提供约束。def inpaint(image, mask, gen, disc, num_steps1000): z torch.randn(1, noise_dim, requires_gradTrue) optimizer torch.optim.Adam([z], lr0.01) for _ in range(num_steps): fake_img gen(z) # context loss已知区域像素误差 context_loss torch.mean(((fake_img - image) * mask) ** 2) # perception loss判别器对整体真实性的惩罚 p_loss -torch.mean(disc(fake_img)) total_loss context_loss 0.1 * p_loss optimizer.zero_grad() total_loss.backward() optimizer.step() return gen(z).detach()这里的mask是固定大小已知区域为1、缺失区域为0。权重0.1不用太较真不同数据集上要微调。关键是理解context loss负责“像素级对齐”perception loss负责“不要生成一个看着不真实的补丁”。这两个目标相互制衡迭代出来的z对应的生成图就能同时满足两部分要求。这个方案对简单图像比如MNIST、人脸缩略图效果可观但对高分辨率图片效果有限。现代工业界的图像修复方案已经演进到扩散模型、ControlNet这类架构但GAN作为入门理解“生成模型如何解决图像补全问题”依然非常合适尤其是那个“固定生成器、优化潜在向量”的管线和很多深度学习问题里的逆问题求解思路完全一致。5.3 接下来可以练什么如果你已经把上面这些代码都跑通了接下来练手的方向可以有这些把MLP换成卷积网络复现DCGAN把生成器换成卷积结构后观察特征图的变化再进一步可以尝试WGAN-GP亲自感受一下Wasserstein距离对训练稳定性的提升。每次改动只动一个变量多做对比知识才会沉淀成你的本能。我个人练GAN代码最大的体会是这类生成模型不像图像分类那样有个明确的准确率指标你必须完全靠损失曲线、生成样例和特征统计来判断模型状态这种训练方式会逼着你把底层原理吃透。如果只是把别人的代码跑通而不去动它你永远只能做一个代码搬运工。把loss换一换、把网络换一换、把报错记下来每一个坑都会变成你后续调参的判断依据。