上一课把话说死了:「梯度告诉每个参数往哪挪」。 可我们那个小模型有 798,720 个参数 —— 那就是 798,720 个方向。 一个一个试?先把这笔账算清楚:按「一次前向 1 毫秒」, 量一轮梯度要 1,597,440 毫秒(26.6 分钟), 跑完第 9 课那 8,000 步是 147.9 天。 这一课讲的,就是把这件事变成 24 秒的那个办法。
先把那张账坐实。第 10 课那个山谷只有一维 —— 一个参数,一个人就能试出来。真实模型不是。 我们第 8 课数出来的模型是 798,720 个参数(798720.000000 —— 这个数在lib/block.ts 里能一行行加出来)。 要知道每个参数该往哪挪,最笨的办法是:把它轻轻推一下,看损失变多少。
一个参数推两次(加一点、减一点),就是 1,597,440 次前向。 一次前向按 1 毫秒算,量完一轮 1,597,440 毫秒; 第 9 课那份记录跑了 8,000 步 —— 乘起来 147.9 天。 而反向传播要的是 3 次前向的代价, 同样 8,000 步只要 24 秒,快 532,480 倍。
所以问题不是「有没有更聪明的试法」,而是:能不能一次就把所有参数的梯度全算出来?能。办法分三步:把函数拆成一张计算图、 用链式法则把导数一层层乘起来、 再从损失那边往回走一遍。下面这张图就是这个过程的全部。
L = (v·ReLU(w1x1 + w2x2 + b) + c − t)² 拆成 12 个节点,每条边记一个「局部导数」。一家连锁店想知道「全国 798,720 家门店, 每家的房租涨 1 块钱,总成本会涨多少」。笨办法是挨家挨户去试一遍。
这就是为什么它叫「反向传播」:账是从总账(损失)往明细(每个参数)方向推的。
页首那张图里的每个圆圈,都是一个大函数里的一次小运算。 乘法、加法、ReLU、平方 —— 每一个单独看都简单得不能再简单, 它们的导数更是小学水平。麻烦只在于它们套在一起。
计算图要做的事就一件:把「套在一起」摊开成「谁喂给谁」。 这样一来,每一步的导数只跟它自己的输入有关,不用去管外面套了多少层。 下面是我们这张图的全部 12 个节点 —— 每个数都能按计算器核对:
有了图,一件重要的事就成立了:前向时顺手把中间结果存下来, 反向时不用重算。图上每条边标的 ×数字,全都是前向时已经算出来的东西 (比如 w1 → z1 那条边上标的是 x1 = 2.000000)。 这就是为什么训练比推理吃显存:推理可以走一个丢一个, 训练必须把这一整张图(术语叫激活值)留着,等反向用。
.grad, 和参数本身一样大。798,720 个参数就是两份账。图上有了边,导数就有了着落。链式法则只有一句话: L 对某个参数的变化率,等于从 L 走到它的路上,每一个中转站贡献的「局部变化率」全部相乘。
从 L 出发到 w1,路上会经过 z2 → a1 → z1 三个站。 每个站贡献一个因子:
对照页首动画:小球在 w1 旁边那个琥珀色小牌子上写的正是 ∇ -24.000000。四个数字,一次乘法链,就是「w1 该往哪拧」的全部答案。
链式法则谁都会背,可它单独用起来是场灾难:每个参数都要把整条路重走一遍。 反向传播的两处「省」,把它从灾难变成一行代码。
现在把页首那张图只走反向那半段再看一遍。这次盯着 b 拖: 把第一层的偏置 b 从 3 拖到 −3,z1 从 1.000000 变成 -5.000000, ReLU 的门就关了。
上面全是推导。推导可以错得很隐蔽 —— 一个符号写反,训练照样跑,只是怎么都学不好。 所以每一次手写反向传播,工程上都要做一件事:用最笨的办法抽查一遍。
这件事之所以成立,是因为我们选了一张小到能手算的图。 真实模型里可没有「推一把」这个选项 —— 那就是第 5 站要算的账。 梯度校验只用来给实现做体检:跑几个样本,量量每个参数的量级对不对, 对上了就关掉它去训练。
把这笔账摊成一张表。前提说清楚:一次前向按 1 毫秒算 (对小模型是合理量级),反向的代价按「前向的 3 倍」估 (业界常用的粗略说法,实际取 2~3 倍)。
反过来说,这也解释了一个工程细节:反向传播既然靠前向存下的中间结果, 那么显存就是这条路的过路费。用不用得起大 batch、 能不能跑长序列,全看这份账本有多大(第 12、15 课会反复回到这里)。
框架把反向传播包装成了 loss.backward() 一行。 第一次看训练循环的人都会卡在那行 optimizer.zero_grad() 上: 为什么要把辛苦算出来的梯度清掉?
optimizer.zero_grad() 不是可有可无的仪式:PyTorch 的 .grad 默认累加, 不清零就会把好几十步的梯度叠在一起。 忘了写它的典型症状是:loss 前几步就炸成 NaN。.grad,所以清零是「把每个参数的账本都划掉」。反向传播不是新数学,是链式法则的工程版: 把复合函数拆成计算图,每条边记一个局部导数; 前向算值的时候顺手把中间结果存下来;然后从 ∂L/∂L = 1 出发往回走一遍, 每个节点把「上游的梯度 × 边上的局部导数」加到自己账上。 一次前向 + 一次反向,798,720 个参数的梯度全出来了 —— 换成最笨的「逐个扰动」要 1,597,440 次前向、147.9 天,反向传播只要 24 秒。 两个必须记住的工程细节:ReLU 把关了门,梯度就一个都传不进来(死区), 以及 .grad 是累加的,每步都要 zero_grad()。