LSGAN:用最小二乘损失解决GAN训练不稳定问题 📅 发布时间:2026/8/29 2:15:28 👁 浏览次数: 1. 从GAN的“不稳定”说起为什么我们需要LSGAN如果你玩过或者研究过生成对抗网络那你一定对它的“训练不稳定”这个老生常谈的问题深有体会。模型要么生成一堆毫无意义的噪声要么判别器早早地就把生成器“打趴下”导致生成质量再也上不去。这背后的一个核心症结就出在原始的GAN所使用的损失函数——交叉熵损失上。想象一下这样一个场景判别器是一个经验老道的古董鉴定师生成器是一个试图制作高仿赝品的学徒。在原始GAN的规则下鉴定师的任务是“二分类”真品1或赝品0。学徒的目标是做出让鉴定师“打眼”误判为真品的赝品。这里有个问题当学徒做的赝品水平太差一眼假时鉴定师会非常自信地给出一个接近0的分数比如0.01。此时对于学徒生成器来说这个0.01和0.001带来的“挫败感”梯度是差不多的都极小。换句话说对于已经被判别器明确判为“假”的样本生成器很难从中学到有效的、指向“更真”的改进信息。它陷入了“梯度消失”的困境不知道朝哪个方向努力才能显著提高分数。LSGAN最小二乘生成对抗网络的提出正是为了根治这个顽疾。它把鉴定游戏从“非黑即白”的二分类变成了一个“打分预测”的回归任务。鉴定师不再说“这是真1或假0”而是给一个赝品打一个分数比如“这个像0.2分那个像0.8分”。学徒的目标也不再是骗过鉴定师而是让自己做出的所有赝品获得的分数都尽可能接近“真品”的分数比如我们设定为1。这个“接近”的程度就用最小二乘Least Squares损失也就是我们熟悉的均方误差MSE来衡量。这么做带来了一个根本性的好处即使一个赝品只得了0.1分离目标1分很远但最小二乘损失会明确地告诉学徒“你离目标还差0.9而且这个差距会产生一个很强的梯度信号指引你朝着分数提高的方向修改你的工艺。” 这相当于给了生成器一个持续的、有方向的“推力”而不是在失败区域陷入迷茫。从实践角度来看LSGAN通常能带来更稳定的训练过程、更快的收敛速度以及很多时候更清晰的生成效果。它没有增加任何复杂的网络结构仅仅通过替换损失函数这一“巧劲”就显著改善了GAN的训练体验这也是它一经提出就备受关注的原因。2. 撕开公式LSGAN损失函数的直观与精妙理解了动机我们来看LSGAN具体是怎么做的。它的核心就是为生成器G和判别器D重新定义了两个损失函数。为了让推导更清晰我们先明确几个符号真实数据分布( p_{data}(x) )我们从数据集中采样的真实图片。生成数据分布( p_{z}(z) )生成器的输入噪声分布如标准正态分布通过生成器 ( G(z) ) 得到生成图片。判别器输出( D(x) )对于输入样本 ( x )判别器输出一个标量值。注意在LSGAN中这个值不再被Sigmoid压缩到(0,1)表示概率而是一个可以超出0-1范围的分数但通常我们仍会约束其范围以获得稳定训练。LSGAN为判别器和生成器分别设定了两个“目标分数”a 我们希望判别器给真实数据打出的分数。b 我们希望判别器给生成数据打出的分数。c 我们希望生成器努力让判别器给生成数据打出的分数。在原始论文中作者提供了一个简单有效的选择a1, b0, c1。这非常直观真实数据的目标分数是1生成数据的目标分数是0而生成器希望自己的数据能被判为1。基于此LSGAN的损失函数定义如下2.1 判别器D的损失函数判别器的目标是成为一个好的“打分员”。对于真实数据它打出的分 ( D(x) ) 应该尽量接近目标 ( a )对于生成数据它打出的分 ( D(G(z)) ) 应该尽量接近目标 ( b )。用均方误差来衡量[ \min_D L_{D} \frac{1}{2} \mathbb{E}{x \sim p{data}(x)}[(D(x) - a)^2] \frac{1}{2} \mathbb{E}{z \sim p{z}(z)}[(D(G(z)) - b)^2] ]这里乘以 ( \frac{1}{2} ) 是为了求导后形式更简洁导数前的系数为1。当 ( a1, b0 ) 时判别器的任务就是把真实图片 ( x ) 的分数 ( D(x) ) 尽量推向1。把生成图片 ( G(z) ) 的分数 ( D(G(z)) ) 尽量推向0。这很好理解判别器要能区分真假。2.2 生成器G的损失函数生成器的目标是“欺骗”判别器它希望自己生成的图片 ( G(z) ) 在判别器那里得到的分数 ( D(G(z)) ) 能尽量接近另一个目标 ( c )。同样使用均方误差[ \min_G L_{G} \frac{1}{2} \mathbb{E}{z \sim p{z}(z)}[(D(G(z)) - c)^2] ]当 ( c1 ) 时生成器的任务就变成了让自己生成的图片在判别器那里获得的分数尽可能接近真实图片的目标分数1。2.3 为什么是“最小二乘”优势何在与原始GAN的交叉熵损失相比最小二乘损失在这里展现了几个关键优势梯度更友好缓解梯度消失这是最核心的一点。对于生成器其损失 ( (D(G(z)) - 1)^2 )。当生成样本很差( D(G(z)) ) 接近0时损失值接近 ( (0-1)^2 1 )其关于 ( D(G(z)) ) 的梯度是 ( 2*(D(G(z))-1) \approx -2 )。这是一个很大且非零的梯度它会强烈地推动生成器去更新参数提高 ( D(G(z)) ) 的值。而在原始GAN中当判别器很自信输出接近0时生成器的交叉熵损失 ( log(1 - D(G(z))) ) 的梯度会变得极其微小梯度消失导致生成器无法学习。惩罚机制更合理最小二乘损失会惩罚那些距离目标很远的样本。即使判别器已经能将某个生成样本判为假分数低只要这个分数离生成器的目标c1很远生成器就会持续收到“你做得还很差”的强信号。这迫使生成器生成所有样本的质量都要向目标看齐。有观点认为这有助于减少原始GAN中可能出现的“模式坍塌”Mode Collapse问题——即生成器只学会生成少数几种样本缺乏多样性。因为LSGAN会惩罚那些分数低的“差生”样本促使生成器去覆盖更多样化的、能获得高分的样本。训练更稳定收敛可能更快由于梯度信号更强、更明确LSGAN在实际训练中通常表现出更好的稳定性。我们不再需要小心翼翼地平衡生成器和判别器的训练轮数例如原始GAN中常说的“判别器训k步生成器训1步”训练过程对超参数的敏感度也有所降低。在许多图像生成任务中LSGAN能更快地收敛到一个视觉质量不错的中间状态。注意虽然LSGAN的损失函数形式简单但在实现时有一个细节至关重要。由于我们使用了最小二乘损失并且目标值如1和0是固定的这就要求判别器的输出值不能无界地增长。在实践中我们通常会在判别器的最后一层不使用Sigmoid激活函数因为它将输出压缩到(0,1)但可能会采用其他正则化手段如梯度惩罚、权重裁剪等或简单的值域约束以防止判别器输出值爆炸导致训练不稳定。一种常见的稳健做法是使用Spectral Normalization谱归一化它不仅能稳定训练还能自然地约束函数空间。3. 实战用PyTorch从零搭建一个LSGAN理论说得再多不如动手跑一遍。下面我们就用PyTorch来实现一个最简单的LSGAN用于生成MNIST手写数字。我们将一步步拆解代码并解释每个关键部分的设计意图。3.1 环境准备与数据加载首先确保你的环境安装了PyTorch和Torchvision。我们使用MNIST数据集它简单且训练快速非常适合演示。import torch import torch.nn as nn import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader import matplotlib.pyplot as plt import numpy as np # 设备配置 device torch.device(cuda if torch.cuda.is_available() else cpu) print(fUsing device: {device}) # 超参数定义 latent_dim 100 # 噪声向量的维度 img_channels 1 # MNIST是灰度图通道数为1 img_size 28 # MNIST图像尺寸为28x28 batch_size 64 lr 0.0002 # 学习率 epochs 50 # 训练轮数 # 数据预处理与加载 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) # 将像素值从[0,1]归一化到[-1,1] ]) train_dataset datasets.MNIST(root./data, trainTrue, downloadTrue, transformtransform) train_loader DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue)这里有几个点需要注意归一化到[-1,1]这是GAN训练的常见技巧。使用Tanh作为生成器最后一层激活函数时其输出范围是(-1,1)与归一化后的数据分布匹配有助于训练稳定。潜在维度latent_dim这是输入生成器的随机噪声z的维度。100是一个常用起点维度越高理论上生成器能建模的分布越复杂但也可能增加训练难度。3.2 构建生成器与判别器网络我们构建一个简单的全连接网络MLP来实现。对于MNIST这种小图MLP已经足够。# 生成器定义 class Generator(nn.Module): def __init__(self, latent_dim, img_channels, img_size): super(Generator, self).__init__() self.img_size img_size self.init_size img_size // 4 # 初始特征图大小经过上采样后变为原图大小 self.l1 nn.Sequential(nn.Linear(latent_dim, 128 * self.init_size ** 2)) # 将噪声映射到足够多的神经元 self.conv_blocks nn.Sequential( nn.BatchNorm2d(128), nn.Upsample(scale_factor2), # 上采样 nn.Conv2d(128, 128, 3, stride1, padding1), nn.BatchNorm2d(128, 0.8), nn.LeakyReLU(0.2, inplaceTrue), nn.Upsample(scale_factor2), nn.Conv2d(128, 64, 3, stride1, padding1), nn.BatchNorm2d(64, 0.8), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, img_channels, 3, stride1, padding1), nn.Tanh() # 输出范围在(-1,1) ) def forward(self, z): out self.l1(z) out out.view(out.shape[0], 128, self.init_size, self.init_size) # 重塑为特征图 img self.conv_blocks(out) return img # 判别器定义 class Discriminator(nn.Module): def __init__(self, img_channels, img_size): super(Discriminator, self).__init__() def discriminator_block(in_filters, out_filters, bnTrue): block [nn.Conv2d(in_filters, out_filters, 3, 2, 1)] if bn: block.append(nn.BatchNorm2d(out_filters, 0.8)) block.append(nn.LeakyReLU(0.2, inplaceTrue)) block.append(nn.Dropout2d(0.25)) return block self.model nn.Sequential( *discriminator_block(img_channels, 16, bnFalse), # 第一层通常不加BN *discriminator_block(16, 32), *discriminator_block(32, 64), *discriminator_block(64, 128), ) # 计算经过卷积块后的特征图尺寸 ds_size img_size // 2 ** 4 # 经过4次stride2的卷积尺寸缩小16倍 self.adv_layer nn.Sequential(nn.Linear(128 * ds_size ** 2, 1)) # 注意最后一层没有Sigmoid输出是一个实数。 def forward(self, img): out self.model(img) out out.view(out.shape[0], -1) validity self.adv_layer(out) return validity # 直接输出分数而非概率网络结构设计的核心要点生成器G使用nn.Linear将噪声向量映射到一个较大的维度然后重塑view成特征图接着通过一系列转置卷积nn.Conv2d配合nn.Upsample或上采样卷积来逐步增大空间尺寸最终得到目标图像。Tanh激活函数确保输出在(-1,1)。判别器D是一个典型的卷积分类器但最后一层是线性层。这是LSGAN与原始GAN在实现上的关键区别之一。原始GAN的判别器最后一层通常是nn.Sigmoid()输出一个概率值。而LSGAN的判别器直接输出一个分数logits这个分数会直接代入最小二乘损失函数计算。我们依靠损失函数本身和目标值a,b,c来约束这个分数的意义。BatchNorm和Dropout在生成器中BatchNorm有助于稳定训练但通常不在输入层使用。在判别器中Dropout可以作为一种正则化防止判别器过强。这些技巧需要根据实际情况调整。3.3 定义LSGAN特有的损失函数与优化器这里就是体现LSGAN精髓的地方。我们使用PyTorch内置的均方误差损失MSELoss。# 初始化网络 generator Generator(latent_dim, img_channels, img_size).to(device) discriminator Discriminator(img_channels, img_size).to(device) # 定义LSGAN的损失函数均方误差 adversarial_loss nn.MSELoss() # 定义优化器 optimizer_G optim.Adam(generator.parameters(), lrlr, betas(0.5, 0.999)) optimizer_D optim.Adam(discriminator.parameters(), lrlr, betas(0.5, 0.999)) # 定义目标标签 real_label 1.0 fake_label 0.0 # 注意这里我们使用简单的1和0。有些实现会使用“软标签”如0.9和0.1来增加鲁棒性。目标标签的设定我们严格遵循论文中的a1, b0, c1。在代码中real_label对应afake_label对应b而生成器损失中的目标对应c同样是real_label即1。3.4 训练循环一步步拆解对抗过程训练循环是GAN的核心我们来看每一步发生了什么。for epoch in range(epochs): for i, (imgs, _) in enumerate(train_loader): # 配置数据 real_imgs imgs.to(device) batch_size real_imgs.size(0) # 创建用于损失函数的标签张量形状与判别器输出匹配 valid torch.full((batch_size, 1), real_label, devicedevice, dtypetorch.float32) fake torch.full((batch_size, 1), fake_label, devicedevice, dtypetorch.float32) # --------------------- # 训练判别器 (D) # --------------------- optimizer_D.zero_grad() # 计算真实图片的损失 real_pred discriminator(real_imgs) d_real_loss adversarial_loss(real_pred, valid) # 希望D(real)接近1 # 计算生成图片的损失 z torch.randn(batch_size, latent_dim, devicedevice) # 采样噪声 gen_imgs generator(z).detach() # 生成图片并detach避免梯度传到G fake_pred discriminator(gen_imgs) d_fake_loss adversarial_loss(fake_pred, fake) # 希望D(fake)接近0 # 判别器总损失 d_loss (d_real_loss d_fake_loss) / 2 d_loss.backward() optimizer_D.step() # --------------------- # 训练生成器 (G) # --------------------- optimizer_G.zero_grad() # 生成一批新图片这里重新采样了噪声也可以复用之前的 z torch.randn(batch_size, latent_dim, devicedevice) gen_imgs generator(z) # 计算生成器损失 # 生成器的目标是让判别器对生成图片的评分接近“真实”标签1 g_pred discriminator(gen_imgs) g_loss adversarial_loss(g_pred, valid) # 希望D(fake)接近1 g_loss.backward() optimizer_G.step() # 打印日志 if i % 200 0: print(f[Epoch {epoch}/{epochs}] [Batch {i}/{len(train_loader)}] f[D loss: {d_loss.item():.4f}] [G loss: {g_loss.item():.4f}])训练步骤详解准备标签为当前批次batch的真实图片和假图片创建对应的目标标签张量全1和全0。训练判别器清零梯度。将真实图片输入判别器计算其输出与“1”的MSE损失d_real_loss。采样新的噪声通过生成器得到假图片。注意这里使用了.detach()意味着在计算判别器对假图片的损失时生成器的参数被冻结梯度不会反向传播到生成器。这是为了防止在更新判别器时影响到生成器。将假图片输入判别器计算其输出与“0”的MSE损失d_fake_loss。判别器的总损失是这两部分损失的平均。反向传播更新判别器参数。训练生成器清零梯度。再次采样噪声或复用之前的通过生成器得到假图片。这次没有使用.detach()。将这批假图片输入刚刚更新过的判别器得到评分。计算生成器损失判别器对假图片的评分与“1”的MSE损失g_loss。这意味着生成器在努力让判别器“误判”假图为真。反向传播更新生成器参数。注意这里的梯度会穿过判别器一直传递到生成器指导生成器如何修改以提升分数。这个“先更新D再更新G”的循环就是GAN训练的基本节奏。LSGAN的损失计算就嵌在这个标准流程里。3.5 可视化与结果保存训练过程中定期查看生成效果至关重要。# 每训练完一个epoch保存一批生成图片 if epoch % 5 0 and i 0: # 每个epoch的第0个batch保存一次 with torch.no_grad(): test_z torch.randn(16, latent_dim, devicedevice) gen_imgs generator(test_z).cpu() # 反归一化将图像从[-1,1]变回[0,1]以便显示 gen_imgs 0.5 * gen_imgs 0.5 fig, axs plt.subplots(4, 4, figsize(8,8)) cnt 0 for i in range(4): for j in range(4): axs[i,j].imshow(gen_imgs[cnt, 0, :, :], cmapgray) axs[i,j].axis(off) cnt 1 plt.savefig(flsgan_images/epoch_{epoch}.png) plt.close()运行这个代码你会观察到随着训练进行生成的数字从最初的随机噪声逐渐变得清晰可辨。LSGAN的训练损失曲线通常也比原始GAN更平滑d_loss和g_loss会围绕一个值相对稳定地振荡而不是剧烈波动或一路飙升/降至零。4. LSGAN的变体、技巧与边界探讨掌握了基础LSGAN后我们来看看它的一些高级玩法和需要注意的边界。4.1 目标标签的“软化”与标签平滑在基础实现中我们使用了硬标签1和0。但在实践中直接使用1和0有时会导致判别器过于自信从而可能产生对抗性的梯度。一种常见的改进技巧是标签平滑Label Smoothing。原理不要求判别器将真实图片的输出严格推向1而是推向一个略小于1的值如0.9不要求将假图片的输出严格推向0而是推向一个略大于0的值如0.1。这相当于给判别器的目标增加了一点噪声防止其过度拟合。在LSGAN中的实现非常简单只需修改real_label和fake_label的值即可。real_label 0.9 # 原来是1.0 fake_label 0.1 # 原来是0.0对于生成器损失的目标c通常仍保持为1.0因为生成器的目标是尽可能“以假乱真”。标签平滑是一种廉价而有效的正则化手段常能提升模型的泛化能力和稳定性。4.2 与其他GAN损失的对比与结合LSGAN是众多GAN变体中的一种。理解它与其它损失的差异能帮助我们在不同场景下做出选择。损失函数类型核心思想优点缺点适用场景原始GAN (最小化JS散度)二分类交叉熵损失判别器输出概率。理论优美是GAN的奠基。训练不稳定易梯度消失/爆炸模式坍塌。理论研究或作为基线模型。LSGAN (最小二乘损失)将判别视为回归用MSE衡量分数与目标差距。梯度更饱和训练稳定收敛快生成样本质量可能更高。可能生成过于“平均”的样本细节锐利度有时稍逊。大多数图像生成任务的优先尝试选择尤其是需要稳定训练的场合。WGAN (Wasserstein距离)通过判别器此时叫Critic输出值的差值来度量分布距离要求判别器是Lipschitz连续的。理论上解决了训练不稳定问题损失值与生成质量相关性强。需要权重裁剪或梯度惩罚来实施Lipschitz约束训练可能较慢。对理论保障要求高或原始GAN/LSGAN训练失败时尝试。WGAN-GP (带梯度惩罚的WGAN)WGAN的改进用梯度惩罚项代替权重裁剪来实施Lipschitz约束。比WGAN更稳定通常能获得更好的效果。计算梯度惩罚需要额外的前向-后向传播计算成本稍高。追求高质量生成效果且计算资源充足时。Hinge Loss GAN使用合页损失鼓励判别器对真假样本的分数有一个“间隔”。训练也相对稳定在有些任务上表现优异。超参数间隔大小可能需要调整。常见于一些现代GAN架构如SAGAN, BigGAN中判别器的损失。如何选择对于新手或大多数应用LSGAN是一个非常好的起点。它实现简单改进直接且能显著提升原始GAN的训练体验。如果LSGAN效果不佳可以尝试WGAN-GP或Hinge Loss。原始GAN的交叉熵损失现在已较少直接用于实践。4.3 LSGAN的局限性它真的是“万能药”吗尽管LSGAN优点突出但它并非没有缺点可能生成“过于平均”的样本最小二乘损失倾向于最小化所有样本的误差平方和。这可能导致生成器为了降低整体损失而去生成那些“安全”的、位于数据分布中间区域的样本从而损失了一些细节和多样性。直观理解就是它可能更倾向于生成一个“轮廓正确但有点模糊”的数字而不是一个“笔画锐利但偶尔出格”的数字。这在一些需要高保真细节的任务中可能成为瓶颈。对离群值敏感MSE损失对大的误差给予非常大的惩罚平方项。如果判别器偶尔对一个样本给出了极端错误的分数这个巨大的损失可能会主导梯度更新导致训练波动。仍需谨慎的网络设计和超参数虽然LSGAN缓解了梯度消失但不意味着可以随意设计网络。生成器和判别器的能力仍需平衡尽管容错性更高学习率、优化器选择等超参数依然影响最终结果。不恰当的网络深度或通道数仍然会导致训练失败。4.4 实战中的调优经验与排坑指南结合我自己的项目经验分享几个让LSGAN跑得更好的技巧判别器别太强也别太弱这是一个永恒的话题。如果判别器太强比如层数太多、通道数太大生成器可能一直无法获得有效的梯度。如果太弱生成器学不到东西。一个实用的启发性原则是让生成器和判别器的参数量或层数保持在同一个数量级。例如生成器是4层卷积判别器也用4层卷积下采样。可以先从简单的对称结构开始。使用谱归一化Spectral Normalization这是稳定GAN训练的“神器”。尤其是在判别器中应用谱归一化可以自动地约束判别器函数的Lipschitz常数防止其梯度爆炸或变得过于尖锐从而让对抗训练更加平稳。在PyTorch中可以用torch.nn.utils.spectral_norm来包装卷积层和线性层。self.conv1 spectral_norm(nn.Conv2d(...))监控损失与可视化并重不要只看损失曲线LSGAN的损失值本身不像WGAN那样与生成质量有明确的单调关系。定期查看生成的样本图片是评估模型状态最直接的方式。如果损失在降但图片质量变差那肯定是出了问题。遇到模式坍塌怎么办如果发现生成器只产出少数几种样本。可以尝试增加噪声向量的维度给生成器更多自由度。在判别器中使用Dropout增加判别器的随机性防止其过快地记住并打击某几种模式。尝试Mini-batch Discrimination这是一种让判别器能够感知一个批次内样本多样性的技术可以有效缓解模式坍塌。检查数据确保你的训练数据本身是多样化的。学习率与优化器Adam优化器(lr0.0002, betas(0.5, 0.999))是GAN训练的黄金标配对于LSGAN同样适用。除非有充分理由否则不要轻易改动。学习率可以尝试小幅调整但0.0001到0.0005是一个比较安全的范围。LSGAN通过一个简单的损失函数替换为GAN的训练稳定性带来了质的飞跃。它就像给一辆难以驾驭的赛车换上了更可靠的轮胎和悬挂系统虽然极限速度未必最高但让绝大多数司机都能更安全、更平稳地开到目的地。理解其原理掌握其实现并知晓其边界和调优技巧你就能在解决图像生成、数据增强乃至跨模态转换等实际问题时多一件强大而顺手的工具。在实际项目中我通常会先搭建一个LSGAN基线模型它快速稳定的特性能让项目前期推进得非常高效在验证想法阶段尤其有用。