反向传播 入门

反向传播在网络里反着走一趟,就把每个参数的梯度都拿到手;而它整个搭在同一张图上:一条曲线,和一条贴着它走的线,那条线的倾角就是梯度。三个主张,每一个都由旁边那张图算出来 —— 反向传播是一串斜率相乘;正是这串乘法让它只花一趟而不是一百万趟;也正是同一串乘法,让深网络会消失、会爆炸。

01

梯度就是一条斜率

不是比喻。是一条线真实的陡峭程度,你可以伸手把它扳过来。

整篇 primer 用的就是这一个网络:一个输入 x = 2、一个 ReLU 隐藏神经元、一个输出,对目标 y = 5 算平方误差,四个参数从 w₁ = 1.5、b₁ = 0.5、w₂ = 0.8、b₂ = 0.1 出发。预测值是 2.90,所以损失是 4.41。随便动其中一个参数,这个数就画出一条曲线 —— 挑一个看看:

沿 w₁ 看,损失 4.41

四条曲线,都穿过同一个点。w₁ 那条在 −0.25 处有个折角:隐藏单元的 ReLU 在那里关掉了,损失从此对 w₁ 毫无反应;b₂ 那条则是一条干净的抛物线。反向传播交还的,是每个参数一个数,而每一个都是那个参数自己那条曲线的性质。

这个性质就是我们此刻站的那一点有多陡。把w₂ 那条曲线拿出来,在上面搭一条切线,再沿曲线拖动那个点,看斜率怎么转:

w₂ = 0.80,梯度 −14.70
w₂ = 0.80,梯度 −14.70

切线就是梯度 —— 除此之外没有别的要算。在初始权重处它向下倾斜 −14.70,这正是 ∂L/∂w₂。往右滑到 它就平了:斜率为零,因为损失到了最低点。「没有梯度」长的就是这个样子。

−14.70 是从哪来的?从导数的定义来的:取两个相邻点上的损失,除以它们之间的距离。把间距收小,看割线怎么甩到切线上:

间距 1.60,割线斜率 4.90

注意初始间距 1.60 时的符号:割线读出 +4.90,跟真相正好相反 —— 因为这么宽的一步直接跨过了最低点。把它收到 ,读数变成 −14.09。torch.autograd.gradcheck做的就是这件事,一次一个参数 —— 它只能当测试,不能当训练循环,因为每个参数都要花掉一整趟前向传播。

斜率存在的意义,是告诉你该往哪边走。它指向上坡,所以我们逆着它走一步 −η · ∂L/∂w₂。学习率一开始是 0,所以这一步还没有长度;把它抬起来,看权重沿曲线往下滑:

η = 0.000,走完这步的损失 4.41

看看越过 之后会发生什么:这一步冲过了最低点,开始往对面的墙上爬。到 时它正好落回出发时的那个损失,再往后每一步都让情况更糟。这个阈值是 2/L″,对这条抛物线就是 1/a² = 0.082 —— 梯度告诉你方向,从来不告诉你距离。

在把斜率串起来之前还有一件事。反向传播需要每个节点的斜率,不只是损失的斜率;而对线性节点来说,这个斜率是白送的。移动输出权重,看这个节点的直线绕着输入转:

w₂ = 0.80,局部斜率 0.80

因为 ŷ = w₂·a + b₂ 是一条直线,它的斜率在每个输入处都一样,而且就是 w₂ —— 一个前向传播本来就已经握在手里的数。框架里每一个算子的局部导数都是这个样子:一条短公式,用的全是前向传播反正都要算出来的值。反向传播便宜,是因为这些公式便宜。

02

链式法则就是斜率相乘

变化率会相乘。把这一件事套到图上每一条边,就是整个算法。

网络是一摞小函数,一个喂给下一个。如果输入端动一下、中间就动两倍,中间动一下、损失就动零点八倍,那么输入端动一下,损失就动 0.8 × 2 = 1.6。背后没有更深的东西。

这张图把那次乘法画成了实物。同一次扰动被画了三遍 —— 一次在输入端,一次在第一段斜率之后,一次在第二段之后 —— 每一条轨都是上面那条乘以这一段的斜率:

