用条件生成对抗网络控制图像生成:从标签注入到可复现实验
2026/8/10 20:49:12 网站建设 项目流程

普通生成对抗网络的目标是学习真实数据分布。以手写数字为例,生成器可能输出任意类别的数字,调用方无法直接指定“生成一个 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", train=True, download=True, transform=transform ) loader = DataLoader(dataset, batch_size=batch_size, shuffle=True, num_workers=0, drop_last=True) def one_hot(labels, classes=num_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], dim=1)) 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], dim=1)).squeeze(1) G = Generator().to(device) D = Discriminator().to(device) criterion = nn.BCEWithLogitsLoss() opt_g = torch.optim.Adam(G.parameters(), lr=lr, betas=(0.5, 0.999)) opt_d = torch.optim.Adam(D.parameters(), lr=lr, 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, device=device) fake_target = torch.zeros(n, device=device) z = torch.randn(n, noise_dim, device=device) 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_none=True) loss_d.backward() opt_d.step() z = torch.randn(n, noise_dim, device=device) fake = G(z, labels) loss_g = criterion(D(fake, labels), real_target) opt_g.zero_grad(set_to_none=True) loss_g.backward() opt_g.step() print(f"epoch={epoch + 1:02d} loss_d={loss_d.item():.4f} " f"loss_g={loss_g.item():.4f}") G.eval() fixed_labels = torch.arange(10, device=device) z = torch.randn(10, noise_dim, device=device) 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], cmap="gray", vmin=-1, vmax=1) ax.set_title(str(index)) ax.axis("off") plt.tight_layout() plt.savefig("cgan_samples.png", dpi=150)

执行步骤

  1. 将代码保存为train_cgan.py
  2. 执行python train_cgan.py,首次运行会下载 MNIST 数据集,因此需要网络访问或提前准备数据缓存。
  3. 观察每轮输出的两个损失值。损失值本身不是图像质量的充分指标,不能仅凭某一轮的数值判断模型优劣。
  4. 训练结束后检查cgan_samples.png。图像标题代表传给生成器的目标类别,应结合视觉结果判断条件是否生效。
  5. 固定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 示例迁移到实际视觉任务时,优先改进数据管线和卷积架构,再处理更复杂的损失函数与评估指标。任何关于收敛速度和生成质量的结论,都应基于具体数据、硬件、随机种子和实验记录,而不能由单次运行的损失值推断。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询