高中数学就够了:把梯度下降手算篇的每一步拆到初高中知识点

《三个词的词表,十二个参数》的数学陪读篇。把原文用到的全部数学拆成七块积木——下标记号、乘加账单、e 和 ln、softmax、斜率、多旋钮下山、链式法则——每块只用初高中知识搭起来,配可按计算器复现的小例子,装完立刻回扣原文对应公式。结尾附逐符号拆解表和自检清单。

《三个词的词表,十二个参数》把梯度下降压到了 12 个参数,但如果 \partialln\ln、链式法则这些符号本身就发虚,那篇再小也硌牙。这篇只做一件事:把原文用到的全部数学拆成七块积木,每一块都从初高中知识搭起——你需要的前置只有:负数、百分比、平均数、平方与开方、y=kxy=kx 的斜率,以及高中的指数、对数和”导数=切线斜率”这个印象。每装好一块,立刻回到原文把对应的那一步亲手按一遍计算器。

先看地图:原文的数学,一共就七块积木

原文分三段:前向(算损失)、反向(算梯度)、更新(改参数)。每段卡人的地方,对应的积木如下:

积木一句话卡的是原文哪一步
⓪ 记号下标、\sum、one-hot所有公式的”读法”
① 乘加账单点乘、矩阵乘法前向②:z=xWz = xW
② e 和 ln一对互逆按钮前向③⑤:eze^{z}lnp-\ln p
③ softmax先按 e 键再算占比前向③:分数变概率
④ 斜率与割线导数、有限差分反向 2.3:数值梯度检验
⑤ 多旋钮下山偏导数、梯度、学习率更新 3.1:θηg\theta - \eta g
⑥ 链式法则影响的传递与分账反向 2.1/2.2:pyp-y 和外积
⑦ 滑动平均与绝对值AdamW 的全部数学更新 3.2:m/vm/\sqrt{v}

本文所有算例的数字都先用脚本验算过(和原文同一条纪律)。下面逐块装。

⓪ 记号扫盲:先能”读出声”

  • 下标是门牌号W02W_{02} 读作”W 的第 0 行第 2 列”——WW 是张表格,下标告诉你去哪个格子取数。zjz_j 里的 jj 是”第 j 个”,jj 可以是 0、1、2。(编号从 0 开始是程序员习惯,数学上无深意。)
  • \sum 是”加总”kezk\sum_k e^{z_k} 读作”把每个 ezke^{z_k} 加起来”,就是 ez0+ez1+ez2e^{z_0} + e^{z_1} + e^{z_2},没有更多内容。
  • one-hot 是”答题卡”y=[0,0,1]y = [0, 0, 1] 表示正确答案是第 2 个词(),对的位置涂 1,其他涂 0。
  • 10310^{-3} 是科学计数法103=0.00110^{-3} = 0.001。学习率 η=103\eta = 10^{-3} 就是 0.001。

回扣原文:现在能读出 LW02\frac{\partial L}{\partial W_{02}} 的字面意思了——“损失 LL 对表格 WW 第 0 行第 2 列那个数的(偏)导数”。\partial 是什么,第⑤块积木装。

① 乘加账单:点乘和矩阵乘法只是”数量×单价再加总”

买 2 个苹果(3 元/个)、1 个梨(2 元/个),总价:

2×3+1×2=82 \times 3 + 1 \times 2 = 8

这个”对应相乘再加总”就是点乘。两个数字清单 [2,1][2,1][3,2][3,2] 点乘得 8。全部线性代数的起点就这一个动作(想多装一点见线性代数与 AI 模型 vol.1)。

矩阵乘法 = 一次算好几份账单。原文的 z=xWz = xWx=[0.3,0.1]x = [0.3, -0.1] 是清单,WW 的每一是一份”单价表”,三列就是三份账单:

z0=0.3×0.5+(0.1)×0.2=0.150.02=0.13z_0 = 0.3 \times 0.5 + (-0.1) \times 0.2 = 0.15 - 0.02 = 0.13