第一段斜率 2.00 · 第二段斜率 0.80

注意任意一个滑块经过 时会发生什么:最下面那条轨整个消失,不管另一段有多大。一段死掉就杀死整条路径 —— 这就是 §05 的问题的缩微版。把两段都推到 1 以上,最下面那条轨反而比最上面还长 —— 梯度在往回走的路上,涨起来和缩下去一样容易。

我们这个网络就是这样的四段:x → z → a → ŷ → L。在任何东西能往回流之前,前向传播必须先跑一遍,在每一根线上留下一个值。按播放,或者直接拖轨道,看四个值按顺序冒出来:

第 0 步,共 5 步 —— x = 2.00

这些数每一个都还要再用一次。激活值 a = 3.50 会变成 ŷ 对 w₂ 的局部导数;输入 x = 2 会变成 z 对 w₁ 的局部导数。这就是前向传播为什么要留着中间结果而不是扔掉 —— §03 会给这个决定标上价钱。

现在把同一张图倒过来走。每一个反向步骤都把前一步的结果乘上一个局部导数,再把结果往左递。拖动轨道,看梯度一根线一根线地出现:

第 0 步,共 5 步 —— L 处的梯度 = 1.00

看箭头上的那些乘数:−4.20,然后 × 0.80,然后 × 1,然后 × 1.50。a 处的梯度是 −3.36,穿过 ReLU 之后还是 −3.36,因为 z = 3.5 是正的,ReLU 在那里的斜率正好是 1。反向传播从不重新推导任何东西;它只是顺着一张表往下乘。

它划算的原因在于复用。每个参数都挂在某一根线上,而那根线的梯度是所有挂在上面的参数共用、只算一次的。逐个走过这四个参数,看路径上共享的那一段纹丝不动,只多出一次乘法:

b₂:在共享路径之上再多 1 次乘法

如果各算各的,这四个梯度要花 2 + 2 + 4 + 4 = 12 次乘法。一起算只要 7 次:三个线上梯度,加上每个参数各一次。把它放大到真实模型上,这个比例就是全部的故事 —— 96 层网络里单干的话要重走 96 层,而现在它只付一次乘法,用的是上一层已经填好的那根线。

03

之所以反着走,是因为便宜

链式法则并没有说该从哪一头开始。是问题的形状说了算。

一串乘法可以从左往右算,也可以从右往左算。从输入端出发,你把某一个输入的影响一路往前推;从损失出发,你把损失的敏感度一路往后推。两种都正确。只有一种便宜,而哪一种便宜取决于形状。

我们的形状是几百万个参数进、一个标量出。把参数个数调大,看前向模式的扫描一层层堆起来,而反向模式那一趟始终只有一遍:

前向模式 · 扫 4 遍

前向模式每个输入要扫一遍;反向模式每个输出扫一遍。就是 12 遍对 1 遍,而一个 1750 亿参数的模型要扫 1750 亿遍。训练之所以可行,全部理由就在这里:损失只是一个数,所以便宜的那一头,是远的那一头。

不过反向模式并不白来 —— 它必须在前向传播离开每根线之后再去看它,所以前向传播不能把中间结果扔掉。把深度拖下去再拖回来,看 GPT-3 那个形状要囤下多少张量:

96 层,2,048 个 token

GPT-3 的 96 层,在一条 2,048 个 token 的序列上,要留下275 GB 的激活值 —— 而卡上只有 80 GB。这个数字出自 Korthikanti 等人,他们的公式还把它拆开了:82.1 GB与序列长度成正比,193 GB是注意力矩阵,与序列长度的平方成正比。

真正扎人的是那个平方项。下面两根轴都是对数轴,所以直线就是幂律,直线的陡峭程度就是那个幂次。逐档加大上下文长度,看总量怎样从线性那一项上拐开:

2,048 个 token,其中注意力占 193 GB

在 处两条曲线相差不到一倍;到 时总量是 13,027 GB,而线性项只有657 GB。长上下文首先是内存问题,其次才是算力问题 —— FlashAttention 就是为此而生的:它根本不把这一项在量的那个注意力矩阵写出来。

