多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

CGAN条件生成对抗网络原理与PyTorch实现:让生成结果可控

CGAN条件生成对抗网络原理与PyTorch实现:让生成结果可控 生成对抗网络这个东西我第一次跑通的时候特别兴奋——手写数字生成得挺好。但紧接着就陷入一个尴尬生成器完全不听话我想让它写一个“7”它偏偏给你一堆“3”。如果你也有过这个经历那 CGANConditional GAN就是你下一步该学的模型。它将类别标签或任意条件信息加入生成器和判别器让生成过程从“随机撒网”变成“按需生成”。这篇文章我会从原理讲到一份可直接运行的 PyTorch 代码再把训练过程中那些最容易出问题的环节逐一拆开尽量让你看完就能自己动手复现。1. 从一个难以控制的GAN说起条件信息为什么能改变一切1.1 原始GAN的“不可控”从何而来原始生成对抗网络做的事情本质上是从一个随机噪声分布去拟合真实数据分布。生成器的输入只有一个随机噪声向量 z它不知道自己在生成什么也没有任何“目标”可以瞄准。打个比方这就像让一个司机发动汽车但没告诉他目的地他只能凭感觉往某个方向开。运气好的时候他可能开到了某个像是城市的地方但大概率是乱转。从数学上看原始GAN的目标函数长这样[ \min_G \max_D \mathbb{E}{x \sim p{data}}[\log D(x)] \mathbb{E}_{z \sim p_z}[\log(1 - D(G(z)))] ]生成器 G 只接收 z判别器 D 只接收图像 x 和生成样本 G(z)。整个系统里没有任何“类别”或者“属性”信息所以生成器学到的是一种“混合分布”——它把所有训练样本的共性都学到了但没法针对其中某一类单独建模。表现在结果上就是生成的图样虽然看起来像手写数字但你无法指定它输出哪个数字。这种现象也被称为“生成不可控”。1.2 CGAN的核心给博弈双方都加上一个“提示词”CGAN 的思路特别简单直接既然没有条件那我就把条件塞进去。具体做法是在生成器的输入里拼上条件变量 c在判别器的输入里也拼上条件变量 c。这个 c 可以是类别标签、属性向量甚至一段文本的编码只要它能够对我们想要控制的语义维度做出区分。修改后的目标函数变成[ \min_G \max_D \mathbb{E}{x \sim p{data}}[\log D(x|c)] \mathbb{E}_{z \sim p_z}[\log(1 - D(G(z|c)|c))] ]注意第二项里的形式G 在接收 z 和 c 的情况下生成假样本D 在接收假样本的同时也要接收同一个 c。这个细节很容易被忽略但极其重要。如果判别器只接收假样本而不接收 c它学到的只是“图片是否真实”而不是“图片是否满足条件”。只有同时输入 c判别器才能学会判断“这张图在类别 c 下是否合理”。这样一来生成器就被迫学会当我输入数字 7 的标签时输出的图像要尽量像真实数字 7否则判别器会依据条件 c7 轻松识破它。1.3 条件本质上是一种“约束维度”有人把CGAN理解成“多加了一个输入”我认为更准确的理解是条件信息给生成过程增加了一个约束维度。原始GAN的生成函数是 ( G: \mathbb{R}^z \rightarrow \mathbb{R}^x )生成器把噪声空间直接映射到图像空间。CGAN的生成函数则是 ( G: \mathbb{R}^z \times \mathbb{R}^c \rightarrow \mathbb{R}^x )映射关系被条件空间“剖分”了。在生成器内部条件是和噪声拼接在一起然后一起通过线性层、卷积层逐步变成图像。这相当于在特征空间里对不同的 c 划分出不同的子区域。训练充分后你在噪声维度上保持固定、只改变条件 c生成结果就会在保持整体风格差异的同时主要变化集中在类别语义上。这也是CGAN最让初学者兴奋的地方——它第一次让你对“生成”有了遥控器。2. 生成器和判别器内部是怎么“接收条件”的2.1 条件编码方式为什么用Embedding而不用One-hot矩阵处理类别标签时最直接的做法是转成 one-hot 向量然后和噪声向量拼接。比如手写数字的类别 0-9one-hot 就是一个长度 10 的向量。这样做的优点是简单缺点是 one-hot 向量非常稀疏包含的信息量极低——它的每个分量要么是0要么是1空间利用率很差。实际代码中更常用的方式是用nn.Embedding。Embedding 本质上是一个可学习的查表操作它会把一个整数标签映射成一个稠密向量。这个向量可以在训练过程中不断更新使得相近的类别在特征空间中天然接近。比如在MNIST上如果类别“1”和“7”在视觉上更相似训练出来的 Embedding 向量也会有更近的距离。从我个人的经验看Embedding 的嵌入维度设置在 32-128 之间都够用。但要注意如果维度设得过大会引入额外的参数量在小数据集上反而容易过拟合。我在MNIST上用 64 维的 Embedding 就已经能取得不错的效果。2.2 生成器从“噪声条件”到图像的逐步上采样生成器的输入是两个向量一个是随机噪声 z维度通常取 100另一个是 Embedding 后的条件向量维度取 64。两者拼接后得到 164 维的向量然后通过一个全连接层映射到足够大的空间尺寸再经过转置卷积逐级上采样到图像尺寸。MNIST 图像是 28×28 单通道。一个典型的生成器结构可以这样安排先把拼接后的 164 维向量映射到 7×7×128 的张量然后接两个转置卷积层分别从 7×7 上采样到 14×14再到 28×28。每一层后面接 BatchNorm 和 ReLU 激活函数。最后输出层用 Tanh把像素值压到 [-1,1] 区间。这里有一个值得注意的细节噪声和条件拼接的位置。早期实现经常把条件向量和噪声向量拼成一个大向量然后一次性过全连接层。这种做法简单有效但缺点是条件信息会和噪声在同一个线性变换里“纠缠”。实际上也有很多实验把条件单独做一个映射再在中间特征图通道维度上拼接这样条件的影响更可控。我在自己代码里选择了前者因为对新手更直观如果你做实验时有更高的形态要求可以尝试后者。2.3 判别器如何让“图像”和“条件”在一个网络里对比判别器要完成的任务是给定一张图像和一个条件判断这对组合是否真实且匹配。所以它必须同时看到图像和条件。图像是二维的条件在预处理阶段是向量两者不能直接拼接。常见的做法是先对条件 c 做 Embedding 和线性映射得到一个与图像尺寸相同的单通道矩阵然后把这张“条件图”与输入图像在通道维度上拼接起来。比如 MNIST 图像是 1×28×28条件映射成 1×28×28拼接后得到 2×28×28再进入卷积网络。另一个思路是把条件映射到多个通道比如 16 个通道让条件信息对后续卷积层的影响更强。不过一个通道在MNIST这种简单任务上已经完全够用。判别器的卷积部分我习惯用 LeakyReLU 激活函数不用 ReLU。因为 ReLU 在负区间的梯度完全为零容易导致神经元“死掉”LeakyReLU 保留了负区间的小梯度更适合训练判别器和生成器对抗的场景。2.4 两种条件注入方式的对比现在主流实现里条件注入方式基本可以分成两类线性拼接和特征层调制。线性拼接就是我们前面讲的把条件向量映射后直接拼到输入或中间特征上。特征层调制则更复杂例如 StyleGAN 用 AdaIN 把条件向量变换成每个特征通道的缩放和偏移量。对于CGAN入门项目线性拼接既简单又稳定没有必要一上来就上复杂的调制方案。我把两种方式的优劣列在下面供你选型时参考注入方式优点缺点适用场景输入层拼接实现简单模型收敛快条件对深层特征影响弱新手复现、类别标签特征层拼接条件影响更深层特征网络结构调整稍复杂多条件、图像级控制调制/AdaIN条件控制能力强生成质量高调参门槛高训练复杂高分辨率、细粒度生成我的建议是先跑通输入层拼接版本的CGAN理解整个训练流程和loss变化规律之后再尝试特征层拼接。直接上调制方案的话出了问题你根本分不清是条件编码的问题还是训练策略的问题。3. 代码逐段拆解从数据装载到训练循环3.1 准备工作与数据集处理代码基于 PyTorch需要安装的依赖只有torch、torchvision和matplotlib。数据集用 MNISTtorchvision可以直接下载。这里有一个关键处理MNIST 图像的像素范围是 [0,1]但生成器输出层用了 Tanh输出范围是 [-1,1]。如果两者不一致判别器会非常容易分辨真假样本——因为真实图像和生成图像的数值分布根本不在一个区间。所以数据预处理中必须加一步归一化把图像像素从 [0,1] 转换到 [-1,1]。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 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset torchvision.datasets.MNIST( root./data, trainTrue, transformtransform, downloadTrue ) train_loader DataLoader(train_dataset, batch_size64, shuffleTrue)Normalize((0.5,), (0.5,))的作用是把每个像素从 [0,1] 线性映射到 [-1,1]。公式是(x - 0.5) / 0.5。这里用 0.5 做均值、0.5 做标准差刚好能把区间对齐到 [-1,1]和生成器输出层的 Tanh 匹配。训练前可以先设置随机种子保证实验可复现torch.manual_seed(42)3.2 生成器完整代码生成器的结构我已经在第二章里画过草稿了这里直接给出可运行代码。我在每个关键层后面加了注释方便对照。class Generator(nn.Module): def __init__(self, z_dim100, num_classes10, embed_dim64): super().__init__() self.embed nn.Embedding(num_classes, embed_dim) # 噪声 条件向量拼接后先映射到 7*7*128 的中间张量 self.fc nn.Linear(z_dim embed_dim, 7 * 7 * 128) self.deconv1 nn.ConvTranspose2d(128, 64, kernel_size4, stride2, padding1) self.deconv2 nn.ConvTranspose2d(64, 1, kernel_size4, stride2, padding1) self.bn1 nn.BatchNorm2d(128) self.bn2 nn.BatchNorm2d(64) self.relu nn.ReLU() self.tanh nn.Tanh() def forward(self, z, labels): # z: (batch, z_dim), labels: (batch,) c self.embed(labels) # (batch, embed_dim) x torch.cat([z, c], dim1) # (batch, z_dim embed_dim) x self.fc(x) # (batch, 7*7*128) x x.view(x.size(0), 128, 7, 7) # (batch, 128, 7, 7) x self.bn1(x) x self.relu(x) x self.deconv1(x) # (batch, 64, 14, 14) x self.bn2(x) x self.relu(x) x self.deconv2(x) # (batch, 1, 28, 28) out self.tanh(x) return out代码里有一个非常容易踩的坑转置卷积之后的尺寸计算。ConvTranspose2d(128, 64, kernel_size4, stride2, padding1)会把 7×7 变成 14×14公式是(H_in - 1) * stride - 2 * padding kernel_size也就是(7-1)*2 - 2 4 14。下一层再从14变28。如果你自己改网络结构务必先按这个公式验算一遍尺寸否则后期会报维度错误。3.3 判别器完整代码判别器接收“图像条件”的组合输入。这里我把条件映射成 1×28×28 的条件图然后和图像拼接成 2 个通道。注意判别器里没有 BatchNorm 的全连接输出层原因后面再说。class Discriminator(nn.Module): def __init__(self, num_classes10, embed_dim64): super().__init__() self.embed nn.Embedding(num_classes, embed_dim) # 把条件向量映射到 28*28变成一张“条件图” self.fc_cond nn.Linear(embed_dim, 28 * 28) # 输入为 2 通道1通道图像 1通道条件图 self.conv1 nn.Conv2d(2, 64, kernel_size4, stride2, padding1) self.conv2 nn.Conv2d(64, 128, kernel_size4, stride2, padding1) self.conv3 nn.Conv2d(128, 256, kernel_size4, stride2, padding1) self.bn2 nn.BatchNorm2d(128) self.bn3 nn.BatchNorm2d(256) self.leaky_relu nn.LeakyReLU(0.2) self.flatten nn.Flatten() # 经过三层卷积后特征图尺寸是 3x3乘256是通道数 self.fc_out nn.Linear(256 * 3 * 3, 1) self.sigmoid nn.Sigmoid() def forward(self, x, labels): # x: (batch, 1, 28, 28), labels: (batch,) c self.embed(labels) # (batch, embed_dim) c self.fc_cond(c) # (batch, 28*28) c c.view(-1, 1, 28, 28) # (batch, 1, 28, 28) x torch.cat([x, c], dim1) # (batch, 2, 28, 28) x self.conv1(x) x self.leaky_relu(x) x self.conv2(x) x self.bn2(x) x self.leaky_relu(x) x self.conv3(x) x self.bn3(x) x self.leaky_relu(x) x self.flatten(x) x self.fc_out(x) out self.sigmoid(x) return out.squeeze(1)需要特别注意卷积层参数和特征图尺寸的匹配。输入 28×28经过 stride2 的卷积28→14→7→3最后是 3×3 的特征图所以全连接层输入是 256×3×3。如果你调整了卷积层数量或步长这里也要对应修改。3.4 训练循环里最重要的几个细节CGAN训练的总体流程和原始GAN相似先训练判别器再训练生成器。但有几个细节直接决定条件是否生效第一判别器训练时真实图像的标签必须使用它本身的真实类别不要随机指定。假图像的标签必须使用生成时输入的那个类别不要换一个标签。这样才能让判别器学会判断“这张图在这个条件下是否真实且匹配”。第二更新生成器时输入条件标签可以在0-9之间随机采样目的是让生成器学习所有类别的生成方式。如果你只固定生成某个类别图片生成器会忘掉其他类别。具体更新过程如下criterion nn.BCELoss() G Generator() D Discriminator() opt_G optim.Adam(G.parameters(), lr2e-4, betas(0.5, 0.999)) opt_D optim.Adam(D.parameters(), lr2e-4, betas(0.5, 0.999)) for epoch in range(50): for real_imgs, real_labels in train_loader: batch_size real_imgs.size(0) real_imgs real_imgs real_labels real_labels # 构造标签真实样本为1假样本为0 valid torch.ones(batch_size, 1) fake torch.zeros(batch_size, 1) # ---------- 训练判别器 ---------- z torch.randn(batch_size, 100) gen_labels torch.randint(0, 10, (batch_size,)) fake_imgs G(z, gen_labels) d_real_loss criterion(D(real_imgs, real_labels), valid) d_fake_loss criterion(D(fake_imgs.detach(), gen_labels), fake) d_loss (d_real_loss d_fake_loss) / 2 opt_D.zero_grad() d_loss.backward() opt_D.step() # ---------- 训练生成器 ---------- z torch.randn(batch_size, 100) gen_labels torch.randint(0, 10, (batch_size,)) fake_imgs G(z, gen_labels) g_loss criterion(D(fake_imgs, gen_labels), valid) opt_G.zero_grad() g_loss.backward() opt_G.step()注意到判别器训练时fake_imgs是用detach()分离的。原因很简单这一步只更新判别器参数生成器的梯度不应该回流。如果不 detach计算图会把生成器的梯度也计算出来虽然opt_D.step()只会更新判别器参数但梯度会累积在生成器的叶子节点上下一轮清空优化器梯度时才会被清掉存在潜在风险。养成习惯凡是在反传前不想更新哪部分网络就对它的输出做detach()。fake_imgs.detach()这一步是很多人写代码时容易漏掉的。漏掉后实验往往也能跑起来但显存占用会更高训练也容易变得不稳定。3.5 训练过程中的监控与可视化为了知道训练有没有正常推进我建议每轮epoch结束后固定几个类别标签生成一组图片并拼接成网格输出到屏幕上或保存成文件。这样你一眼就能看出生成器是否学到了类别区分。def show_generated_images(G, epoch, num_classes10, z_dim100): G.eval() with torch.no_grad(): z torch.randn(num_classes, z_dim) labels torch.arange(num_classes) imgs G(z, labels) imgs (imgs 1) / 2 # 从[-1,1]转回[0,1]方便显示 fig, axes plt.subplots(1, num_classes, figsize(10, 1)) for i in range(num_classes): axes[i].imshow(imgs[i].squeeze(0), cmapgray) axes[i].axis(off) plt.savefig(fcgan_epoch_{epoch}.png) plt.close() G.train()如果你发现某几个类别的生成图像明显比其他的更杂乱这说明生成器对这些类别学习得不够好需要增加训练轮数或者调整学习率。网格可视化虽然简单但在调试阶段比任何数值指标都直观。4. 训练环节最容易翻车的几个地方4.1 症状与排查表先对号入座再动手CGAN和原始GAN一样训练过程并不总是平稳的。但CGAN多了一个条件维度有时候问题会表现得更隐蔽。我把这些年复现CGAN时总结出的常见症状和对应的排查思路整理成表格建议你训练之前先扫一眼现象可能原因我的处理方式生成图模糊各个类别之间几乎看不出差异条件信息没有进入生成器或Embedding参数没有更新查看Embedding上是否有梯度确认拼接维度正确判别器loss快速降到0附近生成器loss不降判别器能力过强生成器的梯度消失调低判别器学习率或给生成器加更多BN层生成器loss降得很低但肉眼效果依然很差生成器“骗过了”判别器但并不是因为图像质量好提高判别器能力或减小batch size训练初期出现NaN学习率过高或网络初始化不当调低学习率检查是否有除零操作生成图像类别正确但形状扭曲条件信息学到的语义还不足增大Embedding维度或增强条件映射层容量4.2 判别器过强导致的梯度消失这是GAN训练中最高频的问题。判别器如果训练得太好给生成器回传的梯度会趋近于零生成器就再也学不到任何东西了。数字上表现为判别器 loss 迅速趋近于0生成器 loss 几乎是一条水平线。解决这个问题的顺序很重要。我一般遵循下面的排查链路第一步调低判别器的学习率。比如从2e-4降到1e-4。让判别器学得慢一些给生成器追赶的时间。第二步给判别器加 Dropout。在卷积层之后、全连接层之前加nn.Dropout(0.3)提升判别器的“容错率”。这种方式在很多任务上比单纯调学习率更管用。第三步更新频率不对等。每更新一次生成器更新两次或三次判别器不是这里恰恰相反。在判别器过强的情况下应该改成每更新两次生成器才更新一次判别器。虽然 PyTorch 实现上需要额外控制迭代次数但效果立竿见影。需要强调一点这三个方法不要一开始就全部用上。每次只改一个参数观察loss曲线的变化否则你根本不知道真正起作用的是哪个。4.3 条件标签失效生成结果不分类别怎么排查这种情况比梯度消失更隐蔽。损失函数数值看起来很健康判别器和生成器的loss都在正常波动但可视化结果里所有类别生成出来的图像几乎一样。这通常意味着条件信息没有对生成过程产生有效影响。我的排查顺序是第一步检查Embedding的梯度。在debug模式下打印G.embed.weight.grad如果梯度为 None 或数值极小说明条件路径没有参与到反向传播。最可能的原因是拼接时维度写错导致条件向量被后续全连接层“忽略”了。第二步检查判别器里的条件图。把c.view(-1, 1, 28, 28)之后的张量拿出来可视化一下看不同类别的条件图是否有明显差异。如果没有差异问题出在fc_cond这一层可能它把不同类别的Embedding向量都映射到了非常接近的特征图。第三步检查训练数据配对的正确性。有人会在构建批次时打乱了标签顺序或者误把gen_labels固定成了一个常数。别看这种错误低端我实际工作中确实遇到过不止一次。4.4 关于“先训练D还是先训练G”的实操经验很多教程会告诉你每个iteration里先后更新D和G顺序并不重要。但实际做实验时我发现先更新D再更新G更稳定。原因在于如果先更新G本轮D还是上一轮迭代的参数梯度方向可能已经过时G的更新会带有“滞后性”先更新D可以保证G更新时面对的是最新的判别器对当前样本的判断。还有一个经验是判别器的总参数量不要比生成器大太多。我见过有人把判别器做得特别复杂卷积层数加到七八层结果生成器完全追不上。在MNIST这种任务上判别器三层卷积、生成器两层转置卷积已经能取得很好的效果。如果做更高分辨率的图片要同步加深两侧网络而不是只加深其中一侧。4.5 训练后的样本多样性陷阱训练完成后有人满心欢喜地检查生成结果却发现生成的图片虽然清晰但每一类的手写数字风格高度相似——所有“7”都是一个写法所有“3”都是一个弯。这就是部分模式崩塌。缓解办法有几个一是把噪声向量的维度调高一些从100调到128或256给生成器更多“自由度”二是确认生成器输入噪声不是固定不变的三是在训练过程中增加判别器的难度比如用真样本和假样本数量不均衡的方式或者偶尔给真实图像加入微小的高斯噪声让判别器不那么容易把真实样本“记住”。5. 生成效果怎么量化评判不能只靠“看着像”来验收5.1 定性质检按类别生成网格图训练结束后我会固定类别标签 0-9对每个类别随机采样多个噪声向量生成一组图片并按行排列。然后逐类检查这一行里的图片是否都属于“该类别”。这个步骤虽然是人工的但能第一时间暴露出类别混淆问题。如果发现某一类生成结果里混入了其他类别的形态说明条件信息还没有完全学到位。这时最直接的方法是延长训练epoch而不是急着调整网络结构。我实验里用50个epoch做基线如果收敛良好15-20个epoch时已经能看到明显的类别区分。5.2 用预训练分类器计算“条件准确率”人工检查存在主观性而且在类别多的时候效率很低。更客观的做法是训练一个简单的MNIST分类器然后把CGAN生成的图片输入进去看分类器输出的类别与生成时设定的条件标签是否一致。这个准确率就是“条件正确率”它能直接衡量CGAN在多大程度上做到了条件可控。我自己复现时在训练30个epoch后条件正确率通常能到90%以上。如果你想用这个指标做定量对比注意生成图片在输入分类器前要把像素从[-1,1]转回[0,1]并且对齐分类器的预处理方式。否则分类器看到的是分布完全不同的输入准确率会大幅下降。5.3 从CGAN往后走FID、ACGAN与条件注入的演进如果你想让生成结果的质量评估更科学可以引入FIDFréchet Inception Distance。FID的基本思想是把真实图片和生成图片都输入预训练的Inception网络取某个中间层的特征然后计算两组特征在特征空间中的距离。距离越小说明生成分布和真实分布越接近。注意FID衡量的是分布相似度不是单张图片的相似度所以生成结果风格统一但缺乏多样性时FID反而会变差。CGAN的进阶方向也有很多。ACGANAuxiliary Classifier GAN在判别器上额外增加一个分类分支要求判别器不仅要判断真假还要判断类别条件约束更强。Projection Discriminator则基于条件与图像特征的内积来建模是Conditional GAN发展到后期的一个重要改进。后面如果你把研究方向推进到文本生成图像会发现如今很多基于扩散模型的工作里条件注入的方式依然能追根溯源到CGAN的设计思路。我个人的经验是做实验时不要满足于“训练完、看效果”有条件就把每次的条件正确率、FID和损失曲线记录下来。你会发现有些模型虽然loss曲线长得差不多但指标差异巨大。定量评估这一步在项目汇报和论文写作中都很关键。最后再分享一个调试CGAN时非常实用的小习惯在代码里写一个简单的“作弊开关”——训练完保存固定条件下的生成结果。每次改动代码或超参后隔一段时间跑一次生成把不同历史节点的输出图放在一起对比。你很快就能看到网络是越学越好还是在某个迭代点开始退化。这个习惯帮我避开了很多“训练了很久才发现模型早就崩了”的坑。
返回列表