从零手写大模型 · 训练篇 02 - Backpropagation(反向传播)
训练篇 · 查看系列一、上集回顾
上一集我们讲了最小二乘法(Least Squares),核心是这么一件事:
想让模型"学"到东西,先得有一把尺子,量出模型现在预测得有多差——这把尺子就是损失函数(Loss Function)。
我们用均方误差(MSE)作为这把尺子:
L = (y_hat - y_true)²
y_hat 是模型的预测值,y_true 是真实标签。L 越小,说明预测越准。
训练的目标,就是不断调整模型里的参数(权重 W、偏置 b),让 L 变小。而"往哪个方向调、调多少",靠的是梯度——loss 对每个参数的偏导数。上一集我们处理的是最简单的情况:只有一层线性变换,梯度可以直接手算出来。
但我们现在手写的 GPT,长什么样?回忆一下推理篇 [EP12](堆叠 Transformer Block),一个输入要经过:
Embedding → Block 1 → Block 2 → ... → Block N → LayerNorm → 输出头 → y_hat
参数散落在每一层里。loss 是在最后算出来的,但需要更新的参数,可能藏在离 loss 很远的第 1 层。loss 对第 1 层参数的梯度,该怎么算?
这就是今天的主角:反向传播。
二、为什么需要反向传播
如果模型只有一层,梯度可以直接写公式硬算,上一集就是这么干的。但只要层数一多,直接硬算的方式会迅速失控——每多一层,公式里就要多嵌套一层,手写公式的复杂度会指数级爆炸。
反向传播要解决的,正是这个问题。它的核心思想其实很朴素:
不要试图一步写出"loss 对某个很深层参数"的完整公式,而是把这个复杂的求导过程,拆成一串简单的、逐层传递的局部求导,再用链式法则把它们乘起来。
前向传播(Forward)是信息从输入一路算到 loss;反向传播(Backward)则是把"loss 有多敏感"这个信号,从输出往回,一层一层地传回给每一层的参数。这也是"反向"这个名字的由来。
三、复习一下链式法则
链式法则(Chain Rule)是高中/大一数学就学过的东西,这里简单复习:
如果 z 是 y 的函数,y 又是 x 的函数,那么 z 对 x 的导数是:
dz/dx = (dz/dy) × (dy/dx)
直白点说:x 的一点变化,先影响 y,y 的变化再影响 z——把这两段"影响力"乘起来,就是 x 对 z 的总影响力。
如果链条更长,比如 z → y → x → w,那就一路乘下去:
dz/dw = (dz/dy) × (dy/dx) × (dx/dw)
神经网络的反向传播,本质上就是把这个链条应用在"每一层的输出是下一层的输入"这个结构上——只不过链条里传的不是标量,而是向量和矩阵。别担心,我们马上用"猫吃鱼"里的具体数字把它拆开揉碎讲清楚。
四、猫吃鱼案例:搭一个迷你两层网络
为了让链式法则的"层层传递"看得见摸得着,这一集我们不用完整的 Transformer Block(那个下一步会单独处理,attention、LayerNorm、残差的反向传播都各有讲究),而是先搭一个最简化的两层网络,把反向传播的机制吃透。
复用我们熟悉的"猫"的 embedding 向量(d_model=4):
x = [1, 2, -1, 0.5] # shape: (4,)
网络结构:
x (4,) → 线性层1 (W1, b1) → h (2,) → ReLU → a (2,) → 线性层2 (W2, b2) → y_hat (标量) → Loss
隐藏层维度取 2,是为了方便手算——原理和 d_model=4 时完全一样,只是数字多一点。
4.1 前向传播:一步步算出具体数字
第一层:线性变换
W1 = [[0.1, 0.2, -0.1, 0.05], # shape: (2, 4)
[0.3, -0.1, 0.2, 0.1 ]]
b1 = [0.1, -0.1] # shape: (2,)
h = W1 @ x + b1 # shape: (2,)
逐项算:
h1 = 0.1×1 + 0.2×2 + (-0.1)×(-1) + 0.05×0.5 + 0.1 = 0.725
h2 = 0.3×1 + (-0.1)×2 + 0.2×(-1) + 0.1×0.5 + (-0.1) = -0.15
h = [0.725, -0.15] # shape: (2,)
ReLU 激活
a = ReLU(h) = [max(0, 0.725), max(0, -0.15)] = [0.725, 0] # shape: (2,)
注意 h2 = -0.15 < 0,被 ReLU 直接清零了。这个"死掉"的分量,等会儿在反向传播里会有个很直观的后果,先记住这一点。
第二层:线性变换,输出预测值
W2 = [0.5, -0.4] # shape: (1, 2)
b2 = 0.2 # shape: (1,)
y_hat = W2 @ a + b2 = 0.5×0.725 + (-0.4)×0 + 0.2 = 0.5625 # shape: 标量
计算 Loss
假设这个简化任务的目标值 y_true = 1.0(可以理解成某种回归目标,不用纠结它具体代表什么,重点是走通反向传播的流程):
L = (y_hat - y_true)² = (0.5625 - 1)² = 0.19140625
前向传播到此结束。整条链路是:
x → (W1,b1) → h → ReLU → a → (W2,b2) → y_hat → L
4.2 反向传播:把链式法则逐层展开
现在反过来,从 L 开始往回走,把梯度一层一层传回去。核心规则只有一条:每往回走一层,就用链式法则乘上这一层的局部导数。
第一步:loss 对 y_hat 的梯度
MSE 求导:d(L)/d(y_hat) = 2 × (y_hat - y_true)
dL/dy_hat = 2 × (0.5625 - 1) = -0.875
第二步:传到第二层参数(W2, b2)
y_hat = W2 @ a + b2,所以 dy_hat/dW2 = a,dy_hat/db2 = 1。套链式法则:
dL/dW2 = dL/dy_hat × a = -0.875 × [0.725, 0] = [-0.634375, 0] # shape: (1,2)
dL/db2 = dL/dy_hat × 1 = -0.875 # shape: (1,)
到这里,第二层的梯度已经算完了——这一步跟上一集单层网络的计算一模一样,没什么新东西。真正的"反向传播"从下一步才开始体现。
第三步:继续往回传,传到 a(第二层的输入)
这一步是关键:a 不是参数,但我们需要知道"loss 对 a 有多敏感",因为 a 是第一层的输出,梯度还要继续往前传。
y_hat = W2 @ a + b2,所以 dy_hat/da = W2:
dL/da = dL/dy_hat × W2 = -0.875 × [0.5, -0.4] = [-0.4375, 0.35] # shape: (2,)
第四步:穿过 ReLU
ReLU 的导数很简单:输入大于 0 的地方导数是 1(原样通过),小于等于 0 的地方导数是 0(直接掐断):
ReLU'(h1=0.725) = 1 (因为 h1 > 0)
ReLU'(h2=-0.15) = 0 (因为 h2 ≤ 0)
dL/dh = dL/da ⊙ ReLU'(h) = [-0.4375×1, 0.35×0] = [-0.4375, 0] # shape: (2,)
(⊙ 表示逐元素相乘)
看到了吗:h2 那条路径的梯度直接变成了 0——前向传播时 ReLU 把 h2 清零了,反向传播时它也"回敬"了一个 0 梯度回去。也就是说,第一层里所有导致 h2 的参数(W1 第二行、b1 第二个分量),这一轮完全学不到任何东西,因为它们对最终 loss 的影响被 ReLU 掐断了。这就是常说的"ReLU 神经元死亡"现象在计算上的真实样子,不是什么玄学,就是链式法则乘上了一个 0。
第五步:传到第一层参数(W1, b1)
h = W1 @ x + b1,所以 dh/dW1 的结构是:h 的每个分量对 W1 对应那一行的导数是 x。展开成外积(outer product):
dL/dW1 = dL/dh 与 x 做外积
第一行(对应 h1):dL/dh1 × x = -0.4375 × [1, 2, -1, 0.5] = [-0.4375, -0.875, 0.4375, -0.21875]
第二行(对应 h2):dL/dh2 × x = 0 × [1, 2, -1, 0.5] = [0, 0, 0, 0]
dL/dW1 = [[-0.4375, -0.875, 0.4375, -0.21875], # shape: (2, 4)
[0, 0, 0, 0 ]]
dL/db1 = dL/dh = [-0.4375, 0] # shape: (2,)
到这里,梯度已经从 loss 一路传回了最底层的 W1、b1。整个反向传播链条走完了:
dL/dy_hat → dL/dW2, dL/db2 → dL/da → dL/dh (穿过ReLU) → dL/dW1, dL/db1
每一步,都只是"当前这一层的局部导数"乘以"从后面传过来的梯度"——没有任何一步需要写出 loss 对 W1 的完整展开公式。这就是反向传播比硬算高明的地方:把一个深层复合函数的求导,拆解成了一串简单的、可以逐层复用的局部计算。
五、一个直观的类比
做后端的同学,大概率跟"责任链"或者"异常堆栈回溯(stack trace)"打过交道:线上出了个 500 错误,你不会指望一眼看穿是底层哪个函数的问题,而是从最外层的报错开始,一层层往里查调用栈——网关报错,往回查是哪个服务返回的异常;服务报错,往回查是哪个函数抛的;函数报错,往回查是哪行代码、传入了什么脏数据。
反向传播干的事情本质上是一回事,只不过它传的不是"哪里出错了",而是"这里的输出该往哪个方向调、调多重要":
- loss 是最外层的"报错入口"
dL/dy_hat是"这次预测错得有多离谱"- 往回传一层,就是把这个"错得离谱"的责任,按照每个参数在这一层里的"权重贡献",分摊给它们
- 传到 ReLU 这种"开关"结构时,如果这一路在前向传播时压根没被激活(就像上面的 h2),反向传播时它自然也分不到任何责任——因为它前向的时候根本没参与"办事"
所以反向传播不是什么高深的黑魔法,它就是"前向传播时谁参与办事了,反向传播时就按参与程度给谁分配责任(梯度)",只不过这个"分配"是用链式法则严格算出来的,不是拍脑袋。
六、PyTorch 实现:手算验证 + autograd 自动求导
6.1 先手写一遍,验证跟手算的数字对不对
import torch
# 前向传播用到的输入和参数,手动设置成和手算完全一样的数字
x = torch.tensor([1.0, 2.0, -1.0, 0.5])
W1 = torch.tensor([[0.1, 0.2, -0.1, 0.05],
[0.3, -0.1, 0.2, 0.1]], requires_grad=True)
b1 = torch.tensor([0.1, -0.1], requires_grad=True)
W2 = torch.tensor([[0.5, -0.4]], requires_grad=True)
b2 = torch.tensor([0.2], requires_grad=True)
y_true = torch.tensor([1.0])
# 前向传播
h = W1 @ x + b1
a = torch.relu(h)
y_hat = W2 @ a + b2
loss = (y_hat - y_true) ** 2
print("h:", h)
print("a:", a)
print("y_hat:", y_hat)
print("loss:", loss)
输出应该和我们手算的一模一样:h = [0.725, -0.15],a = [0.725, 0],y_hat = 0.5625,loss ≈ 0.1914。
6.2 一行代码,交给 autograd
这才是 PyTorch 真正强大的地方——我们上面手算了一整节的链式法则,PyTorch 只需要一行:
loss.backward()
print("dL/dW2:", W2.grad)
print("dL/db2:", b2.grad)
print("dL/dW1:", W1.grad)
print("dL/db1:", b1.grad)
运行一下,你会看到跟我们手算完全一致的结果:
dL/dW2 = [[-0.634375, 0]]
dL/db2 = [-0.875]
dL/dW1 = [[-0.4375, -0.875, 0.4375, -0.21875],
[0, 0, 0, 0]]
dL/db1 = [-0.4375, 0]
backward() 这一行代码背后做的事情,就是我们上面手推的那一整套链式法则——只不过 PyTorch 在前向传播时,会偷偷记录下一张"计算图"(哪个变量是怎么从哪些变量算出来的),backward() 调用时就沿着这张图反着走一遍,自动帮你把每一步的局部导数乘好、传下去。这也是为什么之前几集我们全程用 nn.Module 手写各种层的时候,从没自己写过一行反向传播代码——PyTorch 的 autograd 引擎全程在后台帮我们干了这件事。
七、和系列前面内容的关联
- 训练篇[EP01](最小二乘法):那一集定义了"用什么尺子衡量模型好坏"(loss function),这一集解决的是"算出来的 loss 该怎么变成每个参数的更新方向"。两者再加上参数更新,才构成完整的"模型是怎么学习的"闭环——loss 告诉你差多少,梯度告诉你怎么改。
- 推理篇 [EP12](堆叠 Transformer Block):我们真正要训练的 GPT,是 N 个 Block 堆叠起来的,比今天这个两层 demo 深得多。但原理完全一样:反向传播会从最后的 loss 出发,依次穿过输出头、每一个 Transformer Block、一直传回 Embedding 层,每穿过一层就用链式法则乘一次局部导数。
- 推理篇 [EP07](残差连接)、[EP09](FFN)、多头注意力等:这些结构各自的反向传播都有自己的讲究——比如残差连接的反向传播会把梯度"一分为二"同时传给两条支路,这也是残差连接能缓解深层网络梯度消失问题的关键原因。这些具体到每种结构的反向传播细节,我们会在把完整训练流程接到真正的 GPT 模型上时,逐个拆开讲。
八、下集预告
这一集我们搞懂了"梯度是怎么算出来的"。但算出梯度之后呢?参数具体该怎么用这个梯度去更新?更新的步子迈多大合适?
下一集是训练篇 EP03「梯度下降」:拿到梯度后,实际更新一次参数,补上训练循环的最后一步。随后在 EP04 中介绍 Adam 优化器。