另一条出路是几乎什么都不留,用重算的代价把它补回来。只把每层的输入存成检查点,反向传播在需要激活值之前,先把那一段的前向重跑一遍。在存下来和重算之间切换:

96 层,存下来

因为反向传播的算术量大约是前向的两倍,多跑一趟前向大约让每步多三分之一的活 —— 换来的是从275 GB 掉到7.7 GB。正是这笔交易,让梯度检查点在每个训练框架里都只是一行开关,也让「模型装不下」时第一个动作永远是把它打开。

04

局部导数是从哪儿来的

反向走那一趟上的每一个乘数,都是某条曲线在某一点上的斜率。曲线在这里。

线性层交出来的是它的权重,跟输入无关。激活函数不一样:它的斜率取决于前向传播恰好落在哪里,所以同一个网络在不同样本上交还的乘数并不相同。

四个激活函数,共用一根横轴。挑一个函数,把切线沿着曲线拖动;读数就是反向传播接下来要乘的那个数:

σ 在 z = 0.00,斜率 0.250
σ 在 z = 0.00,斜率 0.250

注意 sigmoid 除了中间那一小段,其余地方有多平。它最陡的位置是原点,可即便在那里斜率也只有 0.250。tanh 的峰值是 1.000, ReLU 在活着的那一半正好是 1、另一半正好是 0, GELU 的上限略高于 1,是 1.129。一百层网络的命运,就由这四个数决定。

sigmoid 的这个上限值得单独画一次。曲线下面是它自己的斜率,画在同一根激活前的横轴上,并把上限横着标出来。移动游标,在两个地方读同一个数:

z = 0.00,斜率 0.250

那条斜率曲线是 σ(z)·(1 − σ(z)),两个和为 1、都落在 [0, 1] 里的数的乘积 —— 所以它在两者相等时最大,也就是 处,正好是四分之一。横轴上没有任何一个输入,能让 sigmoid 交出大于 0.25 的数。这是一个上界,不是一种倾向,而 §05 会把它自己乘自己。

把同一个单元推进尾巴里,这个上界就不重要了,因为真实的数远在它之下。把点滑出去,看斜率三角形怎么压成一条线:

z = 0.00,斜率 0.2500

在 处斜率是 0.0025 —— 最大值的百分之一 —— 到 是 0.0003。这个单元没有坏:它照样输出一个自信的 0.9997。它只是不再能学习了,因为它下游的一切都要乘上万分之三。这就是饱和,而且是一次沉默的失败 —— 损失只是不再往下走了。

ReLU 没有可以饱和的尾巴,但它有个更糟的把戏。它在死的那一侧斜率正好是零,所以一个整批输入都落在折点左边的单元,什么都收不到。把偏置往下拖,看整批越过去:

偏置 0.00,8 个输入里有 6 个还活着
偏置 0.00,8 个输入里有 6 个还活着

8 个输入里一开始有 6 个活着。拖过 之后一个都不剩:每个样本拿到的斜率都是 0,于是 ∂L/∂w 正好是零,权重不动,偏置不动,这个单元永远死了。没有任何东西会抛异常。把这批输入摁在那儿不动,换一下这个单元算什么:

ReLU,偏置 −2.50

在 ReLU 下,8 个斜率加起来正好是 0.000。leaky ReLU 在负半边那个常数 0.01把它变成 0.080;GELU 在那里的斜率很小、还短暂地是负的,加起来是 0.152。这些都不大。但它们全都不是零,而不是零正是唯一要紧的性质:还能收到点什么的单元,就还能爬回来。

05

五十条斜率,乘在一起

反向传播的不变式是:第 i 层的梯度,等于第 i+1 层的梯度乘上一个局部导数。反复迭代,你手里就是一个连乘。

深度的全部困难,一句话就说完了。五十层就是五十个因子相乘,而许多个数的乘积不像求和那样温和:它不是慢慢漂移,它复利。两种失败模式,就是乘积能做的两件事。

下面这些柱子是每一层上梯度的大小,画在对数轴上 —— 往上一格就是十倍,所以一排笔直的柱子意味着指数增长或指数衰减。每一段的系数一开始正好是 1.00,所以既不缩也不涨;把系数往任意一边推一点:

系数 1.00,走过 50 层,到达 1.0e+0

看它需要的偏离有多小。 —— 每层损失一成,听上去人畜无害 —— 走完五十层就只剩 5.2e−3。到 时,抵达第 50 层的梯度是 1.3e−20;到 时是 1.6e+10。只有一个系数值能让乘积原封不动,而没有哪个网络会碰巧坐在上面。

对 sigmoid 来说这个系数由不得你选。 §04 已经证明它的斜率永远不超过四分之一,所以一摞 sigmoid 对「能有多少梯度活下来」有一个保证的上限。加层数,看这排柱子怎么穿过半精度放弃的那两条线:

4 层,到达 3.9e−3

这两条线都是 2 的整数次幂,所以它们正好落在整数层上。0.25⁷ = 2⁻¹⁴ 是 fp16 最小的正规数,所以 就是半精度梯度开始掉比特的地方。0.25¹² = 2⁻²⁴ 是 fp16 手里最后一个数,所以到 时格式已经用尽,再多一层就正好是零。而这还是最好的情况:每个单元都恰好待在自己最陡的那一点上。

对付这件事有个标准的绕法,而且只是一次乘法:在调用反向传播之前把损失放大,图里每一个梯度就都按同样的倍数回来。把这排柱子整个抬离地板:

损失缩放 1

损失缩放取 时, 12 层这一趟落在 9.8e−4,稳稳待在 fp16 的正规数范围里;更新之前优化器再除回 16,384 —— 所以走的这一步没变,变的只是它被表示的方式。注意第二条线:最上面那根柱子必须待在 fp16 自己的最大值以下,这就是 GradScaler为什么要靠不断翻倍、直到某处溢出再退回来的方式去找这个系数。

往上的那堵墙比人们以为的更硬也更近。这里链条有 128 层深,每一段都乘上大于 1 的数,顶上横着的是float32 的最大值。把系数抬高,直到柱子撞上它:

系数 1.60,到达 1.3e+26

在 —— 每层翻一倍 —— 这个乘积在第 128 层越过 3.4e+38,之后每根柱子读出的都是 inf。梯度里出现一个 inf,权重更新里就出现一个 inf,而 inf − inf 是 nan:再过两轮迭代,模型里每个参数都是 nan,损失也不再打印出数字。爆炸至少是响的,这是它唯一的仁慈。

在动手修之前先做一次诚实的更正。一层的系数不是标量 —— 它是一个雅可比矩阵,对不同方向的拉伸程度不一样。转动梯度到来的方向,看出去的那一支贴着椭圆走:

120°
σ_max = 1.00

那个圆是梯度可能到来的所有方向;椭圆是这一层把它们送到的地方。长半轴是 σ_max,短的那条是 σ_min,这里是它的 0.63 倍。所以「系数 1.00」的真正意思是:梯度会被乘上 0.63 到 1.00 之间的某个数,取决于它从哪儿来 —— 而必须待在 1 附近的,是 σ_max。谱归一化约束的正是它。

06

权重从哪儿出发

你交给反向传播什么,它就改进什么。而起始权重有两种交法,会让它什么也改进不了。

最自然的第一反应是把所有权重都设成 0。它对称、无偏、一行代码就写完,而且它在第一步之前就杀死了网络 —— 原因是反向传播自身的性质,跟具体是什么损失无关。

四个隐藏神经元,画成它们各自的输入权重向量,每个尖端挂着它收到的那次更新。它们一开始完全叠在一起。把散开程度打开,看本来只有一个的地方冒出四个:

散开 0.00,1 个彼此不同的神经元

一模一样的权重看到一模一样的输入,于是算出一模一样的输出,于是收到一模一样的梯度,于是走完这一步之后它们还是一模一样。反向传播只会保持交给它的对称性;它没有任何打破对称的机制。这样初始化的一百个神经元的层,就是一个神经元复制了一百份,永远如此 —— 而且它不报错,它只是停在一个单神经元的层能达到的准确率上。

所以权重要随机出发。但那份随机的尺度同样不是白给的,因为 §05 里的那个连乘,对前向的激活值和对反向的梯度一样有效。下面是一个 12 层 ReLU 堆叠里逐层的方差,初始化按 He 配方乘上一个增益。移动这个增益:

增益 1.00,第 12 层的方差是 1.0e+0

增益正好是 1.00 时方差是平的:每一层把收到的原样传下去。到 时,第 12 层只剩 1.9e−4;到 时是 3.2e+3。一个超参数上三成的偏差,复利十二次。这就是为什么初始化有一堆带人名的配方,而不是一个默认值。

这些配方彼此只差一个 2,而这个 2 正好就是 ReLU 扔掉的那一半。下面是抵达某一层的 64 个激活前的值,左边是被 ReLU 归零的那一半。先送层号,再换配方:

第 1 层,Xavier

因为 ReLU 删掉了分布的一半,离开一层的方差就是进入这一层的方差的一半。所以Xavier 的 √(1/n)让信号每层减半 ——,散布只剩 0.022,出发时是 1.000。He 的 √(2/n)把那个 2 补了回去,散布就一动不动。在 GPT-2 的宽度 768 上,这意味着标准差 0.051 对 Xavier 的 0.036 —— 一个能训练的深层 ReLU 网络和一个逐渐淡出的网络,差别全在这里。

07

是什么把这个乘积摁在 1 附近

两个结构性的修法,加一个粗暴的。合起来,就是一百层网络居然能训练的原因。

残差块算的是 x + f(x),而不是 f(x),所以它的局部导数是 1 + f′(x),而不是 f′(x)。那个 1是一条通路,梯度走上去,等于什么都不乘。

24 层,画两遍:普通链条和残差链条,恒等通路正好横在 1 的位置上。移动分支自己的斜率:

分支斜率 0.05 —— 普通 6.0e−32,残差 3.2e+0

在初始斜率 0.05 处 —— 一个弱分支,也就是初始化得当的块该有的样子 —— 普通链条走到6.0e−32,残差链条走到3.2e+0。不过要留意残差链条错在哪一边:把分支抬到 ,它走到 1.7e+4。残差并没有修好这个乘积,它只是把失败从消失换成了增长 —— 而增长是归一化摁得住的那一种。

LayerNorm 夹在块与块之间干的就是这件事:把每个块的输出重新缩回一个固定的方差,这样不管上面再叠多少块,残差流都不会无限膨胀。同一个念头也把输出投影的初始化按 1/√(2N) 缩小(N 是层数),也就是一开始就给一个更小的 f′,而不是事后再纠。

而当这一切在某个倒霉的 batch 上全都失效时,还有一件钝器。一个半径为 1 的球,以及落在球外的原始梯度。把尖端拖到任意位置,看优化器实际收到的是什么:

‖g‖ = 2.61,实际用 1.00
‖g‖ = 2.61,实际用 1.00

裁剪是缩放,不是截断:沿着球外侧拖动尖端,真正被用的梯度方向分毫不变,只是长度被削。把尖端拖,它就什么也不做。本来会打印 nan 的那一步,代价变成一次有偏的更新。‖g‖ ≤ 1.0 就是 GPT-3 用的那个值。

剩下的就是循环本身了。六行,其中有一行是所有人都会忘的那一行:

for x, y in loader:
    opt.zero_grad()          # or grads add up
    loss = mse(model(x), y)  # keeps every
    loss.backward()          #   activation
    clip_grad_norm_(p, 1.0)  # the ball above
    opt.step()               # w -= lr * grad

backward() 是往 .grad 里累加,不是覆盖 —— 好让一个 batch 能拆成几趟反向传播。去掉 opt.zero_grad(),看每一步实际施加的是什么:

第 12 步,在累加

第二步施加的是 g₁ + g₂;第十二步施加十二步的量。方向大致是对的,所以没有任何异常,损失照样在降 —— 错的是尺度,而它随迭代次数线性往上爬,直到一小时后整趟训练发散。,每根柱子都掉回 1 那条线上。

真正要抓住的不变式,就是 §05 那一句:任意一层的梯度,等于后一层的梯度乘上一个局部导数。便宜是因为乘法共享;两种失败是因为乘积复利;要配方是因为唯一安全的乘积,因子都待在 1 附近。