回扣原文:前向②的 z=[0.13,0.13,0.08]z = [0.13, -0.13, 0.08],就是三笔这样的账。请现在亲手把 z1=0.3×(0.3)+(0.1)×0.4z_1 = 0.3 \times (-0.3) + (-0.1) \times 0.4 按出来,确认等于 0.13-0.13。“查行”(x=E[1]x = E[1])更简单:去表格 EE 里取出第 1 行,抄下来,没有任何计算。

② e 和 ln:一对互逆的按钮

计算器上有个 exe^x 键和一个 ln\ln 键。e2.71828e \approx 2.71828 是个常数(和 π\pi 一样是自然界的固定数),这两个键互为逆操作:ln(ex)=x\ln(e^x) = x,就像平方和开方互逆。高中指数对数的全部性质里,本文只用四条:

性质数字验证(可按计算器)
exe^x 永远是正数e2=0.1353>0e^{-2} = 0.1353 > 0,再负也是正
exe^x 单调递增(保排名)e0.13=1.1388>e0.08=1.0833>e0.13=0.8781e^{0.13}=1.1388 > e^{0.08}=1.0833 > e^{-0.13}=0.8781,排名不变
e0=1e^0 = 1
lnab=lnalnb\ln\frac{a}{b} = \ln a - \ln blnex=x\ln e^x = x除法变减法,这条待会儿推导要用

为什么损失是 lnp-\ln p:概率 pp 在 0 到 1 之间,lnp\ln p 是负数,加个负号变正。pp 越接近 1,lnp-\ln p 越接近 0(答对没惩罚);pp 越小,lnp-\ln p 爆炸(答错重罚)。按一下:ln(0.3494)=1.0515-\ln(0.3494) = 1.0515——这就是原文前向⑤那个损失值的全部来历(更完整的动机链在交叉熵为什么用 log)。

③ softmax:先按 e 键,再算占比

把分数 [1,2,3][1, 2, 3] 变成”总和为 1 的概率”,最朴素的想法是直接算占比:[1,2,3]/6=[0.1667,0.3333,0.5][1,2,3]/6 = [0.1667, 0.3333, 0.5]。为什么模型不这么干?因为分数可能是负数——[1,2,3][1, -2, 3] 直接算占比会出现”负概率”,荒谬。

softmax 的解法:先给每个分数按一下 exe^x 键(负数也变正,排名不变——积木②的前两条性质),再算占比:

[1,2,3]ex[2.718,7.389,20.086]占比[0.0900,0.2447,0.6652][1,2,3] \xrightarrow{e^x} [2.718, 7.389, 20.086] \xrightarrow{\text{占比}} [0.0900, 0.2447, 0.6652]

对比直接占比 [0.17,0.33,0.5][0.17, 0.33, 0.5]:softmax 把差距拉大了(3 分的从 50% 涨到 67%)——exe^x 增长快,领先者被放大。这是它的性格,不是 bug。

回扣原文:前向③就是这两步——z=[0.13,0.13,0.08]z=[0.13,-0.13,0.08]exe^x 键得 [1.1388,0.8781,1.0833][1.1388, 0.8781, 1.0833],加总 Z=3.1002Z=3.1002,各自除以 ZZ1.1388/3.1002=0.36731.1388 / 3.1002 = 0.3673。请把剩下两个占比也按出来。

④ 斜率与割线:导数不过是”陡不陡”

高中导数的核心图像:导数 = 切线斜率 = 这一点附近”输入动一点,输出动几倍”f(x)=x2f(x) = x^2x=3x=3 处导数是 f(x)=2x=6f'(x) = 2x = 6xx 挪 0.001,ff 大约挪 0.006。

不会求导公式也能”量”出斜率——取两个很近的点连线(割线):

f(3.001)f(2.999)0.002=9.0060018.9940010.002=6.000000\frac{f(3.001) - f(2.999)}{0.002} = \frac{9.006001 - 8.994001}{0.002} = 6.000000

