1. 从零理解 GAN:它到底在做什么
先把故事情景铺开。想象一个造假币的团伙,和一个验钞机厂家。造假币的每天琢磨怎么把假币做得更像真的,验钞机的任务就是越来越精准地挑出假币。两者天天对着干,互相逼着对方进步。最后的结果是:假币已经能以假乱真,验钞机的鉴别能力也强到离谱。
这就是生成对抗网络(Generative Adversarial Network,简称 GAN)的核心逻辑。2014 年由 Ian Goodfellow 提出,把“造假者”和“验钞机”同时放进一个模型里互相博弈,靠对抗来共同进化。生成器(Generator)就是那个造假币的,它的任务是让生成的数据骗过判别器;判别器(Discriminator)就是那个验钞机,它的任务是区分真实样本和生成样本。
如果你是第一次接触深度学习,只需要抓住一句话:GAN 的基本原理就是让两个网络互相为难,在对抗中把生成能力逼出来。生成器负责无中生有,判别器负责火眼金睛,最后我们留下的是那个能生成以假乱真样本的生成器。
这篇文章面向真正零基础的读者,从目标函数开始推导,把交叉熵里为什么没有负号、JS 散度又是怎么回事、训练时为什么会模式崩塌这些问题全部拆开讲透,再给出一份可运行的代码骨架和调参经验。看完之后,你能手写一个最简单的 DCGAN,也能在训练不收敛时知道自己到底踩了哪个坑。
2. 目标函数拆解:交叉熵公式里的每个符号
2.1 GAN 的目标函数到底长什么样
原始 GAN 的优化目标公式写成:
min_G max_D V(D, G) = E[log D(x)] + E[log(1 - D(G(z)))]其中:
x是来自真实数据集的样本。z是随机噪声向量,通常从标准正态分布采样。G(z)是生成器根据噪声 z 生成的假样本。D(x)是判别器对真实样本的输出,表示“这个东西有多大概率是真的”,输出范围在 0 到 1 之间。D(G(z))是判别器对生成样本的输出,表示“这个东西被判定为真的概率”。
这个公式读起来就是:训练判别器时,我们希望它尽可能把真实样本判成 1,把生成样本判成 0,也就是让 D(x) 逼近 1,让 D(G(z)) 逼近 0。训练生成器时,我们希望它生成的东西能让判别器误判为真,也就是 D(G(z)) 越接近 1 越好。
所以这个目标函数其实写了两个模型的诉求。max_D是判别器的目标,它想让整个表达式变大。min_G是生成器的目标,它想让整个表达式变小。两个模型在同一个函数上较劲,这就是“对抗”二字的根本来源。
2.2 原始 GAN 的交叉熵为什么没有负号
这是很多入门者第一次看公式时最困惑的地方。按常理,交叉熵损失应该长这样:
L = -[y * log(p) + (1 - y) * log(1 - p)]损失函数自带负号,训练时最小化它。但 GAN 公式里没有负号,直接写了E[log D(x)] + E[log(1 - D(G(z)))],这到底是不是交叉熵?
答案其实很简单:它确实是交叉熵,只是换了个等价写法。判别器的目标是最大化E[log D(x)] + E[log(1 - D(G(z)))],也就是让真实样本的 log 概率和生成样本的 log 补概率都尽量大。把整个式子加个负号,就变成了最小化:
L_D = -E[log D(x)] - E[log(1 - D(G(z)))]这就是标准的二元交叉熵形式。交叉熵本质上是一个衡量两个概率分布差异的指标,它的基本形式就是-∑ p * log q。论文里为了让“判别器最大化、生成器最小化”这个博弈关系看起来更对称,直接写成了不带负号的极大化形式。所以不是交叉熵没有负号,而是公式经过了等价变形,负号藏在“max”这个操作里了。
如果写代码实现,实际用的还是带负号的交叉熵损失,因为 PyTorch 里的BCELoss默认就是最小化带负号的交叉熵。公式上写成 max 是为了数学推导直观,代码里写成 min 是为了优化器好用,这两个不矛盾。
2.3 判别器的输出为什么要经过 Sigmoid
判别器的最后一层必须是 Sigmoid 激活函数,把输出压缩到 (0, 1) 区间,用来表示“真实概率”。这个概率再进入交叉熵公式计算损失。
理解这个点以后,再看判别器的输入输出就清楚了:
- 输入一张真实图片,输出的 D(x) 应该接近 1。
- 输入一张生成图片,输出的 D(G(z)) 应该接近 0。
对应到目标函数,D(x) 接近 1 时 log D(x) 接近 0,D(x) 接近 0 时 log D(x) 趋向负无穷,所以判别器会努力让 log D(x) 变大,也就是惩罚那些把真样本判错的行为。生成器那边则正好反过来,它希望 D(G(z)) 接近 1,这样 log(1 - D(G(z))) 就会很小,整个 max 表达式变小,生成器就算赢了。
这里必须注意:生成器训练时只更新生成器的参数,判别器的参数要冻结。否则两个网络同时更新,梯度方向会互相干扰,训练直接乱掉。
2.4 用一次前向传播串起整个公式
画成时间线来理解整个流程。每一步训练,都按这个顺序执行:
- 从真实数据集采样一个 batch 的 x。
- 从噪声分布采样一个 batch 的 z。
- 把 z 输入生成器,得到假图片 G(z)。
- 把真实图片 x 输入判别器,得到 D(x)。
- 把假图片 G(z) 输入判别器,得到 D(G(z))。
- 计算判别器损失,反向传播更新判别器参数。
- 再算一次生成器损失,反向传播更新生成器参数。
注意第 6 步和第 7 步不是一次反向传播搞定两个网络,而是分开两次反向传播。因为 GAN 里两个网络的优化目标不同,必须分别计算损失、分别更新。
很多初学代码的读者在跑 GAN 时经常遇到一个问题:生成器更新的 loss 到底是log(1 - D(G(z)))还是-log D(G(z))。原始论文里生成器目标是log(1 - D(G(z))),意思是让判别器对假样本的判定结果“不那么假”。但这个函数在训练早期梯度很小,生成器学得慢,所以实际工程里更常用的是最大化log D(G(z)),也就是让判别器直接把假样本判定为真。这两种写法的根本逻辑一致,只是梯度特性不同。代码实现时,通常用-mean(log(D(G(z))))作为生成器损失,等价于最大化 log D(G(z))。
3. 核心细节:JS 散度、KL 散度与训练难点
3.1 为什么 GAN 的目标是最小化 JS 散度
如果只把 GAN 理解成“造假和验钞”,那就只看懂了表层。从概率分布的角度看,GAN 真正做的事情是:让生成数据的概率分布 P_G 尽可能接近真实数据的概率分布 P_data。
真实数据分布我们不知道具体函数形式,但我们能采样。生成器给出一个参数化的概率分布 P_G,我们的目标就是让 P_G 逼近 P_data。怎么衡量两个分布的距离?最常用的指标是 KL 散度和 JS 散度。
KL 散度的定义是:
KL(P||Q) = ∫ P(x) * log(P(x) / Q(x)) dx它有几个问题:
- 不对称:KL(P||Q) 不等于 KL(Q||P)。同一个距离,换一下方向结果不同,这很不符合“距离”的直觉。
- 当 P(x) 大于 0 而 Q(x) 趋向 0 时,KL 值趋向无穷大,梯度会很陡峭。
原始 GAN 论文用的是 JS 散度,定义是:
JS(P||Q) = 0.5 * KL(P||(P+Q)/2) + 0.5 * KL(Q||(P+Q)/2)JS 散度是对称的,数值范围被压缩到 [0, log2] 之间,整体性质比 KL 好很多。当判别器训练到最优时,GAN 的目标函数就等价于最小化真实分布和生成分布之间的 JS 散度。这就解释了为什么 GAN 训练的核心目标不是“骗过判别器”这么简单,而是在不断调整生成分布,让它和真实分布在统计意义上越来越近。
3.2 判别器最优解推导:为什么 D(x) = P_data / (P_data + P_G)
这个推导是理解 GAN 绕不开的一步。在固定生成器参数的前提下,把目标函数展开,对 D(x) 求导。令导数为 0,可以得到最优判别器的形式:
D*(x) = P_data(x) / (P_data(x) + P_G(x))当 P_data(x) 远大于 P_G(x) 时,D*(x) 接近 1,说明该位置大概率来自真实分布。当 P_data(x) 和 P_G(x) 相当时,D*(x) 等于 0.5,说明判别器已经完全分不清真假了。
这个结论有什么用?它告诉我们:生成器训练到理想状态时,判别器的输出应该稳定在 0.5 附近。如果你训练时发现判别器的 loss 一直压到很低,或者 D(x) 输出永远接近 1 或者 0,那大概率就是生成器没跟上,导致判别器太容易区分真假了。
3.3 为什么原始 GAN 训练不稳定:JS 散度的致命问题
理论上 JS 散度很完美,对称、有界、非负。但实际使用时会遇到一个致命问题:当真实分布和生成分布完全没有重叠时,JS 散度恒等于 log2,梯度为 0,生成器根本学不到任何信息。
什么是“完全没有重叠”?你可以想象真实数据分布是一条窄长的直线,生成数据的分布是另一条平行的直线,两者间隔很小但不重合。在低维空间里,这种“不重叠”是常态,因为真实图片在高维空间里往往分布在非常低维的流形上,两个低维流形在高维空间里几乎不可能重合。
这就导致训练早期,生成器生成的图片质量很差,和真实图片分布完全不一样,判别器很快就能完美区分真假。此时 JS 散度恒等于 log2,梯度消失,生成器得不到有效的反馈,训练直接陷入停滞。
这个问题的直接后果就是原始 GAN 训练极其不稳定,经常出现模式崩塌或者训练一整天生成器还在输出噪声的情况。
3.4 模式崩塌的本质
模式崩塌是 GAN 训练里最出名的坑,表现为生成器只学会生成少量重复的样本,多样性严重不足。比如生成手写数字,它永远只生成同一类数字甚至同一个数字。
背后的原因是:生成器发现自己生成某一种样本时,最容易骗过判别器,就会拼命往那个方向跑,不断强化这条“成功路线”,最后坍缩到单个模式。判别器看到大量重复样本,虽然能识别出来,但它在真实样本上的分类能力也被反复针对,导致训练进入恶性循环。
解决模式崩塌的方法有很多,后面会讲到的 WGAN 就是一个经典方案,它的梯度性质让生成器不容易“固步自封”。早期工程上还会用 minibatch discrimination、特征匹配、Unrolled GAN 等方法,各有各的代价和收益。
4. 动手实现:一个最小可跑的 GAN 代码骨架
4.1 环境准备与数据集选择
说实话,零基础入门 GAN,不建议一上来就用 ImageNet 那种大规模数据集。建议选择 MNIST 手写数字数据集,单张图片只有 28×28 像素,类别明确,训练速度快,判别器和生成器都不用太深,几层全连接或简单卷积就能跑出效果。
环境方面,Python 3.8 以上版本,PyTorch 2.0 以上,CUDA 可用更好,没有 GPU 用 CPU 跑 MNIST 也能在半小时内看到初步效果。需要安装的库就三个:torch、torchvision、matplotlib。
MNIST 不需要额外下载数据文件,torchvision 自带下载接口。第一次运行会自动从官网下载,如果网络慢,可以手动下载放到对应目录。
4.2 生成器的代码实现
生成器的输入是随机噪声向量 z,输出是一张和真实图片同尺寸的假图片。MNIST 的图片是 28×28 单通道,以下是一个简单的全连接生成器:
import torch import torch.nn as nn class Generator(nn.Module): def __init__(self, latent_dim=100, img_dim=784): super().__init__() self.fc = nn.Sequential( nn.Linear(latent_dim, 128), nn.ReLU(), nn.Linear(128, 256), nn.ReLU(), nn.Linear(256, 512), nn.ReLU(), nn.Linear(512, img_dim), nn.Tanh() ) def forward(self, z): return self.fc(z)几个关键点:
- latent_dim 是噪声向量的维度,默认 100,可以调。维度太低可能表达力不足,维度太高训练成本上去但效果不一定更好。
- 隐藏层用了 ReLU,输出层用了 Tanh。Tanh 的输出范围是 [-1, 1],所以真实图片的像素值也要缩放到 [-1, 1],不能保持原来的 [0, 255] 或者 [0, 1]。
- 不要用 ReLU 作为输出层激活,因为它输出非负,图片的像素值范围不对称,很难拟合真实分布。
4.3 判别器的代码实现
判别器输入一张图片,输出一个标量概率。零基础版本用全连接网络实现即可:
class Discriminator(nn.Module): def __init__(self, img_dim=784): super().__init__() self.fc = nn.Sequential( nn.Linear(img_dim, 512), nn.LeakyReLU(0.2), nn.Linear(512, 256), nn.LeakyReLU(0.2), nn.Linear(256, 1), nn.Sigmoid() ) def forward(self, x): return self.fc(x)判别器隐藏层用 LeakyReLU 而不是 ReLU,是为了避免 ReLU 把负值全部截断为 0,导致梯度稀疏,判别器学得太慢。LeakyReLU 的负斜率参数设为 0.2 是 GAN 里的常见经验值。
最后输出层必须用 Sigmoid,把结果压到 (0, 1),表示“图片为真的概率”。
4.4 训练循环的完整实现
训练循环是 GAN 代码里最重要的部分,前面的网络定义只是堆模块,真正容易出问题的全在这里:
import torch.optim as optim from torchvision import datasets, transforms from torch.utils.data import DataLoader latent_dim = 100 batch_size = 128 epochs = 50 lr = 0.0002 transform = transforms.Compose([ transforms.ToTensor(), transforms.Normalize([0.5], [0.5]) ]) dataset = datasets.MNIST(root="./data", train=True, transform=transform, download=True) dataloader = DataLoader(dataset, batch_size=batch_size, shuffle=True) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") generator = Generator(latent_dim).to(device) discriminator = Discriminator().to(device) g_opt = optim.Adam(generator.parameters(), lr=lr, betas=(0.5, 0.999)) d_opt = optim.Adam(discriminator.parameters(), lr=lr, betas=(0.5, 0.999)) criterion = nn.BCELoss() for epoch in range(epochs): for batch_idx, (real_imgs, _) in enumerate(dataloader): real_imgs = real_imgs.view(-1, 784).to(device) batch_size_current = real_imgs.size(0) # 样本标签:真实样本标 1,生成样本标 0 real_labels = torch.ones(batch_size_current, 1).to(device) fake_labels = torch.zeros(batch_size_current, 1).to(device) # 1. 训练判别器 z = torch.randn(batch_size_current, latent_dim).to(device) fake_imgs = generator(z) d_loss_real = criterion(discriminator(real_imgs), real_labels) d_loss_fake = criterion(discriminator(fake_imgs.detach()), fake_labels) d_loss = d_loss_real + d_loss_fake d_opt.zero_grad() d_loss.backward() d_opt.step() # 2. 训练生成器 z = torch.randn(batch_size_current, latent_dim).to(device) fake_imgs = generator(z) g_loss = criterion(discriminator(fake_imgs), real_labels) g_opt.zero_grad() g_loss.backward() g_opt.step() if batch_idx % 200 == 0: print(f"Epoch [{epoch}/{epochs}] Batch {batch_idx} " f"D Loss: {d_loss.item():.4f} G Loss: {g_loss.item():.4f}")这里有几个必须讲清楚的细节:
第一,训练判别器时,输入的假图片要调用detach()。原因是在计算判别器损失之前,假图片是通过生成器生成的,如果不对假图片做 detach,梯度会同时传到生成器和判别器里。但这一轮我们只想更新判别器,生成器的梯度应该冻结。detach 之后,计算图断裂,梯度只在判别器内部传播。
第二,训练生成器时,判别器的参数要冻结。怎么冻结?代码里没有显式冻结,但细想一下,当计算g_loss = criterion(discriminator(fake_imgs), real_labels)时,fake_imgs 依然依赖生成器参数,反向传播时梯度会穿过判别器传到生成器。判别器本身的参数在这个 backward 里也会被计算梯度,但随后我们只调用g_opt.step(),只更新生成器参数,判别器的梯度即便被算出来了也会被丢弃。
第三,为什么生成器的标签是 1 而不是 0。生成器的目标是让判别器把假图片判定为真,所以它的目标是让 D(G(z)) 接近 1,因此标签是 real_labels。从公式上理解,就是最大化 log D(G(z)),在 BCELoss 里等价于让 D(G(z)) 的预测值逼近 1。
第四,优化器用 Adam,学习率设置成 0.0002,beta1 设置为 0.5。这是 DCGAN 论文中验证过的经典配置。普通 Adam 的 beta1 默认是 0.9,但 GAN 训练时 0.9 容易导致震荡,0.5 会让优化过程更平稳。这个参数直接决定生成器和判别器之间是否容易出现“一方把另一方打趴下”的状态。
4.5 训练结果评估:怎么判断有没有收敛
GAN 没有像分类任务那样的准确率指标,判断训练好坏需要看生成图片。每训练一个 epoch,可以固定一批随机的 z,喂给生成器生成图片,再保存成网格查看。代码里加上保存图片的逻辑:
import matplotlib.pyplot as plt import torchvision.utils as vutils def save_generated_images(generator, epoch, sample_z, output_path): generator.eval() with torch.no_grad(): fake = generator(sample_z).view(-1, 1, 28, 28) vutils.save_image(fake, f"{output_path}/epoch_{epoch}.png", nrow=8, normalize=True) generator.train()观察生成的图片:
- 前几个 epoch 应该是模糊的噪声块,逐渐出现类似数字的轮廓。
- 10 个 epoch 左右,应该能看到明显的手写数字轮廓,但边缘可能粗糙。
- 20 个 epoch 以后,数字应该比较清晰,整体多样。
如果你训练到 30 个 epoch,生成出来的图片永远只有几个固定的数字,那就说明模式崩塌了,需要后面 WGAN 的方法来改善。
5. 训练 GAN 的常见问题与排查清单
GAN 的训练过程不像普通神经网络那样 loss 稳步下降,它更像是两个人在拔河,tension 一直存在。下面这张表总结了我在实践中踩过的坑和对应解法。
| 问题 | 现象 | 常见原因 | 解决思路 |
|---|---|---|---|
| 生成器loss很低但图片质量差 | D(G(z)) 很高,但图片很糊 | 判别器没有给足梯度信号,生成器在骗过判别器时钻了空子 | 增强判别器容量;尝试不同架构 |
| 判别器loss趋近0 | 判别器很快完美区分真假 | 生成器太弱,判别器太强 | 调低判别器学习率;加深生成器;给判别器加 Dropout |
| 生成器loss震荡剧烈 | 生成的图片时好时坏 | 学习率太高;优化器参数不合适 | 降低学习率;调整 beta1;考虑梯度惩罚 |
| 模式崩塌 | 生成图片多样性差,重复同一类 | JS 散度梯度消失;生成器过于贪心 | 换用 WGAN;增加噪声;使用标签平滑 |
| 训练中期loss崩盘 | 数值变为 NaN | 梯度爆炸 | 梯度裁剪;降低学习率;检查是否有除零 |
5.1 判别器太强怎么处理
这是我个人最开始跑 GAN 时最常遇到的问题。判别器训练得太猛,把真实样本和生成样本分得太清,生成器拿不到有效梯度。最直观的解决方法:
- 把判别器的学习率降低,比如从 0.0002 降到 0.0001,让判别器“慢半步”。
- 给判别器加 Dropout,削弱它的判断能力,让生成器有机会骗过它。
- 训练判别器时少更新几次,比如每训练两次生成器,再训练一次判别器。
- 使用标签平滑:把真实样本的标签从 1 改成 0.9。这样判别器不会对训练数据过于自信,目标函数的梯度更温和,能有效缓解训练不稳定。
标签平滑的原理很直白——真实标签给 1 时,判别器会被迫输出极端概率,一旦遇到稍微不一样的样本就会剧烈调整。把标签压缩到 0.9 以后,判别器的输出不需要那么极端,梯度更平稳,网络不容易被一个异常样本“带节奏”。
5.2 生成器 loss 降到很低但图片依然差
这种情况最容易让人怀疑人生:生成器明明赢了判别器,输出图片却完全没法看。原因通常是生成器找到了判别器的“盲区”,在这个盲区里 D(G(z)) 很高,但图片跟真实图片没有任何关系。
解决办法是不要让判别器一次把所有信息都学完,用 Early Stopping 控制判别器训练次数,同时观察生成图片而不是依赖 loss 数值来判断好坏。我在项目里会把生成图片实时可视化,每隔固定步数保存一次网格图,通过眼睛而不是数字来判断训练状态。
5.3 BN 层的坑与使用建议
想在 GAN 里用 BatchNorm 需要格外小心。生成器和判别器的 BatchNorm 行为不一样:生成器里 BN 有助于稳定训练,因为每一层输出的分布被归一化了,生成样本的幅度不会剧烈波动;判别器里 BN 反而容易出问题,尤其是 batch size 较小时,BN 的统计量不稳定,判别器会随样本批次波动。
经验做法是生成器使用 BN,判别器尽量少用或者不用,如果一定要用,使用 LayerNorm 或 InstanceNorm 代替。虽然 DCGAN 论文里生成器和判别器都推荐 BN,但那是基于特定设置下的经验总结,在小数据集或者 batch size 较小时,BN 的副作用可能超过收益。
5.4 为什么我的损失函数不下降
GAN 的 loss 不下降不等于训练失败,要分情况判断:判别器 loss 保持在 0.69 附近(log2 左右),说明判别器无法区分真假,生成的图片很可能已经不错了;判别器 loss 不断下降到 0.1 以下,说明区分能力太强,生成器需要加强;判别器 loss 上下剧烈波动,说明两个网络在激烈对抗,需要调整学习率或网络容量。
千万别只盯着 loss 曲线就断言模型失败。先保存几个 epoch 的生成图片,用眼睛看,再结合 loss 曲线综合判断。这是 GAN 训练和普通分类训练最大的区别。
6. 从原始 GAN 到 WGAN:解决训练不稳定的关键改进
6.1 WGAN 做了什么
原始 GAN 用 JS 散度衡量两个分布的差距,这个问题在分布不重叠时会导致梯度消失。WGAN 的核心思路是:把衡量分布差距的指标从 JS 散度换成了 Wasserstein 距离,也叫推土机距离。
Wasserstein 距离的直观解释是:把真实分布的一堆土搬到生成分布,需要的最少搬运量。不管两个分布是否重叠,Wasserstein 距离都能提供一个有意义、可微的梯度信号。这从底层解决了原始 GAN 梯度消失的问题。
WGAN 的改动并不大:
- 判别器输出层去掉 Sigmoid,输出的是一个实数,表示样本的“真实程度评分”,不再限制在 (0, 1)。
- 判别器的损失函数不再是交叉熵,而是真实样本评分减去生成样本评分,目标是把两者的评分差拉大。
- 生成器的目标是让判别器对生成样本的评分尽可能高。
- 为了保证 Wasserstein 距离的有效性,判别器的参数更新后要裁到一个小范围内,比如 [-0.01, 0.01],这就是重量裁剪。
6.2 WGAN-GP 进一步改进
重量裁剪的问题在于会让判别器参数集中在两端,拟合能力下降。后续改进版 WGAN-GP(Gradient Penalty)把重量裁剪换成了梯度惩罚,在判别器损失里加了一个正则项,让判别器在真实分布和生成分布之间的梯度范数接近 1。这个改进让训练过程更稳定,生成质量也更好。
WGAN-GP 的代码实现比原始 GAN 复杂一些,需要额外计算插值样本处的梯度:
def compute_gradient_penalty(discriminator, real_imgs, fake_imgs, device): alpha = torch.rand(real_imgs.size(0), 1).to(device) interpolated = alpha * real_imgs + (1 - alpha) * fake_imgs interpolated.requires_grad_(True) d_interpolated = discriminator(interpolated) gradients = torch.autograd.grad( outputs=d_interpolated, inputs=interpolated, grad_outputs=torch.ones_like(d_interpolated), create_graph=True, retain_graph=True )[0] gradients = gradients.view(gradients.size(0), -1) gradient_penalty = ((gradients.norm(2, dim=1) - 1) ** 2).mean() return gradient_penalty使用 WGAN-GP 时,判别器最后一层去掉 Sigmoid,损失函数不直接用 BCELoss,而是用评分差值结合梯度惩罚。当时我踩过不少坑,最深刻的一条是:WGAN-GP 对学习率非常敏感,0.0002 在 MNIST 上表现很好,换到更大的模型上可能直接崩,此时降低学习率到 0.0001 往往能救回来。
6.3 从易到难的实践路线
如果你从零开始接触 GAN,我的建议是先跑通原始 GAN 的 mnist 示例,体验一下训练过程的颠簸,然后再切到 WGAN-GP 试试训练稳定性。这样做的好处是你能从对比中真正理解为什么 WGAN 要用距离度量替代散度,而不是背一个结论。
接下来可以按这条路线循序渐进:
- 原始 GAN 全连接版本,跑 MNIST,理解基础流程。
- DCGAN 卷积版本,跑 Anime Face 或 CIFAR-10,理解卷积操作在 GAN 中的作用。
- WGAN-GP,替换损失函数和优化策略,对比原始 GAN 的训练体验。
- 条件 GAN(CGAN),在生成时额外输入类别标签,控制生成结果。
- CycleGAN,处理图像风格迁移任务,理解多个生成器和判别器如何协同。
到第 4、5 步之后,你就已经完全脱离“零基础”的位置,可以进入特定领域的实战了。
7. GAN 能做什么,以及往后学什么
GAN 的应用场景比很多人想象的要广。图像生成是最常见的,比如生成人脸、动漫头像、艺术作品。图像转换领域有 CycleGAN,能把马变成斑马、把照片变成油画、把夏天变成冬天。超分辨率重建领域,SRGAN 能把低分辨率图像补到高清,GAN 的对抗损失在这里比单纯 L2 损失更容易生成纹理细节丰富的图像。语音合成领域也有 GAN 的影子,它让合成语音更自然。医疗影像上的数据增强更直接,很多医学数据集样本数量稀缺,GAN 生成逼真的病灶样本帮助模型提升了泛化能力。
但 GAN 也有自己的边界。生成图片的分辨率很难在无限增大时保持所有细节真实,训练高分辨率 GAN 的成本和难度指数级上升。而且 GAN 在大规模文本生成任务上的效果不如 Transformer 系模型,它的主战场还是在连续数据尤其是图像数据上。
很多人口中的“深度学习入门”,其实走的是监督学习路线——给数据打标签,训练分类或回归模型。这当然重要,但 GAN 是另一条完全不同的路线,它不需要标签,而是从无标注数据中学会数据本身的分布。理解这种“无监督生成式学习”的逻辑,会让你对模型的本质有更深的把握。
后续如果你想系统地深入,建议直接读原始论文《Generative Adversarial Nets》,以及后来的《Unsupervised Representation Learning with Deep Convolutional Generative Adversarial Networks》(DCGAN)和《Improved Training of Wasserstein GANs》(WGAN-GP)。这几篇论文写的都比较清晰,公式不多,配着代码看非常容易理解。
资源和工具方面,PyTorch 官方有一个pytorch/tutorials仓库,里面的 DCGAN 示例代码注释非常详细。GitHub 上搜索GAN zoo或是pytorch-GAN,能找到几十种经典 GAN 变体的实现,代码风格都很接近,适合刷源码对比各个模型之间的差异。
最后分享几个基于实操的教训
第一次跑通 GAN 代码的时候,我当时用的是全连接版本,训练了五十个 epoch,生成出来的数字轮廓模糊,但已经能看出来是数字了。那种从深陷黑暗到看到一丁点亮光的过程,确实奇妙。但真正把它搞清楚,是在后面踩了很多坑之后。
先说一个最常见的坑:训练 GAN 的火候比训练普通网络难掌握得多。普通分类网络过拟合可以早停,GAN 是过拟合和欠拟合之间找平衡。我试过把判别器训练得很强,结果生成器彻底罢工。也试过把生成器调得太猛,结果图片反复横跳。
再分享一个小技巧:在训练 GAN 的时候,把生成器的输入噪声维度、学习率、beta1 这三个参数作为一组超参来调,一次只动一个。不要同时改两个以上,否则出了问题你根本不知道是哪一个调整导致的。我当时为了提高生成质量,同时改了噪声维度、网络层数和学习率,结果训练三天不出效果,回到只动学习率之后很快就找到规律了。
最后一点心得:GAN 的训练更像是在练手感,而不是在推公式。理解了损失函数和网络结构之后,剩下大量的时间花在观察曲线、看生成图片、调超参数上。如果你想零基础快速上手,可以先跑通代码,再回头啃数学原理,这个过程比我反过来的路径要顺利得多。
希望你也能体会到,看着一张完全由噪声生成的空白图片,一点点长出清晰的轮廓,最终变成一张可辨识的数字图片时的满足感。