共计 2379 个字符,预计需要花费 6 分钟才能阅读完成。
背景痛点
传统 GAN(生成对抗网络)虽然能生成逼真图像,但在需要精确控制生成内容时显得力不从心。比如我们想生成特定数字的手写体,传统 GAN 无法保证生成的图像符合我们的预期类别。这就是 CGAN(条件生成对抗网络)要解决的问题。

初学者在实现 CGAN 时常会遇到以下典型问题:
- 模式坍塌(Mode Collapse):生成器只学会生成少数几种样本,缺乏多样性
- 梯度消失(Gradient Vanishing):判别器太强导致生成器无法获得有效梯度
- 训练不稳定:损失函数波动剧烈,难以收敛
- 标签融合不佳:条件信息未能有效影响生成结果
- 评估困难:缺乏有效的指标衡量生成质量
技术对比
CGAN 与 DCGAN(深度卷积生成对抗网络)、WGAN(Wasserstein 生成对抗网络)的主要区别在于:
- 结构差异:
- DCGAN 使用卷积网络,适合图像生成
- WGAN 改进了损失函数,训练更稳定
-
CGAN 增加了条件输入,实现可控生成
-
标签嵌入方式:
- 拼接(Concat):将条件标签直接拼接到输入向量
- 投影(Projection):使用嵌入层将标签映射到特定维度
核心实现
下面是用 PyTorch 实现 CGAN 的关键代码(基于 PyTorch 1.10+):
import torch
import torch.nn as nn
# 条件生成器
class Generator(nn.Module):
def __init__(self, latent_dim, num_classes, img_shape):
super().__init__()
# 标签嵌入层(投影方式)self.label_embedding = nn.Embedding(num_classes, num_classes)
# 网络主体
self.model = nn.Sequential(nn.Linear(latent_dim + num_classes, 256),
nn.LeakyReLU(0.2),
nn.Linear(256, 512),
nn.LeakyReLU(0.2),
nn.Linear(512, int(torch.prod(torch.tensor(img_shape)))),
nn.Tanh())
def forward(self, z, labels):
# 将噪声 z 和标签嵌入拼接
c = self.label_embedding(labels)
x = torch.cat([z, c], dim=1)
img = self.model(x)
return img.view(img.size(0), *img_shape)
# 带梯度惩罚的损失函数
def compute_gradient_penalty(D, real_samples, fake_samples, labels):
# 计算梯度惩罚项(WGAN-GP)alpha = torch.rand(real_samples.size(0), 1, 1, 1)
interpolates = (alpha * real_samples + ((1 - alpha) * fake_samples)).requires_grad_(True)
d_interpolates = D(interpolates, labels)
gradients = torch.autograd.grad(
outputs=d_interpolates,
inputs=interpolates,
grad_outputs=torch.ones_like(d_interpolates),
create_graph=True,
retain_graph=True,
)[0]
gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean()
return gradient_penalty
训练监控
使用 TensorBoard 监控训练过程:
from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter()
# 在训练循环中添加记录
for epoch in range(epochs):
# ... 训练代码...
writer.add_scalar('Loss/Discriminator', d_loss.item(), epoch)
writer.add_scalar('Loss/Generator', g_loss.item(), epoch)
# 定期保存生成样本
if epoch % 10 == 0:
with torch.no_grad():
fake_imgs = generator(z_sample, label_sample)
writer.add_images('Generated_images', fake_imgs, epoch)
正常训练过程中,生成图像应该逐渐变得清晰,损失函数会呈现周期性波动。如果出现以下情况说明训练异常:
- 生成图像始终模糊 → 可能是模式坍塌
- 判别器损失快速趋近 0 → 判别器过强
- 损失值剧烈波动 → 学习率可能过高
避坑指南
生产环境部署需注意:
- 批量归一化设置:生成器最后一层不要用 BN(批量归一化),否则可能导致颜色异常
- 学习率调整:使用 Adam 优化器时,betas 参数建议设为(0.5, 0.999)
- 判别器强度:判别器更新步数不宜过多,通常生成器: 判别器 =1:1 到 1:5
如果判别器过强,可以尝试:
- 降低判别器学习率
- 减少判别器层数
- 添加梯度惩罚(如 WGAN-GP)
延伸思考
改进方向建议:
- 加入注意力机制:在生成器中添加自注意力层,提升细节质量
- 多模态生成:支持一个标签对应多种生成风格
扩展实验建议:
- 尝试在 CIFAR-10 数据集上实现类别条件生成,比较不同标签嵌入方式的效果差异
结语
通过本文的实践,我们完成了从理论到实现的完整 CGAN 构建过程。关键是要理解条件信息的融合方式,以及如何平衡生成器和判别器的训练。在实践中多观察中间结果,及时调整超参数,就能逐渐掌握 CGAN 的训练技巧。
正文完
