ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

GAN零基础入门:生成对抗网络原理、损失函数与PyTorch实战

GAN零基础入门:生成对抗网络原理、损失函数与PyTorch实战 1. GAN 到底在解决什么问题1.1 一个让零基础也能秒懂的例子先说个故事。假设你是个做假钞的人你花了很多年练习伪造技术目标是造出能以假乱真的钞票。与此同时银行请了一个非常厉害的验钞员他的工作就是专门识别你造的假钞。一开始你的手艺很烂验钞员一眼就能识破。但你每次被抓到后都会根据“哪里被识破”去改进把水印做得更像一点、手感做得更接近一点。而验钞员也不是吃素的他也在不断升级自己的识别手段。这样一来一回假钞越做越真验钞员也越来越刁钻最后假钞几乎和真钞一模一样甚至连验钞员都可能看走眼。这就是 GAN 的全部思想。GAN 的全称是 Generative Adversarial Network中文叫生成对抗网络里面有两个角色一个是生成器Generator就是那个造假的另一个是判别器Discriminator就是那个验钞的。两者互相博弈、互相促进最终生成器能造出极其逼真的数据。我当年第一次接触这个概念时觉得它简单到不像一个深度学习方法但深入了解后发现这个“简单”的背后藏着一套非常精巧的数学逻辑也是这几年生成式 AI 爆发的基石。你看到的 AI 换脸、动漫头像生成、老照片修复、超分辨率重建底层很多都能追溯到 GAN 这套框架。1.2 GAN 能做什么凭什么这么火在 GAN 出现之前深度学习的主流方向是判别式模型给你一张图问你是不是猫给你一段音频问是不是你的声音。这类任务本质上是“分类”模型学的是一个边界把千千万万的输入划分到有限的标签里。但生成式模型不一样。它要做的是“凭空造数据”这比“判断数据”难得多。你得先真正理解一张脸长什么样、五官怎么分布、肤色和光线怎么配合才能从一堆随机噪声里画出一张像人的脸。GAN 就是第一个能让这种“凭空造物”稳定落地的框架。它的应用场景多得吓人图像生成输入 100 维随机向量输出高清人脸图比如 StyleGAN 系列。图像翻译把草稿变成油画、把白天变黑夜、把马变成斑马代表是 CycleGAN。超分辨率与修复把模糊的小图变清晰给老照片上色、去噪代表是 SRGAN。语音与音乐生成合成语音、谱曲、降噪也有很多 GAN 的身影。数据增强当手上样本不够时用 GAN 生成更多训练样本提高模型的泛化能力。所以很多招聘要求里写“熟悉 GAN、VAE 等生成模型”不是虚的。它确实是工业界和学术界都非常看重的方向。1.3 学习 GAN 需要什么基础标题写了“零基础必看”我得负责任地说一句零基础完全能看懂 GAN 的核心思想但要真正动手训练一个能用的 GAN大概需要准备三样东西懂一点 Python至少会写循环和类会调用 PyTorch 或 TensorFlow 的 API。懂一点神经网络基础知道全连接层、卷积层、激活函数是什么不需要非常深。懂一点点概率和二分类交叉熵这是理解损失函数的关键我会在后面用大白话讲清楚。如果你前两样还不太熟建议先用一个下午看看神经网络基础再回来读这篇效果会好很多。如果你已经会训练分类网络了那这篇对你来说应该非常轻松唯一的新概念就是“两个网络打架”。2. GAN 的核心架构拆解2.1 生成器从噪声到数据的“魔术师”生成器是一个神经网络输入是一段随机噪声输出是一张图片或其他类型的数据。在原始 GAN 论文里生成器输入是一个 100 维的向量每个元素服从标准正态分布。这个向量我们通常叫 z理解成“随机的灵感种子”。这个 z 里的每个数字本身没有意义但它组合在一起就编码了生成图片的某些特征。比如某一维可能影响了发色某一维可能影响了脸型某一维可能影响了背景亮度。当然在原始 GAN 里这种编码是完全无序的后来 StyleGAN 才把这种隐空间做得很可控。生成器本质上在做的事是把一个低维的随机分布映射到一个高维的、非常复杂的数据分布上。这句话听着玄乎其实可以换个角度理解它就像一个画家读了一段“灵感”画出一张画。一开始画得乱七八糟但随着训练它慢慢摸清了真实长什么样画得越来越像。我在实际写代码时生成器通常用转置卷积ConvTranspose2d或者上采样加普通卷积来实现把 100 维向量逐步放大成 64×64 或 128×128 的图片。每个中间层之间要加 BatchNorm 和 ReLU 激活函数这是稳定训练的关键。2.2 判别器眼睛毒辣的“鉴定师”判别器也是一个神经网络输入是图片输出是一个 0 到 1 之间的分数。这个分数表示“图片是真实的概率”。1 代表百分百确定是真的0 代表百分百确定是假的。训练时判别器会同时看到两类图一类是数据集里的真实图标签为 1另一类是生成器造出来的假图标签为 0。它的任务就是尽可能把这俩分清楚。这本质就是一个二分类问题所以判别器的输出层用 Sigmoid 激活函数损失用交叉熵。跟普通图像分类网络唯一的区别是它的对手不是固定的数据集而是一个随时在变强的生成器。判别器在结构上通常就是一个普通的分类网络可能用几层卷积加全连接。不需要特别深因为它的任务是“判断真假”而不是精细分类特征提取层够了就行。2.3 两者如何配合一场零和博弈如果只训练判别器它很快就变成火眼金睛但我们不需要一个只会鉴定的网络。如果只训练生成器而没有判别器的反馈它就不知道该往哪个方向优化只能在原地打转。GAN 的精髓在于交替训练。每一步先拿一批真实数据和一批生成数据去更新判别器的参数让它变得更眼尖然后再拿生成数据去更新生成器的参数但这时候的标签是“真图”目标是骗过刚变强的判别器。这个交替过程就是一场零和博弈。生成器的收益就是判别器的损失判别器的收益就是生成器的损失两者的期望完全相反。训练到最后理想状态下会出现一个均衡点生成器生成的数据分布完全等于真实数据分布判别器无论怎么判断正确率都只能停留在 50%跟扔硬币一样。这就是 Nash 均衡在神经网络里的体现。当然在实际操作中我们几乎永远到不了这个完美均衡点只要能生成足够逼真的图就算成功。注意虽然这里强调“交替训练”但不是说每一步都要严格一比一。很多时候生成器训练一步判别器要训练两步甚至三步这个比例需要根据训练情况动态调整。后面我会详细讲。3. 损失函数与数学原理精讲3.1 二分类交叉熵的推导要理解 GAN 的损失函数得先搞明白二分类交叉熵。这个东西本质上衡量的是“模型预测的概率分布”和“真实标签的概率分布”之间的距离。假设真实标签 y取值 0 或 1。模型输出的概率是 p它表示模型认为样本属于类别 1 的概率。那么交叉熵损失写成公式就是L -[ y * log(p) (1 - y) * log(1 - p) ]如果样本是正类y1损失就是 -log(p)p 越接近 1损失越接近 0如果样本是负类y0损失就是 -log(1-p)p 越接近 0损失越接近 0。这个公式里的负号很多人会困惑为什么前面有个负号因为 log 函数在 0 到 1 之间是负数log(0.5) ≈ -0.693如果不加负号损失就是负的不符合“损失越小越好”的习惯。函数线画出来是个从左到右递减的形状我们希望 p 越大越好时p 大原本包含的信息量应该小加负号只是为了让损失值和“我们的直觉”对齐。更本质地说交叉熵来自信息论里的 KL 散度。它衡量两个分布之间的差异差异越大熵越高需要的信息量越多损失自然越大。3.2 原始 GAN 的损失函数逐项拆解2014 年 Goodfellow 那篇名为 “Generative Adversarial Nets” 的论文里把 GAN 的整体目标函数写成了这样一个形式min_G max_D V(D, G) E_x~p_data[log D(x)] E_z~p_z[log(1 - D(G(z)))]这个式子分成两部分左边是判别器的期望收益右边是生成器的期望损失。仔细拆开来分析对于判别器来说它的目标是“最大化” V(D, G)。对真实数据 x它希望 D(x) 接近 1所以 log(D(x)) 接近 0最大值。对生成数据 G(z)它希望 D(G(z)) 接近 0所以 log(1 - D(G(z))) 接近 0也是最大值。两个加在一起判别器的最大化目标就变得很明确。对于生成器来说它的目标是“最小化” V(D, G)。但由于 V(D, G) 里只有第二项跟生成器有关所以生成器只需要最小化 E[log(1 - D(G(z)))]。这个式子越小越好意味着 D(G(z)) 越接近 1也就是生成器成功骗过了判别器。于是一个 min-max 优化问题就出现了。两个网络各怀鬼胎一个想最大一个想最小最后在博弈中找到平衡。3.3 为什么热搜里都在问“没有负号”“原始 GAN 公式的交叉熵为什么没有负号”这个我在备课查资料时才注意到原来有非常多初学者在问这个问题。先说结论这个公式本身是从“最大化似然”的角度出发写的所以没有显式加负号如果你把它改写成“损失函数”的形式负号自然就出现了。我们重新看一下判别器面对的那两个期望E[log(D(x))] E[log(1 - D(G(z)))]把这两项看作判别器的“收益”——越高越好所以是最大化问题。但是写损失函数的时候我们习惯让网络做最小化。于是只需要把上面加个负号L_D -E[log(D(x))] - E[log(1 - D(G(z)))]这时负号就出现了。所以并不是“交叉熵没有负号”而是论文在表达时站在了“最大化”视角不是在写最小化损失。初学者如果直接拿损失函数的标准公式去对照自然会觉得哪里不对劲。另外还有一层更深的视角把整个 GAN 当作一个极大似然估计问题。判别器的输出 D(x) 可以理解为“x 来自真实分布”这个事件的后验概率。我们希望最大化整个模型对真实数据的似然。对于真实数据 xlikelihood 就是 D(x)对于生成数据 G(z)likelihood 就是 1 - D(G(z))。把所有样本的 likelihood 相乘再取对数变成相加就得到了上面的公式。取对数的好处是把小概率相乘变成负数相加数值稳定性更好。3.4 全局最优解生成器完美收敛的证明很多人在看 GAN 原论文时会对那个“全局最优解证明”一头雾水。这个证明其实是一个很漂亮的数学推导核心结论是当且仅当生成数据的分布 p_g 等于真实数据分布 p_data 时整个博弈达到全局最优。大致思路是固定生成器 G 时我们可以写出最优判别器 D* 的显式表达。对目标函数里的被积函数求导并令导数为零设 y D(x)那么目标项是 p_data(x) * log(y) p_g(x) * log(1 - y)。对 y 求导p_data(x) / y - p_g(x) / (1 - y) 0解得y p_data(x) / (p_data(x) p_g(x))这就是最优判别器的形式。当 p_g 和 p_data 相等时y 0.5也就是判别器完全分不出来输出概率永远是一半一半。再把 D* 代回目标函数可以证明此时的目标函数值恰为 -2log2。如果 p_g 不等于 p_data目标函数值一定会大于 -2log2。所以这个点就是全局最小值的位置对应生成器已经完美学会了数据分布。这个证明过程中没有任何需要“调参”的部分非常优雅。但是注意它假设生成器和判别器都有无限容量且每一步都能更新到最优这在现实里都不成立所以 GAN 训练才那么脆弱。4. 从零训练一个 GANPyTorch 实战4.1 环境准备与数据加载纸上谈兵差不多了咱们动手写代码。环境我建议用 PyTorch原因是它的自动求导和动态图机制让 GAN 的实现非常直观。如果你机器上没有 GPU用 CPU 也能跑通 MNIST 这个例子只是慢一点。先装依赖pip install torch torchvision matplotlib然后加载 MNIST 数据集并做预处理。MNIST 是 28×28 的灰度手写数字图片类别 0 到 9。我们不需要标签因为 GAN 是无监督的只要图片本身。import torch import torch.nn as nn import torch.optim as optim import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader import matplotlib.pyplot as plt # 不用归一化到 [-1, 1]后面用 Tanh 输出时最好统一 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) batch_size 128 dataset torchvision.datasets.MNIST( root./data, trainTrue, transformtransform, downloadTrue ) dataloader DataLoader(dataset, batch_sizebatch_size, shuffleTrue)这里Normalize([0.5], [0.5])的作用是把像素值从 [0,1] 映射到 [-1,1]因为生成器最后用的 Tanh 激活函数输出范围是 [-1,1]保持一致能让训练稳定很多。数据集准备就绪后就可以定义生成器和判别器了。MNIST 是 28×28 灰度图比较简单用全连接层也能取得不错效果但为了更接近实际工程我用了转置卷积和卷积结构。4.2 生成器与判别器网络定义生成器输入是 100 维噪声向量输出是 1×28×28 的图。我走的路线是全连接层 Reshape 转置卷积 Tanh。class Generator(nn.Module): def __init__(self, latent_dim100): super().__init__() self.fc nn.Linear(latent_dim, 64 * 7 * 7) self.conv_layers nn.Sequential( nn.BatchNorm1d(64 * 7 * 7), nn.ReLU(), # 把 [batch, 64*7*7] 变形为 [batch, 64, 7, 7] nn.Unflatten(1, (64, 7, 7)), nn.ConvTranspose2d(64, 32, kernel_size4, stride2, padding1), nn.BatchNorm2d(32), nn.ReLU(), nn.ConvTranspose2d(32, 1, kernel_size4, stride2, padding1), nn.Tanh() ) def forward(self, z): return self.conv_layers(self.fc(z))转置卷积的参数需要自己算一下输出尺寸。假设输入尺寸是 H_in卷积核 k、步长 s、填充 p输出尺寸就是 (H_in - 1) * s - 2 * p k。从 7×7 经过 stride2 的转置卷积后变成 14×14再经过一次变成 28×28正好匹配 MNIST。判别器就用普通卷积网络输入是一张 28×28 图输出是一个标量概率。class Discriminator(nn.Module): def __init__(self): super().__init__() self.conv_layers nn.Sequential( nn.Conv2d(1, 32, kernel_size4, stride2, padding1), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(32, 64, kernel_size4, stride2, padding1), nn.BatchNorm2d(64), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(64, 1, kernel_size4, stride2, padding0), ) self.fc nn.Sequential( nn.Linear(1 * 3 * 3, 1), nn.Sigmoid() ) def forward(self, img): x self.conv_layers(img) x x.view(x.size(0), -1) return self.fc(x)判别器里用了 LeakyReLU 而不是 ReLU避免在负区间死掉。这是一个非常重要的小细节。4.3 训练循环与超参数设定准备好两个网络后定义优化器。生成器和判别器分别用独立的 Adam 优化器。学习率一般取 1e-4 到 3e-4 之间太高容易震荡太低收敛慢。Beta1 设置成 0.5 能降低训练初期的波动这在 GAN 中是一个常用技巧。latent_dim 100 lr 2e-4 G Generator(latent_dim) D Discriminator() G.to(device) D.to(device) opt_G optim.Adam(G.parameters(), lrlr, betas(0.5, 0.999)) opt_D optim.Adam(D.parameters(), lrlr, betas(0.5, 0.999)) criterion nn.BCELoss()训练循环里每一步都做这样几件事取一批真实图片标签置 1计算判别器在真实图上的损失。生成一批噪声标签置 0计算判别器在假图上的损失。把两个损失相加反向传播更新判别器。再生成一批新噪声但标签置 1训练生成器目的是让生成器学会“骗人”。每过一定轮数生成一张图看看效果。num_epochs 50 for epoch in range(num_epochs): for i, (imgs, _) in enumerate(dataloader): batch_size imgs.size(0) real_imgs imgs.to(device) real_labels torch.ones(batch_size, 1, devicedevice) fake_labels torch.zeros(batch_size, 1, devicedevice) # 训练判别器 z torch.randn(batch_size, latent_dim).to(device) fake_imgs G(z).detach() loss_D_real criterion(D(real_imgs), real_labels) loss_D_fake criterion(D(fake_imgs), fake_labels) loss_D loss_D_real loss_D_fake opt_D.zero_grad() loss_D.backward() opt_D.step() # 训练生成器 z torch.randn(batch_size, latent_dim).to(device) fake_imgs G(z) loss_G criterion(D(fake_imgs), real_labels) opt_G.zero_grad() loss_G.backward() opt_G.step() print(fEpoch [{epoch1}/{num_epochs}] Loss_D: {loss_D.item():.4f} Loss_G: {loss_G.item():.4f})这里必须注意一个细节训练判别器时传给它的假图要用fake_imgs.detach()也就是切断梯度传播。我们只是想训练判别器去识别假图而不是让生成器收到来自判别器的梯度。如果不 detach梯度会一路传回生成器干扰生成器的优化。训练生成器时则不能 detach要让梯度从判别器流回生成器。4.4 训练结果观察训练 50 个 epoch 后你可以把生成器输出可视化。正常情况下你会看到从最初的纯噪声逐步变成模糊的数字轮廓最后变成清晰可辨的手写数字。我用这套代码跑出来的效果大约在第 5 个 epoch 就能看到数字雏形20 个 epoch 后已经很清晰。下面是保存生成图片的代码def save_samples(generator, epoch, pathsamples): generator.eval() z torch.randn(16, latent_dim).to(device) imgs generator(z).cpu().detach() imgs (imgs 1) / 2 # 从 [-1,1] 还原到 [0,1] fig, axes plt.subplots(4, 4, figsize(6, 6)) for idx, ax in enumerate(axes.flatten()): ax.imshow(imgs[idx].squeeze(), cmapgray) ax.axis(off) plt.savefig(f{path}_epoch{epoch}.png) plt.close() generator.train()注意保存前要eval()保存后再切回train()。不然 BatchNorm 的统计量会受 eval 状态影响虽然在我们这个简单的例子里影响不大但养成好习惯很重要。5. 训练 GAN 必须掌握的 6 个调参技巧与细节5.1 BCE 损失里的标签设置我在代码里用了criterion nn.BCELoss()它要求输入的概率和标签都在 0 到 1 之间。判别器输出经过 Sigmoid 后满足这个条件。标签方面常见做法是真实图用 1假图用 0。但这里有个被很多人忽略的点标签不要都用严格的 0 和 1最好用平滑过的标签比如真实标签用 0.9假标签用 0.1。这就是所谓的标签平滑Label Smoothing。它可以防止判别器过于自信给生成器留下更多梯度信息。如果判别器的输出老是极度接近 1 或极度接近 0梯度会迅速消失生成器学不到东西。我实测下来标签平滑在 GAN 里真的能明显提升训练稳定性强烈建议尝试。5.2 学习率与优化器选择GAN 的优化问题不是一个凸优化问题常规的带动量 SGD 容易陷入震荡。Adam 凭借自适应学习率成为 GAN 训练的首选。学习率不要用 1e-3 这种偏大的值判别器和生成器之间容易互相踩脚。我个人习惯是先在 2e-4 起步如果 Loss 震荡厉害降到 1e-4。Adam 的 beta1 默认是 0.9但 GAN 里推荐设成 0.5甚至 0.3。原因在于一阶矩的长期累积会带入较早的梯度信息生成器在快速变化时会被“惯住”。减小 beta1 等于让优化器更快遗忘过去的梯度更适应 GAN 这种动态博弈场景。5.3 BatchNorm 与激活函数生成器里几乎每个中间层都要用 BatchNorm。它的好处是让每一层的输入分布保持稳定避免生成器输出波动过大。判别器里加一层 BatchNorm 也可以但要注意不要在输出层前加会影响概率输出。激活函数方面生成器隐层用 ReLU输出层用 Tanh。ReLU 不会饱和梯度传导好Tanh 把输出限制在 [-1,1]和归一化的图片范围一致。判别器隐层用 LeakyReLU负斜率设置为 0.2这样即使输入落入负区间梯度也不会变成 0。如果你用 ReLU 作为判别器的激活函数输入一旦为负神经元就死了梯度永远为 0判别器完全失去学习能力。这个坑我踩过当时检查了半天才发现是激活函数的问题。5.4 训练比例的平衡理论上判别器和生成器应该各训练一步交替进行。但如果你发现判别器 loss 降到极小比如 0.01说明它太强了生成器完全骗不过它。这时应该让生成器多训练几步或者给判别器加大学习率惩罚。反过来如果生成器 loss 降到极小、判别器 loss 一直不下降说明判别器太弱需要多训练判别器。一个比较常用的策略是每训练一次生成器训练 k 次判别器k 在 1 到 3 之间调。没有万能比例只能观察损失曲线动态调整。5.5 损失曲线的正确解读方式初学 GAN 的人最容易问好的 GAN 训练曲线长什么样答案可能让你意外——好的 GAN 训练过程中生成器和判别器的损失不一定是稳定下降的它们可能是来回波动的。因为这是一个零和博弈一方变强另一方损失就变大非常正常。你应该关注的是生成图片的视觉效果和多样性而不是盯着 loss 数字追求“越小越好”。判别器的 loss 如果长期徘徊在 0.69 附近log2说明判别器已经基本分不清真假了这在 GAN 里反而是好消息。注意loss 从 0.69 附近开始是正常的。因为训练刚开始时判别器面对真假各一半的样本如果它完全乱猜期望损失就是 -0.5log(0.5) - 0.5log(0.5) log2 ≈ 0.693。如果开场 loss 远低于这个值说明判别器一上来就学会了明显差异这不是坏事但要注意它可能太快变强。5.6 评估生成效果的艺术GAN 不像分类问题那样有一个明确的准确率指标。怎么判断训练得好不好最基本的办法就是人工看图。但人眼有主观性所以后来有了 FIDFréchet Inception Distance这类指标。FID 用 Inception V3 网络提取特征计算真实图和生成图在特征空间上的分布距离。得分越低说明两个分布越接近。对初学者来说先把眼睛练好能肉眼判断生成图片清晰度、多样性、有没有明显重复就已经比死盯 loss 强多了。6. 常见问题与排查实录6.1 模式崩溃生成的都是同一类图模式崩溃Mode Collapse是 GAN 训练里最经典的问题。表现就是生成器找到一条“捷径”输出的图几乎一样或者只有少数几种失去了多样性。我遇到过一次生成器只输出几个固定数字怎么调都绕不开。后来排查发现是生成器过于强大已经能够在判别器面前“蒙混过关”但它是靠重复生成某一类样本来做到的而不是真正学到了整个数据分布。解决办法可以从这几方面入手降低生成器学习率让它不要甩开判别器太远。增强判别器的能力比如加深层数或者让判别器多训练几步。用 Mini-Batch Discrimination 或引入正则化手段让生成器考虑样本之间的多样性。6.2 判别器 loss 直接归零如果你看到判别器的 loss 在训练初期就掉到 0.001 以下这是一个危险信号。它说明判别器已经完全碾压生成器生成器产出的图像一眼假判别器毫不费力就能识别。这个时候生成器的梯度几乎是 0完全没有学习动力。解决办法是让生成器“喘口气”。具体操作包括把生成器的学习率调大一点点让它更新幅度更大。减小判别器的学习率。训练判别器时使用标签平滑降低它的自信程度。让生成器每训练两次判别器才训练一次。我见过不少人一看到判别器 loss 很低就觉得训练很成功其实恰恰相反这往往是灾难的前奏。6.3 生成器 loss 不下降生成器 loss 一直在零点几徘徊生成的图片始终是模糊的一团。这个情况通常意味着判别器给出的梯度信号太微弱或者生成器本身没有足够容量去建模复杂分布。你可以做几件事把噪声向量的维度加大一点比如从 100 变成 256。增强生成器的网络宽度比如把中间层通道数翻倍。换成更好的损失函数比如 WGAN 的 Wasserstein 距离它对生成器的梯度要平滑得多。检查是不是数据归一化没有做好生成器 Tanh 输出 [-1,1] 但真实图还在 [0,1]两者范围不一致会影响训练。6.4 训练过程震荡loss 像过山车训练过程中损失剧烈波动生成图片时好时坏这种情况其实是 GAN 的常态不用太焦虑。如果波动实在太大可以用下面几个手段压一压降低学习率给双方“冷静一下”的空间。使用批量归一化之外的 LayerNorm有时能稳定很多。使用梯度惩罚Gradient Penalty让判别器的梯度变化更平滑。WGAN-GP 就是经典方案。6.5 上手建议速查表症状可能原因实践方案生成图模糊不清生成器容量不足增大网络宽度或深度提高噪声维度生成图缺乏多样性模式崩溃降低生成器 lr引入多样性惩罚判别器 loss 几乎为 0判别器太强标签平滑减小判别器 lr增加生成器训练步数生成器 loss 不降生成器梯度消失换 WGAN 或使用 LeakyReLU检查归一化训练震荡剧烈学习率过高降低 lr调整 beta1 至 0.5生成器输出全是 0 或边界值Tanh 饱和降低判别器 lr使用 BatchNorm 稳定生成器最后再说一个我自己踩过的坑训练 GAN 遇到问题很多人第一反应是换架构、换损失函数但我现在会先检查一件事数据预处理和网络输出范围是不是一致。我早期做实验的时候真实图片归一化到了 [0,1]生成器输出层用的 Tanh 输出 [-1,1]结果判别器很容易通过“值的绝对值大小”判断真假根本不需要学图像特征。这算是一个特别隐蔽的 bug但它会把训练节奏彻底打乱。所以不管你用什么数据、什么网络结构先确认输入范围、输出范围、损失函数三者匹配再谈调参。GAN 的学习曲线确实比分类网络陡峭不少但当你第一次看到生成器从纯噪声里慢慢画出清晰的数字时那种“入门”的感觉比训练一百个分类模型还要爽。这篇的代码和思路全部基于原始 GAN 以及我自己的实战经验你可以直接照着写一遍。跑通了之后再去尝试 DCGAN、WGAN、StyleGAN会顺畅很多。
返回列表