📑 本页目录(点开跳转)
08 · 反向传播
⏱ 25 分钟 | ⭐⭐ 手写一遍就永远忘不掉 | 🔨 必须动手
🎯 一句话
反向传播 = 链式法则 + 从后往前算,把中间结果复用起来。 它不是什么高深理论,就是求导时别做重复劳动。
🧠 一、先建立直觉:谁该为错误负责
💡 人话:一个公司业绩差了,要追责到每个部门、每个人。 反向传播就是从结果开始,一层层往回追责——而且追责时会复用上一层已经算好的结论,不重复劳动。
⛓️ 二、链式法则:唯一需要的数学
如果 loss ← a ← z ← w (w 影响 z,z 影响 a,a 影响 loss)
那么 ∂loss/∂w = ∂loss/∂a × ∂a/∂z × ∂z/∂w
↑ 一路乘回去
💡 人话:「w 变一点点,z 变多少;z 变一点点,a 变多少;a 变一点点,loss 变多少」——三个比例连乘。
为什么要从后往前算:因为 ∂loss/∂a 这一段,前面所有参数都要用。
从后往前算一遍存下来,所有参数复用 → O(n)。
从前往后每个参数各算一遍 → O(n²)。这就是反向传播的全部优化。
📊 这个"复用"能省多少
朴素做法(数值梯度):
对每个参数 ±ε 各跑一次前向 → 2P 次完整前向传播
P = 100 万参数 → 200 万次前向 💀 完全不可行
反向传播:
1 次前向 + 1 次反向 ≈ 2 次前向的代价
【一次拿到全部参数的梯度】 ⭐
🔑 所以 1986 年反向传播的普及是个里程碑—— 想法不新(链式法则有几百年了), 但它让训练大网络在计算上第一次变得可行。
💡 一个值得记住的推论: 反向传播的代价大约是前向的 2 倍。 所以「训练一个 epoch 的成本 ≈ 推理 3 次全部数据」—— 这是估算训练时间时最好用的心算公式。
🔨 三、动手:纯 numpy 手写一个能学会异或的网络
这段代码是本章的核心。跑一遍,然后逐行对照下面的注解。
import numpy as np
X = np.array([[0,0],[0,1],[1,0],[1,1]], dtype=float)
y = np.array([[0],[1],[1],[0]], dtype=float) # XOR
def sigmoid(z): return 1/(1+np.exp(-z))
rng = np.random.RandomState(0)
# 网络结构:2 → 8 → 1
W1 = rng.randn(2, 8) * 0.5; b1 = np.zeros(8)
W2 = rng.randn(8, 1) * 0.5; b2 = np.zeros(1)
lr = 0.5
for step in range(3001):
# ============ 前向传播:算出预测和损失 ============
z1 = X @ W1 + b1 # 线性
a1 = np.tanh(z1) # 非线性(第 7 章)
z2 = a1 @ W2 + b2 # 线性
a2 = sigmoid(z2) # 输出概率
loss = -np.mean(y*np.log(a2+1e-9) + (1-y)*np.log(1-a2+1e-9)) # 交叉熵
# ============ 反向传播:从后往前追责 ============
n = len(X)
dz2 = (a2 - y) / n # ⭐ sigmoid+交叉熵的梯度就这么干净:预测−真实
dW2 = a1.T @ dz2 # z2 = a1@W2 → ∂z2/∂W2 = a1
db2 = dz2.sum(0)
da1 = dz2 @ W2.T # 责任传给上一层
dz1 = da1 * (1 - a1**2) # 穿过 tanh:tanh 的导数是 1−tanh²
dW1 = X.T @ dz1
db1 = dz1.sum(0)
# ============ 更新参数:朝梯度反方向走一小步 ============
W2 -= lr*dW2; b2 -= lr*db2
W1 -= lr*dW1; b1 -= lr*db1
if step % 1000 == 0:
print(f"step {step:<5} loss={loss:.4f}")
print("预测:", (a2>0.5).astype(int).ravel(), " 真实:", y.astype(int).ravel())
实测输出:
step 0 loss=0.8049
step 1000 loss=0.0038
step 2000 loss=0.0016
step 3000 loss=0.0009
预测: [0 1 1 0] 真实: [0 1 1 0] ✅ 学会异或了
你刚刚从零实现了一个神经网络的完整训练。 大约 25 行。
🔍 四、四个值得盯着看的细节
① dz2 = (a2 - y) —— 干净得可疑
Sigmoid + 交叉熵组合时,梯度正好就是「预测 − 真实」,所有复杂项全约掉了。
🔗 这就是第 3 章说的「分类为什么不用 MSE」的另一面—— 这个组合是被数学选中的,不是随便配的。
② 反向的每一步都对应前向的一步
前向:z1 → a1 → z2 → a2 → loss
反向:dz1 ← da1 ← dz2 ←────── loss
(每一步都在问:上一步的输出变一点,最终 loss 变多少)
③ 穿过激活函数时乘它的导数
tanh: dz = da * (1 - a²)
ReLU: dz = da * (z > 0) ← 负区间直接置 0
sigmoid: dz = da * a * (1-a)
💡 梯度消失的根源就在这里:如果激活函数的导数小于 1(如 sigmoid 最大 0.25), 层数一多,连乘就指数级衰减。ReLU 导数是 1,所以不衰减。
④ .sum(0) 是因为 batch
一个 batch 里每个样本都对偏置有贡献,要加起来。
⑤ 每一行梯度的形状怎么推出来(不用背)
⭐ 一条规则走天下:【梯度的形状 = 它对应那个量的形状】
W1 是 (2, 8) → dW1 必须是 (2, 8)
X 是 (4, 2) → X.T 是 (2, 4)
dz1 是 (4, 8)
→ dW1 = X.T @ dz1 形状是 (2,4)@(4,8) = (2,8) ✅ 对上了
💡 卡住时不用推数学,用形状凑:
要得到 (2,8),手上有 X(4,2) 和 dz1(4,8)—— 只有
X.T @ dz1这一种乘法能对上形状。 这个技巧在手写任何层的反向传播时都管用。
✅ 五、梯度检验:怎么确认自己写对了
手写反向传播极易出错,而且错了不报错,只是"学得慢"。 专业做法是数值检验:
# 数值梯度:把某个参数挪动一点点,看 loss 变多少
eps = 1e-5
W1p = W1.copy(); W1p[0,0] += eps
W1m = W1.copy(); W1m[0,0] -= eps
num_grad = (forward_loss(W1p,...) - forward_loss(W1m,...)) / (2*eps)
# 和你的解析梯度对比
print(f"数值梯度 {num_grad:.8f}")
print(f"解析梯度 {dW1[0,0]:.8f}")
实测结果:
数值梯度 -0.00007294
解析梯度 -0.00007294
相对误差 1.49e-08 ✅ 实现正确
🔑 判断标准:相对误差 < 1e-6 就算对。 自己实现任何自定义层/损失函数时,梯度检验是必做的一步——第 18 章挑战 A 会大量用到。
⚠️ 梯度检验的四个坑:
| 坑 | 说明 |
|---|---|
| 必须用中心差分 | (L(+ε)−L(−ε))/(2ε) 误差是 O(ε²);单边差分是 O(ε),差两个数量级 ⭐ |
| ε 不能太小 | 1e-5 左右最好。更小会被浮点舍入误差淹没,反而更不准 |
| ReLU 的折点 | z≈0 处数值梯度天然不准,这是函数本身不可导,不是 bug |
| 必须关掉随机性 | Dropout 和 BN 会让两次前向的网络不一样,比较毫无意义 ⭐ |
💡 相对误差怎么算:
|数值−解析| / (|数值|+|解析|+1e-12)。 用相对误差而不是绝对误差,是因为梯度本身的量级可能很小或很大。
🚀 六、现代框架帮你做了什么
# 上面 25 行的反向传播,在 PyTorch 里是这样:
loss.backward() # 就这一行
optimizer.step()
框架做的事叫「自动微分」:
- 前向计算时,悄悄记录下一张计算图(谁是谁算出来的)
- .backward() 时沿着图从后往前,自动套用链式法则
🔬 想看严格的推导(任意层数、矩阵形式的四个方程)? 见《机器学习的数学原理》第 12 章—— 那里还解释了残差连接在公式上到底多了什么。
💡 所以你以后不用手写反向传播 —— 但手写过一遍,你就知道: - 梯度爆炸/消失是怎么发生的(连乘) - 为什么要
optimizer.zero_grad()(梯度默认累加) - 自定义损失函数为什么可能不收敛(可导性)第 18 章挑战 A 会让你从零造一个 mini-torch,把这套自动微分自己实现出来。
🔗 和站内其他章的关系
| 相关的地方 | 根因在这 |
|---|---|
| 推荐算法第 5 章手写 MF 的 SGD 更新 | 同一套梯度下降,只是网络更简单 |
| 全景导论第 15 章「求导交给框架」 | 框架做的就是本章的自动微分 |
| 梯度消失/爆炸 | 反向连乘导致的指数衰减/增长 |
✅ 检查点
- 为什么反向传播要"从后往前"算?复杂度差别是多少?
- 反向传播的代价大约是前向的几倍?这个数字有什么实用价值?
- Sigmoid + 交叉熵的梯度是什么?为什么这么干净?
- 手写时梯度的形状怎么快速推出来?
- 梯度穿过 ReLU 时怎么变?穿过 sigmoid 呢?
- 梯度消失的根源是什么?
- 梯度检验怎么做?判断标准是什么?有哪四个坑?
- 为什么用相对误差而不是绝对误差?
👀 答案
- 因为靠后的那段梯度(∂loss/∂a)是前面所有参数共用的,从后往前算一遍存下来复用是 O(n);从前往后每个参数各算一遍是 O(n²)。数值梯度要 2P 次前向(P 是参数量),反向传播只要 1 次前向 + 1 次反向。
- 约 2 倍。所以「训练一个 epoch ≈ 推理 3 次全部数据」——估算训练时间最好用的心算公式。
- 正好是 (预测 − 真实)。因为这个组合在数学上互相约掉了复杂项——它们是被设计成配对使用的。
- 梯度的形状 = 它对应那个量的形状。要得到 dW1(2,8),手上有 X(4,2) 和 dz1(4,8),只有
X.T @ dz1能对上形状。卡住时用形状凑,不用推数学。 - ReLU:
dz = da * (z>0),负区间直接置 0;sigmoid:dz = da * a * (1-a),最大值只有 0.25。 - 反向传播时每穿过一层要乘激活函数的导数,如果导数小于 1(sigmoid 最大 0.25),层数一多连乘就指数级衰减。ReLU 导数为 1 所以不衰减。
- 把某参数 ±ε 各算一次 loss,用
(L(+ε)−L(−ε))/(2ε)得数值梯度,和解析梯度对比,相对误差 < 1e-6 算对。四个坑:①必须中心差分(误差 O(ε²) vs O(ε))②ε 不能太小(1e-5 最好,太小被浮点误差淹没)③ReLU 折点处天然不准(函数本身不可导,不是 bug)④必须关掉 Dropout 和 BN。 - 因为梯度本身的量级可能很小或很大,绝对误差没有统一的判断标准;相对误差
|数值−解析|/(|数值|+|解析|)才是尺度无关的。
🛑 可以停在这里
⚡ 走神救援
反向传播=链式法则+从后往前算以复用中间结果(数值梯度要2P次前向,反向传播只要1次前向+1次反向)——1986年它的普及是里程碑,因为它让训练大网络第一次在计算上可行;代价约是前向的2倍 → "训练一个epoch≈推理3次全部数据"是估时间的心算公式。手写25行让网络学会异或。五细节:①sigmoid+交叉熵梯度=预测−真实(干净是因为被数学选中)②反向每步对应前向每步③穿过激活要乘它的导数(ReLU乘(z>0),sigmoid乘a(1-a)最大0.25)→ 这就是梯度消失的根源④sum(0)是因为batch ⑤⭐梯度形状=对应量的形状,卡住时用形状凑不用推数学。梯度检验:中心差分,ε≈1e-5,相对误差<1e-6 才算对;四坑:必须中心差分、ε别太小、ReLU折点天然不准、先关Dropout和BN。框架的
.backward()=自动微分(记录计算图再反向套链式法则)。
下一节 👉 09-优化器与学习率.md