会加减乘除,就能训会一个神经网络
用 0~9 的全部 100 道乘法,从头搭一个最小的神经网络, 亲手把它训到满分:先用只有加减乘除的笨办法真训对, 再讲透反向传播为什么快了几百倍。
这是「乘法表」系列的上篇;下篇《注意力究竟做了什么》 拿本篇训出的模型当对照物。
全文只追一条因果链:
题目 (a,b) → 20 格输入 x → 24 格 z → GELU 后的 h → 82 个分数 u → 概率 p → 罚分 L 参数稍微变化 → L 怎样变化 → 得到梯度 → 参数反向挪一步 → 下一轮概率改变
每出现一个名字,先指出它由哪些已知数算出、形状是多少、交给下一步什么;最后再把整条链倒着走一遍。模型能答对 100 道题,只说明它拟合了这 100 个训练样本,不自动证明它学会了对未见数字进行乘法推理。
传统做法:用一个 MLP 学会 0~9 的全部 100 道乘法
MLP(多层感知机)是最基础的那种神经网络: 输入一次性铺开,一层层往前算,中间没有任何分支。 今天它仍是深度学习的基本积木——每个 Transformer block 里的 FFN 就是 MLP。
00任务
输入两个个位数 a、b,输出 a×b。
模型里没有「乘法」这个概念——它只有一堆权重和加权求和。 乘法表得从 100 个例子里学出来。
挑 0~9 乘法表,是因为它小到能把每一根连线都画出来: 100 条数据、两千多个参数。下面每张图都从真模型上现算,没有示意图。
01模型结构与推理
三层:左边 20 维输入——两个数字各占 10 格,只有自己那一格是 1,
其余全 0(这叫 one-hot,独热编码);
中间 24 个隐藏单元;
右边 82 维输出——0×0=0 到 9×9=81,
乘法表的全部可能答案,谁的分数最高就答谁。
每根连线上挂一个数,叫权重。一个隐藏单元的值 =
20 个输入各乘自己那根线的权重、求和,再加偏置
(bias,每个单元自己的一个常数)。把 24 个单元一起写出来,
就是 z = W1·x + b1;再算一次到输出层,就是 u = W2·h + b2。
所以这个模型要训练的东西一共就四块:两个矩阵、两个偏置。 它们各有多少个数,加起来是多少:
W1 长什么样
矩阵乘法在这里退化成了「取两列,相加」。x 是 one-hot,20 项里 18 项乘的是 0。
所以:W1 第 a 列那 24 个数, 就是模型对数字 a 的全部表示。做一道题=把两个数的表示相加。
注意这条路径是静态的:a 走第 a 列,由它被放在哪一维决定,和它是几无关。
中间那道非线性
z 算完不能直接往下传,每一维还要单独过一次激活函数 (activation,作用在单个数上的非线性函数)。这里用 GELU (Gaussian Error Linear Unit)。公式就一行:
GELU(t) = 0.5 · t · ( 1 + tanh( 0.7978845608 · ( t + 0.044715 · t³ ) ) )
里面只有三样东西:
t³就是t · t · t;tanh是个现成的函数,作用是把任何数压进 −1 到 +1 之间: 数很大时它趋近 +1,很小(很负)时趋近 −1,0 附近近似原样。 于是括号外面那个1 + tanh(…)就在 0 到 2 之间, 乘上0.5就成了一个 0 到 1 的开关;0.7978845608和0.044715是两个写死的常数 (前者是√(2÷π)),不参与训练—— ③-1 数 2,554 个参数时它们不在里面。
连起来读:GELU(t) = t × 一个 0 到 1 的开关,
而这个开关的大小由 t 自己决定。t 越大开关越接近 1
(原样放行),t 越负开关越接近 0(压扁)。拖滑块看:
W2 长什么样
过完非线性得到 h,再乘第二个矩阵,得到 82 个候选答案的分数:
「退化成两列」是 one-hot 输入的特殊待遇,不是矩阵乘法的普遍规律—— h 不是 one-hot,24 项一个都不能省。
推理的完整代码
把上面几步连起来,一次推理就是这六行。左边的代码可以直接改, 右边是当前这一行算出来的数——单步走一遍,或者动手拆一下看它怎么坏。
末尾的判决是错的:参数还全是随机数,它只能瞎猜。 第 2 节把它训出来之后,翻回这一节,同一个位置就会变成 56。
这台机器的六个「为什么这么设计」 — 作者视角
01 节从头到尾路过了六个没解释的决定,把推理逐个摆出来:
为什么 one-hot,不直接把 7 喂进去?——模型只会乘和加。
直接喂数字,「7 比 3 大、挨着 8」这些大小关系会被硬塞进每个算式,
可乘法表里 7×8 和 6×8 的答案毫无「相邻」可言
⇒ one-hot 给每个数字独立的一列参数,互不牵连——数字成了「查表的键」,不是「量」
为什么要中间那层?——先看去掉非线性会怎样(乘法结合律,两行):
u = W2·( W1·x ) = ( W2·W1 )·x # 两张表乘好合成一张,两层塌成一层
⇒ 没有 GELU,中间层加了白加;先有那道闸门,中间层才存在——闸门买什么,下篇 FFN 节当场反证
为什么是 24 个隐藏单元?——试出来的旋钮,和学习率同类,没有神秘含义
为什么每个单元还要加个 b?——z = 两列之和 + b:没有 b,闸门拐点焊死在原点;
b 是每个闸门可学的门槛,把「压扁区」挪到对这道题有用的位置
为什么随机初始化,不全填 0?——全 0 时 24 个单元式子相同、输入相同
⇒ 前向输出相同 ⇒ 斜率也相同 ⇒ 挪完还相同 ⇒ 永远是 24 个复制品,等于只有 1 个单元
随机数唯一的作用:把它们错开,让分工有机会长出来
为什么 82 个输出「选一个」,不直接输出 56 这个数?——输出一个数,罚分只能用「差多少」:
56 猜成 55 罚 1、猜成 12 罚 44——可乘法表错 1 和错 44 一样是错;
更要命的是 softmax + −ln 那套「为训练编排」的性质全没了(02 章的核心)
⇒ 82 个候选各占一格,「对」才有明确的一格,−ln(py) 那招才使得出来
02训练
上一节结束时模型答错了:7 × 8 它答 62。这一节把它训对——
只改那 2,554 个数,一根线都不动。
难的不是这一道题,是 100 道题同时做对。
100 道题互相拉扯:7×8 想把 W1 第 7 列往这边拽,7×3 想往那边拽。
先钉住一个参数:100 道题怎样共同决定它往哪里走
挑 W1[k,7]:它是隐藏单元 k 接收“第一个乘数为 7”这格输入时使用的那一个权重。
因为 x 是 one-hot,只有 7×0…7×9 这 10 道题会用到它;其余 90 道对它的单题梯度严格为 0。
第 b 道题给它的责任: ∂L7×b/∂W1k,7 = dzk(7,b) × x7 = dzk(7,b) 100 道题平均后的梯度: ∂L/∂W1k,7 = [ dzk(7,0) + dzk(7,1) + … + dzk(7,9) ] ÷ 100 真正更新: W1k,7 ← W1k,7 − 学习率 × ∂L/∂W1k,7
“训练朝合理方向走”在这里没有比喻:10 道题若大多给同号梯度,就共同推动这个参数;若有正有负,就在求和时抵消。优化器听到的是总和,不知道哪道题“更有道理”。所谓可复用规律,就是同一参数在许多样本上反复收到相容的更新。
注意分母仍是 100,因为本文的 L 定义为 100 道题的平均;那 90 个零也属于这个平均。若改成只抽一批题,分母才换成 batch 大小。
①算一遍:得到 p
这一步第 1 节讲完了,只补最后那行 p = softmax(u)。
u 是模型给 82 个候选答案打的分,可正可负、没上下限,没法当「把握」读。
82 个候选是编了号的(第 0 个代表答案 0,……第 81 个代表 81),于是:
c = 随便哪一个候选答案的编号(0…81) # 占位的编号,说「对任意 c 成立」就是 82 个都成立 uc = 第 c 个候选拿到的分数 pc = 第 c 个候选的概率 pc = e^uc ÷ (e^u₀ + e^u₁ + … + e^u₈₁) y = 正确答案的编号 # 这道题 y = 56,因为 7×8=56 py = 模型给正确答案的概率 # 这道题 p₅₆ = 0.0041 —— 几乎没往它上面想
每个分数先取 e 的指数(一定为正),再除以所有 82 个指数的总和(于是加起来正好等于 1)。
82 个非负、和为 1 的数,就是 p。
②量一下差多少:loss
要让模型自己变好,先得有一个数说清「它现在有多差」。
直接用 py 有两处不趁手:习惯上要一个越小越好的罚分;
而且 py 从 0.50 掉到 0.49、和从 0.02 掉到 0.01,在它眼里都是「掉了 0.01」,
可后者是概率掉一半。取对数再取负号,正好同时解决:
ln x 回答一个问题:e 的几次方等于 x。
(e ≈ 2.718;ln = 以 e 为底的 log,叫自然对数——不是中学默认的以 10 为底。
本文数学式一律写 ln;但代码里都写 log,numpy、PyTorch 的 log 就是 ln。)
ln(1) = 0 因为 e 的 0 次方 = 1
ln(0.1) = −2.30 ln(0.01) = −4.61
概率每掉一个数量级,固定加 2.30
概率都不超过 1,所以 ln 全是负数或 0——前面加个负号,罚分才是正的。 这就是 −ln 里那个负号的来历。
一道题的罚分 L = −ln( py ) # 100% → 0 分;10% → 2.30;1% → 4.61;趋近 0 → 趋于无穷
100 道题合成一个数 L = ( L0×0 + L0×1 + … + L9×9 ) ÷ 100
这个「÷ 100」比看上去大。100 个要求没法一个个去调和; 一个数却只有高低,永远知道该往哪边走——开篇那个「互相拉扯」,解法就藏在这里。
③算出每个参数该往哪挪:梯度
这是唯一有难度的一步。整节只做一件事:把「L 该往哪边挪」拆成 2,554 个 「这一个数该调大还是调小」。
先问清楚:L 是谁的函数
要求导,先得知道对谁求。把 L 的来历倒着捋一遍:
L ← 100 道题的罚分平均 ← 每道题的 p # 由 W1 b1 W2 b2 和这道题的两个数算出来 ← 题目 (a, b) 和答案 a×b # 7×8=56 永远是 56,动不了 ← W1 b1 W2 b2 # ← 只有这些能动 W1 24 × 20 = 480 b1 24 = 24 W2 82 × 24 = 1,968 b2 82 = 82 ⇒ 合计 2,554 个能拧的数
题目和答案是给定的常数,x、z、h、u、p 是算出来的中间结果,不能直接改。
整条链上唯一能拧的旋钮,就是那 4 张表里的数。
所以 L 是一个吃 2,554 个数、吐 1 个数的函数。
把它完整写出来,就是第 1 节那 6 行
x = onehot2(a, b) # 题目摊成 20 个数:第 a 个和第 10+b 个是 1 z = matvec(W1, x, b1) # 24 个数 用掉 W1 的 48 个 + b1 全部 24 个 h = gelu(z) # 24 个数 0 个参数 u = matvec(W2, h, b2) # 82 个数 用掉 W2 全部 1,968 个 + b2 全部 82 个 p = softmax(u) # 82 个数 0 个参数 L = −ln( p[a×b] ) # ← 只有这一行是新的
前 5 行一个字没改,第 6 行换掉了:那边是 answer = argmax(p),
挑概率最大的当答案,这是推理;这边拿正确答案那一格的概率算罚分,这是训练。
同一个前向,用完之后一个用来答题,一个用来打分。
把 py 的定义代进去(除法在 ln 里变减法,而 ln(e^uy)
就是 uy),再把 100 道题平均起来,整个 L 完整写出来就是这三行:
L = 1/100 · Σa=0…9 Σb=0…9 [ −ua·b + ln( e^u₀ + e^u₁ + … + e^u₈₁ ) ]
其中 uc = Σk=0…23 W2c,k · GELU( W1k,a + W1k,10+b + b1k ) + b2c
其中 GELU(t) = 0.5·t·( 1 + tanh( 0.7978845608·( t + 0.044715·t³ ) ) )
Σ 就是「把下面那些加起来」:两个 Σ 套在一起 = a、b 各取 0~9,100 道题各算一遍
就这些,没有别的了。式子里的字母只有两类:
① W1 b1 W2 b2——那 2,554 个能拧的数;
② a、b——题目,常数。
剩下全是加、乘、tanh、e 的幂、ln。
「深度学习模型」这五个字底下,就是这么一个式子。训练要做的事一句话说完: 找一组 w₁…w₂₅₅₄,让它最小。
推一点点,量比值——斜率就是这么测出来的
2,554 个未知数想不出来,先退到只有一个未知数:那条大家都见过的抛物线
y = x²。在 x = 3 处把 x 推大 0.01,y 从 9.0000 涨到 9.0601——
涨的量 ÷ 推的量 = 6.01;推 0.001 得 6.001,推 0.0001 得 6.0001。
这个越推越稳的数 6,就是 y = x² 在 x = 3 处的斜率,
也叫导数。它只说一件事:在这儿把 x 推一点点,y 会涨这一点点的 6 倍。
回到 2,554 个未知数。办法一样,只要先把其余 2,553 个按住不动—— 它们不动,L 就只剩一条随这个参数变的一元曲线,照样推一点点、量斜率。
这就是「偏」字的全部含义:只推一个,其余按住。
这样量出来的斜率叫偏导数,记作 ∂L/∂w——读作「w 推一点点,L 涨多少倍」。
斜率为正就往小了调,为负就往大了调——反着走就对了。这一句就是训练的全部动作。
2,554 个都问一遍,挪一步——真的把它训对
对每个参数问一次得一个斜率,2,554 个各问一遍得 2,554 个斜率,排成一列 就叫梯度。它指向 L 上升最快的方向,所以往它的反方向走:
wi ← wi − 学习率 × gi # gi 是这个参数的偏导;学习率是自己定的旋钮,④ 再细说
方向藏在 g 自己的正负号里:g 为正,减去正数自动变小;g 为负,减去负数自动变大。 一行公式不用任何判断。剩下的只是重复:量一遍 2,554 个偏导 → 全体各挪一小步 → 再量一遍。 不需要链式法则,不需要反向传播,不需要矩阵。
先把记号焊死——上面是数学写法 wi、gi,
代码里它们叫别的名字。本页所有代码窗格只用下面这一套,不再换词:
数学写法 代码里的名字(lab.js 里就是这么写的) ──────────────────────────────────────────────────────────── w₁ … w₂₅₅₄ 那 2,554 个数 → ps = [W1, b1, W2, b2] # 四张表,摊平就是这 2,554 个(第 122 行) wi 现在的值 → p.d[i] # d = data,表里那个数本身 wi 的斜率 gi → p.g[i] # g = grad,跟 p.d 一样大的一块,专门装斜率 「2,554 个挨个过一遍」 → each((p, i) => …) # 就是套两层 for,写一次,下面反复用 L(100 道题的平均罚分) → 平均罚分() # lab.js 里叫 fullLoss() 学习率 → 代码里写死的那个 5
于是办法一整个写出来,一共 9 行:
为什么各挪各的,整体能一路降到 100 道全对?——因为因果是反的: 不是训练办法碰巧配上了 L,是 L 从头就是为这个训练办法编排的。 这才是深度学习的核心。
训练只有一招:沿斜率,各挪一小步。回头看 L 的每一处设计,全是在配这一招—— softmax 把分数变成能罚的概率;−ln 让错得越离谱罚得越重,且处处光滑、有底; 100 道取平均,把 100 个要求压成一个数。 「降到底正好 100 分」,上面亲手摇已经实测过了;「为什么一定降」,下面两行算给你。
上面那句「因果是反的」,两行算得出来。每个 gi 的定义都是
「其他参数不动,只动这一个」,可代码里是 2,554 个同时动——
凭什么它们不互相拆台?
把总变化算一算。每个参数挪了 Δwi = −学习率 × gi,
各自让 L 变 gi × Δwi。只要每一步都足够小,
总变化就是它们直接相加:
ΔL ≈ g₁·Δw₁ + g₂·Δw₂ + … + g₂₅₅₄·Δw₂₅₅₄ = −学习率 × ( g₁² + g₂² + … + g₂₅₅₄² ) = −学习率 × Σg²
看括号里:全是平方,每一项都 ≥ 0,整个和一定非负,前面带个负号 → ΔL 一定 ≤ 0,loss 一定降。
不会起瓢的原因就在这儿:每个参数贡献的都是 gi²,全是「帮忙」,谁也抵消不了谁。 换任何别的方向,这些平方项就会变成有正有负,真的可能互相抵消——这正是非要走负梯度方向的原因。
这条证明对 L 的要求只有两个字:光滑——处处有斜率。 我们的 L 由加、乘、tanh、e 的幂、ln 拼成,处处光滑,所以够格。 这不是巧合——上面说过:L 就是为配这条证明编排的。
办法一到此闭环:真的训对了。代价也看见了——每挪一步,全量数据重算 2,555 遍。
先算账:笨办法慢在哪
慢的根子是浪费:量第 1 个参数那一遍算出来的中间结果 z、h、u,
量第 2 个参数时全扔了、原样重算——同样的东西被推倒重来 2,554 次。
省下这份浪费,就是快的全部来源。
一个规则背到底:链式法则
反向传播要做到:整份数据只走 1 遍,2,554 个斜率一次全出来。靠的只有一个规则。 先看清依赖关系——就一条链:
真正要的梯度只有 4 个:∂L/∂W1、∂L/∂b1、∂L/∂W2、∂L/∂b2——只有它们是参数。
∂L/∂u、∂L/∂h、∂L/∂z 是路过的:算它们不是目的,是为了把误差从右边搬到左边,搬完就扔。 ∂L/∂x 根本不用算——x 是输入,改不了。
链式法则:A 通过 B 影响 C,那么「A 对 C 的影响」=「A 对 B 的影响」×「B 对 C 的影响」。 把 A 推一点点,B 跟着动,C 再跟着动——两段效应相乘。
照着那条链,从右往左套一遍:
∂L/∂u = 直接算 # 就是五行的第 1 行 ∂L/∂W2 = ∂L/∂u × ∂u/∂W2 # 上游 × 这一层的局部导数 ∂L/∂h = ∂L/∂u × ∂u/∂h ∂L/∂z = ∂L/∂h × ∂h/∂z # 换成新的上游,继续 ∂L/∂W1 = ∂L/∂z × ∂z/∂W1 ∂L/∂b1 = ∂L/∂z × ∂z/∂b1
看加粗的部分:每一行的第一个因子,都是上一行刚算出来的。 这就是「反向传播」四个字的全部含义:从 L 出发倒着走,每一步复用上一步的结果, 只算这一层新增的那个局部导数。不复用的话,每个参数都得从 L 一路推到底——那就是笨办法。
从偏导的定义,把五行推出来
全部推导只用一个式子——偏导的定义:
把 x 推大 ε,看比值: ( f(x+ε) − f(x) ) ÷ ε
导数 ∂f/∂x = 这个比值在 ε 缩到 0 时停住的那个数
拿前面那条抛物线套一遍:f(x) = x²,不代具体数字,整条按字母算:
( (x+ε)² − x² ) ÷ ε = ( x² + 2xε + ε² − x² ) ÷ ε = ( 2xε + ε² ) ÷ ε = 2x + ε
ε 缩到 0 ⇒ x² 的导数 = 2x # 一条对每个 x 都成立的新函数;代 x=3 得 6,正是前面量到的 6.01、6.001
照它办事就三步:把 w 换成 w+ε → 上下两式相减(不含 w 的项一字不差,全消掉) → 除以 ε、把 ε 缩到 0。五行全是这三步。
几个量同时被牵动时,各自贡献相加——理由:一个一个地挪, 第二个挪动落在「第一个已挪过」的位置上,带来的偏差是 ε×ε 级,除以 ε 后还剩一个 ε,缩到 0 就没了。
下面五行,推的是一道题。
用的 L 是那一道题的罚分 −uy + ln S——
没有 1/100,也没有那两个 Σ。为什么这样就够?还是照定义,两行:
( (f+g)(w+ε) − (f+g)(w) ) ÷ ε = ( f(w+ε)−f(w) )÷ε + ( g(w+ε)−g(w) )÷ε # 拆开就是
ε 缩到 0 ⇒ 和的导数 = 导数的和;同理 (c·f) 的导数 = c × f 的导数
⇒ ∂/∂w [ ( L0×0 + L0×1 + … + L9×9 ) ÷ 100 ] = ( ∂L0×0/∂w + … + ∂L9×9/∂w ) ÷ 100
所以:每道题各推各的,把 100 份斜率加起来、除以 100,就是真正那个 L 的斜率。
代码里就是 W2.g += …(累加)和 p.g[i] /= 100(取平均)那两行——
两个办法训练时都是 100 道全上,一道也没少。
先立两个记号,五行全靠它们:
W2i,j = W2 这张 82 行 × 24 列的表里,第 i 行、第 j 列的那一个数
ui = u 这排 82 个数里的第 i 个(hj、zi、xj 同理)
onehot(y) = 82 格里只有第 y 格是 1、其余全 0 的一排数 # y=正确答案的编号,①里立过
第 1 行 du = p − onehot(y) 82 个数
先备三件工具——都只是数「乘了几份」,唯一限定 底数为正(负数开方出鬼;这里全为正):
规矩1 x³·x² = (x·x·x)·(x·x) = x⁵ # 同底相乘=指数相加 ⇒ e^a·e^b = e^(a+b)
规矩2 (x³)² = (x·x·x)·(x·x·x) = x⁶ # 幂的幂=指数相乘 ⇒ (e^p)ⁿ = e^(np)
e ≝ ( 1 + 1/n )^n n=1 → 2 n=10 → 2.59374… n=1000 → 2.71692… ⇒ e = 2.71828…
ln 就是 e^ 的反查表:e^y = w 这一张表,从 y 查 w 叫 e^,从 w 查 y 叫 ln。
ln e = 1 # e¹ = e,反查回去就是 1
两条规矩翻到 ln 这边(记 a = e^p、b = e^q,即 p = ln a、q = ln b):
ln(a/b) = ln a − ln b # a/b = e^p ÷ e^q = e^(p−q) (规矩1)
ln(xⁿ) = n·ln x # xⁿ = (e^p)ⁿ = e^(np) (规矩2)
引理一 ln w 的导数 = 1/w 照定义办,三步到底: ( ln(w+ε) − ln w ) ÷ ε = ln( (w+ε)/w ) ÷ ε = ln( 1 + ε/w ) ÷ ε # 除法变减法 记 n = w/ε,即 ε = w/n(ε 缩到 0 ⇔ n 变大): = (n/w)·ln( 1 + 1/n ) = (1/w) · ln( (1+1/n)ⁿ ) # 系数进指数 n 变大时 (1+1/n)ⁿ 挤向 e ⇒ 那个 ln 挤向 ln e = 1 # 这一步借了「表连着查」:被查的数挪一点点,查出来的也只挪一点点 ⇒ ln w 的导数 = 1/w # 验 w=3、ε=0.001:n=3000,(1+1/n)ⁿ=2.71783,ln 它=0.99983,整个比值=0.333278 → 1/3 ✓
反过来那半张表白送——同一段台阶,横竖对调,比值就上下颠倒:
引理二 e^t 的导数 = e^t # t 是任意一点,后面代成 uc 同一段台阶:t → t+δ w = e^t → w+ε # 即 t = ln w、t+δ = ln(w+ε);一头缩到 0 另一头跟着缩 ln 那边(推 w 看 t):( ln(w+ε) − ln w ) ÷ ε = δ ÷ ε → 1/w # 引理一 e^ 那边(推 t 看 w):( e^(t+δ) − e^t ) ÷ δ = ε ÷ δ = 1 ÷ (δ÷ε) # 同一对数,上下颠倒 ⇒ e^t 的导数 = w = e^t # 验 t=ln3、δ=0.001:ε÷δ=3.001500 → 3 ✓
要证:∂L/∂uc = pc − onehot(y)c L(u) = −uy + ln S,其中 S = e^u₀ + … + e^u₈₁ # 「L 是谁的函数」里推过的形式 照定义办:把 uc 换成 uc+ε(其余 81 个不动),上下相减,除以 ε。 第一项 −uy,上下相减: c = y:( −(uy+ε) ) − ( −uy ) = −ε ÷ ε = −1 c ≠ y:式里没有 uc,推它前后一字不差 ⇒ 差 = 0 ÷ ε = 0 第二项 ln S——先看 S 的差,82 项逐项相减,只有 e^uc 那一项变了: ΔS = e^(uc+ε) − e^uc ΔS ÷ ε → e^uc # 这正是引理二的比值,代 t = uc;ε 缩到 0 时 ΔS 跟着缩到 0 再看 ln 的差——拆成两个比值相乘(就是链式法则那一招): ( ln(S+ΔS) − ln S ) ÷ ε = ( ln(S+ΔS) − ln S ) ÷ ΔS × ΔS ÷ ε ε 缩到 0:前一个停在 1/S(引理一,把 w 换成 S),后一个停在 e^uc ⇒ ∂(ln S)/∂uc = (1/S) × e^uc = e^uc/S = pc # 这正是 softmax 的定义 两项相加: c = y:∂L/∂uy = py − 1 c ≠ y:∂L/∂uc = pc 82 个合成向量:du = p − onehot(y) 82 个数 ✓ 证毕
第 2 行 dW2 = du ⊗ h , db2 = du 82×24 + 82
要证:∂L/∂W2i,j = dui × hj
先把 u 的第 i 行摆全(01 节那次矩阵乘法的其中一行,24 次乘法加 1 个偏置):
ui = W2i,0·h0 + W2i,1·h1 + … + W2i,j·hj + … + W2i,23·h23 + b2i
把 W2i,j 换成 W2i,j+ε;h 由 W1 和 x 算出,纹丝不动。
u 差多少——82 行逐行上下相减(下面的 … 就是上头那些不含 W2i,j 的项):
第 i 行:( … + (W2i,j+ε)·hj + … ) − ( … + W2i,j·hj + … ) = ε·hj # 其余各项一字不差,全消掉
其余 81 行:式里没有 W2i,j ⇒ 差 = 0
L 差多少:整个 u 里只有 ui 挪了,挪量 ε·hj
第 1 行已证 dui =「ui 每挪 1,L 挪多少」:
ΔL = dui × ε·hj
÷ ε,ε 缩到 0 ⇒ ∂L/∂W2i,j = dui·hj 证毕
1,968 格摆回表:第 i 行第 j 列放 dui·hj——记作 dW2 = du ⊗ h(外积) 82×24 ✓
b2i 同法:Δui = ε ⇒ ΔL = dui·ε ⇒ db2 = du
第 3 行 dh = W2ᵀ · du 24 个数
要证:∂L/∂hj = Σi dui·W2i,j——写成矩阵,就是 W2ᵀ·du 把 hj 换成 hj+ε。u 差多少——82 行逐行上下相减: u₀ :( … + W20,j·(hj+ε) + … ) − ( … + W20,j·hj + … ) = W20,j·ε u₁ :同法 = W21,j·ε ⋮ # 这次每一行都含 hj——82 个全动了 u₈₁:同法 = W281,j·ε L 差多少:82 个 u 同时各挪一点,各自贡献相加(定义那条): ΔL = du₀·W20,j·ε + du₁·W21,j·ε + … + du₈₁·W281,j·ε ÷ ε,ε 缩到 0 ⇒ ∂L/∂hj = du₀·W20,j + … + du₈₁·W281,j = 把 W2 的第 j 列和 du 逐位相乘再相加 证毕 24 个 j 各用一列;可矩阵乘法拿的从来是「行」⇒ 把 W2 翻个身,行列对调:dh = W2ᵀ·du 24 ✓ # 转置不是规定,是「hj 同时喂给 82 个 u」逼出来的
第 4 行 dz = dh ⊙ GELU′(z) 24 个数
要证:∂L/∂zj = dhj × GELU′(zj) 把 zj 换成 zj+ε。h 差多少:GELU 逐维各算各的,只有 hj 动: Δhj = GELU(zj+ε) − GELU(zj) = GELU′(zj)·ε # GELU′ 的定义就是这个比值——那条曲线的斜率 L 差多少:只有 hj 挪了:ΔL = dhj × GELU′(zj)·ε ÷ ε,ε 缩到 0 ⇒ ∂L/∂zj = dhj·GELU′(zj) 证毕 24 维各管各的 ⇒ dz = dh ⊙ GELU′(z)(⊙=逐位乘) 24 ✓ # zj 很负时斜率 ≈ 0:它当时没出力,责任也传不进去
第 5 行 dW1 = dz ⊗ x , db1 = dz 24×20 + 24
要证:∂L/∂W1i,j = dzi × xj——和第 2 行一模一样的推法 把 W1i,j 换成 W1i,j+ε。z 逐行上下相减: 第 i 行:( … + (W1i,j+ε)·xj + … ) − ( … + W1i,j·xj + … ) = ε·xj 其余 23 行:不含 W1i,j ⇒ 差 = 0 L 差多少:只有 zi 挪了 ε·xj ⇒ ΔL = dzi × ε·xj ÷ ε,ε 缩到 0 ⇒ ∂L/∂W1i,j = dzi·xj ⇒ dW1 = dz ⊗ x 24×20 ✓ 证毕 b1 同法:Δzi = ε ⇒ db1 = dz 附赠:x 是 one-hot,20 格里 18 格是 0 xj = 0 ⇒ 那一列每格 = dzi·0 = 0——整列梯度为零 ⇒ 一道 a×b 只改第 a、第 10+b 两列——第 1 节「连线焊死」,证毕于此
加上 db1 = dz,到这儿
2,554 个梯度全齐了——一遍前向、一遍反向。
回头看「亲手摇一遍」那节的账:笨办法量满 2,554 个偏导要把 100 道题重算 2,554 遍; 这五行做同一件事,只走 1 遍。它们算出来的是同一批数。
装进循环:省下的浪费,就是一个变量
和亲手摇同一个起点、同一个挪法,只是斜率换成这五行来算。 代码里多两个名字,正好就是「先算账」那节说的那份被扔掉的中间结果:
记号约定 · 第二段(第一段在「亲手摇一遍」那节):
五行里的写法 代码里的名字
────────────────────────────────────────────────────────
前向那 5 行 fwd(a, b) 算完把中间结果打包吐出来
那一包中间结果 act = { z, h, u, p }
反向那 5 行 bwd(act)
dW2 db2 dW1 db1(要的) W2.g b2.g W1.g b1.g # 不另开变量,直接落进 .g
du dh dz(路过的) 临时变量,用完就扔 # 链式法则那节说过:搬运工,不是目的
正确答案编号 y a×b
为什么是 W2.g += … 而不是 dW2 = …?
因为 L 是 100 道题的平均:每道题算出自己那份 dW2,累加到同一块 W2.g,
100 道跑完再除以 100。那个加号就是「取平均」的前半步。
五行要用的中间结果逐个都在 act 里:第 1 行取 act.p、第 2 行取 act.h、
第 4 行取 act.z(u 已被 softmax 用掉,反向用不着)。
「省下那份浪费」在代码里,就是 act 这一个变量。
两拍到此结束:两条路算出的是同一批斜率,差别只有快慢。 四件事还剩最后一件——挪多远。
④挪一小步:学习率和 Adam
梯度只说了往哪边走,没说走多远。 那个「多远」就是学习率——两拍的代码里都是它,笨办法那 9 行写死 5,反向传播那 13 行也是 5。 太小走不动,太大会冲过头(第一拍那条证明里的「只要每一步都足够小」,就是它的代价)。 同一个模型、三个学习率,各训一遍看:
一个学习率要管 2,554 个参数,可它们的梯度差着几个数量级。Adam 补的就是这个: 攒动量抵消抖动,再按各方向自己的历史幅度归一化—— 相当于每个参数有了自己的学习率。本质还是沿负梯度挪一小步,只是步长不再一刀切。
梯度下降保证什么,不保证什么
ΔL ≤ 0 只说了「下一步更低」,没说会走到最低点。
梯度下降从头到尾只用当前这一点的信息。那它停在哪儿?
换个随机起点重训一遍就知道了——要是最低处只有一个,从哪儿出发都该滑到同一个底:
成绩一样,参数完全不同。「让 100 道题全对」的参数组合多得是, 梯度下降只是就近找了一个。它从来不是在找「那个最优解」,是在找「一个够用的解」。
每步不必算满 100 道:batch
到这里为止,每一步都把 100 道题全算了一遍——L 的定义就是那个平均,一步没偷工: 笨办法每量一个参数重算一遍 100 道,反向传播每步也把 100 道各跑一遍前向和反向。
可这很贵。而那个平均抽几条就能估出来—— 估得糙一点,方向大体还在。抽出来的一小撮叫一个 batch,条数叫 batch size。
最终版:上面那段循环,只改两处
不是新东西。「装进循环」那节你亲手走过的循环,把 Adam 和 batch 装上去, 就是所有人真在用的训练代码。逐行比,只有两行变了(开篇那张图的四件事 ①②③④,标在右边):
for (let step = 1; ; step++) {
each((p, i) => p.g[i] = 0) # 斜率清零
for (const [a, b] of 抽出来的这一批) { ← 改动 1:原来是 100 道全上,现在抽一批
const act = fwd(a, b) # ① 算一遍
loss = −ln act.p[a×b] # ② 量一下差多少——只是拿来看的,训练不用它
bwd(act) # ③ 那五行原样折成一句,斜率累加进 p.g
}
each((p, i) => p.g[i] /= 这一批的条数) # 取平均
adam(ps, 0.02, step) ← 改动 2:原来是 p.d[i] −= 5 × p.g[i]
}
② 那行 loss,整段训练不需要它。
五行里从头到尾没出现过 L——第 1 行 du = act.p − onehot(a×b) 直接就是斜率的起点。
源码就是证据:lab.js 算全部梯度的那个函数(gradOnly(),182–191 行)
里只有一行 this.bwd(this.fwd(a, b));——连返回值都没接,
bwd 顺手算的 loss 算完就扔,2,554 个斜率照样全齐。
真正驱动训练的从来只有那 2,554 个斜率。
训给你看
这个模型跑在你自己的浏览器里,没有服务器,也没有预训练权重。 先点「训练 1 步」看 loss 动一下,再点「一直训到全做对」,约 100 步收敛。
训完往上翻回第 1 节那张图。结构一个格子都没动,只是 2,554 个数换了一遍——
判决从 62 ✗ 变成 56 ✓,把握 100%。
W1 第 7 列那些数不再是噪声,而是模型自己长出来的、对数字 7 的表示。
而 7 走 W1 的第 7 列,是它被放在第 7 维决定的,跟它是几毫无关系。 换任何数字摆在第一个位置,走的都是同一批连线。这条路径是出厂焊死的。
下篇:连线焊死会卡在哪、transformer 多出来的那一步到底在干什么—— 《注意力究竟做了什么》。本篇训到满分的这台 MLP,就是那边的对照物。
控制台可以直接玩:inside.mlp 就是这台模型,inside.redraw() 刷新所有图。