RockDesk

会加减乘除,就能训会一个神经网络

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任务

输入两个个位数 ab,输出 a×b

模型里没有「乘法」这个概念——它只有一堆权重和加权求和。 乘法表得从 100 个例子里学出来

挑 0~9 乘法表,是因为它小到能把每一根连线都画出来: 100 条数据、两千多个参数。下面每张图都从真模型上现算,没有示意图。


01模型结构与推理

三层:左边 20 维输入——两个数字各占 10 格,只有自己那一格是 1, 其余全 0(这叫 one-hot,独热编码); 中间 24 个隐藏单元右边 82 维输出——0×0=09×9=81, 乘法表的全部可能答案,谁的分数最高就答谁。

模型结构 — 点中间那列的圆点换一个隐藏单元 数字都是这道题此刻的真实值

每根连线上挂一个数,叫权重。一个隐藏单元的值 = 20 个输入各乘自己那根线的权重、求和,再加偏置 (bias,每个单元自己的一个常数)。把 24 个单元一起写出来, 就是 z = W1·x + b1;再算一次到输出层,就是 u = W2·h + b2

所以这个模型要训练的东西一共就四块:两个矩阵、两个偏置。 它们各有多少个数,加起来是多少:

2,554 个参数」是怎么加出来的 这就是模型的全部家当

W1 长什么样

x × W1 + b1 = z — 点任意一行,下面摊开那一行的算式 高亮的两列=x 里那两个 1

矩阵乘法在这里退化成了「取两列,相加」。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³ ) ) )

里面只有三样东西:

连起来读:GELU(t) = t × 一个 0 到 1 的开关, 而这个开关的大小由 t 自己决定。t 越大开关越接近 1 (原样放行),t 越负开关越接近 0(压扁)。拖滑块看:

GELU — 拖滑块看它把一个数变成什么 虚线=恒等

W2 长什么样

过完非线性得到 h,再乘第二个矩阵,得到 82 个候选答案的分数:

h × W2 + b2 = u — 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 的几次方等于 xe ≈ 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 个能拧的数; ab——题目,常数。 剩下全是加、乘、tanh、e 的幂、ln

「深度学习模型」这五个字底下,就是这么一个式子。训练要做的事一句话说完: 找一组 w₁…w₂₅₅₄,让它最小。

第一拍 凭直觉的笨办法 —— 会加减乘除,就能把 2,554 个参数真训对

推一点点,量比值——斜率就是这么测出来的

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 涨多少倍」。

切一刀:只动一个参数,loss 怎么变 曲线每个点都把 100 道题重跑一遍;白点是当前值,橙线是那一点的切线

斜率为正就往小了调,为负就往大了调——反着走就对了。这一句就是训练的全部动作。

2,554 个都问一遍,挪一步——真的把它训对

对每个参数问一次得一个斜率,2,554 个各问一遍得 2,554 个斜率,排成一列 就叫梯度。它指向 L 上升最快的方向,所以往它的反方向走:

wi ← wi − 学习率 × gi      # gi 是这个参数的偏导;学习率是自己定的旋钮,④ 再细说

方向藏在 g 自己的正负号里:g 为正,减去正数自动变小;g 为负,减去负数自动变大。 一行公式不用任何判断。剩下的只是重复:量一遍 2,554 个偏导 → 全体各挪一小步 → 再量一遍。 不需要链式法则,不需要反向传播,不需要矩阵。

先把记号焊死——上面是数学写法 wigi, 代码里它们叫别的名字。本页所有代码窗格只用下面这一套,不再换词:

数学写法                    代码里的名字(lab.js 里就是这么写的)
────────────────────────────────────────────────────────────
w₁ … w₂₅₅₄ 那 2,554 个数  →  ps = [W1, b1, W2, b2]      # 四张表,摊平就是这 2,554 个(第 122 行)
wi 现在的值               →  p.d[i]                     # d = data,表里那个数本身
wi 的斜率 gip.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 遍。

第二拍 反向传播 —— 同样那批斜率,一遍全出来

先算账:笨办法慢在哪

办法一的账单 — 从 99 乘法表放大到手写数字 左列当场实测,右列=左列 × 纯算术倍数

慢的根子是浪费:量第 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,jdui·hj 证毕
  1,968 格摆回表:第 i 行第 j 列放 dui·hj——记作 dW2 = du ⊗ h(外积) 82×24 ✓
  b2i 同法:Δui = ε ⇒ ΔL = dui·ε ⇒ db2 = du
三步各验一遍 — 公式 vs 真推一下 换行号列号、换题目,结论都成立

第 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.zu 已被 softmax 用掉,反向用不着)。 「省下那份浪费」在代码里,就是 act 这一个变量。

反向传播的训练 — 亲手走,或一口气到全对 每步全量只走 2 遍(1 遍算斜率、1 遍复核 L)——办法一是 2,555 遍
乘法账本 — 反向传播一步 vs 办法一一步 全按当前模型的形状数:24×20、82×24

两拍到此结束:两条路算出的是同一批斜率,差别只有快慢。 四件事还剩最后一件——挪多远。

挪一小步:学习率和 Adam

梯度只说了往哪边走,没说走多远。 那个「多远」就是学习率——两拍的代码里都是它,笨办法那 9 行写死 5,反向传播那 13 行也是 5。 太小走不动,太大会冲过头(第一拍那条证明里的「只要每一步都足够小」,就是它的代价)。 同一个模型、三个学习率,各训一遍看:

同一个模型,三种学习率 现场训,约十几秒

一个学习率要管 2,554 个参数,可它们的梯度差着几个数量级。Adam 补的就是这个: 攒动量抵消抖动,再按各方向自己的历史幅度归一化—— 相当于每个参数有了自己的学习率。本质还是沿负梯度挪一小步,只是步长不再一刀切。

梯度下降保证什么,不保证什么

ΔL ≤ 0 只说了「下一步更低」,没说会走到最低点。 梯度下降从头到尾只用当前这一点的信息。那它停在哪儿? 换个随机起点重训一遍就知道了——要是最低处只有一个,从哪儿出发都该滑到同一个底:

三个不同的随机起点,各训一遍 现场训 3 个模型,比它们最后的 2,554 个数

成绩一样,参数完全不同。「让 100 道题全对」的参数组合多得是, 梯度下降只是就近找了一个。它从来不是在找「那个最优解」,是在找「一个够用的解」。

每步不必算满 100 道:batch

到这里为止,每一步都把 100 道题全算了一遍——L 的定义就是那个平均,一步没偷工: 笨办法每量一个参数重算一遍 100 道,反向传播每步也把 100 道各跑一遍前向和反向。

可这很贵。而那个平均抽几条就能估出来—— 估得糙一点,方向大体还在。抽出来的一小撮叫一个 batch,条数叫 batch size

每步 1 条 / 8 条 / 100 条,同时训三个 现场训,约十几秒

最终版:上面那段循环,只改两处

不是新东西。「装进循环」那节你亲手走过的循环,把 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() 刷新所有图。