微积分 入门

训练循环真正跑的那点微积分,全部搭在三张图上:一条被直线贴住的曲线、一只从上往下看、上面插着一支箭头的碗,以及一张让导数倒着走回去的计算图。灵敏度、链式法则、梯度、步长,以及它失败的四种方式 —— 页面上每一个数字都由旁边那张图算出来,所以你可以把它拖到任何状态,它依然成立。不讲积分:训练一个十亿参数的模型,从头到尾用不到一次积分。

01

斜率就是灵敏度

一个数,回答优化器唯一关心的问题:把这个输入拨动一点点,输出会跟着走多远?

训练就是一场搜索。每走一步,都要对十亿个旋钮里的每一个问同一个问题 ——把它拧动一丝,损失是变好还是变坏,变了多少?答案是每个旋钮一个数。下面所有机器,都是为了便宜地算出这些数,而且全部搭在同一条曲线上。沿着曲线拖动那个点,看它的两个坐标跟着走:

x = 1.80,f(x) = 1.14
x = 1.80,f(x) = 1.14

函数不过是把一个数变成另一个数的规则,而这个点就是这条规则被看见了一次。它还不是变化率:变化率需要两次读数,以及两者之间的距离。往前挪 h 取第二个点,把两点连起来,得到的这条线就有了可以算的斜率。缩小步长,看它远端那道缝隙怎么合上:

h = 1.40,割线斜率 2.05

注意读数收敛得有多快。h = 1.40 时割线读数是 2.05; 时是 0.21; 时是 0.05。割线正走向 0,这个极限就是导数 —— 写作 f′(1),想把两个量都点名时就写 df/dx。同一个东西,三种写法,论文里三种都会出现。

取极限之后,割线变成一条只在一处碰到曲线、并且和曲线在那里同向的直线。那就是切线,它的斜率就是该点的导数。拖着点走,看这条线跟着倾斜:

x = −1.90,f′(x) = 2.61
x = −1.90,f′(x) = 2.61

注意,每一个 x 都对应一个斜率,所以斜率本身就是一个函数:f′(x) = x² − 1。曲线上升处它为正,下降处它为负,在两个灰色刻度处 —— 和 —— 恰好为零,那里曲线走平了。这些平坦处正是优化器在找的东西。

既然是函数,就可以画出来。上面是曲线,下面是它自己的斜率,共用一条输入轴;随便拖一下,两边一起动:

x = −1.60,f′(x) = 1.56
x = −1.60,f′(x) = 1.56

下面那条曲线穿过零点的位置,正是上面那条掉头的位置。有一点要说清楚:两个框并不是按同一比例画的 —— 斜率请读读数,不要靠目测角度。这张图真正说的是:一个公式一次性带上了每一点的上升速率,这就是导数值得用符号推导、而不是一根根量割线的原因。

关于切线还有一种更强的说法,后面所有内容都依赖它:只要离那个点足够近,曲线和直线就是同一个东西。把窗口对半缩小,看两者之间最大的那道缝会怎样:

窗口 ±1.200,最大偏差 2.23

窗口每对折一次,缝隙就大致缩到 四分之一 —— 时是 0.112, 时是 0.027 —— 所以越放大,线性模型赢得越彻底。这就是可微的全部含义:f(x + h) = f(x) + f′(x)·h + (误差),而误差比 h 死得更快。所有框架里的每一次梯度更新,都是在一个很小的 h 上信任这条直线。

于是失效模式就是:某个函数无论放大到什么程度都不会变直。 ReLU 在原点正是如此,而它是深度学习里用得最多的激活函数。把点滑到上,读一读两侧的单边斜率:

x = 1.10,斜率 1.00

在拐角处,从左边看答案是 0,从右边看是 1,所以没有唯一的切线,f′(0) 不存在。而且没有任何东西会报错:relu'(0) 在 PyTorch、TensorFlow、JAX 里都返回 0 —— JAX 是 2022 年才用自定义 JVP 把它钉死的,在那之前返回 1。每个库都悄悄替你了结了这场分歧,其中一个还改过主意。[0, 1] 里任何值都是合法的次梯度,所以不要紧 —— 但最好在你跑去别处找 bug 之前就知道。

02

一次面对很多个旋钮

模型的输入不是一个,是几十亿个。办法是把单变量那个问题,对每个变量各问一遍 —— 同时想清楚「把其余的按住不动」到底代价是什么。

给损失两个参数而不是一个,它就不再是一条曲线,而变成一片地形:每一对 (w₁, w₂) 都有一个高度 L。我们从正上方往下看,于是每一圈等高线串起了代价相同的那些点。拖动那个点,读一读它站在多高的地方:

w₁ = 0.50,w₂ = 1.60,L = 4.63
w₁ = 0.50,w₂ = 1.60,L = 4.63

这就是平方误差损失在极小点附近真实的样子,要看的两件事都跟这些等高线有关。它们是椭圆而不是正圆 —— 曲面横穿山谷的方向比顺着山谷的方向硬得多 —— 而且在陡的地方挤在一起。圈心处它们收敛到 L = 0,那正是训练要找的极小点。

现在把 w₂ 钉住,只让 w₁ 动。这相当于从地形上切下一条曲线,而在这条曲线上,我们又回到了第 01 节的单变量情形。沿着切面滑动,再换一换你站在哪一条切面上:

w₁ = 1.20,∂L/∂w₁ = −0.76

那条切线的斜率就是偏导数,记作 ∂L/∂w₁,用弯的 ∂ 而不是直的 d。这个弯钩不是什么新数学,它只是给读者的一条注记:说明哪些变量被按住了。算法也完全一样:把其余变量当常数,然后用第 01 节的规则。

站在一个点上,这样的问题有两个,每根轴一个,答案也就有两个。挪动那个点,把两个都读出来 —— 它们是从同一点出发的两支箭头 ——沿 w₁ 的变化率和沿 w₂ 的变化率:

∂L/∂w₁ = −2.34,∂L/∂w₂ = 6.52
∂L/∂w₁ = −2.34,∂L/∂w₂ = 6.52

在初始点上,第一个是 −2.34,第二个是 6.52:把 w₁ 往上推会降低损失,把 w₂ 往上推则会让损失陡增。一个 D 元函数有 D 个这样的数,每个变量一个,而配方从来没变过 —— 变多的只是记账。

关于「按住不动」有一点常常绊住人,值得亲手感受一下,而不是读过就算。∂L/∂w₁ 本身也是所有变量的函数。把 w₁ 原地钉死,只动另外那一个:

w₂ = 1.60,∂L/∂w₁ = −0.76

看好了:没人碰过 w₁,切线照样在倾斜。w₂ = 1.60 时斜率读数是 −0.76; 时是 2.70; 时是 6.16。它甚至会变号,就在 w₂ = 1.25 处。所以梯度是关于某个点的事实,而不是关于某个参数的事实:网络里任何一个别的权重一动,你算好的每一个偏导数就都过期了。

03

变化率是乘起来的

神经网络就是函数套函数。把导数穿过整摞函数的规则,是每一环做一次乘法 —— 而且全程不会出现比第 01 节更难的东西。

两环:x 进 g 得到 u,u 进 h 得到 y。下面两个框共用中间那根轴,所以从左边进去的一个扰动,出到右边时已经被缩放了两次。把扰动从初始宽度往小里收,盯住那三段区间一起合拢:

Δx = 0.800,Δy/Δx = 0.155

注意哪些数字对得上。在初始扰动下 Δu/Δx = 1.80,Δy/Δu = 0.086,它们的乘积 0.155 正好就是 Δy/Δx 的读数。这个等式不是近似,也不要求扰动很小 —— Δu 被约掉了,就像任意两个首尾相接的分数约掉共同项一样。

把扰动缩到零,这三个比值各自变成导数,等式原封不动地活过了取极限。下面是两个局部斜率画成的切线;拖动输入,看它们俩一起变:

x = 1.40,dy/dx = 0.198
x = 1.40,dy/dx = 0.198

这就是链式法则:dy/dx = h′(g(x)) · g′(x),或者写成能看出约分的形式 dy/dx = dy/du · du/dx。在初始点上,两个因子是 0.142 和 1.40,答案是 0.198。N 环就是 N 个因子,配方永远不会更难。

里面有一个坑,而且几乎人人都会踩一次:h′ 必须在前向传播留下的那个值处取,不是在 x 处。把输入拖到 ,图上读数是 0.083;要是写成 h′(x)·g′(x),得到 0.190,大了 2.3 倍,而且悄无声息地错。正确的取值点由前向传播提供 —— 框架先正着跑一遍才能倒着求导,原因就在这里。

因子既然是乘起来的,它们的大小就会复利式累积。一个压缩型的环节贡献一个小于 1 的因子。logistic 曲线是最典型的例子:沿着它拖动,看下面它自己的斜率:

x = 1.60,σ′(x) = 0.140
x = 1.60,σ′(x) = 0.140

下面那条曲线的峰值恰好是 0.25,出现在 ;到了 就已经不到 0.007。这个数就是全部故事:σ′ = σ(1−σ) 是关于 σ 的抛物线,峰值在 σ = ½,所以任何地方的任何一个 sigmoid 层,贡献的因子都不可能超过四分之一。

