
生成对抗网络GAN现在已经是深度学习里绕不开的一个名字了而CGANConditional GAN条件生成对抗网络则是让GAN从“能画图”进化到“按指定要求画图”的关键一步。我在实际项目里第一次用CGAN跑手写数字生成那种“让模型生成数字7它就生成7”的体验比纯粹无条件的GAN直观太多。这篇博文就用PyTorch从零实现一个基于MNIST的CGAN代码逐段拆开讲原理尽量说人话适合想真正把对抗生成网络原理和代码一起吃透的读者。1. 从GAN到CGAN先搞清楚生成对抗网络在解决什么问题1.1 造假画师与鉴定师的博弈GAN的基本盘讲CGAN之前必须先讲清楚GAN的基础逻辑。GAN的核心思想可以粗暴理解成一场“造假画师”和“鉴定师”之间的博弈。造假画师生成器负责画图一开始画得稀烂而鉴定师判别器负责判断眼前的画到底是真迹还是赝品。画师每被拆穿一次就偷偷改进一次画技鉴定师每被骗一次也逼着自己变得更刁钻。两个人互相“卷”最后卷到画师画出来的东西能乱真鉴定师已经分不清真假。从这个故事能提炼出GAN的两个核心角色生成器GGenerator和判别器DDiscriminator。G接收一个随机噪声向量z输出一张假图G(z)D接收一张图像输出一个0到1之间的概率表示“这张图是真图的概率”。D的目标是把真图判为1、假图判为0G的目标是让自己的假图尽可能被判为1。两人的博弈在数学上写成一个极小极大问题min_G max_D V(D,G) E[log D(x)] E[log(1 - D(G(z)))]其中x来自真实数据分布z来自噪声分布。理想状态下当G完全骗过D时D(G(z))的输出会稳定在0.5也就是鉴定师彻底放弃治疗开始随机瞎猜。原始的GAN有个很扎心的问题生成完全不可控。训练完一个GAN之后你往G里丢一个随机噪声它确实能生成一堆看起来像样的图但你决定不了它具体生成什么。你运气好丢进去一个噪声出来是数字“3”再丢一个可能就变成“8”了。这就好比你请了个画家但画家画什么全凭心情你没法跟他说“帮我画一只猫”或者“画个穿红衣服的人”。在真实应用里这种不可控性几乎是致命的。1.2 CGAN多了什么给生成过程加一个“命题”CGAN的论文是2014年由Mirza和Osindero提出的标题就叫Conditional Generative Adversarial Nets。它的核心改动非常简洁给生成器和判别器都额外输入一个条件y。这个y可以是类别标签、属性向量、文本描述甚至可以是另一张图像。对于MNIST手写数字任务来说y就是0到9的数字类别。有了y之后生成器的输入从单纯的噪声z变成了拼接后的(z, y)。你告诉它“画个7”它就必须往7的方向去画。判别器也一样输入变成了(x, y)它不再只判断“这张图是不是真图”而是判断“这张图在条件y下是不是真图”。换句话说D收到的命题是“这张图是不是一张真实的7”而不是笼统的“这张图是不是真实图片”。如果给它一张真实的“3”配上标签“7”它会毫不犹豫判假因为这张图不符合“真实的7”这个命题。对应的损失函数变成min_G max_D V(D,G) E[log D(x|y)] E[log(1 - D(G(z|y)))]单看公式变化就是每个输入后面多了一个条件y。但这个“多一嘴”的意义非常深远它把“生成什么东西”这件事从模型自带的神秘意志里剥离了出来变成一个可以由人明确控制的变量。我在第一次跑通CGAN的时候最大的感受就是“终于有了掌控感”。同样是玩MNIST数字生成原始GAN生成出来的图是一锅乱炖你根本不知道下一个batch是哪几个数字CGAN则可以把固定噪声配上不同的标签让同一批噪声生成出完全不同的数字形态。这种感觉非常像从“拆盲盒”变成了“按菜单点菜”。1.3 CGAN能用在哪从手写数字到图像编辑CGAN的价值不光体现在MNIST这类玩具数据集上。它的条件机制实际上是整个条件生成模型家族的基石后来的很多方法都把条件思想用到了更深的地方。表情/属性控制给定一张人脸图配上“戴眼镜”“金发”“微笑”等属性标签生成对应的人脸编辑结果。图像翻译pix2pix系列把输入图像当作条件输入边缘图生成真实照片输入白天街景生成夜晚街景。文字转图像用文本编码器把一句话编码成向量作为条件输入生成器让模型根据文字描述画出对应图像。数据增强在样本类别不平衡时指定少数类别的标签生成该类别的额外训练样本。所以别把CGAN当成一个只能跑MNIST的玩具。搞明白它“条件是往哪里塞、怎么塞”这件事后面看pix2pix、CycleGAN、StyleGAN里的条件注入方式都会顺畅很多。这也是我坚持用MNIST代码带你上手的原因——模型简单、数据轻量、条件机制清晰所有注意力都可以集中在理解核心思想上。2. 网络结构与损失函数CGAN改动了哪些核心部件2.1 整体架构与信息流先画一下CGAN的数据流。MNIST图像是28x28的灰度图我们把像素铺平就是784维向量。在CGAN里生成器G的输入是一个拼接向量 [z(100维随机噪声), y(10维one-hot标签)]总共110维。G输出一个784维的向量reshape成1x28x28就是一张生成的数字图。判别器D的输入是 [真实图像或生成图像(784维), y(10维one-hot标签)]总共794维。D输出一个标量表示输入图像在当前条件下为真的概率。需要注意一个关键细节条件y在训练和推理时必须同时提供给G和D。在G里条件是控制生成方向的“指令”在D里条件是判定真假的“考卷题目”。如果你只在G里注入条件D完全不知道你要求的是什么数字那么D的反馈就退化成“这张图像不像MNIST”生成器学到的条件关联会很弱。反过来只给D注入条件、不给G注入条件那更是直接废了。我在代码里就曾经漏掉过一次D的条件拼接结果训练出来的生成图完全不受标签控制排查了很久才发现是这个低级错误。2.2 生成器里“条件”是怎么注入的生成器的核心任务把噪声和标签融合成一张符合标签语义的图像。常见的条件注入方式主要有两种第一种也是最简单的直接在输入层拼接。噪声z是100维向量标签转成10维one-hot向量在特征维度上cat到一起变成110维。后面的全连接层会自己学着从这110维里提取“哪些维度对应数字7、哪些维度对应笔画走向”的信息。第二种用Embedding层把标签映射成稠密向量再拼接。one-hot是稀疏表示向量里大部分维度都是0Embedding则可以让模型在训练中学习到一个更适合当前任务的标签向量表示。从效果上看Embedding往往能稍微提升生成质量同时也可以降低标签向量维度避免条件信息被100维的噪声“淹没”。在MNIST这种简单数据集上两种方式差别不大。但有一个细节我一直很在意条件向量的维度如果太小比如Embedding维度设成8、16标签信息在拼接后很容易被噪声信息稀释生成结果会表现出标签控制力变弱。我实测下来噪声100维配Embedding 32维或64维效果比较稳。生成器内部结构我用的是三层全连接加BatchNorm加LeakyReLU最后一层用Tanh。关于Tanh多说一句MNIST像素经过Normalize处理后值域在[-1, 1]Tanh的输出正好匹配这个范围。如果你把最后一层换成Sigmoid那范围是[0, 1]数据预处理也得跟着改成不归一化否则训练很难收敛。很多人抄代码喜欢随手换激活函数但换之前务必想清楚输出范围和数据分布是否匹配。2.3 判别器里“条件”是怎么注入的判别器的结构跟生成器是镜像关系但更简单把图像展平成784维拼接上10维one-hot标签变成794维然后通过几层全连接压缩到1个输出节点。中间激活函数用LeakyReLU这里有个值得注意的点判别器里我一般不用BatchNorm。为什么判别器要同时处理真实图像和生成图像两个不同分布的数据BatchNorm会强制把每个batch的数据都归一化到均值为0、方差为1的分布。当真实图和生成图混在一个batch里时BN统计出来的均值方差会来回横跳反而让训练不稳定。这个现象在MNIST任务上可能不算严重但换到CIFAR或者更高分辨率的图像上你就知道有多难受了。生成器里用BN没问题因为G每次只处理自己生成的数据分布相对可控。D里面我更推荐用LeakyReLU负斜率0.2配合轻微Dropout实测下来稳定性明显更好。判别器的最后一层很多老代码会写Sigmoid然后再用BCELoss来计算损失。但我在PyTorch里更推荐的做法是最后一层不接Sigmoid直接输出logits配合nn.BCEWithLogitsLoss使用。原因是BCEWithLogitsLoss内部把Sigmoid和交叉熵融合在一起数值计算上更稳定能有效避免log(0)这类问题梯度回传也更干净。2.4 损失函数与优化器选择的细节CGAN的损失本质还是二分类交叉熵。判别器的目标是真图配真标签时输出尽量接近1假图配标签时输出尽量接近0。生成器的目标是假图配上标签后让判别器输出尽量接近1。用PyTorch的BCEWithLogitsLoss判别器的损失由两部分组成真实图像在对应标签下的损失目标值为1加上生成图像在对应标签下的损失目标值为0。生成器的损失则是把生成的假图再喂给判别器目标是让输出接近1。代码形式会很直观D损失 BCE(D(real_img, label), real_target) BCE(D(fake_img, label), fake_target)G损失 BCE(D(fake_img, label), real_target)注意G损失里的目标也是real_target也就是1而不是fake_target。生成器的目的就是欺骗判别器所以它希望D对假图输出越接近1越好。优化器这里有个非常重要的细节Adam的beta1参数。PyTorch里Adam默认的beta1是0.9这个值在普通分类任务上很好用但在GAN上会导致损失震荡得非常厉害。标准做法是把beta1调到0.5相当于给动量“降温”让梯度更新不那么冲。学习率一般设成2e-4这也是GAN训练里经过大量实践验证的“甜点区”。我早期用默认的1e-3训练D的loss经常在1个epoch内就掉到接近0而G的loss直接飙到十几生成图全是噪声后来改成2e-4加beta10.5之后整个训练过程明显平稳。还有一个实用小技巧标签平滑。把真实标签从1改成0.9让判别器不要太自信可以缓解生成的图片过于单一的问题。这个技巧对训练稳定性有一定改善后面代码里我会直接用上。3. 基于PyTorch的CGAN完整代码实现3.1 环境准备与数据加载开始写代码之前先把环境准备好。建议直接用Python 3.8以上的环境PyTorch 2.x版本都可以需要torchvision和matplotlib这两个配套库。装好之后直接跑pip install torch torchvision matplotlib数据加载用的是torchvision自带的MNIST第一次运行会自动下载到本地。预处理里我用了两个操作ToTensor把像素范围从0-255转到0-1然后Normalize把范围调整到[-1, 1]匹配生成器Tanh的输出范围。import torch import torch.nn as nn import torch.optim as optim import torch.nn.functional as F import torchvision import torchvision.transforms as transforms from torch.utils.data import DataLoader transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) train_dataset torchvision.datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) train_loader DataLoader( train_dataset, batch_size128, shuffleTrue, drop_lastTrue )这里drop_lastTrue是刻意加的。MNIST训练集有60000张图除以128会剩一点尾巴最后一个batch如果太小BatchNorm的统计量就会不稳。既然用了batch size 128干脆把不完整的尾巴丢掉。标签需要转成one-hot向量PyTorch里直接用F.one_hot函数labels_onehot F.one_hot(labels, num_classes10).float()labels的形状是(batch_size,)转完之后是(batch_size, 10)。后面在模型里会把这个向量和图像特征或噪声向量直接拼接。3.2 生成器与判别器的模型定义生成器我用全连接结构110维输入经过三层全连接映射到784维reshape成(1, 28, 28)。class Generator(nn.Module): def __init__(self, noise_dim100, num_classes10): super().__init__() self.model nn.Sequential( nn.Linear(noise_dim num_classes, 256), nn.BatchNorm1d(256), nn.LeakyReLU(0.2), nn.Linear(256, 512), nn.BatchNorm1d(512), nn.LeakyReLU(0.2), nn.Linear(512, 784), nn.Tanh() ) def forward(self, z, labels): x torch.cat([z, labels], dim1) out self.model(x) return out.view(-1, 1, 28, 28)forward里的输入z是随机噪声labels是one-hot标签cat之后就是一个(batch_size, 110)的拼接向量。BatchNorm1d接在全连接层后面是标准操作注意BatchNorm在训练和推理两种模式下行为不同推理时要记得调用model.eval()。判别器的输入是图像和标签图像先view成(batch_size, 784)再和标签拼接class Discriminator(nn.Module): def __init__(self, num_classes10): super().__init__() self.model nn.Sequential( nn.Linear(784 num_classes, 512), nn.LeakyReLU(0.2), nn.Linear(512, 256), nn.LeakyReLU(0.2), nn.Linear(256, 1) ) def forward(self, x, labels): x x.view(x.size(0), -1) x torch.cat([x, labels], dim1) return self.model(x).squeeze(1)注意判别器最后输出的是(batch_size,)的logits不是概率。配合BCEWithLogitsLoss使用模型内部不显式加Sigmoid。3.3 训练循环判别器和生成器的交替更新训练循环是CGAN代码的重头戏我拆成两个阶段看。先定义优化器和损失函数。优化器两个网络分开定义Adam的学习率和beta1都是关键参数前面已经说过原因device torch.device(cuda if torch.cuda.is_available() else cpu) G Generator().to(device) D Discriminator().to(device) criterion nn.BCEWithLogitsLoss() g_opt optim.Adam(G.parameters(), lr2e-4, betas(0.5, 0.999)) d_opt optim.Adam(D.parameters(), lr2e-4, betas(0.5, 0.999))然后是核心循环。判别器每个batch更新一次生成器每个batch更新一次两者交替进行epochs 60 for epoch in range(epochs): for i, (imgs, labels) in enumerate(train_loader): batch_size imgs.size(0) imgs imgs.to(device) labels_onehot F.one_hot(labels, num_classes10).float().to(device) # 真实标签用0.9做平滑生成标签用0 real_target torch.ones(batch_size, devicedevice) * 0.9 fake_target torch.zeros(batch_size, devicedevice) # ---------- 第一步训练判别器 ---------- z torch.randn(batch_size, 100, devicedevice) fake_imgs G(z, labels_onehot) d_real_logits D(imgs, labels_onehot) d_fake_logits D(fake_imgs.detach(), labels_onehot) d_loss criterion(d_real_logits, real_target) criterion(d_fake_logits, fake_target) d_opt.zero_grad() d_loss.backward() d_opt.step() # ---------- 第二步训练生成器 ---------- z torch.randn(batch_size, 100, devicedevice) fake_imgs G(z, labels_onehot) fake_logits D(fake_imgs, labels_onehot) g_loss criterion(fake_logits, real_target) g_opt.zero_grad() g_loss.backward() g_opt.step() print(fEpoch [{epoch1}/{epochs}] D loss: {d_loss.item():.4f}, G loss: {g_loss.item():.4f})这里有几个细节值得单独拿出来说。第一个是fake_imgs.detach()。训练判别器时生成器产生的fake_imgs只用来给D提供输入不需要反向传播到G的参数所以必须detach掉。如果不detach梯度会顺着fake_imgs传到G的参数里等于你更新D的时候顺带改了G两个网络搅在一起训练必然乱套。第二个是生成器训练时重新采样了一次噪声z而不是复用训练判别器时的那一批。其实复用也没问题但重新采样能让每个batch里G接收到更多样化的噪声输入多少能增加些生成多样性。我个人的习惯是重新采样训练过程会更平稳。第三个是为什么生成器训练时也要用real_target0.9而不是1.0。配合标签平滑生成器不会陷入“必须让D的输出无限接近1”的极端这在某种程度上能抑制生成样本单一化。3.4 训练结果观察与可视化训练过程光看loss没什么意义必须把生成的图像存下来肉眼看。GAN这个领域有一条铁律loss下降不代表生成质量好loss震荡也不代表训练失败。真正靠谱的评估方式就是直接看生成图像。我在训练时写了一个简单的保存函数每轮epoch结束后用固定的噪声和0-9标签各生成一张图拼成一个网格保存def save_samples(G, device, epoch): G.eval() with torch.no_grad(): fixed_noise torch.randn(10, 100, devicedevice) fixed_labels torch.eye(10, 10, devicedevice) fake_imgs G(fixed_noise, fixed_labels) grid torchvision.utils.make_grid(fake_imgs, nrow10, normalizeTrue) torchvision.utils.save_image(grid, fcgan_epoch_{epoch:03d}.png) G.train() # 训练循环每轮结束后调用 save_samples(G, device, epoch)这个函数有意思的地方在于“固定的噪声”和“固定的标签”fixed_noise每轮都一样fixed_labels则是0到9各一个one-hot。这样你连续看几十轮保存下来的图片就能清晰看到同一个噪声向量在标签控制下是怎么一步步长出对应数字的轮廓的。我实测的观察是这个节奏前5个epoch生成图像基本是一团模糊的随机纹理几乎看不出数字但不同标签列已经隐隐有些灰度差异10到20个epoch数字的轮廓开始出现能勉强分辨出某些标签对应的数字30个epoch以后大部分数字已经比较清晰但部分类别比如“4”“9”可能还会有些歪50到60个epoch图像质量趋于稳定每列基本就是对应数字的正常手写体。如果你训练到后期发现某一列标签对应的图像和标签不匹配比如标签“7”那列生成的图看起来像“1”别急着慌。先看是不是固定噪声那几行的问题有些特定噪声向量生成出的结果确实会有歧义。更准确的验证方法是随机抽样让G配合随机噪声和标签“7”生成几十张图如果大部分都像“7”说明条件控制是生效的。4. 训练CGAN的实战心得常见问题与排查技巧4.1 损失震荡、模式崩塌最常遇到的两座大山GAN训练最折磨人的两个问题一个是损失剧烈震荡一个是模式崩塌。损失震荡的表现是D的loss在0.1到2之间来回蹦极G的loss也跟着上蹿下跳。这时候人很容易慌但其实GAN训练本身就是一个动态博弈过程D和G交替变化的损失本来就该有波动。真正需要警惕的是某一方的loss长期一边倒比如D的loss长期小于0.1说明D把真假图完美区分开了G一点活路都没有此时生成的图像通常全是噪声。遇到这种情况先把判别器的学习率调小比如从2e-4降到1e-4同时给判别器加一点Dropout让D别那么“卷”。反过来如果G的loss长期压着D打D的loss一直居高不下那可能是生成器能力过强或者判别器太弱需要给G降学习率或者增加D的容量。模式崩塌的典型表现是生成器学会了疯狂产出同一个或几个固定模式的样本多样性极差。在MNIST CGAN里模式崩塌可能表现为所有标签列生成的图像都长得差不多。条件GAN比普通GAN好在标签信息本身能强迫生成器区分不同类别所以全类别塌缩的情况相对少见更常见的是某一两个数字崩了比如“0”和“6”长得一模一样。这时候我一般先做两件事一是调大噪声向量维度从100加到128或256给G更多表达空间二是检查标签拼接是否正确确认条件信息确实在每个batch都传进去了。4.2 条件没生效检查这几个地方条件信息“没生效”是CGAN新手最容易踩的坑。生成出来的图像不受标签控制所有标签列看起来都像同一个数字或者图跟标签明显对不上。遇到这个问题按下面这个清单逐项排查。第一检查生成器的输入拼接。在forward里打印一下z和labels的形状确保cat之前两个张量的第一个维度都是batch_size。如果labels的形状是(batch_size, 1, 10)cat之后张量维度就变成了(batch_size, 1, 110)全连接层肯定会报维度错误。如果没报错但确实拼错了那就要看一下你的拼接逻辑。第二检查判别器是否也接收了条件。有些简化版的CGAN代码只在G里拼接标签D里只接收图像训练出来的效果就是标签控制力很弱。因为D完全不知道条件的存在它只会判断图像像不像MNIST而不是判断图像在当前标签下的合理性。G缺少一个“按标签区分真伪”的监督信号自然懒得分不同数字去画。第三用小实验验证条件注入确实有效。把一张真实的“3”配标签“3”给DD输出应该接近真实再把同一张“3”配标签“5”给DD输出应该明显下降。如果这两种输出的差距很小说明D没有好好利用标签信息要么是条件维度被其他特征淹没了要么是训练还没收敛。这个实验我每次调完模型都会跑一遍非常直观比看loss有用得多。第四检查one-hot标签的构造。PyTorch的F.one_hot要求输入是整数张量输出是Float类型才方便拼接。很多人忘了.float()导致cat的时候因为dtype不一致直接报错或者用Hot编码时类别数设错标签维度不匹配。4.3 一组可以长期使用的调试习惯训练GAN的过程里我总结了一套比较趁手的调试流程分享给你参考。每5个epoch保存一批固定噪声和固定标签的生成图。这里的重点是“固定”如果每次都换新的噪声你就没法对比同一个噪声在不同训练阶段的变化也就看不出模型学习的过程。连续翻看这些图片你能很清楚地看到数字的轮廓是否越来越清晰、不同标签的区分度是否越来越大。每个epoch记录四个值D的real输出均值、D的fake输出均值、D的loss、G的loss。理想情况下D的real输出均值在0.8左右fake输出均值在0.2左右两者随着训练慢慢向0.5靠拢。如果real和fake的输出差距迅速拉大到0.9对0.1说明D太强了赶紧调参不要等训练崩了再动手。还有一个我很喜欢的小技巧固定噪声向量集合但改变标签顺序看模型对标签是否敏感。比如拿16个固定噪声配合一组乱序的标签生成64张图按标签值排列观察同一噪声在不同标签下的差异。如果差异明显说明标签信息确实在干预生成过程。5. 把CGAN改一改条件注入方式的几种进阶玩法5.1 用Embedding层替代one-hot拼接one-hot编码虽然简单但有个问题是标签之间的“距离”都是平等的模型很难学到数字之间的“相似关系”。在MNIST里“3”和“8”在笔画上其实比较接近但one-hot编码看不出这层关系模型只能自己从数据里摸索。用Embedding层能多少缓解这个问题。Embedding会把每个类别映射成一个可学习的稠密向量模型在训练中可以自己调整这些向量之间的距离让相似的类别在向量空间里也靠得近一些。改动起来非常小生成器可以改成这样class GeneratorEmbed(nn.Module): def __init__(self, noise_dim100, num_classes10, embed_dim64): super().__init__() self.embed nn.Embedding(num_classes, embed_dim) self.model nn.Sequential( nn.Linear(noise_dim embed_dim, 256), nn.BatchNorm1d(256), nn.LeakyReLU(0.2), nn.Linear(256, 512), nn.BatchNorm1d(512), nn.LeakyReLU(0.2), nn.Linear(512, 784), nn.Tanh() ) def forward(self, z, labels): cond self.embed(labels) x torch.cat([z, cond], dim1) out self.model(x) return out.view(-1, 1, 28, 28)注意forward里的labels直接传整数索引不需要再转one-hot。判别器也可以做同样的替换但我建议先只改生成器来观察效果差异。我实测Embedding维度过大比如128维时标签信息和噪声会在拼接后出现一定程度的干扰反而没有32到64维清爽。5.2 把标签换成数值型条件CGAN的条件不一定是离散类别标签也可以是连续数值。举个例子假设你想控制生成数字的笔画粗细可以给每个样本定义一个0到1之间的粗细系数作为条件向量拼接到生成器输入里。这种连续条件的引入方式和分类标签本质相同只是在数据处理上要做归一化。数值型条件的坑在于量纲问题。如果粗细系数的范围在0到1之间而噪声是从标准正态分布采样的拼接后两者数值范围大体匹配问题不大。但如果某个条件的数值是100、1000这种量级它会在全连接层的加权求和里完全压过噪声信号生成器会过度依赖这个条件而丢掉了随机多样性。所以数值型条件务必归一化到和噪声相近的尺度一般在0到1或-1到1之间。我曾尝试用CGAN控制生成数字的旋转角度角度值从0到359如果直接用原始度数作为条件训练出来的效果很差生成图像要么倾斜过度要么不倾斜。把角度除以359归一化成0到1之后训练马上就正常了。这类经验听起来很小但非常影响实际效果。5.3 把条件换成图像或文本向量CGAN条件的抽象程度可以不断升级。在pix2pix这类模型里条件不再是一个简单标签而是一整张图像——比如输入一张边缘图输出对应的真实照片。这种情况下“条件注入”的方式也从简单的特征拼接变成了编码器网络的信息融合。理解了CGAN的拼接思路后你会发现这些进阶模型本质上还是同一套逻辑把条件编码成向量和生成器的输入融合让生成过程按照条件的语义去走。区别只在于条件的编码方式更复杂了条件注入的位置也不一定只在输入层。对于想深入做生成模型的朋友我建议从CGAN这个小而美的模型出发先彻底搞懂“条件是怎么影响生成的”再去看pix2pix、CycleGAN、StyleGAN都不会太吃力。条件注入是生成模型里最核心的骨架之一把这个骨架搭稳了后面加什么都顺理成章。写在最后的一点体会从跑通第一个GAN到真正理解CGAN我对“条件”这件事的感受越来越深它不只是给模型多看一个标签而是把整个生成问题变成了“在约束下解难题”。约束给得好模型学得又快又稳约束给得模糊模型的输出也会含糊。MNIST上的CGAN是我见过最适合用来体会这种差别的小实验几十轮训练几分钟跑完可以反复试、反复改。最后分享一个我自己的小习惯训练CGAN时不要只盯着最终的生成图多去看看固定噪声在训练过程中的逐轮演化。图像从一团噪声慢慢长成一个个对应标签的数字这个过程会把很多抽象的原理具象化也会让你对生成模型的把握感强很多。等你能一眼看出“这轮训练D太强了该调参了”恭喜你这个模型你就真吃透了。