一、原理解释:对抗即学习
GAN 的核心不是"拟合数据",而是博弈。两个网络共享同一片数据战场:
| 网络 | 输入 | 输出 | 心理活动 |
|---|---|---|---|
| 判别器 D | 真实样本 x 或 生成样本 G(z) | 概率 D(x) ∈ 0,1 | "我要练出一双火眼金睛" |
| 生成器 G | 随机噪声 z ~ p(z) | 假样本 G(z) | "我要让 D 看不出破绽" |
训练动态 :D 给 G 提供"梯度信号"------G(z) 被 D 识破得越多,G 就知道该往哪个方向改进。最终达到纳什均衡:G 完美复刻真实分布,D 只能随机猜测(准确率 50%)。
二、数学公式:从目标函数到最优解
1. 价值函数( minimax 博弈 )
minGmaxDV(D,G)=Ex∼pdata(x)logD(x)+Ez∼pz(z)log(1−D(G(z)))\min_G \max_D V(D, G) = \mathbb{E}{x \sim p{data}(x)}\\log D(x) + \mathbb{E}_{z \sim p_z(z)}\\log(1 - D(G(z)))GminDmaxV(D,G)=Ex∼pdata(x)logD(x)+Ez∼pz(z)log(1−D(G(z)))
- D 的目标 :最大化 V(D,G)V(D,G)V(D,G),即让真实样本得分高、假样本得分低
- G 的目标 :最小化 V(D,G)V(D,G)V(D,G),即让 D(G(z))D(G(z))D(G(z)) 接近 1(骗过 D)
2. 判别器的最优解(固定 G 时)
对任意给定的 G,最优判别器为:
DG∗(x)=pdata(x)pdata(x)+pg(x)D^*G(x) = \frac{p{data}(x)}{p_{data}(x) + p_g(x)}DG∗(x)=pdata(x)+pg(x)pdata(x)
其中 pgp_gpg 是生成器隐式定义的分布。此时 D 是在做贝叶斯最优分类。
3. 全局最优(纳什均衡)
当 pg=pdatap_g = p_{data}pg=pdata 时,D∗(x)=12D^*(x) = \frac{1}{2}D∗(x)=21,此时价值函数达到全局最小值:
V(D∗,G∗)=−log4V(D^*, G^*) = -\log 4V(D∗,G∗)=−log4
4. 生成器的等价目标(实际训练用)
原始公式中 log(1−D(G(z)))\log(1-D(G(z)))log(1−D(G(z))) 在训练初期梯度极小(因为 D 很容易识破 G),导致 G 学习缓慢。因此实际采用非饱和目标:
maxGEz∼pz(z)logD(G(z))\max_G \mathbb{E}_{z \sim p_z(z)}\\log D(G(z))GmaxEz∼pz(z)logD(G(z))
这等价于最小化 pdatap_{data}pdata 与 pgp_gpg 之间的 Jensen-Shannon 散度。
三、举例说明:伪造名画的博弈
想象两个角色:
- 画家(生成器 G):从未见过真迹,只从鉴定专家那里得到反馈。起初他画的是涂鸦,专家一眼看穿。画家不断调整笔触、色彩、构图。
- 专家(判别器 D):见过大量真迹,也见过画家的赝品。起初赝品太假,专家轻松识别。但随着画家技艺精进,专家必须更仔细地观察笔触细节、颜料氧化痕迹。
博弈过程:
- 第 1 轮:画家画了一只四不像的"猫",专家笑出声 → 画家知道要画得像猫
- 第 50 轮:画家能画出猫的轮廓,但眼睛位置不对 → 专家指出"眼神不对"
- 第 200 轮:画家画的猫与照片几乎无差,专家只能抛硬币猜测 → 均衡达成
对应到 GAN:
- 画家的"画布"就是神经网络的输出
- 专家的"鉴定报告"就是判别器的输出概率
- 画家的"改画方向"就是生成器损失的反向传播梯度
四、伪代码:训练循环的本质
python
# 超参数
lr = 0.0002
batch_size = 64
z_dim = 100
k = 1 # 每训练 G 一次,训练 D 的次数
# 网络
G = Generator(input_dim=z_dim, output_dim=data_dim) # 生成器
D = Discriminator(input_dim=data_dim) # 判别器
opt_G = Adam(G.parameters(), lr=lr)
opt_D = Adam(D.parameters(), lr=lr)
for epoch in range(num_epochs):
for real_data in dataloader: # real_data ~ p_data(x)
batch_size = real_data.size(0)
# ========== 训练判别器 D ==========
for _ in range(k):
z = sample_noise(batch_size, z_dim) # z ~ p(z)
fake_data = G(z).detach() # 假样本,不计算G的梯度
d_real = D(real_data) # D(x)
d_fake = D(fake_data) # D(G(z))
# 最大化 log(D(x)) + log(1-D(G(z))) → 最小化负值
loss_D = -mean(log(d_real) + log(1 - d_fake))
opt_D.zero_grad()
loss_D.backward()
opt_D.step()
# ========== 训练生成器 G ==========
z = sample_noise(batch_size, z_dim)
fake_data = G(z)
d_fake = D(fake_data) # D(G(z))
# 非饱和目标:最大化 log(D(G(z)))
loss_G = -mean(log(d_fake))
opt_G.zero_grad()
loss_G.backward() # 梯度通过 D 传回 G
opt_G.step()
# 打印状态
print(f"D_loss: {loss_D:.3f} | G_loss: {loss_G:.3f}")
关键细节:
fake_data.detach():训练 D 时,G 的参数不更新loss_G的梯度流经 D 再回传 G,因此 D 必须是可微的- 实际实现中常用 BCEWithLogitsLoss 替代手工 log,数值更稳定
五、应用场景:从实验室到产业
1. 图像生成与编辑
- StyleGAN / StyleGAN2:生成 1024×1024 级超逼真的人脸(thispersondoesnotexist.com)
- DALL-E / Stable Diffusion:文本到图像生成,底层思想源于 GAN 的对抗训练
2. 图像到图像翻译
- Pix2Pix:有监督的图像翻译(素描 → 照片、卫星图 → 地图)
- CycleGAN:无监督的域迁移(马 ↔ 斑马、夏天 ↔ 冬天、照片 ↔ 莫奈风格)
3. 数据增强(稀缺数据场景)
- 医疗影像:为罕见病灶生成合成 CT/MRI,扩充训练集
- 自动驾驶:生成极端天气、罕见交通场景的合成数据
4. 超分辨率与修复
- SRGAN:将 64×64 低清图重建为 256×256 高清图,PSNR 指标外视觉质量更优
- 图像修复:填补缺失区域(老照片修复、水印去除)
5. 其他生成领域
- WaveGAN:生成原始音频波形,用于音乐合成
- MolGAN:生成满足化学约束的分子结构,用于药物发现
- TextGAN:生成离散文本序列(因文本不可微,需配合 Gumbel-softmax 或强化学习)
总结
GAN 的精髓在于用对抗代替监督:没有"标准答案"告诉生成器该怎么做,只有一个越来越挑剔的裁判。生成器在欺骗裁判的过程中,学会了真实数据的本质结构。