用条件生成对抗网络控制图像生成:从标签注入到可复现实验
普通生成对抗网络的目标是学习真实数据分布。以手写数字为例生成器可能输出任意类别的数字调用方无法直接指定“生成一个 7”。如果业务需要按类别、属性或文本条件生成样本就必须把条件信息纳入生成过程这正是条件生成对抗网络Conditional GAN简称 cGAN解决的问题。cGAN 并不等于“给 GAN 加一个标签参数”这么简单。标签必须同时影响生成器和判别器生成器要根据标签改变输出判别器则要判断“图像是否真实”以及“图像是否符合给定标签”。否则模型可能忽略条件退化成普通 GAN。本文使用 MNIST 作为演示数据集。示例只用于说明训练流程和工程结构不预设固定的生成质量、收敛速度或最终准确率实际结果会受到硬件、随机种子、依赖版本和超参数影响。原理拆解设随机噪声为z类别标签为y真实图像为x。cGAN 的生成器学习G(z, y) - x_fake判别器接收图像和标签D(x, y) - [0, 1]其中输出值通常被解释为图像在给定条件下为真实样本的概率。训练时判别器需要区分两类正样本和负样本(真实图像, 真实标签)应判为真。(生成图像, 指定标签)应判为假。生成器则试图让(生成图像, 指定标签)被判为真。采用二元交叉熵时常见目标可以写成L_D BCE(D(x, y), 1) BCE(D(G(z, y), y), 0)L_G BCE(D(G(z, y), y), 1)条件信息的注入有多种方式。对于简单的类别生成任务可以把标签转换为独热向量再与噪声拼接也可以使用嵌入层把类别映射为连续向量。判别器同样可以把图像特征与标签向量拼接后进行判断。独热编码实现直观嵌入方式则更容易扩展到大量类别。实验准备准备 Python 环境后安装 PyTorch、torchvision 和 Matplotlib。不同平台的 PyTorch 安装命令可能不同尤其是 CPU 与 CUDA 构建版本建议按照目标平台的官方安装说明选择对应命令。下面的代码假设这些包已经可正常导入。建议先确认设备和数据目录权限python -c import torch, torchvision; print(torch.__version__); print(torch.cuda.is_available())示例使用全连接网络便于观察条件输入的形状变化。对于更高分辨率图像应改用卷积结构例如 DCGAN 风格的生成器和判别器全连接模型不适合作为通用图像生成架构。完整示例下面代码训练一个按数字类别生成 MNIST 风格图像的 cGAN。标签通过one_hot转为 10 维向量并分别送入生成器和判别器。为避免把模型输出直接当作概率判别器最后一层保留 logits损失函数使用BCEWithLogitsLoss。import random import numpy as np import torch from torch import nn from torch.utils.data import DataLoader from torchvision import datasets, transforms import matplotlib.pyplot as plt SEED 42 random.seed(SEED) np.random.seed(SEED) torch.manual_seed(SEED) device torch.device(cuda if torch.cuda.is_available() else cpu) batch_size 128 noise_dim 100 num_classes 10 epochs 20 lr 2e-4 transform transforms.Compose([ transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ]) dataset datasets.MNIST( root./data, trainTrue, downloadTrue, transformtransform ) loader DataLoader(dataset, batch_sizebatch_size, shuffleTrue, num_workers0, drop_lastTrue) def one_hot(labels, classesnum_classes): return torch.nn.functional.one_hot(labels, classes).float() class Generator(nn.Module): def __init__(self): super().__init__() self.net nn.Sequential( nn.Linear(noise_dim num_classes, 256), nn.LeakyReLU(0.2), nn.Linear(256, 512), nn.LeakyReLU(0.2), nn.Linear(512, 28 * 28), nn.Tanh() ) def forward(self, z, labels): condition one_hot(labels).to(z.device) return self.net(torch.cat([z, condition], dim1)) class Discriminator(nn.Module): def __init__(self): super().__init__() self.net nn.Sequential( nn.Linear(28 * 28 num_classes, 512), nn.LeakyReLU(0.2), nn.Dropout(0.3), nn.Linear(512, 256), nn.LeakyReLU(0.2), nn.Linear(256, 1) ) def forward(self, images, labels): condition one_hot(labels).to(images.device) return self.net(torch.cat([images, condition], dim1)).squeeze(1) G Generator().to(device) D Discriminator().to(device) criterion nn.BCEWithLogitsLoss() opt_g torch.optim.Adam(G.parameters(), lrlr, betas(0.5, 0.999)) opt_d torch.optim.Adam(D.parameters(), lrlr, betas(0.5, 0.999)) for epoch in range(epochs): G.train() D.train() for real, labels in loader: real real.view(real.size(0), -1).to(device) labels labels.to(device) n real.size(0) real_target torch.ones(n, devicedevice) fake_target torch.zeros(n, devicedevice) z torch.randn(n, noise_dim, devicedevice) fake G(z, labels) d_real D(real, labels) d_fake D(fake.detach(), labels) loss_d criterion(d_real, real_target) criterion(d_fake, fake_target) opt_d.zero_grad(set_to_noneTrue) loss_d.backward() opt_d.step() z torch.randn(n, noise_dim, devicedevice) fake G(z, labels) loss_g criterion(D(fake, labels), real_target) opt_g.zero_grad(set_to_noneTrue) loss_g.backward() opt_g.step() print(fepoch{epoch 1:02d} loss_d{loss_d.item():.4f} floss_g{loss_g.item():.4f}) G.eval() fixed_labels torch.arange(10, devicedevice) z torch.randn(10, noise_dim, devicedevice) with torch.no_grad(): samples G(z, fixed_labels).view(-1, 28, 28).cpu() fig, axes plt.subplots(2, 5, figsize(8, 4)) for index, ax in enumerate(axes.flat): ax.imshow(samples[index], cmapgray, vmin-1, vmax1) ax.set_title(str(index)) ax.axis(off) plt.tight_layout() plt.savefig(cgan_samples.png, dpi150)执行步骤将代码保存为train_cgan.py。执行python train_cgan.py首次运行会下载 MNIST 数据集因此需要网络访问或提前准备数据缓存。观察每轮输出的两个损失值。损失值本身不是图像质量的充分指标不能仅凭某一轮的数值判断模型优劣。训练结束后检查cgan_samples.png。图像标题代表传给生成器的目标类别应结合视觉结果判断条件是否生效。固定fixed_labels和随机噪声后可重复生成同一批样本若只固定标签而不固定噪声每次输出仍可能不同这是生成模型保留多样性的正常结果。如何验证条件是否生效最直接的检查是建立固定标签网格每一列使用相同标签每一行使用不同噪声。若同一列的类别特征基本一致同时不同样本仍有笔画差异说明模型同时保留了条件一致性和一定多样性。更严格的验证可以使用独立的数字分类器对生成图像进行分类再统计预测类别与输入标签的一致性。但这个指标会受分类器分布、阈值和预处理影响不能单独代表生成质量。还应检查重复样本、模糊程度和类别覆盖情况。常见问题1. 生成器为什么会忽略标签常见原因包括判别器没有接收标签、标签拼接位置错误、训练不足或者类别信息相对于图像特征过弱。应先打印z、独热向量和拼接结果的形状确认生成器与判别器使用的是同一套类别编码。将不同标签输入同一个固定噪声比较输出差异也能帮助定位条件是否被使用。2. 判别器损失迅速接近零怎么办这通常说明判别器暂时过强但仅凭损失不能确定具体原因。可以检查数据归一化是否与生成器末端激活匹配。本例使用Tanh所以真实图像被归一化到大致[-1, 1]。此外还可以降低判别器学习率、调整网络容量或采用卷积结构改善图像建模能力。每次只改变一个因素便于判断影响。3. 输出全黑、全灰或高度重复先确认推理阶段调用了eval()并在torch.no_grad()中生成再检查保存图像时是否正确反归一化或设置显示范围。若样本高度重复可能是模式崩溃。可从降低学习率、调整判别器正则化、增加数据多样性和改用更稳定的 GAN 目标函数开始排查但不同数据集的有效方案并不相同。4. 为什么BCEWithLogitsLoss前不能再加 Sigmoid该损失函数内部已经包含对 logits 的数值稳定处理。若模型末端再加Sigmoid会改变预期输入形式可能带来梯度和数值稳定性问题。若确实需要输出概率应在评估或展示时单独调用torch.sigmoid。5. CPU 运行很慢是否代表代码错误不一定。生成对抗训练需要反复更新两个网络CPU 速度通常取决于处理器、批量大小和数据加载方式。可以减少epochs进行流程验证再按设备能力调整批量大小。num_workers的最佳值与操作系统和存储环境有关示例设为0是为了降低跨平台启动问题不代表所有环境的最优配置。工程化建议真实项目中应把超参数、数据路径和输出目录放入配置文件或命令行参数并保存模型检查点。检查点至少应包含生成器、判别器和两个优化器的状态这样中断后才能较完整地恢复训练。数据预处理必须在训练和评估阶段保持一致类别编码也应固定并记录。如果模型用于业务数据还需要关注训练数据的授权、敏感信息泄露和生成内容的审查。生成图像可用于数据增强但不能默认替代真实样本合成数据进入下游训练前应验证其是否引入类别偏差或重复模式。总结cGAN 的关键不是单纯增加标签而是让条件同时进入生成器和判别器并在训练目标中约束“图像是否符合条件”。一个可执行的实验应包括统一的数据归一化、明确的标签编码、独立的生成与判别更新以及固定标签网格验证。从全连接 MNIST 示例迁移到实际视觉任务时优先改进数据管线和卷积架构再处理更复杂的损失函数与评估指标。任何关于收敛速度和生成质量的结论都应基于具体数据、硬件、随机种子和实验记录而不能由单次运行的损失值推断。