ARTICLE DETAIL

资讯详情

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

对抗生成网络架构拆解:生成器判别器、损失函数与训练排查

对抗生成网络架构拆解:生成器判别器、损失函数与训练排查 对抗生成网络这几年从学术圈的玩具变成了不少业务里的常规工具做数据增强、做图像修复、做风格迁移、做合成样本补足长尾类别都能见到它的身影。但真正上手写过的人都知道对抗生成网络最劝退的地方不是代码量而是它的架构设计逻辑跟普通的分类网络完全不是一回事你训练的是一个会互相拆台的系统损失曲线好看不代表生成质量好判别器太强反而会把生成器饿死。这篇内容我打算把对抗生成网络的架构原理从头拆一遍包括生成器与判别器各自该怎么设计、张量尺寸怎么推、损失函数怎么配、训练过程中崩掉之后怎么排查。不管你是刚接触生成模型的学生还是准备把合成数据接进业务管线的工程师都能从这套拆解里拿到可以直接抄的结构和参数。1. 先搞清楚对抗生成网络在博弈什么很多人第一次看对抗生成网络的论文会被那一串极小极大公式劝退其实它的核心思想非常朴素让一个人造假让另一个人鉴假两个人水平同步上涨。造假的叫生成器鉴假的叫判别器两者共享同一个数据分布作为目标。理解这个博弈关系比背公式重要得多因为后面所有的架构调整、损失函数变体、训练技巧本质上都是在调这两个角色的强弱平衡。1.1 生成器与判别器的角色分工生成器的输入是一段随机噪声通常是从标准正态分布里采样的向量维度常见的有 100、128、256。它要做的事情是把这段毫无结构的噪声映射成一张看起来像真实样本的图片。注意这里的映射必须是确定的函数同样的噪声输入必须得到同样的输出否则训练过程就没法收敛。生成器内部没有任何标签信息它唯一的学习信号来自判别器给出的反馈。判别器的输入是一张图片输出是一个标量表示这张图来自真实数据的概率。它的训练数据一半来自真实数据集一半来自生成器当前的输出。判别器的目标是尽可能把两者分开生成器的目标是尽可能让判别器分错。这就是对抗的来源一方的损失下降往往意味着另一方的处境变难。我常跟新人打个比方生成器像是一个临摹字帖的学生判别器像是一个只看笔迹的鉴定师。学生一开始乱涂鉴定师一眼就认出是假的学生慢慢学会模仿笔画走势鉴定师被迫去看更细的结构比如笔锋、间距、墨色。双方在互相逼迫中一起进步。但如果鉴定师太强学生怎么改都被否定就干脆摆烂只写一个最像的字反复交差这就是后面要讲的模式崩塌。从参数量上看两个网络的规模一般不需要对等。生成器负责从低维到高维的扩张计算量通常更大判别器做的是降维分类结构可以更浅。我自己的习惯是判别器参数量控制在生成器的三分之一到二分之一之间这个比例在 32×32 到 128×128 的分辨率区间都比较好用。1.2 用生活场景理解极小极大目标函数原始对抗生成网络的目标函数写成这样min_G max_D V(D, G) E_{x~p_data}[log D(x)] E_{z~p_z}[log(1 - D(G(z)))]拆开看就两件事。判别器 D 要最大化这个值对真实样本 x它希望 D(x) 接近 1所以 log D(x) 接近 0对生成样本 G(z)它希望 D(G(z)) 接近 0所以 log(1 - D(G(z))) 也接近 0。生成器 G 要最小化这个值也就是让 D(G(z)) 接近 1让 log(1 - D(G(z))) 变得很小。注意生成器只能影响公式的第二项第一项跟它无关。这也是为什么在代码实现里生成器的损失通常只计算假样本那一部分不需要把真实样本再喂一遍。理论上当判别器达到最优时生成器最小化这个目标等价于最小化真实分布和生成分布之间的 JS 散度。JS 散度有个致命问题当两个分布几乎没有重叠时它是常数梯度为零。而训练初期生成分布和真实分布几乎必然不重叠这就导致生成器拿不到有效梯度训练原地不动。这个推导解释了一个反直觉的现象判别器训得太好反而害了生成器。所以在实操中我们经常故意削弱判别器比如降低它的学习率、减少它的更新次数、加 Dropout、加标签噪声。这些做法从纯理论角度看是不严谨的但从工程角度看是必要的。1.3 从 JS 散度到梯度消失判别器太强的代价我在早期做手写数字生成的时候遇到过一种很典型的场面判别器损失迅速掉到 0.01 以下D(x) 稳定在 0.99D(G(z)) 稳定在 0.001而生成器的损失一路飙到十几。当时我以为是生成器学习率太小调大之后情况更糟生成器直接输出全黑图片。后来才想明白问题出在梯度上。当判别器对假样本的判定非常自信时它对输入的梯度趋近于零生成器通过链式法则拿到的梯度也就趋近于零。生成器不是学不会而是根本没有信号告诉它该往哪改。解决思路有三条我都实际用过降低判别器的学习率让它和生成器保持接近的收敛速度比如生成器用 2e-4判别器用 1e-4。对判别器使用标签平滑把真实样本的标签从 1.0 改成 0.9 或 0.95避免它输出极端置信度。换用 Wasserstein 距离替代 JS 散度这也是 WGAN 系列的起点后面第 5 节会详细展开。这里还有个容易忽视的点判别器最后一层千万不要加 Sigmoid 再配 BCELoss数值上容易溢出。我一般直接用单输出接 BCEWithLogitsLoss它内部做了 log-sum-exp 的稳定化处理实测比手动 Sigmoid 稳定不少尤其是混合精度训练的时候。2. 架构选型的几个关键决策点动手写代码之前有几个结构性的选择必须先定下来。这些选择决定了你的模型是能在几个小时内出效果还是调一周都看不到有意义的图。我把这些决策点按重要性排了序从最影响成败的开始讲。2.1 全连接还是卷积我为什么劝新手别用 MLP 版 GAN原始论文里的对抗生成网络用的是全连接层输入 100 维噪声经过几层全连接升到 1024 维再 reshape 成 28×28。这个结构在 MNIST 上能跑通但只限于极低分辨率。原因在于全连接层把空间结构完全打平了图像里相邻像素的相关性被丢掉生成出来的图会带有明显的块状噪声而且参数量爆炸一张 128×128 的 RGB 图有 49152 个像素值第一层全连接就要吃掉几千万参数。卷积结构天然保留了局部相关性同时通过权重共享把参数量压下来。DCGAN 这篇工作把这一点讲得很透它给出了一组后来被广泛沿用的骨架规则生成器用转置卷积做上采样步长设为 2卷积核设为 4。判别器用普通卷积做下采样同样是步长 2、卷积核 4。生成器里用 BatchNorm 加 ReLU最后一层用 Tanh 把输出压到 [-1, 1]。判别器里用 BatchNorm 加 LeakyReLU斜率 0.2最后一层不加激活。提示Tanh 输出配合把真实图像归一化到 [-1, 1]是标配组合。如果你用 Sigmoid 输出把图像归一化到 [0, 1] 也能跑但实测收敛会慢一些因为 Sigmoid 在两端的梯度太小。BatchNorm 在判别器上的使用一直有争议。有观点认为它会让判别器对批次内其他样本产生依赖导致训练不稳定。我在小批量比如 64 以上的场景里用 BatchNorm 没遇到明显问题但如果做 WGAN-GP就必须去掉 BatchNorm换成 LayerNorm 或者 InstanceNorm因为梯度惩罚项要求判别器对每个样本独立。2.2 生成器的上采样路径该怎么排生成器的结构本质上是一条从低维到高维的路径。以 32×32 灰度图为例我的常规排法是z(100, 1, 1) - ConvT(100-512, k4, s1, p0) - (512, 4, 4) - ConvT(512-256, k4, s2, p1) - (256, 8, 8) - ConvT(256-128, k4, s2, p1) - (128, 16, 16) - ConvT(128-1, k4, s2, p1) - (1, 32, 32)每一层的尺寸推导公式是out (in - 1) * stride - 2 * padding kernel_size以第二层为例输入空间尺寸 4步长 2填充 1卷积核 4算出来 (4-1)×2 - 2×1 4 8正好翻倍。第一层比较特殊输入是 1×1步长 1填充 0卷积核 4算出来 (1-1)×1 - 0 4 4一步到位得到 4×4 的特征图。这里有个经验第一层直接一步放大到 4×4 比逐层放大更稳。我试过把第一层改成步长 2 的转置卷积得到 2×2然后再翻倍结果训练前期生成器更容易崩因为 2×2 的特征图容量太小后面几层要承担过重的信息扩张任务。通道数一般从 512 开始逐层减半最后按输出通道收尾。如果你的分辨率到 128×128可以在尾部多加一层通道序列变成 512、256、128、64、3。我不建议在最前面堆到 1024除了显存吃紧之外也没有明显收益反而会因为参数太多导致训练初期震荡。2.3 判别器越弱越容易入门新手最容易犯的错是按做分类任务的思路去堆判别器。加层、加宽、加残差连接结果判别器 AUC 冲到 0.99生成器彻底躺平。判别器的定位是陪练不是考官它的能力只需要略高于生成器当前的造假水平。我的常规配置是三层卷积通道 64、128、256每层步长 2最后接一个 4×4 的卷积把空间维度压到 1×1输出一个标量。这个结构在 32×32 到 64×64 上都够用。判别器里我会加 Dropout比例 0.3作用相当于给判别器加噪声防止它记住训练样本。这个技巧在处理小数据集的时候尤其有效因为真实样本数量少判别器很容易过拟合一旦过拟合它就会对训练集里的图给高分对生成图一律给低分生成器就学不到东西。另外判别器输入层之后可以加一点高斯噪声标准差 0.05 到 0.1这招来自 StyleGAN 的实践本质上是让判别器的决策边界更平滑。我实测在样本量低于一万的时候用上FID 能降两到三个点。3. 手搓一个可复现的 DCGAN完整代码与参数推导光讲原理容易飘下面把一整套能跑的代码写出来以 32×32 的灰度手写数字为例。选这个数据集是因为它小、下载快、单卡几分钟就能看到结果适合验证整套流程是否正确。3.1 环境与数据准备依赖只需要 PyTorch、torchvision 和 NumPy。我用的版本是 PyTorch 2.xCUDA 版本按你的显卡驱动选这里不展开。import torch import torch.nn as nn import torchvision import torchvision.transforms as T from torch.utils.data import DataLoader device torch.device(cuda if torch.cuda.is_available() else cpu) transform T.Compose([ T.Resize(32), T.ToTensor(), T.Normalize((0.5,), (0.5,)), # 单通道压到 [-1, 1] ]) dataset torchvision.datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) loader DataLoader( dataset, batch_size128, shuffleTrue, num_workers4, drop_lastTrue, pin_memoryTrue )drop_lastTrue这个参数值得单独说。如果最后一个批次只有几个样本BatchNorm 统计出来的均值和方差噪声极大会让整个训练步产生异常梯度。我在早期没加这个参数训练到后期偶尔会出现生成图片突然糊掉加了之后就再没出现过。注意如果你用的是 InstanceNorm 或 LayerNorm对批次大小不敏感drop_last就不是必须的。但只要用了 BatchNorm强烈建议加上。3.2 尺寸推导每一层张量长什么样把网络里的张量形状写清楚调试的时候对着打印比对能省掉大量排查时间。阶段输入形状操作输出形状G 第 1 层(128, 100, 1, 1)ConvT k4 s1 p0(128, 512, 4, 4)G 第 2 层(128, 512, 4, 4)ConvT k4 s2 p1(128, 256, 8, 8)G 第 3 层(128, 256, 8, 8)ConvT k4 s2 p1(128, 128, 16, 16)G 第 4 层(128, 128, 16, 16)ConvT k4 s2 p1(128, 1, 32, 32)D 第 1 层(128, 1, 32, 32)Conv k4 s2 p1(128, 64, 16, 16)D 第 2 层(128, 64, 16, 16)Conv k4 s2 p1(128, 128, 8, 8)D 第 3 层(128, 128, 8, 8)Conv k4 s2 p1(128, 256, 4, 4)D 输出层(128, 256, 4, 4)Conv k4 s1 p0(128, 1, 1, 1)卷积输出尺寸的公式和转置卷积略有不同out floor((in 2 * padding - kernel_size) / stride) 1以判别器第一层为例(32 2×1 - 4) / 2 1 16。最后输出层 (4 0 - 4) / 1 1 1得到一个标量。整个判别器参数量大约 200 万生成器大约 350 万比例合适。3.3 生成器与判别器代码class Generator(nn.Module): def __init__(self, z_dim100, base512, out_ch1): super().__init__() self.net nn.Sequential( nn.ConvTranspose2d(z_dim, base, 4, 1, 0, biasFalse), nn.BatchNorm2d(base), nn.ReLU(inplaceTrue), nn.ConvTranspose2d(base, base // 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(base // 2), nn.ReLU(inplaceTrue), nn.ConvTranspose2d(base // 2, base // 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(base // 4), nn.ReLU(inplaceTrue), nn.ConvTranspose2d(base // 4, out_ch, 4, 2, 1, biasFalse), nn.Tanh(), ) def forward(self, z): return self.net(z) class Discriminator(nn.Module): def __init__(self, in_ch1, base64): super().__init__() self.net nn.Sequential( nn.Conv2d(in_ch, base, 4, 2, 1, biasFalse), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout2d(0.3), nn.Conv2d(base, base * 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(base * 2), nn.LeakyReLU(0.2, inplaceTrue), nn.Dropout2d(0.3), nn.Conv2d(base * 2, base * 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(base * 4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(base * 4, 1, 4, 1, 0, biasFalse), ) def forward(self, x): return self.net(x).view(-1)几个细节值得说明。生成器每层都加了biasFalse因为后面紧跟 BatchNorm偏置项会被归一化抵消属于冗余参数。判别器第一层没有 BatchNorm因为输入是原始图像直接归一化反而会破坏像素级的分布信息。判别器输出层用.view(-1)拉平成 (B,) 的形状配合BCEWithLogitsLoss使用。判别器最后一个卷积是 4×4 卷积核直接压到 1×1没有用全局平均池化。这两种做法我都试过全局平均池化会让判别器更弱一些训练更稳但收敛慢直接压到 1×1 判别能力更强需要靠 Dropout 和标签平滑来压制。你可以根据实际情况二选一。3.4 损失函数、标签平滑与优化器配置z_dim 100 G Generator(z_dim).to(device) D Discriminator().to(device) criterion nn.BCEWithLogitsLoss() opt_G torch.optim.Adam(G.parameters(), lr2e-4, betas(0.5, 0.999)) opt_D torch.optim.Adam(D.parameters(), lr1e-4, betas(0.5, 0.999))优化器的 betas 参数是我踩坑最多的地方。默认的 (0.9, 0.999) 在 GAN 上经常导致训练震荡改成 (0.5, 0.999) 之后明显平稳。原因是一阶动量系数太大时历史梯度的影响持续时间太长而 GAN 的梯度和目标在训练过程中一直在变旧梯度反而是干扰。学习率生成器用 2e-4、判别器用 1e-4 是我的默认起点。如果你发现判别器损失掉得太快就把判别器的学习率再降一档如果生成器迟迟不收敛就把它的学习率升到 3e-4 试试但不要超过 5e-4。标签平滑的处理方式是这样real_label 0.9 fake_label 0.0我不用 1.0 和 0.0 这组极端值。真实样本标签给 0.9假样本给 0.0让判别器的输出不会饱和到极端值。有时候我还会给假样本标签加一点噪声在 0.0 到 0.1 之间随机取效果是让判别器的决策边界更平滑。实测下来标签平滑能让训练稳定运行的 epoch 数明显增加。3.5 训练循环里两个容易被忽略的细节完整的训练循环长这样fixed_z torch.randn(64, z_dim, 1, 1, devicedevice) for epoch in range(50): for i, (real, _) in enumerate(loader): real real.to(device, non_blockingTrue) bsz real.size(0) # ---- 更新判别器 ---- opt_D.zero_grad(set_to_noneTrue) z torch.randn(bsz, z_dim, 1, 1, devicedevice) fake G(z) d_real D(real) d_fake D(fake.detach()) loss_d_real criterion(d_real, torch.full_like(d_real, real_label)) loss_d_fake criterion(d_fake, torch.full_like(d_fake, fake_label)) loss_d loss_d_real loss_d_fake loss_d.backward() opt_D.step() # ---- 更新生成器 ---- opt_G.zero_grad(set_to_noneTrue) d_fake_for_g D(fake) loss_g criterion(d_fake_for_g, torch.full_like(d_fake_for_g, 1.0)) loss_g.backward() opt_G.step() if i % 100 0: print(fepoch {epoch} step {i} fd_loss {loss_d.item():.3f} g_loss {loss_g.item():.3f} fD(x) {torch.sigmoid(d_real).mean().item():.3f} fD(G(z)) {torch.sigmoid(d_fake).mean().item():.3f}) with torch.no_grad(): samples G(fixed_z) torchvision.utils.save_image( samples, fout/epoch_{epoch:03d}.png, nrow8, normalizeTrue )第一个细节是fake.detach()。更新判别器的时候假样本必须断开梯度否则反向传播会一路传到生成器把生成器的参数也更新一遍。我见过有人忘记加这个训练出来的模型看起来能出图但生成器实际上是被判别器的目标在优化质量差一大截。第二个细节是zero_grad(set_to_noneTrue)。这是 PyTorch 1.7 之后推荐的写法把梯度置为 None 而不是零张量能省一点显存也略微提速。在 GAN 这种每步都要清两次梯度的场景里累积起来的效果还不错。固定噪声fixed_z的作用是画训练进度图。用同一组噪声在每个 epoch 结束时生成图像你可以直观看到生成器是在进步还是在崩塌。这比看损失曲线有效得多因为 GAN 的损失值没有绝对意义只能看趋势。提示如果你的显存紧张可以把判别器的更新和生成器的更新放在两个torch.no_grad()块里分别计算避免同时保留两份中间激活。实测能省 20% 到 30% 显存。3.6 训练监控与评估指标判断 GAN 训练是否健康我一般看三个信号。第一个是 D(x) 和 D(G(z)) 的均值。健康状态下D(x) 在 0.7 到 0.9 之间D(G(z)) 在 0.1 到 0.3 之间两者有明显间隔但不极端。如果 D(x) 稳定在 0.99、D(G(z)) 稳定在 0.01说明判别器太强了。如果两者都在 0.5 附近晃说明判别器太弱生成器也没有明确的学习信号。第二个是生成样本的视觉质量。这是最直接的判断方式。我习惯每个 epoch 存一张 8×8 的网格图训练结束后做一次快速回放看看从第几个 epoch 开始出形状、第几个 epoch 开始出现纹理、第几个 epoch 开始退化。第三个是 FID。FID 的计算方式是先用预训练网络提取真实样本和生成样本的特征各自拟合成多元高斯分布然后计算两个分布之间的 Frechet 距离FID ||mu_r - mu_g||^2 Tr(Sigma_r Sigma_g - 2 * sqrt(Sigma_r * Sigma_g))FID 越低越好但它需要足够数量的样本才稳定一般要一万张以上。我在小规模实验里用得不多因为算一次要几分钟而且样本量不够时波动很大。做正式对比实验的时候我会在训练结束后统一算一次不放进训练循环。还有一个低成本指标是最近邻检查从生成样本里随机抽一批在真实训练集里找它们的最近邻。如果生成图和最近邻几乎一模一样说明模型在记忆训练数据泛化性差如果最近邻在语义上相似但细节不同说明模型学到了分布规律。这个检查用预训练特征算余弦距离就行不需要额外训练。4. 训练崩掉的典型症状与排查手册GAN 的训练失败方式有很多种但症状就那么几类。下面按症状分类把可能的原因和对应的处理办法整理出来这套排查顺序是我这几年积累下来的基本能覆盖九成以上的问题。4.1 模式崩塌模式崩塌的表现很典型生成器只会生成少数几类样本比如训练手写数字时只出 1 和 7其他数字完全看不到。原因是生成器发现只要骗过判别器的那几个安全样本就够拿高分没必要覆盖整个数据分布。处理模式崩塌我有几个常用手段。第一个是加小批量判别让判别器不只评估单张图还能看到一批样本之间的差异从而对重复样本给低分。第二个是换 WGAN-GPWasserstein 距离在分布不重叠时依然能提供有效梯度生成器不会因为某一类样本难学就放弃它们。第三个是调整判别器的更新频率改成每更新一次判别器就更新两次生成器逼迫生成器更快地覆盖分布。注意模式崩塌有时是数据不平衡导致的。如果你的训练集里某一类样本特别少生成器自然倾向于忽略它。这种情况下先去处理数据再调模型。我实际排查时还会做一个实验固定判别器只训练生成器若干步看生成样本的多样性是否提升。如果提升明显说明问题在判别器压制得太狠如果没变化说明生成器的容量或结构有问题。4.2 损失震荡与梯度消失损失剧烈震荡一会儿 d_loss 是 0.2一会儿跳到 4.0通常和批次大小、学习率、BatchNorm 有关。批次太小时 BatchNorm 的统计量不稳我的处理办法是把批次加到 128 以上或者把判别器里的 BatchNorm 换成 InstanceNorm。梯度消失的表现是生成器损失几乎不动生成图像在一个较差的水平上停滞。除了前面提到的降低判别器强度还可以在生成器损失里加一点 L1 或 L2 正则让生成器有一个稳定的辅助目标。我试过在生成器损失里加上对生成图像总变分的惩罚权重给 1e-5 量级能缓解生成图像里的高频噪声但权重不能大否则图像会糊。梯度爆炸相对少见一般出现在学习率过大或者没用 BatchNorm 的配置里。如果看到损失突然变成 NaN先检查学习率再检查有没有对输入做正确的归一化。梯度裁剪也是一个保险手段把判别器的梯度范数裁到 1.0。4.3 棋盘伪影与颜色失真转置卷积有一个结构性问题当卷积核大小不能被步长整除时输出图像会出现周期性的棋盘格伪影。比如卷积核 3、步长 2 的组合就会出问题因为不同的输出像素被覆盖的次数不一致。解决方式有两个。一是使用卷积核 4、步长 2 的组合4 能被 2 整除每个输出像素的覆盖次数一致。这也是我在前面结构里一直用 k4s2 的原因。二是先上采样再卷积用最近邻插值把特征图放大两倍然后接一个普通卷积做特征整合。这种方式计算量稍大但完全没有棋盘伪影StyleGAN 系列就是用的这个方案。颜色失真一般是训练不充分或者数据归一化不一致导致的。检查两点真实图像的归一化范围是否和生成器输出层匹配判别器是否在训练早期就过拟合。如果生成图像整体偏绿或者偏紫多半是生成器在某个通道上产生了偏置可以检查一下生成器最后一层的初始化方式用正态分布初始化并缩小标准差会有帮助。4.4 常见问题速查表症状可能原因处理方式D(x)≈1D(G(z))≈0判别器过强降低 D 学习率加 Dropout标签平滑只生成少数类别模式崩塌加小批量判别换 WGAN-GP调整更新比例生成图像全灰或全黑梯度消失或输出层饱和检查 Tanh 前的 BatchNorm降低学习率损失出现 NaN学习率过大或数值溢出降学习率用 BCEWithLogitsLoss加梯度裁剪图像有网格状纹路转置卷积核与步长不匹配改用 k4s2 或先上采样再卷积训练后期质量突然下降判别器过拟合或批次过小加数据增广检查 drop_last加噪声生成图与训练图几乎相同判别器记忆训练集加 Dropout加输入噪声减少 D 参数量损失一直不降数据归一化不一致检查真实数据与生成器输出的值域是否匹配这张表我建议贴在显示器旁边遇到问题先对号入座能省掉大量盲调的时间。5. 主流变体架构什么场景换什么刀基础版 DCGAN 能解决低分辨率的生成问题但一旦遇到训练不稳、需要条件控制、需要高分辨率就得换架构。下面按问题类型梳理几个主流方向每个方向说清楚它改了什么、为什么这么改、什么时候该用。5.1 WGAN 与 WGAN-GPWGAN 的核心改动是把 JS 散度换成 Wasserstein 距离同时去掉判别器最后一层的 Sigmoid判别器不再输出概率而是输出一个实数评分。这样一来即使两个分布完全不重叠Wasserstein 距离依然能给出有意义的梯度。为了满足 Wasserstein 距离要求判别器满足 Lipschitz 连续条件原始 WGAN 用的是权重裁剪每次更新后把判别器参数裁到 [-0.01, 0.01]。这个做法简单但容易导致参数被压到边界上训练效果受裁剪阈值影响很大。WGAN-GP 改用梯度惩罚项loss E[D(fake)] - E[D(real)] lambda * E[(||grad D(x_hat)||_2 - 1)^2]其中 x_hat 是在真实样本和生成样本之间做随机插值得到的点lambda 一般取 10。梯度惩罚比权重裁剪稳定得多代价是每次判别器更新要额外算一次梯度训练速度大概慢 20% 到 30%。注意WGAN-GP 里判别器不能用 BatchNorm因为梯度惩罚是逐样本计算的BatchNorm 会让同批次样本产生耦合破坏惩罚项的含义。换成 LayerNorm 或 InstanceNorm 都可以。我在 64×64 以上分辨率的任务里基本都用 WGAN-GP训练曲线平稳模式覆盖也更完整。代价是调参要多一个 lambda不过 10 这个默认值很少需要改。5.2 条件式与图像翻译CGAN、Pix2Pix、CycleGAN基础 GAN 只能随机生成无法控制生成内容。CGAN 的做法是把标签信息拼接到生成器输入和判别器输入上。生成器这边把类别标签做嵌入后和噪声向量拼接判别器那边把标签扩展到空间维度后和图像在通道维度拼接。Pix2Pix 把这种条件控制扩展到图像到图像的翻译输入不是噪声而是另一张图生成器采用 U-Net 结构判别器采用 PatchGAN也就是输出一个特征图而不是单个标量每个位置对应输入图像的一块区域。PatchGAN 的好处是能关注局部纹理同时参数量比全局判别器少得多。CycleGAN 解决的是无配对数据的翻译问题它用两个生成器和两个判别器构成循环加上循环一致性损失约束两个方向的映射互为逆运算。这个架构在风格迁移、季节转换、材质替换这类任务里特别实用因为收集配对数据成本太高。我做过一个把草图转成实物照片的小项目用的就是 Pix2Pix。经验是 U-Net 的跳跃连接对细节保留非常关键去掉之后边缘会明显变糊。判别器的感受野也要和任务匹配做局部纹理转换用 PatchGAN做整体风格转换用全局判别器。5.3 高分辨率路线StyleGAN 与 EMA 权重到 256×256 以上前面这些架构就开始吃力了。StyleGAN 系列的关键改动有三个。一是把噪声和风格解耦用一个映射网络把噪声映射到中间隐空间再通过自适应实例归一化注入到生成器的每一层。二是逐层注入随机噪声让人脸的发丝、皮肤纹理这类细节更加自然。三是使用 Equalized Learning Rate让每一层的更新尺度一致。StyleGAN2 进一步去掉了自适应实例归一化带来的水滴状伪影改用权重解调。同时引入路径长度正则化让隐空间更加平滑便于插值编辑。工程上还有一个通用技巧是高分辨率 GAN 都会用的EMA 权重。训练过程中维护一份生成器参数的指数滑动平均评估和推理时用这份权重而不是当前权重。EMA 的衰减系数一般取 0.999 或 0.9999实测能把 FID 降好几个点因为这相当于对训练后期的参数做了集成平滑掉了随机波动。这个技巧和具体架构无关任何 GAN 都能用。5.4 现在还要不要学 GAN这几年扩散模型在图像生成上声量很大很多人问 GAN 是不是过时了。我的判断是在追求极致生成质量和多样性的场景里扩散模型确实更占优势但 GAN 有两个扩散模型短期难以替代的特点一是推理速度快只需要一次前向传播而扩散模型要迭代几十步二是它的隐空间结构更规整做属性编辑和插值的时候更可控。实际业务里实时人脸特效、视频帧级的风格化、移动端的轻量生成这些场景对延迟敏感GAN 依然是首选。而且理解对抗训练这套机制对理解扩散模型的训练目标、理解各种生成模型的评估方式都有帮助。所以我的建议是先把 GAN 的架构吃透再去学扩散模型路会顺很多。6. 从训练到落地显存、速度与部署模型训出来只是第一步真正接进业务管线还要过显存、速度、部署这几关。这部分内容网上讲得少我把实际项目里的经验整理一下。6.1 显存与批次的选择32×32 分辨率、批次 128 的配置生成器加判别器前向反向大概占 1.5GB 显存加上优化器状态和中间激活单卡 8GB 足够跑。64×64 分辨率、批次 64 大概占 4GB。128×128 分辨率、批次 32 大概占 8GB。这组数字是基于我前面给的结构如果你把通道数翻倍显存也要翻倍。显存不够的时候优先级顺序是先降批次再降分辨率最后才考虑减通道。降通道对生成质量的影响最大因为通道数直接决定特征表达能力。混合精度训练能省 30% 到 40% 显存但 GAN 用 AMP 要小心判别器的梯度惩罚和 BatchNorm 在 FP16 下容易出数值问题。我的做法是只在生成器和判别器的前向计算用 autocast梯度惩罚部分强制转回 FP32。批次大小对 GAN 的影响比普通分类网络大。批次太小判别器对单批次内的样本过拟合会加剧模式崩塌。我的经验是判别器里的 BatchNorm 决定了批次下限一般是 32。如果显存实在有限把 BatchNorm 换成 InstanceNorm批次可以降到 8 甚至 4但训练会更不稳定需要更小的学习率和更多的训练步数。6.2 导出与推理推理阶段只需要生成器判别器可以直接丢掉。导出流程是先把生成器设为 eval 模式然后用 torch.jit.trace 或 torch.onnx.export 导出。需要注意的是生成器的输入是四维张量 (B, z_dim, 1, 1)不是二维的很多人在导出时会忘记这一点导致 shape 不匹配。G.eval() dummy torch.randn(1, 100, 1, 1, devicedevice) torch.onnx.export( G, dummy, generator.onnx, input_names[z], output_names[image], dynamic_axes{z: {0: batch}, image: {0: batch}}, opset_version13, )导出后一定要做数值对齐验证同样一组噪声PyTorch 和 ONNX Runtime 的输出差异应该在 1e-5 以内。我遇到过一次因为 BatchNorm 的 running stats 没正确导出导致 ONNX 输出和 PyTorch 差了一大截生成图像完全不一样。推理延迟方面32×32 灰度图在单张消费级显卡上生成器的单次前向大概 2 到 5 毫秒批次 64 大概 15 毫秒。如果部署到边缘设备可以进一步做通道剪枝和 INT8 量化但要重新评估生成质量因为 GAN 的生成器对量化误差比分类网络敏感。提示部署时务必把生成器里的 BatchNorm 固定在 eval 模式避免上线后因为输入统计量变化导致输出漂移。如果你用的是 WGAN-GP推理时不需要任何额外处理直接用生成器即可。6.3 我的实战心得最后分享几个我在实际项目里反复验证过的经验。第一GAN 的项目周期里数据质量占的比重远超模型结构。我做过一次数据增强的项目前期花了两周调架构FID 卡在 35 下不去后来花三天清洗了训练集把标注错误和重复样本去掉同样的架构 FID 直接降到 22。所以如果你的生成质量不理想先去检查数据不要急着改网络。第二训练日志里除了损失一定要记录判别器输出的均值和固定噪声的生成图。损失是滞后的、间接的指标判别器输出和可视化图像才是直接信号。第三别迷信单次训练的结果。GAN 对随机种子敏感同一个配置跑两次FID 可能差三到五个点。做对比实验的时候每个配置至少跑三次取中位数否则很容易得出错误结论。第四把训练配置完整记录下来包括随机种子、库版本、超参数、数据划分。GAN 的复现难度比分类任务高得多没有完整记录的话过两周你自己都复现不出来当时的那个好结果。这套架构拆解和实操流程我从最早的 MNIST 一路用到 256×256 的人脸生成中间换过很多次结构细节但核心的博弈逻辑和张量推演方式一直没变。真正把生成器和判别器的每一步形状、每一个超参数的作用弄清楚遇到问题的时候就不需要靠猜了直接对着症状找原因效率会高很多。
返回列表