回扣原文:反向 2.3 的”数值梯度检验”公式 L(θ+ε)L(θε)2ε\frac{L(\theta+\varepsilon) - L(\theta-\varepsilon)}{2\varepsilon},和上面这条割线是同一个式子ε=0.001\varepsilon = 0.001 换成 10510^{-5} 而已)。所谓”生产级的 gradcheck”,就是用高中的割线斜率去验证公式算出的切线斜率。神秘感应该在这里死掉一半。

⑤ 多旋钮下山:偏导数、梯度、学习率

偏导数(\partial:函数有多个输入时,“固定其他,只动一个”算出的斜率。f(x,y)=x2+3yf(x, y) = x^2 + 3y,把 yy 当常数,对 xx 的偏导 fx=2x\frac{\partial f}{\partial x} = 2x;把 xx 当常数,fy=3\frac{\partial f}{\partial y} = 3没有新数学,只是”轮流单独考察每个旋钮”\partial 这个符号读”偏”,作用和 dd 一样,只是提醒你屋里还有别的变量被按住了。

梯度:把每个旋钮的偏导数排成一排的清单。原文模型有 12 个旋钮,梯度就是 12 个斜率。

梯度下降:每个旋钮朝自己斜率的反方向拧一点(斜率为正→函数在涨→往回拧)。拧多少由学习率 η\eta 控制:θθηg\theta \leftarrow \theta - \eta \cdot g。用 f(x)=x2f(x)=x^2x=3x=3 出发(η=0.1\eta = 0.1)亲手走三步:

斜率 g=2xg = 2xx=x0.1gx = x - 0.1gf(x)f(x)
16.002.40005.7600
24.801.92003.6864
33.841.53602.3593

ff 从 9 一路降下来,而且越接近谷底斜率越小、步子自动变小——不需要任何智能,纯机械。(为什么”沿负梯度”是局部最快下降方向,严格版见泰勒展开和大模型。)

回扣原文:3.1 节的 SGD 表格就是这张表的 12 旋钮版;loss 从 1.0515 → 0.8894 → … → 0.0125 与这里 9 → 5.76 → 3.69 是同一件事。

⑥ 链式法则:影响的传递,加一条”多路分账”

这是原文反向传播一节唯一真正的”新”数学,分两步装。

第一步:串联传递(高中链式法则)。汇率类比:1 元 = 0.14 美元,1 美元 = 150 日元,那么 1 元 = 0.14×150=210.14 \times 150 = 21 日元——影响沿链条相乘。函数版:f(x)=(2x+1)2f(x) = (2x+1)^2,外层”平方”的斜率是 2×(2x+1)2 \times (2x+1),内层"2x+12x+1"的斜率是 2,在 x=1x=1 处总斜率 =2×3×2=12= 2 \times 3 \times 2 = 12(割线验证:f(1.001)f(0.999)0.002=12.0000\frac{f(1.001)-f(0.999)}{0.002} = 12.0000 ✓)。

还需要高中导数公式表里的三条(都可用积木④的割线自行验证,我验过后两条:割线值 0.500000 和 2.71828):

(x2)=2x(lnx)=1x(ex)=ex(x^2)' = 2x \qquad (\ln x)' = \frac{1}{x} \qquad (e^x)' = e^x

第二步:多路分账(超出高中的唯一新规则)。如果一个旋钮通过好几条路影响结果,每条路各自”沿链相乘”,然后把几条路加起来。就这一句。

现在慢速重推原文最重要的公式:L/z=py\partial L/\partial z = p - y

原文 2.1 节两行推完,这里逐步标注用了哪块积木:

  1. L=lneztZL = -\ln\frac{e^{z_t}}{Z},用积木②”除法变减法”:L=(lneztlnZ)=zt+lnZL = -\big(\ln e^{z_t} - \ln Z\big) = -z_t + \ln Z
  2. zjz_j 求偏导(积木⑤:其他 zz 按住不动),两项分开算:
    • zt-z_t:它就是条直线。若 j=tj = t,斜率 1-1;若 jtj \neq t,这项里根本没有 zjz_j,斜率 0。
    • lnZ\ln Z:链式串联(积木⑥第一步)。外层 ln\ln 的斜率是 1Z\frac{1}{Z};内层 Z=ez0+ez1+ez2Z = e^{z_0}+e^{z_1}+e^{z_2}zjz_j 的斜率——三项里只有 ezje^{z_j}zjz_j,而 (ex)=ex(e^x)'=e^x,所以是 ezje^{z_j}。相乘:ezjZ\frac{e^{z_j}}{Z}——这恰好是 softmax 的定义,即 pjp_j
  3. 合并:Lzj=pj(答题卡上第 j 位)=pjyj\frac{\partial L}{\partial z_j} = p_j - (\text{答题卡上第 } j \text{ 位})= p_j - y_j

按计算器核对原文数字:j=2j=2(正确答案位):0.34941=0.65060.3494 - 1 = -0.6506 ✓;j=0j=00.36730=0.36730.3673 - 0 = 0.3673 ✓。

再用”多路分账”拿下原文 2.2 的两笔账

  • W02W_{02} 的梯度z2=x0W02+x1W12z_2 = x_0 W_{02} + x_1 W_{12}。把别的都按住,它就是 z2=0.3W02+常数z_2 = 0.3 \cdot W_{02} + \text{常数}——一条 y=kx+by = kx + b 直线,斜率就是系数 0.30.3!串联:LW02=(p21)0.6506×0.3x0=0.1952\frac{\partial L}{\partial W_{02}} = \underbrace{(p_2 - 1)}_{-0.6506} \times \underbrace{0.3}_{x_0} = -0.1952 ✓。WW 的每个格子只有一条路通向 LL(只经过自己那一列的 zjz_j),所以不用分账——这就是”外积”的全部秘密:格子 (i,j)(i,j) 的梯度 = xi×(pjyj)x_i \times (p_j - y_j)
  • x0x_0 的梯度x0x_0 出现在 z0,z1,z2z_0, z_1, z_2 三条账单里——三条路,分账相加:

Lx0=0.5×0.3673+(0.3)×0.2832+0.1×(0.6506)=0.0336 \frac{\partial L}{\partial x_0} = 0.5 \times 0.3673 + (-0.3) \times 0.2832 + 0.1 \times (-0.6506) = 0.0336 \ ✓

每条路都是”这条路上的系数 × 这条路终点的误差”。原文说”没参与前向的参数分不到账”(的嵌入行梯度为 0),现在你能看出为什么:它一条路都没有。

⑦ AdamW 的全部数学:加权平均、平方、开方、绝对值

原文 3.2 的表格只需要初中工具。

滑动平均m0.9m+0.1gm \leftarrow 0.9m + 0.1g 就是”旧印象占 9 成、新消息占 1 成”的加权平均。偏差修正是为了救第一步:mm 从 0 起步,第一次只装进了 0.1g0.1g,规模缩了 10 倍,除以 (10.9)=0.1(1 - 0.9) = 0.1 补回来。数字验证:g=10g=10m1=1m_1 = 1,修正 1/0.1=101/0.1 = 10,恢复原规模 ✓。

第一步为什么恰好等于 η×sign(g)\eta \times \mathrm{sign}(g)——初中代数三行证完:

m^=g(修正后,如上)v^=g2(同理)m^v^=gg2=gg=±1\hat m = g \quad (\text{修正后,如上}) \qquad \hat v = g^2 \quad (\text{同理}) \qquad \frac{\hat m}{\sqrt{\hat v}} = \frac{g}{\sqrt{g^2}} = \frac{g}{|g|} = \pm 1

g2=g\sqrt{g^2} = |g| 是初中的:(3)2=3\sqrt{(-3)^2} = 3。)梯度大小被自己除掉了,只剩方向——原文 3.2 观察 1 的”自适应归一化”,本质是一次除法。

