CGAN损失函数详解:从理论到PyTorch实战

1次阅读
没有评论

共计 1951 个字符,预计需要花费 5 分钟才能阅读完成。

image.webp

为什么需要 CGAN?原始 GAN 的局限性

传统 GAN(Generative Adversarial Network)虽然能生成逼真数据,但存在两个致命缺陷:

  • 缺乏定向控制:无法指定生成特定类别的内容(比如想要生成 ” 数字 7 ″ 却随机输出)
  • 模式崩溃(Mode Collapse):生成器只学会生成部分样本,多样性急剧下降

CGAN(Conditional GAN)通过引入条件变量 y(通常是类别标签)完美解决了这些问题。在 MNIST 手写数字的例子中,我们可以用条件标签控制生成 ” 指定数字 ”。

CGAN 损失函数数学原理

原始 GAN 的损失函数:
$$\min_G \max_D V(D,G) = \mathbb{E}{x\sim p[\log(1-D(G(z)))]$$}(x)}[\log D(x)] + \mathbb{E}_{z\sim p_z(z)

CGAN 在此基础上增加条件变量 y,核心改进体现在:

  1. 判别器损失:需要同时判断 ” 数据真实性 ” 和 ” 条件匹配性 ”
    $$L_D = -\mathbb{E}{x,y\sim p[\log(1-D(G(z|y)|y))]$$}}[\log D(x|y)] – \mathbb{E}_{z\sim p_z, y\sim p_y

  2. 生成器损失:既要欺骗判别器,又要满足条件约束
    $$L_G = -\mathbb{E}_{z\sim p_z, y\sim p_y}[\log D(G(z|y)|y)]$$

关键改进点:所有输入输出都串联了条件信息 y,相当于给模型加了一个 ” 目标指引 ”。

PyTorch 实战代码

import torch
import torch.nn as nn

# 判别器结构(加入条件信息拼接)class Discriminator(nn.Module):
    def __init__(self):
        super().__init__()
        self.label_embedding = nn.Embedding(10, 10)  # MNIST 有 10 类
        self.model = nn.Sequential(nn.Linear(794, 1024),  # 784+10=794
            nn.LeakyReLU(0.2),
            nn.Dropout(0.3),
            nn.Linear(1024, 1),
            nn.Sigmoid())

    def forward(self, x, labels):
        c = self.label_embedding(labels)
        x = torch.cat([x.view(x.size(0), -1), c], 1)
        return self.model(x)

# 损失函数定义
criterion = nn.BCELoss()
def generator_loss(fake_output):
    return criterion(fake_output, torch.ones_like(fake_output))

def discriminator_loss(real_output, fake_output):
    real_loss = criterion(real_output, torch.ones_like(real_output))
    fake_loss = criterion(fake_output, torch.zeros_like(fake_output))
    return (real_loss + fake_loss) / 2

超参数调优经验

根据我们在 MNIST 上的实验,推荐以下参数组合:

  • 学习率:Generator 建议 0.0002,Discriminator 建议 0.0001(判别器通常需要更小的学习率)
  • Batch Size:64-256 之间(太小会导致模式崩溃,太大会降低生成多样性)
  • 优化器:Adam 优于 SGD,β1 设为 0.5 能稳定训练
  • 标签平滑:将真实样本标签从 1.0 改为 0.9~1.0 随机值,可防止判别器过强

三大训练陷阱及解决方案

  1. 模式崩溃(Mode Collapse)
  2. 现象:生成器只输出几种固定模式
  3. 解决:增加 mini-batch 判别层、适度降低学习率

  4. 梯度消失(Gradient Vanishing)

  5. 现象:判别器 loss→0 且不再更新
  6. 解决:改用 Wasserstein GAN(WGAN)的损失函数

  7. 条件失效(Condition Ignoring)

  8. 现象:生成结果与输入条件无关
  9. 解决:加强条件信息的串联方式(如用乘法代替拼接)

效果验证(MNIST 示例)

条件标签 生成结果
3 CGAN 损失函数详解:从理论到 PyTorch 实战
7

可以看到模型能根据输入标签准确生成对应数字,且笔画风格多样。

延伸练习

  1. 尝试修改损失函数权重:给条件匹配项增加系数(如 1.2 倍)
  2. 用 CIFAR-10 数据集测试彩色图像生成效果
  3. 比较拼接 (concat) 和乘法 (element-wise product) 两种条件融合方式

通过本文的代码和调参技巧,你应该能快速实现一个可用的 CGAN 模型。关键要理解:CGAN 的损失函数不仅是真假判断,更是条件约束下的对抗优化。

正文完
 0
评论(没有评论)