把这些因子叠起来,看看一条长链会把变化率变成什么样。设定每层贡献的增益和一共多少层,看每一节从传到它那里的东西上咬掉多少:

增益 0.75 连乘 24 层:0.0010

增益 0.75 —— 只是轻微压缩 —— 24 层就只剩千分之一了,链条还没走完,光束早就贴在框底上;时只剩 3.19×10⁻⁸。把增益降到 sigmoid 的最好情况 ,光是 10 层就只剩 9.54×10⁻⁷。这就是梯度消失,而且它是无声失败的:不报错,不出 NaN,只是靠近输入的那些层的更新被舍入成零,而损失一动不动。反过来把增益推到 1 以上,同一个乘法就会爆炸 —— 那倒是会自报家门,以 NaN 的形式。

04

把所有偏导数打包成一支箭头

把每个变量的答案装进一个向量,它就获得了单个分量都没有的性质:指向最陡的上坡方向。

把第 02 节的两个偏导数写成同一个向量的两个分量,就得到 ∇L = (∂L/∂w₁, ∂L/∂w₂),也就是梯度。它和参数住在同一个空间里,所以可以直接画在地图上当作一支箭头。挪动那个点,把箭头和它所站的那圈等高线对照着看:

∇L = [−2.34, 6.52],‖∇L‖ = 6.92
∇L = [−2.34, 6.52],‖∇L‖ = 6.92

箭头处处与等高线成直角,而这是被逼出来的,不是碰巧:沿着等高线走,L 根本不变,所以那个方向上的变化率是零,所以梯度在那个方向上没有分量。在初始点上,∇L 的读数是 [−2.34, 6.52],长度 6.92。

说它是最陡方向,这是一个可以验证的断言。沿任意单位方向 û 的上升速率是 ∇L · û —— 也就是梯度投在那个方向上的影子。让影子绕一整圈,找出它的最大值:

0°,上升速率 −2.34

注意影子在哪里最长 —— 在 处,读数是 6.92 —— 正是梯度自己的长度,在梯度自己的方向上。在 处它只剩 0.03 —— 那是沿着等高线的方向,L 在那里根本不变;在正对面它取到最负值,那就是最陡的下坡。投影不可能超过被投影者的长度,证明就这一句。

梯度不是一条路线。它是无穷小一步的最佳方向,而不是通往极小点的方位角;碗一旦不圆,两者就会分道扬镳。把碗拉长,看下坡方向和极小点方向之间的夹角:

κ = 6,偏离 37°

在 时等高线是正圆,两支箭头完全重合,夹角 0°。到了 κ = 6 就差 37°,而 时是 44°。最陡下降会横穿山谷跑掉,而不是顺着谷底走;而那个数 —— 曲面最硬的曲率与最软的曲率之比,也就是它的条件数 —— 正是下一节要讲的东西。

05

那一步

一行式子 —— w ← w − η ∇L(w) —— 以及两种出错的方式,两种你都可以在这里亲手造出来。

梯度指向上坡,我们要往下走,所以减掉它。走多远不是微积分能决定的:导数只在极限意义下成立,任何一次真实的步长都是在赌切线还能诚实多久。学习率 η 就是这个赌注。把它设好,看一步会落在哪里:

η = 0.100,L 4.63 → 1.23

由于曲面会从切线旁边弯开,步子大并不等于更好。从 L = 4.63 出发,η = 0.10 的一步落到 1.23; 落到 0.53,那才是从这里出发单步能达到的最好结果 —— 图里标出了它:在那里,落点连成的那条线只是擦过它能碰到的最小的那条等高线,而不是穿过去,对应 η★ = ∇ᵀ∇ / ∇ᵀH∇ = 0.171。再往前,这一步又爬了出来: 又回到 0.64,而 落到 5.52 —— 比出发时还高。

训练就是把这件事反复做下去:在你当前所在的位置算梯度,走一步,再来一遍。跑起来,把这笔权衡的两头都拨一拨 ——步长和走多少步:

η = 0.100,走了 0 步后 L = 4.63

注意那道之字形。路径横穿山谷而不是顺着谷底走,原因正是第 04 节说过的那个 —— 而把它摆在那里的是碗的形状,不是步长。把 η 调小,同时压住了摆幅、也压慢了下降:12 步之后, 走到 L = 0.223,η = 0.10 走到 0.0606, 走到 0.0036。三者都还在这只碗给出的上限之下,所以小步长并不是更安全的那个,只是更慢的那个 —— 这里最快的是, 12 步走到 6.6×10⁻⁴。

