这些年我陆陆续续带过不少新人入门深度学习,发现一个特别有意思的现象:很多人接触到“深度学习”核心算法列表时,一看到“GAN 对抗网络”这几个字就开始心里发怵。大家普遍觉得这个模型的名字太“硬核”,什么“生成器”“判别器”“对抗训练”,听起来像是要在两台神经网络之间打一场拳击赛。可实际上,GAN 的原理拆开来看,反而比很多经典分类网络更好理解,因为它的核心思想极其贴近我们的生活直觉——造假与验伪。
这篇文章就是写给那些真正零基础、或者学过一些深度学习但始终没把 GAN 原理吃透的读者。我会把 GAN 对抗网络的来龙去脉、数学目标、训练过程、以及那些“劝退”味十足的公式,全部用大白话和一个从业者的实测视角重新讲一遍。如果你正准备学习生成模型,或者学校课程里正讲到 GAN,又或者你只是想弄明白“AI 是怎么凭空画出不存在的人脸”这件事,那么这篇内容就是为你准备的。
1. 为什么造不出“新图片”?先看清生成模型的困境
1.1 判别模型和生成模型是两码事
很多新手在学 GAN 之前,其实并没有意识到一个问题:他们之前接触的绝大部分深度学习任务,都属于判别模型,而 GAN 属于生成模型。这两者之间的鸿沟,正是理解 GAN 价值的关键。
判别模型做的事情,用一句话概括就是“看了输入给结论”。比如给一张猫的图片,判断它是猫还是狗;给一段评论文本,判断它是好评还是差评;给一段语音,判断它说的是哪几个字。判别模型学习的是一条分界线,把不同类别的数据分开。它不需要真的理解“猫是什么样的”,只需要知道“猫和狗在哪些特征上有区别”。
生成模型则完全反过来了,它的目标是“根据学到的规律造出新的东西”。比如看了一百万张猫的图片之后,让模型自己去画一张全新的猫。这听上去好像只是换了个任务,实际上难度天差地别。判别模型只需要建模条件概率 P(类别|数据),而生成模型需要建模数据的真实分布 P(数据)——也就是说,它要理解“全世界的猫都能长成什么样、毛色有哪些组合方式、姿态空间有多大”。这个分布极其复杂,高维、非结构化、充满各种相关性。
我在带新人的时候经常做一个类比:判别模型相当于一个门卫,他只需要认识本单位员工的几张脸就能工作;生成模型则相当于一个画师,他见过无数张脸之后,必须能凭空画出一张“看上去合理”的新面孔。门卫可以只记住几个特征,画师却必须内化一整类数据的规律。所以生成模型远比判别模型难搞,这也是深度学习发展早期,生成方向一直落后于判别方向的原因。
1.2 老式生成思路为什么吃力不讨好
在 GAN 出现之前,主流的生成模型走的基本是最大似然估计这条路。思路很直接:先假设数据服从某个带参数的分布族,比如高斯混合模型,然后调整参数,让当前这批真实数据出现的概率最大。问题是,真实世界的图片、语音、文本,其分布根本不是简单的高斯分布能描述的。如果强行用弱分布模型去拟合强数据分布,最后得到的生成结果会非常模糊,像蒙了一层雾。
后来出现了自回归模型和变分自编码器(VAE),情况有所改善,但依然有各自的麻烦。自回归模型把一张图片拆成逐像素预测,生成一张图要成千上万步,速度慢得感人;VAE 通过引入隐变量来间接建模分布,但它优化的是对数似然的下界,生成的图片天然偏模糊。这些模型不是不能用,只是每一步都在“绕远路”,而且对分布的表达能力始终有上限。
GAN 的提出,直接换了一个思路:与其费尽心力去拟合一个复杂的概率密度函数,不如让两个网络互相“踢屁股”,踢着踢着,生成出来的数据就自然接近真实数据了。这个想法最早由 Ian Goodfellow 在 2014 年提出,论文标题就叫《Generative Adversarial Nets》。当时深度学习圈子里很多人看完都拍大腿,觉得这个思路怎么早没人想到。
1.3 GAN 的整体思路可以用一句大实话概括
GAN 干的事情,翻译成大实话就是:有一个造假者,努力造出以假乱真的伪钞;有一个鉴定者,努力分辨收到的钞票是真的还是假的。造假者每造一轮,鉴定者就检查一轮。鉴定者的分辨能力越来越强,反过来逼着造假者不断改进工艺。经过无数轮博弈,造假者造出来的东西越来越像真的,直到鉴定者完全分不清真假为止。
在数学上,造假者就是生成器网络(Generator,简称 G),鉴定者就是判别器网络(Discriminator,简称 D)。生成器的输入是一个随机噪声向量,输出是一张图片(或一段文本、一条语音等);判别器的输入是一张图片,输出是一个 0 到 1 之间的实数,表示“这张图片是真货的概率”。
这个对抗式的结构,就是 GAN 对抗网络的灵魂。它不直接去建模概率密度,而是通过网络之间的动态博弈,间接逼近真实分布。这一点,是理解后面所有公式的大前提。明白了这一点,后面那些看似吓人的 min-max 目标函数,其实都是在用数学语言描述这个造假与验伪的过程。
2. 造假者与鉴定者:GAN 的两个网络,到底谁更“聪明”
2.1 生成器 G 和判别器 D 的分工细节
先看生成器 G。它的输入是一个固定维度的随机噪声向量,通常服从标准正态分布或者均匀分布,记为 z。这个 z 本身没有任何语义含义,就是一堆随机数。G 的内部通过反卷积(转置卷积)或者上采样操作,把这堆随机数一步步变成一张完整的图片。本质上,G 学习的是一个映射函数:从噪声空间映射到图像空间。
你可以把 z 理解为生成器的“灵感种子”。同一个 G,输入不同的 z,就能得到不同的生成结果。而且这里面有一个很微妙的性质:如果 z 空间中的两个点离得近,那么它们生成的图片往往也比较相似。这意味着 G 隐隐约约在噪声空间和图像语义之间建立了一种对应关系,这也是后来很多人做“潜空间插值”实验的基础——从一个 z 线性过渡到另一个 z,生成的图片会连贯地变化。
再看判别器 D。它的结构通常就是一个标准的卷积分类网络,输入一张图片,输出一个标量。刚开始学的时候,很多人会以为 D 是在给图片“打分”,分数越高越像真的。这么理解没有大问题,但要注意,D 的输出并不是概率分布,它只是通过 sigmoid 函数压缩到 (0, 1) 区间的一个数值,解释成“真”的程度。
有意思的地方来了:D 的“智商”不是固定的,它的判别能力是在和 G 的对抗中不断增长的。G 刚训练时造出来的图片一团糟,D 很容易分辨;但 D 不能因为“很好分辨”就躺在功劳簿上,因为 G 也在不断进化。D 必须持续提升自己的鉴别力,否则下一轮就会被 G 骗过去。这正是对抗训练最本质的动力学过程:两个网络都在一刻不停地追赶对方,谁也不许停下来。
2.2 零和博弈结构:一方所得即另一方所失
GAN 的对抗结构,在博弈论里有一个准确的名字叫零和博弈。游戏里双方的总收益恒定为一个常数,你多拿一分,我就少拿一分,正好对冲。放到 GAN 里,D 和 G 的收益函数互为相反数。
D 的目标很单纯:看到真实图片时,输出尽量接近 1;看到 G 生成的假图片时,输出尽量接近 0。G 的目标也很单纯:让 D 看到自己生成的假图片时,输出尽量接近 1——也就是骗过 D。
这两个目标在字面上是直接冲突的。D 想把真假分得越来越清楚,G 想把 D 越来越糊涂。这种朴素的矛盾,就是整个 GAN 运转的驱动力。有一句话我觉得说得很到位:GAN 里的 D 不是一个评判老师,而是 G 的陪练对手。G 的每一次进步,都是因为 D 变得更难骗之后被逼出来的。
当然,对抗训练的设计要讲究平衡。如果 G 太强,D 完全认不出假图,那 D 会输出一个无意义的 0.5 左右,梯度信息近乎为零,G 就失去了学习方向;反过来,如果 D 太强,一眼看穿 G 所有的把戏,G 的梯度又可能爆炸或者消失,同样学不下去。如何让两个网络在博弈中保持“势均力敌”,是 GAN 训练中最重要的技巧,后面我会专门展开。
2.3 纳什均衡与“理想状态”到底是什么样
深度学习里有很多概念是从物理、数学或者其他领域借来的,GAN 里的“均衡”概念就来自博弈论,具体的名字叫纳什均衡。要理解这个,我们得先回答一个问题:GAN 训练到最后,到底会发生什么?
理想情况下,G 生成的图片分布完全等于真实图片分布,也就是说 p_g(x) = p_data(x)。这时候,D 无论看到什么图片,都无法判断它来自真实数据集还是来自 G——因为两个分布已经重合了。D 的最优策略只能是什么都不区分,对所有输入都输出 0.5,即“我看不出来”。此时 D 和 G 谁也无法再单方面进步,这个状态就是纳什均衡点。
需要注意的是,GAN 的理论分析都建立在理想假设上,比如两个网络都有无限的表达能力、训练过程中能精确地找到每一步的最优解。实际训练中,我们根本不可能达到这个完美均衡,甚至连“接近”都不容易。真实情况是,训练过程中 G 和 D 的损失曲线像两条咬得很紧的波浪线,来回震荡。G 一时骗过了 D,D 又反应过来重新压制 G,如此反复。
我在实践中经常跟新人说,别看论文里画的那种漂亮的平滑收敛曲线,那是理想插图。你拿真实数据集训练一个原生 GAN,打开 TensorBoard 看到的损失曲线绝大多数时候都是上下乱窜的。不震荡的 GAN 是少数,震荡是常态,关键在于震荡范围是否可控、生成质量是否在波动中提升。
3. 目标函数与数学推导:那些“劝退”公式到底在说什么
3.1 min-max 目标函数逐项拆解
现在进入很多人最头疼的部分:数学公式。但我要说,GAN 的目标函数只要拆开来看,含义其实非常直白。原论文中的目标函数长这样:
min_G max_D V(D, G) = E_{x~p_data}[log D(x)] + E_{z~p_z}[log(1 - D(G(z)))]先看外层结构。max_D 是说,站在判别器 D 的角度,它希望让整个 V(D, G) 尽量大。V 由两项组成。第一项 E[log D(x)],x 来自真实数据分布 p_data,D(x) 越大,log D(x) 越接近 0,这一项越大。第二项 E[log(1 - D(G(z)))],z 是随机噪声,G(z) 是生成的假图片,D(G(z)) 表示判别器给假图片的判断结果。如果 D 能成功识破假图片,D(G(z)) 接近 0,那么 log(1 - 0) = log(1) = 0,这一项也是最大的。所以 D 的目标就是把真图判为真、假图判为假,让两项都尽可能大。
再看 min_G,站在生成器 G 的角度,它只能控制第二项。G 希望 D(G(z)) 接近 1,也就是希望判别器把自己生成的图片错判成真图。这样一来,log(1 - D(G(z))) 里的 D(G(z)) 接近 1,整体就接近 log(0),趋近负无穷。G 想让这一项越小越好,所以外面套着 min_G。
把这两层放在一起读,就是:D 拼命放大这个目标值,G 拼命缩小这个目标值。两个目标相反的网络,共同优化同一个目标函数,这就是对抗训练的数学表达。当你读懂了这个 min-max 结构,你其实已经掌握了 GAN 的核心,后面的零散公式都是在为这个框架服务。
3.2 最优判别器:D 的“最佳状态”可以精确算出来
理解了目标函数之后,下一个里程碑式的推导是:在给定 G 的情况下,D 的最优解到底长什么样。这个推导不需要高深数学,但结果非常重要,它是连接 GAN 与散度理论的关键桥梁。
对于一个固定的 x,目标函数里关于 D 的内部是:
p_data(x) * log(D(x)) + p_g(x) * log(1 - D(x))这里 p_data(x) 是真实数据分布在 x 处的密度,p_g(x) 是生成数据分布在 x 处的密度。把它们当成常数,只把 D(x) 当成变量。对 D(x) 求导并令导数为 0,可以得到:
D*(x) = p_data(x) / (p_data(x) + p_g(x))这就是最优判别器的公式。它的含义很直观:如果某个位置 x 上真实数据的密度远大于生成数据,D*(x) 就接近 1,判别器倾向于认为它是真图;反过来就接近 0。如果真实密度和生成密度相等,D*(x) 恰好等于 0.5,完全无法区分。
这个公式还有一个副产品:在训练初期,p_g(x) 和 p_data(x) 几乎不重叠,最优判别器可以轻松做到接近 1 或 0 的判断。这也是为什么 GAN 训练初期 D 的准确率会迅速冲到接近 100%——不是 D 多聪明,而是 G 的太菜实在太明显了。
3.3 代入最优 D 之后:GAN 在最小化 JS 散度
求出 D* 之后,把它代入原目标函数,会发现一个惊人的结果。经过一番整理,目标函数可以变成:
V(G, D*) = -2log2 + 2 * JSD(p_data || p_g)其中 JSD 表示JS 散度(Jensen-Shannon Divergence),它是衡量两个概率分布之间差异大小的指标。JS 散度有一个很好的性质:当两个分布完全相同时,它等于 0;两个分布差异越大,它的值越大。
所以这句话实际上是在说:当判别器处于最优状态时,生成器要最小化的那个目标函数,本质上就是在最小化“真实分布”和“生成分布”之间的 JS 散度。GAN 训练,从底层逻辑上看,就是在想尽一切办法让两个分布重合。这个结论非常重要,因为后续很多 GAN 的改进(比如 WGAN),就是在替换这两个分布之间的度量方式。
这里要插一句给零基础读者的建议:如果你第一遍看不懂这个推导,不要焦虑,多花点时间把 D* 的公式自己手推一遍,然后再代入一次,这个过程值得。我在带新人时经常说,GAN 原理的“顿悟时刻”,往往就发生在你亲手代入公式、看到“-2log2 + 2*JSD”跳出来的那几秒。
3.4 热搜问题解读:原始 GAN 公式里的交叉熵为什么没有负号
热搜词里有这样一个问题:“原始 GAN 公式的交叉熵为什么没有负号?”问的人显然是把 GAN 的目标函数和普通分类任务的交叉熵损失放一起比较了。这个问题其实问得很好,因为它暴露了一个常见的混淆点。
交叉熵损失的标准形式是这样的:
L = -[y * log(p) + (1 - y) * log(1 - p)]注意最前面有一个负号,因为神经网络的训练目标是最小化损失。而机器学习库(比如 PyTorch 的 BCELoss)里实现的就是这个带负号的版本。但 GAN 论文里的目标函数是 V(D, G) = E[log D(x)] + E[log(1 - D(G(z)))],前面没有负号。于是有人就懵了:同样使用了 log 函数,为什么这里不加负号?
答案是:因为 GAN 写的是“最大化目标函数”的形式。max_D V(D, G) 本质上是最大化上式的内部部分,等价于最小化带负号的交叉熵。换句话说,负号不是消失了,而是被挪到了不等号的另一边。你拿 PyTorch 的 BCELoss 去实现 GAN 的判别器时,代码里依然在算交叉熵损失,而且带负号,与你追随论文写出的将目标函数“最大化”是同一个意思。符号的方向只是同一个数学目标的不同写法而已,不代表公式里少了东西。
我在给别人做代码 review 的时候经常遇到这种困惑,其实只要记住一句话就不容易绕晕:写代码时看库函数的约定(通常是最小化损失),看论文时看图作者的约定(通常是最大化目标),两者只是正负号之差。理解了这一点,你再看很多网上互相矛盾的 GAN 代码实现,就不会觉得有理解障碍了。
4. 手把手过一遍训练流程:到底谁先更新,谁让着谁
4.1 经典训练循环伪代码:为什么是五步 D 一步 G
很多新手在理解 GAN 原理时,数学公式看明白了,但一到代码层面就开始糊涂:到底先训练 D 还是先训练 G?一个 epoch 里各训练几次?步调不一致会不会崩?这些问题在论文里其实有标准答案,但我发现大多数教程都只是快速掠过,导致新人动手时容易迷路。
原始 GAN 论文给出的训练流程是这样的:
- 每个训练迭代内部,重复 k 次(默认 k=1,但常用 5)判别器的更新;
- 然后更新 1 次生成器;
- 训练若干轮,直到满足停止条件。
用伪代码表示大概是:
for epoch in range(epochs): for batch in real_data_loader: # 第一步:用真实图片和当前 G 生成的假图片训练 D z = sample_noise(batch_size) fake_images = G(z) d_loss = discriminator_loss(real_batch, fake_images) d_optimizer.zero_grad() d_loss.backward() d_optimizer.step() # 第二步:固定 D,训练 G z = sample_noise(batch_size) fake_images = G(z) g_loss = generator_loss(fake_images) g_optimizer.zero_grad() g_loss.backward() g_optimizer.step()为什么先训 D 且要让 D 多学几步?道理在上一节的数学推导里已经埋下伏笔:G 的优化目标只有在 D 处于最优状态时,才严格等于最小化 JS 散度。如果 D 太弱,G 随便生成一点东西就能骗过 D,G 会以为自己已经天下无敌,停止进步,整个训练也就失去了意义。所以训练时经常让 D 多走几步,确保它足够“犀利”,这样 G 收到的反馈信号才有价值。
4.2 从目标函数到 PyTorch 代码:损失函数怎么写
了解框架之后,实际代码就水到渠成了。判别器的损失是一个标准二分类交叉熵,把真实图片标签设为 1,生成图片标签设为 0 即可:
import torch import torch.nn as nn bce = nn.BCEWithLogitsLoss() def discriminator_loss(real_output, fake_output): real_loss = bce(real_output, torch.ones_like(real_output)) fake_loss = bce(fake_output, torch.zeros_like(fake_output)) return real_loss + fake_loss生成器的损失有两种常见写法。第一种是原始论文里的 min log(1 - D(G(z))),对应代码:
def generator_loss_v1(fake_output): return bce(fake_output, torch.zeros_like(fake_output))注意,这里把 fake_output 的标签设为了 0,意思是“判别器认为假图是真的概率越低,损失越大”,正好逼迫生成器生成更高的 D(G(z))。不过我在实践中很少用这种写法,因为早期梯度太小,训练很慢。
第二种是 Goodfellow 在论文里也提到过的非饱和损失,直接把生成器损失改为最小化 -log(D(G(z))),也就是让生成器想办法让判别器输出接近 1:
def generator_loss_v2(fake_output): return bce(fake_output, torch.ones_like(fake_output))两个版本表面看只差一个标签设置,实际梯度行为差很远。v1 在 D 太强时梯度趋近 0,G 学不到东西;v2 在同样的情况下梯度依然比较大,能让 G 继续更新。所以我的建议是,无论你看的是哪篇教程,动手写代码时都优先用非饱和版本。当年我带着第一批项目跑的时候,就是因为一直用 v1 写法,整整卡了三天,后来换成 v2 才跑出肉眼可见的生成效果。
4.3 训练动态曲线:如何判断 G 和 D 是否在有效博弈
光有代码还不够,训练的时候你总得知道自己调的模型跑得对不对。判别器和生成器的损失曲线是零基础读者最需要学会看的“仪表盘”。
训练的早期阶段,D 的损失会下降得很快,因为 G 还在原地踏步,D 基本上看一眼就能分辨真假。与此同时,G 的损失会慢慢上升,甚至涨到一个看上去很吓人的高值。这时候不要慌,更不要盲目调低学习率。只要 G 的损失在一段时间震荡后开始回落,同时 D 的损失开始反弹,就说明两个网络开始进入了博弈节奏。
如果长时间 G 的损失一直居高不下,同时生成图片仍然是一团模糊的噪点,那大概率是 D 强势压制了 G,梯度没有有效回流到生成器。处理手段通常是:降低 D 的学习率、增加 G 的训练步数、或者修改 D 的结构让它“弱”一点。如果 D 的损失迅速降到 0 附近且长期保持不变,而 G 的损失也在 0 附近徘徊,那么要警惕模型已经“崩溃”了——这个现象叫模式崩溃,下一节我会详细讲。
我的经验值是这样的:一个结构相对简单的 DCGAN 在 MNIST 上训练时,判别器的损失大致会稳定在 0.6 到 1.2 之间来回震荡。如果你看到 D 的损失稳定在 0.1 以下,那大概率不是 G 太强,而是你已经陷入了某种训练陷阱。用损失曲线的数值范围来反推训练是否正常,是每个 GAN 玩家都必须掌握的直觉。
5. 实战中的经典翻车现场:不收敛、模式崩溃、梯度消失
5.1 模式崩溃:为什么 GAN 只会画同一种猫
几乎所有实际训练过 GAN 的人,都遇到过同一个噩梦——生成器生成的所有图片都惊人地相似。在 MNIST 上,它可能只会写一个数字“1”;在人脸上,它只会生成同一张面孔的不同角度。这个现象在业内被称为模式崩溃。
模式崩溃的成因不复杂:G 发现只要反复输出一个特定的、能骗过当前 D 的样本,就可以在最小化损失上偷懒。因为 G 不需要覆盖整个真实分布,它只需要找到一个能让 D 失守的“漏洞”然后无限输出。D 虽然能发现这个规律并修复漏洞,但 G 的搜索速度往往比 D 修复漏洞的速度快得多,于是 G 会在不同模式下跳来跳去,或者干脆锁定单一模式。
应对模式崩溃,传统招数有几类。第一类是小批量判别,让 D 同时看一整批图片,如果这批图片太过相似就判假,能有效打断 G 的偷懒。第二类是Unrolled GAN,在更新 G 之前先“预演”几步 D 的更新,让 G 对 D 的反应更有前瞻性。第三类是调整训练节奏,降低 D 的学习率,让 G 有更多空间探索分布的其他模式。
但在我的实测经验里,真正对零基础用户最有效的方式,反而是先检查两个非常简单的东西:一是生成器的输出层是否用了合适的激活函数,二是 z 向量的维度是不是设得太低了。z 维度过低时,G 的“表达空间”太小,它根本无法覆盖复杂分布,生成结果的多样性必然受限。如果你把 z 从 100 加到 512 之后,模式崩溃现象有缓解,那大概率就是容量不够的问题。
5.2 不收敛与振荡:D 和 G 在“打地鼠”
第二种经典翻车场景是训练过程陷入剧烈的振荡,损失曲线像过山车一样上下翻滚,生成质量也跟着忽好忽坏。这种现象背后是 GAN 的对抗本质:D 学会了识破一种造假手段,G 立刻换一种手段;G 换完手段,D 再重新适应。整个过程像一个高强度的“打地鼠”游戏,永远没有尽头。
如果只是小幅振荡,其实可以接受,因为 GAN 本身就是一个动态博弈过程。但如果振荡幅度很大,导致生成结果始终不稳定,就要考虑干预了。常用的止损手段包括:降低学习率、使用 Adam 优化器时调整 beta1 参数(推荐设到 0.5 左右而不是默认的 0.9)、增大 batch size、给 D 的输入添加少量噪声(标签平滑也算其中一种)。
我自己的经验,给真实标签加上标签平滑处理是最省事也最有效的策略之一。做法很简单,把真实图片的标签从 1 改成 0.8 到 0.95 之间的随机数,不要把话说死,这样 D 即使看到真实图片也不会输出绝对置信,梯度更有韧性,训练稳定性会有肉眼可见的提升。对应到代码里,就是把 torch.ones_like(real_output) 改成 torch.full_like(real_output, 0.9) 之类的操作。
5.3 梯度消失:D 太强,G 根本学不到东西
第三种常见问题,通常会从前面的 D 损失极低、G 损失纹丝不动这些状态中体现出来。原理回到前面讲过的最优判别器公式:当 D 处于最优状态时,D*(x) 在真实图片的位置接近 1、在生成图片的位置接近 0,两个分布几乎没有重叠。此时 G 的梯度会在 JS 散度面前失效,生成器得不到有效的学习信号。
这不是 G 的结构问题,而是度量方式的问题。JS 散度在两个分布不重叠时不能提供平滑的梯度,导致 G 陷入“无论我怎么改,D 都觉得我是一坨垃圾”的状态。解决思路有很多,最著名的是WGAN系列,把 JS 散度换成 Wasserstein 距离,两个分布即使不重叠,也能给出有意义的梯度方向。
但对零基础读者,我给的第一个建议是:不用急着学 WGAN 那一大堆理论,先回到最简单的操作层面检查一下。生成器的信息是不是被 D 完全碾压了?如果是,把 D 的结构变浅一点、参数变少一点,制造一个“不公平”的对抗环境,反而更容易训出东西。我见过不少新手在 MNIST 上卡住,最后发现原因就是 D 写得太深,一个只有两三层卷积的 G 根本骗不过一个六层卷积的 D。
5.4 我的调参避坑清单
最后,把我在多次实操中踩过的坑总结成一张可以直接照抄的清单,给刚入门的读者一个抓手:
- 优化器用 Adam,学习率建议 G 和 D 都从 2e-4 起步,尽量不要直接用默认的 1e-3。
- Adam 的 beta1 设置成 0.5,而不是 PyTorch 默认的 0.9,训练稳定性会好很多。
- 不要用 SGD 训 GAN,除非你对调参有充分的把握。
- 输入噪声 z 的维度不要太小,通常从 100 起步。
- 激活函数方面,生成器内部用 ReLU,输出层用 Tanh;判别器内部用 LeakyReLU,输出层用 Sigmoid 或者直接用带 logits 的 BCE 损失。
- Batch size 不要太小,不要小于 32,否则 BN 层的行为会很不稳定。
- G 和 D 的网络深度不要悬殊,D 明显强于 G 时,训练几乎必然失败。
- 每训练一轮,保存一次生成器输出的样例图片,方便回溯是哪个 epoch 开始崩溃的。
6. 从 Vanilla GAN 出发:给零基础者的学习路线与实操建议
6.1 不建议一上来就学各种 GAN 变体
现在网上关于 GAN 的资料浩如烟海,有 DCGAN、WGAN、CycleGAN、StyleGAN、BigGAN……标题一个比一个酷炫。我看到太多零基础读者的学习路径是先找 WGAN 的论文,结果看到 Wasserstein 距离的推导直接劝退。这完全本末倒置了。
我的强烈建议是:先耐住性子把 Vanilla GAN(原始 GAN)的原理吃透一遍。它是所有变体的地基,理解了原始 GAN 的 min-max 博弈、判别器最优解、分布度量方式,你再看 WGAN 的时候会豁然开朗:它不过是把 JS 散度换成了 Wasserstein 距离,顺带把判别器换成了评价器(Critic),改动核心只有几行。而 CycleGAN 本质上就是在两个域之间各放一个 GAN 再拼起来。
一定不要在基础不稳的时候追求最前沿。带过太多新人,看过太多例子:一上来就想学 StyleGAN 的人,花了一个月连代码都跑不通,最后回头补原始 GAN 的基础;而老老实实从 Vanilla GAN 起步的人,反而在两周内就能训练出自己的第一个能生成人脸的模型。学习顺序,往往比你投入的时间更重要。
6.2 自己的第一个 GAN 项目:数据集和评价指标怎么选
想亲手训练一个 GAN,数据集的选择决定了你的入门体验。我的推荐顺序是:MNIST 或者 Fashion-MNIST 入门、CIFAR-10 进阶、自己的人脸数据集或 CelebA 再往后。
MNIST 几乎是完美的入门数据集:图片只有 28×28 单通道,生成器网络不用很深,训练一个 epoch 也非常快。当年我第一次跑通 GAN 时,在 MNIST 上训练 50 个 epoch 大概只需要几分钟(单卡 CPU 也能慢慢跑),而 CIFAR-10 上跑同样的网络结构可能要几十分钟。入门阶段你需要的不是 GPU 算力,而是快速迭代反馈,MNIST 能帮你做到。
评价生成效果也是一个难点。很多零基础读者会问:怎么判断 GAN 好不好?这里我给出一个朴素的方案:不用追求量化指标,直接肉眼观察加上把生成图和真实图并排对比。理论上虽然有一些量化指标,比如 Inception Score(IS)和 Frechet Inception Distance(FID),但零基础阶段看这些数字意义不大,而且计算 FID 还需要额外下载预训练的特征提取模型。我的一致建议是:先学会用眼睛判断“生成得像不像”“多样性强不强”,等到你对自己的模型有了信心,再引入 FID 作为辅助指标。
6.3 实验管理习惯:日志、模型保存、可视化三件套
最后想分享一个很少被教程提及、但实际工程里极其重要的经验:在训练 GAN 时,实验管理习惯直接决定你的调试效率。GAN 训练本质上是不稳定的,你可能跑了 100 个 epoch,发现真正好的结果只出现在第 80 到第 90 个 epoch 之间。如果你没有定期保存模型权重和生成图片的习惯,你会为这个失误懊恼很久。
我的做法是,写训练代码时就固定住三个习惯:
- 每 10 个 epoch 保存一次生成器 G 的权重,文件名带上 epoch 编号,比如
G_epoch_80.pth; - 每次保存权重时,同时保存当前 epoch 生成的 64 张拼接图片,形成一张
grid_epoch_80.png,方便翻看整个训练过程的生成质量变化; - 用 TensorBoard 或者 Weights & Biases 记录每一轮的 D 损失、G 损失和真实/生成的判别输出均值。前两个看博弈状态,后面两个能帮你判断 D 是不是已经“饱和”了。
这套习惯看着不起眼,但在你调了几天参、发现模型开始崩溃的时候,它是唯一能帮你定位问题出在哪个 epoch 的救命稻草。我见过太多新手,辛辛苦苦训了一整夜,第二天起来发现模型崩了,但同时也没有任何保存记录,只能从头再来。这种低级代价,完全可以避免。
回头看,GAN 的很多“劝退点”其实不过是用复杂术语包装了一个简单的思想。零基础学 GAN,最忌讳的是被公式和变体带乱节奏。抓住“造假者与鉴定者博弈”这条主线,理解目标函数在做什么、最优判别器长什么样、模型崩溃和梯度消失的原因,你就能拥有一个足够扎实的起点。剩下的所有 GAN 变体,都是在承认 Vanilla GAN 不完美的前提下,一点一点修补它的短板。先把那个最初的、有点粗糙的 GAN 亲手训起来,让它在你的屏幕上画出一张张清晰又逼真的小图,你一定会比读任何教程都学得更快。