共计 1726 个字符,预计需要花费 5 分钟才能阅读完成。
背景痛点:小样本数据增强的困境
在深度学习项目中,数据量不足常导致模型性能瓶颈。传统数据增强方法(如旋转、裁剪、颜色变换)虽然简单易用,但存在明显局限性:

- 仅能产生线性变换的衍生样本,无法增加数据分布的多样性
- 对图像局部语义理解有限(如医学影像中的病灶位置)
- 文本数据增强时容易破坏语法结构(如 NLP 中的同义词替换)
技术对比:GAN 家族演进路线
| 模型类型 | 条件控制 | 生成质量 | 训练稳定性 | 适用场景 |
|---|---|---|---|---|
| Vanilla GAN | 无 | 中等 | 低 | 简单图像生成 |
| CGAN | 类别标签 | 高 | 中 | 分类条件生成 |
| DCGAN | 无 | 较高 | 中高 | 通用图像生成 |
核心实现:PyTorch 实战 CGAN
生成器架构设计
class Generator(nn.Module):
def __init__(self, latent_dim, num_classes, img_channels):
super().__init__()
self.label_embedding = nn.Embedding(num_classes, latent_dim) # 标签嵌入层
self.model = nn.Sequential(# 输入:latent_dim + latent_dim ( 噪声 + 标签)
nn.ConvTranspose2d(latent_dim*2, 512, 4, 1, 0, bias=False),
nn.BatchNorm2d(512),
nn.ReLU(True),
# 后续层省略...
)
def forward(self, noise, labels):
# 将标签嵌入与噪声拼接
label_embed = self.label_embedding(labels).unsqueeze(2).unsqueeze(3)
gen_input = torch.cat((label_embed, noise), dim=1)
return self.model(gen_input)
关键参数说明:
– latent_dim=100:噪声向量维度,影响生成多样性
– nn.LeakyReLU(0.2):负斜率防止梯度消失
判别器条件处理
class Discriminator(nn.Module):
def __init__(self, num_classes, img_channels):
super().__init__()
self.label_embedding = nn.Embedding(num_classes, img_size*img_size) # 展平后的尺寸
self.model = nn.Sequential(# 输入:img_channels + 1 ( 图像 + 标签 map)
nn.Conv2d(img_channels+1, 64, 4, 2, 1, bias=False),
nn.LeakyReLU(0.2, inplace=True),
# 后续层省略...
)
def forward(self, img, labels):
# 将标签嵌入转为空间 map
label_map = self.label_embedding(labels).view(labels.size(0), 1, img_size, img_size)
d_in = torch.cat((img, label_map), dim=1)
return self.model(d_in)
避坑指南:稳定训练技巧
- 模式崩溃解决方案
- Mini-batch 判别:在判别器最后层添加特征统计量计算
- 标签平滑:将真实样本标签从 1 调整为 0.9-1.0 随机值
-
学习率比例:保持 D:G 的学习率比例在 1:4 到 1:10 之间
-
质量评估方法
- Inception Score:基于分类器的多样性和清晰度评估
- FID(Frechet Inception Distance):比较真实 / 生成分布的距离
生产环境建议
- 分布式训练时采用
torch.nn.parallel.DistributedDataParallel - 避免棋盘伪影:使用 PixelShuffle 替代转置卷积
实践资源
- Colab 完整代码
- 延伸阅读:
- Mirza M, Osindero S. Conditional Generative Adversarial Nets. arXiv:1411.1784
- Gulrajani I, et al. Improved Training of Wasserstein GANs. NeurIPS 2017
正文完