再往上有一个硬天花板,而且它不是品味问题。沿每个曲率方向,一步会把误差乘以 1 − ηλ,只有在 |1 − ηλ| < 1 时误差才会缩小。把最硬那个方向和最软那个方向的这个倍数都画出来,拖动 η,找出它在哪里穿过 1:

η = 0.100:每步误差 × 0.40

在初始学习率下,硬方向每步把误差缩到 ×0.40,软方向只缩到 ×0.90,所以真正耗时间的是软方向。硬方向那条曲线恰好在 2/6 = 0.333 处到达 1,越过之后每一步都会让那个方向更糟: 时 12 步落到 9.92, 时落到 1.31×10⁶。在真实训练里,这就是跑了几百步之后损失变成 NaN —— 属于会大声报错、也容易诊断的那种失败。

安静的那种失败,是天花板本身在动。λ_max 是曲面的性质,所以一个在某个方向上很窄的碗,会把小 η 强加给所有方向,包括那些本来需要大步长的方向。把碗拉长,数一数步数:

κ = 6:21 步

既然每次都用该形状下最优的 η 来跑,这里量的就是纯粹的条件数效应,学习率已经调好了: 需要 7 步,κ = 6 需要 21 步, 需要 70 步 —— 与 κ 成正比地涨。对训练好的图像分类网络实测 Hessian 谱,最大特征值在数百量级,而绝大多数特征值挤在零附近(Ghorbani、Krishnan、Xiao,2019),所以真实训练面对的比值远不止 6 —— 这就是没人真的直接用朴素梯度下降的原因:动量、Adam、层归一化,归根到底都是在把碗变圆一点。

06

反向传播

这一节没有任何新的数学。它就是第 03 节用在一张图上,而且只走那个能让它便宜下来的方向。

下面是能说明问题的最小网络:z₁ = w₁x,接着 a = tanh(z₁),接着 z₂ = w₂a,最后 L = ½(z₂ − y)²。正着跑一遍,把各个值填好;然后从损失出发,让导数一个结点一个结点地往回走:

反向第 0 步,共 4 步

每一个反向步骤,都是乘上一个局部导数,取值就取在前向传播留在那里的数上。∂L/∂z₂ = z₂ − y 得到 0.834;再乘 w₂ 得到 ∂L/∂a = 1.334;再乘 tanh′ = 1 − a² 得到 ∂L/∂z₁ = 0.407;再乘 x 得到 ∂L/∂w₁ = 0.407。四次乘法,用不到任何超出第 01 节的微积分。

方向才是全部诀窍。链式法则正着倒着都成立,但一个网络只有一个损失,却有几百万个参数,这种不对称决定了一切。在同一张有六个权重的图上,把每个参数走一遍和所有参数只走一遍对比一下 —— 把滑块一格一格推上去,数一数两个方向各自要付多少:

正向:在计算图上走 0 遍

注意六次拖动换来了什么,而一次拖动换来了什么。正着走,你传播的是某一个输入的影响,所以每推一格就在主干上多铺一条路线,只点亮一个答案:有 n 个参数,就要遍历 n 次。倒着走,你传播的是某一个输出的灵敏度,一次遍历就把所有 ∂L/∂wᵢ 一次性交到你手上。对一个 70 亿参数的模型来说,这就是「一次反向」和「70 亿次正向」的差别。

反向模式换来的代价是内存:前向传播产生的每一个中间量,都必须一直活到反向传播走到它为止。设定层数,把全部留住和只留几个、其余重算对比一下:

4 层:留住 4 个,用检查点则是 4 个

到 时,朴素反向传播要留住 48 个张量,用检查点的只留 14 个 —— √n 个检查点,加上当前正在重放的那一段里的 √n 个激活,所以峰值是 2√n,代价是多花大约 30% 的计算量(Chen 等,2016)。正是这笔交易,让 batch size 由激活值内存决定。

07

全部就这些

六条规则、一个循环、四种破法。

  d/dx [c]      = 0            d/dx [eˣ]    = eˣ
  d/dx [xⁿ]     = n·xⁿ⁻¹       d/dx [ln x]  = 1/x
  d/dx [f + g]  = f′ + g′      d/dx [f(g)]  = f′(g)·g′

  loss = forward(w)            # 每个中间量都要留着
  g    = backward(loss)        # 一遍,拿到全部 ∂L/∂wᵢ
  w   -= eta * g               # eta < 2 / 最大曲率

全部都靠同一条性质:在一点附近,函数和它的切线吻合到比一阶更好的程度。破坏它的有四件事 —— 拐角、一串小增益的乘积、越过 2/λ_max 的步长、条件数很差的曲面 —— 而只有第三种会主动告诉你。