weight decayθθηλθ\theta \leftarrow \theta - \eta\lambda\theta,按现值收比例税往 0 拉,纯乘法。原文的 103×0.1×0.1=10510^{-3} \times 0.1 \times 0.1 = 10^{-5},按一下就有。

带得走的东西:原文公式逐符号拆解表

原文公式符号读法/含义靠哪块积木
x=E[1]x = E[1]E[1]E[1]表格 EE 的第 1 行⓪ 门牌号
z=xWz = xWxWxW三份”数量×单价”账单① 乘加
pj=ezj/kezkp_j = e^{z_j}/\sum_k e^{z_k}ezje^{z_j}exe^x 键(变正、保排名)
同上k\sum_k加总求分母 ZZ
同上整体先 e 后占比 = softmax
L=lnptL = -\ln p_tln-\ln答对趋 0、答错爆炸的罚分
Lzj=pjyj\frac{\partial L}{\partial z_j} = p_j - y_j\partial按住其他,只动 zjz_j 的斜率
同上yjy_j答题卡第 jj 位(0 或 1)
同上推导ln\ln 拆减法 + 链式串联②⑥
LWij=xi(pjyj)\frac{\partial L}{\partial W_{ij}} = x_i(p_j - y_j)外积单路串联:直线斜率×终点误差
Lxi=jWij(pjyj)\frac{\partial L}{\partial x_i} = \sum_j W_{ij}(p_j - y_j)j\sum_j三条路分账相加⑥ 多路
L(θ+ε)L(θε)2ε\frac{L(\theta+\varepsilon)-L(\theta-\varepsilon)}{2\varepsilon}整体割线斜率量导数
θθηg\theta \leftarrow \theta - \eta gη\eta学习率:一步拧多少
m=β1m+(1β1)gm = \beta_1 m + (1-\beta_1)gmm9:1 加权平均(惯性)
m^/(v^+ϵ)\hat m/(\sqrt{\hat v} + \epsilon)v^\sqrt{\hat v}首步 =g=\|g\|,除完只剩 ±1
ηλθ\eta\lambda\thetaλ\lambda比例税率(往 0 拉)

