news 2026/9/24 18:53:36

GAN代码实战:从损失函数到训练调参的完整指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
GAN代码实战:从损失函数到训练调参的完整指南

第一次动手写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或者CelebA,MNIST能让你的调试迭代周期短很多。

数据加载部分直接用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", train=True, transform=transform, download=True) dataloader = DataLoader(dataset, batch_size=128, shuffle=True, drop_last=True)

这里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_dim=784): super().__init__() self.model = nn.Sequential( nn.Linear(img_dim, 256), nn.LeakyReLU(0.2, inplace=True), nn.Linear(256, 256), nn.LeakyReLU(0.2, inplace=True), 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_dim=100, img_dim=784): super().__init__() self.model = nn.Sequential( nn.Linear(noise_dim, 256), nn.ReLU(inplace=True), nn.Linear(256, 256), nn.ReLU(inplace=True), 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(), lr=2e-4, betas=(0.5, 0.999)) g_optim = torch.optim.Adam(gen.parameters(), lr=2e-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, nrow=8, normalize=True) 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,生成器发现了某个“以假乱真”的捷径,于是只输出这一种样本,放弃了对整个数据分布的覆盖。

我当时的排查流程是这样的:

  1. 先看判别器的loss是不是下降得非常快。如果判别器很快就收敛到接近0,说明它对真实样本和生成样本区分得太容易了,生成器完全没有机会学到东西。
  2. 检查生成器输出的标准差。对固定一批噪声生成的图片计算像素标准差,如果标准差非常小,说明所有图片几乎一样,基本可以确定模式崩溃。
  3. 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_dim=784): 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_feat=False): 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_feat=True) _, fake_feat = disc(fake_imgs, return_feat=True) fm_loss = torch.mean((real_feat.mean(dim=0) - fake_feat.mean(dim=0)) ** 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_steps=1000): z = torch.randn(1, noise_dim, requires_grad=True) optimizer = torch.optim.Adam([z], lr=0.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换一换、把网络换一换、把报错记下来,每一个坑都会变成你后续调参的判断依据。

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/24 18:52:56

WorkBuddy 十大技能实战:从代码脚手架到跨工具协同的效率提升指南

1. 为什么 WorkBuddy 的技能体系值得认真拆解WorkBuddy 这类工具型产品&#xff0c;最怕的就是“装完即吃灰”。我见过太多人兴冲冲下载、安装、登录&#xff0c;然后对着工作台发呆——不知道从哪下手&#xff0c;也不知道哪些功能真正能省时间。问题不在工具本身&#xff0c;…

作者头像 李华
网站建设 2026/9/24 18:52:47

Spring Boot Maven插件not found报错:原因排查与解决方案

1. 问题现象与初步定位1.1 报错出现的典型场景先说说最常见的踩坑现场。你在IDEA里新建了一个Spring Boot项目&#xff0c;可能是从Spring Initializr生成的&#xff0c;也可能是直接在Maven项目里手动加的依赖。一切看起来都很正常&#xff1a;pom.xml里依赖声明也写了&#x…

作者头像 李华
网站建设 2026/9/24 18:50:01

YOLO车道线虚线检测数据集:标签格式与训练实战解析

简介&#xff1a;面向目标检测学习者与YOLO系列算法实践者&#xff0c;这份数据集专为车道线与虚线检测任务打造&#xff0c;涵盖1659张已标注图像&#xff0c;标签完整&#xff0c;并已划分好训练集与验证集&#xff0c;可直接用于YOLOv5、YOLOv7、YOLOv8、YOLOv9、YOLOv10、Y…

作者头像 李华
网站建设 2026/9/24 18:49:59

AWS云计算术语中英文对照:从基础到实战的完全指南

要搞清楚 AWS 这一堆云计算术语&#xff0c;光背单词没用&#xff0c;得知道每个词背后对应的是什么场景、什么服务&#xff0c;以及中文资料里最常被翻译成什么。我今天把这些年实操和带团队时反复用到的高频 AWS 云计算词汇整理了一份中英文对照&#xff0c;不是简单罗列字典…

作者头像 李华
网站建设 2026/9/24 18:49:11

程序员起点:从零搭建购物车系统完整实战指南

1. 程序员的起点&#xff1a;先想清楚这三件事最早开始带新人那阵子&#xff0c;几乎每周都会收到类似的私信&#xff1a;"想转行做程序员&#xff0c;该从哪里开始&#xff1f;""Java 和 Python 到底选哪个&#xff1f;""培训班学了半年能找到工作吗…

作者头像 李华
网站建设 2026/9/24 18:47:48

mscomm32.ocx 报错修复:从注册机制到 64 位系统兼容实践

前阵子帮一个客户收拾一台工控电脑&#xff0c;Windows 10 64 位系统&#xff0c;一开生产管理系统就弹窗&#xff1a;Component mscomm32.ocx not correctly registered&#xff0c;file is missing or invalid。点几次确定之后程序直接闪退。这是一台刚升级完系统的老机器&am…

作者头像 李华