自检清单

  1. W12W_{12} 指哪个格子?kezk\sum_k e^{z_k} 用原文数字展开是哪三个数相加?
  2. 不查原文,手算 z2=0.3×0.1+(0.1)×(0.5)z_2 = 0.3 \times 0.1 + (-0.1) \times (-0.5),和前向表对得上吗?
  3. 为什么把分数变概率前要先按 exe^x 键?直接算占比会死在哪种输入上?
  4. 用割线法(两点相距 0.002)量 f(x)=x2f(x)=x^2x=5x=5 的斜率,和 2x2x 给的 10 差多少?
  5. \partialdd 的区别是什么?“按住其他旋钮”按住的是原文里的哪些量?
  6. 重推 L/z1\partial L/\partial z_1(非正确答案位):哪一项直接是 0?剩下那项等于什么?
  7. x0x_0 的梯度为什么是三项相加而 W02W_{02} 的梯度只有一项?
  8. g2=g\sqrt{g^2} = |g| 说明:AdamW 第一步的步长为什么和梯度大小无关?

八题全能答,就可以回原文把 12 个参数的完整一步重新手算一遍——这次每个符号都应该是透明的。

诚实的提醒

  • 亲手验算过的:本文全部数字例子(ee 的各值、softmax 两例、割线三例、x2x^2 下山三步、链式 12、外积/分账两笔、滑动平均修正)都先用一次性 Python 脚本跑过,与文中一致。
  • 表述范围声明:“在高中导数公式表里”指中国高中数学选修内容中的基本初等函数求导公式((xn)(x^n)'(ex)(e^x)'(lnx)(\ln x)' 属于此列);不同教材版本编排有差异,未逐版核对。“多路分账”(多元链式法则)确实超出高中范围,本文按规则直接给出并用数字验证,未给严格证明——严格版本对应大学多元微积分,前置链条见从算术到前沿的 AI 数学地图
  • 最低成本验证:全文任何一个数字都能用手机计算器(带 exe^xln\ln 键)在一分钟内复现——这是本篇特意设计的验证门槛。

参考与回链

三个词的词表,十二个参数(原文) · 大模型为什么能学习,知识怎么存 · 交叉熵为什么用 log · 泰勒展开和大模型 · 线性代数与 AI 模型 vol.1 · 从算术到前沿的 AI 数学地图