引言:前向传播把路铺好了,这一章让梯度"倒着走"
先花 30 秒回顾一下上一篇(第十六篇)。我们把 RNN 的前向传播真正"跑"了起来:
- 函数落地 ------
rnn_forward(x, W_xh, W_hh, W_hy, b_h, b_y, h0)用一行for t in range(S)统治了所有时间步,返回(h_seq, y_seq, cache); - cache 四件套 ------每个时间步存下 xt/ht−1/zt/htx_t / h_{t-1} / z_t / h_txt/ht−1/zt/ht,堆叠成 (S,B,⋅)(S,B,\cdot)(S,B,⋅) 的"口粮",就是为了今天;
- 维度验证 ------输出 (10,32,256)(10,32,256)(10,32,256) 与第二章推演表逐位一致,前向没写错;
- S=1S=1S=1 退化实验 ------RNN 退化成"tanh 隐藏层 + 线性输出层"的全连接网络,误差 0.00.00.0。
上一篇结尾我们说了:cache 是第四章 BPTT 的粮仓。 今天就是"开仓放粮"的日子------前向传播把每个时间步的中间量都算好、存好了,反向传播的任务,就是沿着时间轴倒着走回去,把损失对每个权重的"贡献"(梯度)一笔一笔算出来。
但你可能还有这些疑问:
- 前向是 h0→h1→⋯→hTh_0 \to h_1 \to \cdots \to h_Th0→h1→⋯→hT 一路往前,梯度凭什么能倒着走?
- 为什么 WhhW_{hh}Whh 的梯度要写成 ∑t\sum_t∑t 的累加和?累加的是什么?
- 一个时间步的误差,怎么知道该"怪"哪个权重、怪多少?
- 传说中 RNN 的梯度消失/爆炸,数学上到底是怎么回事?
- BPTT 和普通全连接网络的反向传播(BP),差在哪?
学完本章,你要能独立完成三件事:
🎯 本章三大目标
- 在纸上完整推导出三个梯度公式:∂L/∂Why\partial L/\partial W_{hy}∂L/∂Why、∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh、∂L/∂Wxh\partial L/\partial W_{xh}∂L/∂Wxh;
- 用自己的话讲清楚:为什么 WhhW_{hh}Whh 的梯度是"沿时间轴的累加和",中间变量 δt\delta_tδt 是怎么递推出来的;
- 用代码把推导"验算"一遍:递推公式 vs 数值梯度,误差应该在 10−1010^{-10}10−10 量级。
本篇路线图 :换个视角看反向(第一节)→ 损失函数定下来(第二节)→ 最简单的梯度 ∂L/∂Why\partial L/\partial W_{hy}∂L/∂Why(第三节)→ 硬骨头 ∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh 与 δt\delta_tδt 递推(第四节)→ 梯度消失/爆炸的数学根源(第五节)→ 四大实操验算(第六节,含本章里程碑)→ 常见坑与 FAQ(第七节)→ 小结与下章预告(第八节)。
沿着上一篇铺好的路,我们把梯度"倒着走"一遍。
一、换个方向看问题:前向是"顺着走",反向是"倒着追"
1.1 前向传播走过的路,就是反向要退的路
前向传播的本质是一条单向链条:
xt ⟶ zt ⟶ ht=tanh(zt) ⟶ yt x_t \;\longrightarrow\; z_t \;\longrightarrow\; h_t = \tanh(z_t) \;\longrightarrow\; y_t xt⟶zt⟶ht=tanh(zt)⟶yt
再加上时间维度的串联:
ht−1 → Whh ht → Whh ht+1 → Whh ⋯ h_{t-1} \;\xrightarrow{\;W_{hh}\;}\; h_t \;\xrightarrow{\;W_{hh}\;}\; h_{t+1} \;\xrightarrow{\;W_{hh}\;}\; \cdots ht−1Whh htWhh ht+1Whh ⋯
每一个 hth_tht 都"记得"前面所有的信息(通过 WhhW_{hh}Whh 的接力),所以越靠后的时间步,它的误差能追溯到的时间越长 。反向传播要做的事恰好反过来:从损失 LLL 出发,沿着这条链条一步步往回退,每退一步,就把"这笔账"记到途经的权重头上。
这和我们第二章讲的全连接网络 BP 是同一个原理,只不过 RNN 的链条既在"深度"上延伸(z→h→y),又在"时间"上延伸(t→t+1→t+2)。
1.2 BPTT = 沿时间轴展开的 BP
第一章我们把 RNN 按时间步"展开"成一张大图:SSS 个时间步,就是 SSS 个"长得一模一样"的全连接层串在一起 ,共享同一套权重。如果忘记"共享"这件事,把展开图当成一个 SSS 层的普通网络,那么反向传播就是现成的:
- 在展开图上做普通的 BP,一层一层往回传;
- 因为 SSS 层共享同一套权重,每个权重矩阵的梯度 = 它在每个时间步上的梯度之和。
这就是 BPTT(Backpropagation Through Time,随时间反向传播) 的全部秘密:BPTT 不是新算法,它就是在展开图上做的 BP,多了一个"梯度沿时间轴跨步累加"的步骤。
💡 一句话:普通 BP 的"深度"是层,BPTT 的"深度"是时间步。RNN 的层数 SSS 有多深,BPTT 就要往回传多深。
1.3 先把损失函数定下来:没有"目标",就没有"误差"
反向传播的起点是损失函数 LLL。RNN 的输出 yty_tyt 是"原始分数"(没有过激活函数,可正可负、无上界),所以我们定义每个时间步都参与计分 的损失。本章手算用的是最直观的均方误差(MSE):
L=12∑t=1S∥yt−tt∥2 L = \frac{1}{2} \sum_{t=1}^{S} \big\| y_t - t_t \big\|^2 L=21t=1∑S yt−tt 2
其中 ttt_ttt 是第 ttt 步的目标值 (ground truth)。系数 12\frac{1}{2}21 是故意的------一会求导会出来一个 ×2\times 2×2 抵消掉,数字更干净。
如果你做的是分类任务 (第十四章 IMDB 实战就是),更常用的是交叉熵 :先把 yty_tyt 过 softmax 变成概率 ptp_tpt,再算 −logpt-\log p_t−logpt。它的输出层误差长得更漂亮:
L=−∑tlogpt,δyt=∂L∂yt=pt−onehot(tt) L = -\sum_t \log p_t, \qquad \delta_{y_t} = \frac{\partial L}{\partial y_t} = p_t - \text{onehot}(t_t) L=−t∑logpt,δyt=∂yt∂L=pt−onehot(tt)
📌 无论用哪种损失,"输出层误差 = 预测 − 目标"这个结构是不变的:MSE 是 yt−tty_t - t_tyt−tt,交叉熵是 pt−ttp_t - t_tpt−tt。下面推导统一用 MSE,公式最清爽。
二、先算最简单的那个:∂L/∂Why\partial L / \partial W_{hy}∂L/∂Why
2.1 链式法则第一课:输出层误差 δyt\delta_{y_t}δyt
WhyW_{hy}Why 只出现在一个地方:yt=Whyht+byy_t = W_{hy} h_t + b_yyt=Whyht+by。LLL 对 yty_tyt 求导(MSE 下):
δyt ≜ ∂L∂yt=yt−tt \delta_{y_t} \;\triangleq\; \frac{\partial L}{\partial y_t} = y_t - t_t δyt≜∂yt∂L=yt−tt
这就是输出层误差 ------"预测偏了目标多少,就往反方向拽多少"。每个时间步都有一个 δyt\delta_{y_t}δyt,彼此独立。
2.2 把梯度"拼"出来
yty_tyt 对 WhyW_{hy}Why 的偏导:∂yt/∂Why=ht\partial y_t / \partial W_{hy} = h_t∂yt/∂Why=ht(行向量右乘权重,回忆第三章 2.2 节"内层对齐")。于是链式法则给出:
∂L∂Why=∑tht⊤⋅δyt \frac{\partial L}{\partial W_{hy}} = \sum_{t} h_t^{\top} \cdot \delta_{y_t} ∂Why∂L=t∑ht⊤⋅δyt
为什么是求和 ?因为 WhyW_{hy}Why 被 SSS 个时间步共用 ------每一个时间步的 yty_tyt 都贡献了一份梯度,改 WhyW_{hy}Why 会同时影响所有 yty_tyt,所以要把 SSS 份贡献全部加起来。这就是"参数共享"在梯度上的直接后果:谁的权重被共享,谁的梯度就是累加和。
偏置同理:∂L∂by=∑tδyt\displaystyle \frac{\partial L}{\partial b_y} = \sum_t \delta_{y_t}∂by∂L=t∑δyt(对批次求和就是对每个样本求和)。
2.3 为什么 WhyW_{hy}Why 的梯度"最简单"?
因为它没有时间依赖 :yty_tyt 只依赖 hth_tht,而 hth_tht 是前向时已经算好、缓存好的现成值。算 ∂L/∂Why\partial L/\partial W_{hy}∂L/∂Why 不需要"往回追",拿着 cache 里的 hth_tht 和刚算的 δyt\delta_{y_t}δyt 乘一下、加一下就行。
真正的麻烦在下一节:WhhW_{hh}Whh 出现在所有时间步的接力里。
三、硬骨头:∂L/∂Whh\partial L / \partial W_{hh}∂L/∂Whh ------ 梯度沿时间轴传播
3.1 问题出在哪:hth_tht 出现在所有"后来"的时间步里
WhhW_{hh}Whh 被用于每一步:zt=Wxhxt+Whhht−1+bhz_t = W_{xh} x_t + W_{hh} h_{t-1} + b_hzt=Wxhxt+Whhht−1+bh。改动 WhhW_{hh}Whh,会同时影响 h1,h2,...,hTh_1, h_2, \dots, h_Th1,h2,...,hT------而每个 hth_tht 又通过 WhyW_{hy}Why 影响 yty_tyt(进而影响 LLL),还通过 WhhW_{hh}Whh 影响 ht+1,ht+2,...h_{t+1}, h_{t+2}, \dotsht+1,ht+2,...。
于是 ∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh 的链式展开会长成"一串":
∂L∂Whh=∑t∂L∂zt⏟δt⋅∂zt∂Whh⏟ht−1⊤=∑tδt⋅ht−1⊤ \frac{\partial L}{\partial W_{hh}} = \sum_{t} \underbrace{\frac{\partial L}{\partial z_t}}{\delta_t} \cdot \underbrace{\frac{\partial z_t}{\partial W{hh}}}{h{t-1}^{\top}} = \sum_t \delta_t \cdot h_{t-1}^{\top} ∂Whh∂L=t∑δt ∂zt∂L⋅ht−1⊤ ∂Whh∂zt=t∑δt⋅ht−1⊤
问题来了:δt=∂L/∂zt\delta_t = \partial L/\partial z_tδt=∂L/∂zt 不是独立的 ------ztz_tzt 影响 hth_tht,hth_tht 影响 zt+1z_{t+1}zt+1 和 yty_tyt,所以 δt\delta_tδt 里藏着"来自未来"的梯度。要把它算出来,只能从最后一刻倒着推。
3.2 核心递推:δt\delta_tδt 的定义与反推
定义隐藏层误差信号:
δt ≜ ∂L∂zt \delta_t \;\triangleq\; \frac{\partial L}{\partial z_t} δt≜∂zt∂L
看 hth_tht 的"身后"有两条路:
- 直接路径 :ht→yt→Lh_t \to y_t \to Lht→yt→L,贡献 δyt⋅Why⊤\delta_{y_t} \cdot W_{hy}^{\top}δyt⋅Why⊤;
- 时间路径 :ht→zt+1→Lh_t \to z_{t+1} \to Lht→zt+1→L(因为 zt+1=⋯+Whhht+⋯z_{t+1} = \cdots + W_{hh} h_t + \cdotszt+1=⋯+Whhht+⋯),贡献 δt+1⋅Whh⊤\delta_{t+1} \cdot W_{hh}^{\top}δt+1⋅Whh⊤。
两条路在 hth_tht 处汇合相加 ,再乘上 ht=tanh(zt)h_t = \tanh(z_t)ht=tanh(zt) 的导数(回忆第二章:tanh\tanhtanh 的导数是 1−tanh21 - \tanh^21−tanh2):
δt=(δt+1⋅Whh⊤+δyt⋅Why⊤)⊙(1−ht2) \delta_t = \Big( \delta_{t+1} \cdot W_{hh}^{\top} + \delta_{y_t} \cdot W_{hy}^{\top} \Big) \odot (1 - h_t^2) δt=(δt+1⋅Whh⊤+δyt⋅Why⊤)⊙(1−ht2)
其中 ⊙\odot⊙ 是逐元素相乘(Hadamard 积)。边界条件是最后一步没有"未来" :δT+1=0\delta_{T+1} = \mathbf{0}δT+1=0,所以:
δT=δyT⋅Why⊤⊙(1−hT2) \delta_T = \delta_{y_T} \cdot W_{hy}^{\top} \odot (1 - h_T^2) δT=δyT⋅Why⊤⊙(1−hT2)
这就是第四章的灵魂递推式 :从 t=Tt = Tt=T 开始,倒着走回 t=0t = 0t=0,每一步都用"后一刻的 δt+1\delta_{t+1}δt+1 + 本步的输出误差 δyt\delta_{y_t}δyt",共同决定当前的 δt\delta_tδt。
💡 记忆口诀:"直接路径走 WhyW_{hy}Why,时间路径走 WhhW_{hh}Whh,汇合之后过 tanh\tanhtanh 导数闸门 (1−ht2)(1-h_t^2)(1−ht2)。"
3.3 一个细节:tanh\tanhtanh 导数为什么是 1−ht21 - h_t^21−ht2
tanh\tanhtanh 的导数是 1−tanh2(z)1 - \tanh^2(z)1−tanh2(z)。而我们缓存里存的是 ht=tanh(zt)h_t = \tanh(z_t)ht=tanh(zt) 本身,所以:
dhtdzt=1−tanh2(zt)=1−ht2 \frac{d h_t}{d z_t} = 1 - \tanh^2(z_t) = 1 - h_t^2 dztdht=1−tanh2(zt)=1−ht2
不用再费劲去算 ztz_tzt ------cache 里现成的 hth_tht 一平方一减就完事。这就是第三章拼命强调"cache 必须存 ztz_tzt 和 hth_tht"的原因:hth_tht 存着是为了算 yty_tyt 的梯度,ztz_tzt 存着是为了在第五章代码里直接算 tanh\tanhtanh 导数(1−tanh2(zt)1 - \tanh^2(z_t)1−tanh2(zt),两者数值等价)。
另外注意 1−ht2≤11 - h_t^2 \le 11−ht2≤1 恒成立(因为 ∣ht∣<1|h_t| < 1∣ht∣<1),而且 ∣ht∣|h_t|∣ht∣ 越接近 111(饱和),这个因子越接近 000。这个"永远不大于 1 的闸门",是梯度消失的第一位帮凶------第五节细讲。
3.4 三个梯度公式,一次集齐
有了 δt\delta_tδt,把"各回各家"的账目清点一遍:
∂L∂Why=∑tht⊤ δyt,∂L∂Whh=∑tht−1⊤ δt,∂L∂Wxh=∑txt⊤ δt \boxed{\; \frac{\partial L}{\partial W_{hy}} = \sum_{t} h_t^{\top} \,\delta_{y_t}, \qquad \frac{\partial L}{\partial W_{hh}} = \sum_{t} h_{t-1}^{\top} \,\delta_t, \qquad \frac{\partial L}{\partial W_{xh}} = \sum_{t} x_t^{\top} \,\delta_t \;} ∂Why∂L=t∑ht⊤δyt,∂Whh∂L=t∑ht−1⊤δt,∂Wxh∂L=t∑xt⊤δt
配套的偏置(对批次求和):
∂L∂bh=∑tδt,∂L∂by=∑tδyt \frac{\partial L}{\partial b_h} = \sum_t \delta_t, \qquad \frac{\partial L}{\partial b_y} = \sum_t \delta_{y_t} ∂bh∂L=t∑δt,∂by∂L=t∑δyt
三个公式的结构一模一样 ,只是"谁陪 δ\deltaδ 相乘"不同:
| 梯度 | 误差信号 | 与谁相乘(cache 里的现成值) | 直觉 |
|---|---|---|---|
| ∂L/∂Why\partial L/\partial W_{hy}∂L/∂Why | δyt\delta_{y_t}δyt | hth_tht(输出前一刻的隐藏状态) | "这一步的产出" |
| ∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh | δt\delta_tδt | ht−1h_{t-1}ht−1(上一步的隐藏状态) | "上一步的记忆" |
| ∂L/∂Wxh\partial L/\partial W_{xh}∂L/∂Wxh | δt\delta_tδt | xtx_txt(当前输入) | "这一步的输入" |
💡 记住一句话:每个时间步的 δ\deltaδ 先"各回各家"(分别陪 ht−1h_{t-1}ht−1、xtx_txt、δy\delta_yδy 相乘),再把 SSS 个时间步的贡献全部累加。 累加不是可选项------因为权重被共享,少加一个时间步,梯度就少算了一份"历史责任"。
四、梯度消失与爆炸:连乘的数学根源
4.1 从递推式看本质:误差要走"很多个闸门"
把递推式里的"时间路径"单独拎出来看。δt\delta_tδt 里来自未来的部分,要经过每一步的 Whh⊤W_{hh}^{\top}Whh⊤ 和 (1−ht2)(1-h_t^2)(1−ht2) 闸门:
δt 中的未来分量 =δt+1 Whh⊤⊙(1−ht2) \delta_t \;\text{中的未来分量}\; = \delta_{t+1} \, W_{hh}^{\top} \odot (1-h_t^2) δt中的未来分量=δt+1Whh⊤⊙(1−ht2)
把这一步递归展开 kkk 层,就是一串连乘:
δT−k ∝ δT⋅(Whh⊤diag(1−hT−12))⋯(Whh⊤diag(1−hT−k2))⏟k 个因子连乘 \delta_{T-k} \;\propto\; \delta_T \cdot \underbrace{\Big( W_{hh}^{\top} \mathrm{diag}(1-h_{T-1}^2) \Big) \cdots \Big( W_{hh}^{\top} \mathrm{diag}(1-h_{T-k}^2) \Big)}_{k\ \text{个因子连乘}} δT−k∝δT⋅k 个因子连乘 (Whh⊤diag(1−hT−12))⋯(Whh⊤diag(1−hT−k2))
写成更经典的形式------从第 1 步一路传到第 TTT 步的雅可比连乘:
∂hT∂h1=∏t=1T−1Whh⊤ diag(1−ht2) \frac{\partial h_T}{\partial h_1} = \prod_{t=1}^{T-1} W_{hh}^{\top} \; \mathrm{diag}(1 - h_t^2) ∂h1∂hT=t=1∏T−1Whh⊤diag(1−ht2)
这就是 RNN 梯度问题的数学根源:时间轴越长,这个连乘的因子越多。
4.2 为什么是指数的:谱范数
对任意矩阵 A,BA, BA,B,有范数不等式 ∥AB∥≤∥A∥⋅∥B∥\|AB\| \le \|A\| \cdot \|B\|∥AB∥≤∥A∥⋅∥B∥。对上面的连乘取范数:
∥∂hT∂h1∥ ≤ ∏t=1T−1∥Whh⊤∥⋅∥diag(1−ht2)∥ ≤ ρ(Whh)T−1 \Big\| \frac{\partial h_T}{\partial h_1} \Big\| \;\le\; \prod_{t=1}^{T-1} \Big\| W_{hh}^{\top} \Big\| \cdot \Big\| \mathrm{diag}(1-h_t^2) \Big\| \;\le\; \rho(W_{hh})^{T-1} ∂h1∂hT ≤t=1∏T−1 Whh⊤ ⋅ diag(1−ht2) ≤ρ(Whh)T−1
其中 ρ(Whh)\rho(W_{hh})ρ(Whh) 是 WhhW_{hh}Whh 的谱范数 (最大奇异值,直观上就是"权重矩阵在能量上最多能把向量放大多少倍"),而 ∥diag(1−ht2)∥≤1\|\mathrm{diag}(1-h_t^2)\| \le 1∥diag(1−ht2)∥≤1 已经把 tanh 闸门算进去了。于是:
- ρ(Whh)<1\rho(W_{hh}) < 1ρ(Whh)<1 :ρT−1\rho^{T-1}ρT−1 随 TTT 指数衰减 ------回传 30 步,梯度缩到 10−910^{-9}10−9 量级,早到不了远处的时间步。这就是梯度消失(vanishing gradient);
- ρ(Whh)>1\rho(W_{hh}) > 1ρ(Whh)>1 :ρT−1\rho^{T-1}ρT−1 随 TTT 指数爆炸 ------回传 30 步,梯度放大十万倍,更新一步就冲出天际。这就是梯度爆炸(exploding gradient);
- ρ≈1\rho \approx 1ρ≈1 是分水岭:比 1 小一点点就消失,比 1 大一点点就爆炸,几乎没有"刚好合适"的区间------这就是普通 RNN 难训练的根本原因。
💡 直观类比:把时间轴想成一串"收费站",每个站都要收"过路费"。收费比例(ρ\rhoρ)小于 1 时,钱越走越少(消失);大于 1 时,越走越多(爆炸)。tanh 闸门 (1−ht2)(1-h_t^2)(1−ht2) 还额外"再刮一层"------所以消失比爆炸更容易发生。
4.3 数值实验:指数规律
光看公式不过瘾,我们用代码把这条指数曲线画出来。实验设计:取一个 16×1616\times1616×16 的随机矩阵,把它缩放到指定的谱范数 ρ\rhoρ ,把"最后一步的误差"当成单位向量,沿着时间轴往回传 kkk 步,记录梯度范数。

左图是纯连乘 (把 tanh 摘掉,梯度范数严格等于 ρk\rho^kρk)------对数坐标下,四条线就是四条笔直的斜线:ρ=0.5\rho=0.5ρ=0.5 直线坠落,ρ=1.5\rho=1.5ρ=1.5 直线起飞。回传 30 步:
- ρ=0.5\rho = 0.5ρ=0.5:梯度缩水到 9.3×10−109.3\times10^{-10}9.3×10−10,十亿分之一,相当于"信息传到 30 步前就完全失联";
- ρ=1.5\rho = 1.5ρ=1.5:梯度放大 191,751191{,}751191,751 倍,19 万倍,更新一步直接溢出;
- 最讽刺的是 ρ=1.05\rho = 1.05ρ=1.05------只比 1 大 5%,30 步后也放大 4.3 倍。指数就是这么快。
右图是真实 tanh RNN :多了 (1−ht2)(1-h_t^2)(1−ht2) 因子后,hhh 很快被压进饱和区,(1−ht2)→0(1-h_t^2)\to 0(1−ht2)→0,于是"爆炸"被拖慢甚至压平,而"消失"变得更彻底。所以实际训练中你见到的大多是梯度消失,偶尔是爆炸------但根源都是这个连乘。
五、实操验算
纸上推导完了,我们写代码把每一步都验证一遍。用的还是第三章实操 B 那套"小到能心算"的模型(I=2,H=3,O=2,B=1,S=2I=2, H=3, O=2, B=1, S=2I=2,H=3,O=2,B=1,S=2),权重、输入、h0h_0h0 全部原封不动,只是补上了目标值 和损失:
python
import numpy as np
# ---- 第3章实操B的同一套小模型 ----
w_xh = np.array([[0.5, -0.2, 0.3], [-0.1, 0.4, 0.2]]) # (2,3)
w_hh = np.array([[0.8, 0.1, 0.0], [0.2, 0.7, -0.1], [0.0, 0.3, 0.6]]) # (3,3)
w_hy = np.array([[0.5, -0.4], [0.2, 0.3], [-0.1, 0.6]]) # (3,2)
bh = np.array([0.1, -0.1, 0.05])
by = np.array([0.0, 0.0])
x = np.array([[[1.0, -1.0], [0.5, 0.5]]]) # (1,2,2)
h0 = np.zeros((1, 3))
target = np.array([[[1.0, -1.0], [-0.5, 0.5]]]) # 目标值 (1,2,2)
def rnn_forward(x, W_xh, W_hh, W_hy, b_h, b_y, h0):
"""与第3章完全一致的前向,返回 (h_seq, y_seq, cache)"""
S = x.shape[1]
h_prev = h0
h_seq, y_seq = [], []
cache = {'x': [], 'h_prev': [], 'z': [], 'h': []}
for t in range(S):
x_t = x[:, t, :]
z_t = np.dot(x_t, W_xh) + np.dot(h_prev, W_hh) + b_h
h_t = np.tanh(z_t)
y_t = np.dot(h_t, W_hy) + b_y
for k, v in (('x', x_t), ('h_prev', h_prev), ('z', z_t), ('h', h_t)):
cache[k].append(v)
h_seq.append(h_t); y_seq.append(y_t)
h_prev = h_t
h_seq = np.stack(h_seq, axis=0) # (S,B,H)
y_seq = np.stack(y_seq, axis=0) # (S,B,O)
for k in cache:
cache[k] = np.stack(cache[k], axis=0)
return h_seq, y_seq, cache
def mse_loss(y_seq, target):
"""L = 1/2 * Σ_t ||y_t - target_t||^2"""
tgt = target.transpose(1, 0, 2) # (B,S,O) -> (S,B,O),对齐 y_seq
diff = y_seq - tgt
return 0.5 * np.sum(diff * diff)
5.1 实操 A:手算反向,一步一个脚印
前向结果和第三章一模一样(zt/ht/ytz_t/h_t/y_tzt/ht/yt 全对得上),现在定义了目标值,损失是:
L=12(∥y0−t0∥2+∥y1−t1∥2)=0.968125 L = \frac{1}{2}\Big( \|y_0 - t_0\|^2 + \|y_1 - t_1\|^2 \Big) = 0.968125 L=21(∥y0−t0∥2+∥y1−t1∥2)=0.968125
第 0 步:输出层误差。 每个时间步独立算:
δy0=y0−t0=−0.8336, 0.6663,δy1=y1−t1=0.6864, −0.5713 \delta_{y_0} = y_0 - t_0 = -0.8336,\\ 0.6663, \qquad \delta_{y_1} = y_1 - t_1 = 0.6864,\\ -0.5713 δy0=y0−t0=−0.8336, 0.6663,δy1=y1−t1=0.6864, −0.5713
第 1 步:最后一步 t=1t=1t=1,没有未来可回传。 只走直接路径,再过 tanh 闸门:
δz1=(δy1Why⊤)⊙(1−h12)=0.5717, −0.0341, −0.4114⊙0.6635, 0.9053, 0.8222=0.3793, −0.0309, −0.3383 \delta_{z_1} = \big( \delta_{y_1} W_{hy}^{\top} \big) \odot (1 - h_1^2) = 0.5717,\\ -0.0341,\\ -0.4114 \odot 0.6635,\\ 0.9053,\\ 0.8222 = 0.3793,\\ -0.0309,\\ -0.3383 δz1=(δy1Why⊤)⊙(1−h12)=0.5717, −0.0341, −0.4114⊙0.6635, 0.9053, 0.8222=0.3793, −0.0309, −0.3383
第 2 步:t=0t=0t=0,两条路径合流。 直接路径来自 y0y_0y0,时间路径来自"未来的 δz1\delta_{z_1}δz1 通过 WhhW_{hh}Whh 传回来":
δz0=(δy0Why⊤⏟直接路径 −0.6833, 0.0332, 0.4831+δz1Whh⊤⏟时间路径 0.3004, 0.0881, −0.2122)⊙(1−h02)=−0.2431, 0.0769, 0.2649 \delta_{z_0} = \big( \underbrace{\delta_{y_0} W_{hy}^{\top}}{\text{直接路径}\ -0.6833,\\ 0.0332,\\ 0.4831} +\underbrace{\delta{z_1} W_{hh}^{\top}}_{\text{时间路径}\ 0.3004,\\ 0.0881,\\ -0.2122} \big) \odot (1 - h_0^2) = -0.2431,\\ 0.0769,\\ 0.2649 δz0=(直接路径 −0.6833, 0.0332, 0.4831 δy0Why⊤+时间路径 0.3004, 0.0881, −0.2122 δz1Whh⊤)⊙(1−h02)=−0.2431, 0.0769, 0.2649
看到没?t=0t=0t=0 的误差里,有 0.3004 这个分量是从 t=1t=1t=1 穿越回来的 ------h0h_0h0 影响了 h1h_1h1,所以 y1y_1y1 的错误也得"怪" h0h_0h0 一点。这就是"梯度沿时间轴传播"最直白的证据。
第 3 步:拼装三个梯度。 以 WhhW_{hh}Whh 为例,累加两项(t=0t=0t=0 的贡献是零矩阵,因为 h−1=0h_{-1} = \mathbf{0}h−1=0):
\\frac{\\partial L}{\\partial W_{hh}} = \\underbrace{h_{-1}\^{\\top} \\delta_{z_0}}_{\\mathbf{0}} * \\underbrace{h_0\^{\\top} \\delta_{z_1}}_{\[0.2292,\\ -0.0187,\\ -0.2044\]} = \\begin{bmatrix} 0.2292 \& -0.0187 \& -0.2044 \\ -0.2292 \& 0.0187 \& 0.2044 \\ 0.0565 \& -0.0046 \& -0.0504 \\end{bmatrix}
逐位核对第一行:h00×δz1=0.6044×0.3793,−0.0309,−0.3383=0.2292,−0.0187,−0.2044h_00 \times \delta_{z_1} = 0.6044 \times 0.3793, -0.0309, -0.3383 = 0.2292, -0.0187, -0.2044h00×δz1=0.6044×0.3793,−0.0309,−0.3383=0.2292,−0.0187,−0.2044,分毫不差。
完整输出长这样(其余梯度同理拼装):
text
==========================================================================
实操 A:小模型手算反向 ------ 从损失出发,把梯度一步步"追"出来
==========================================================================
前向结果(与第3章实操B逐位一致):
t=0: z_t = [ 0.7 -0.7 0.15]
h_t = [ 0.6044 -0.6044 0.1489]
y_t = [ 0.1664 -0.3337]
t=1: z_t = [ 0.6626 -0.318 0.4498]
h_t = [ 0.5801 -0.3077 0.4217]
y_t = [ 0.1864 -0.0713]
目标值 target: t=0: [ 1. -1.] t=1: [-0.5 0.5]
损失 L = 1/2*Σ||y_t - target_t||^2 = 0.968125
[第0步] 输出层误差 δ_{y_t} = y_t - target_t(MSE 的导数):
δ_y0 = y_0 - target_0 = [-0.8336 0.6663]
δ_y1 = y_1 - target_1 = [ 0.6864 -0.5713]
[第1步] 最后一步 t=1:δ_{z_1} = (δ_{y_1} @ W_hy^T) ⊙ (1 - h_1^2)
δ_{y_1} @ W_hy^T = [ 0.5717 -0.0341 -0.4114]
(1 - h_1^2) = [0.6635 0.9053 0.8222]
δ_{z_1} = [ 0.3793 -0.0309 -0.3383]
[第2步] t=0:δ_{z_0} = (δ_{y_0} @ W_hy^T + δ_{z_1} @ W_hh^T) ⊙ (1 - h_0^2)
直接路径 δ_{y_0} @ W_hy^T = [-0.6833 0.0332 0.4831]
时间路径 δ_{z_1} @ W_hh^T = [ 0.3004 0.0881 -0.2122]
两条路径相加,再 ⊙ (1 - h_0^2)
δ_{z_0} = [-0.2431 0.0769 0.2649]
[第3步] 用 δ 拼出三个梯度(这就是 BPTT 的终点站):
∂L/∂W_hy = Σ_t h_t^T δ_{y_t}
= [[-0.1056 0.0713]
[ 0.2926 -0.2269]
[ 0.1653 -0.1417]]
∂L/∂W_hh = Σ_t h_{t-1}^T δ_t
= [[ 0.2292 -0.0187 -0.2044]
[-0.2292 0.0187 0.2044]
[ 0.0565 -0.0046 -0.0504]]
∂L/∂W_xh = Σ_t x_t^T δ_t
= [[-0.0534 0.0615 0.0958]
[ 0.4327 -0.0924 -0.434 ]]
∂L/∂b_h = [ 0.1362 0.0461 -0.0734]
∂L/∂b_y = [-0.1472 0.095 ]
5.2 实操 B:递推公式 vs 数值梯度
"手算对上了"只能说明我们自己没算错 ,不能证明公式本身正确 。金标准是数值梯度 :把某个 ztz_tzt 的元素微微拨动 ε\varepsilonε(比如 10−610^{-6}10−6),用中心差分估计真实梯度:
∂L∂zt ≈ L(zt+ε)−L(zt−ε)2ε \frac{\partial L}{\partial z_t} \;\approx\; \frac{L(z_t + \varepsilon) - L(z_t - \varepsilon)}{2\varepsilon} ∂zt∂L≈2εL(zt+ε)−L(zt−ε)
然后和递推公式算出的 δt\delta_tδt 对比。注意:扰动 z0z_0z0 时,z1z_1z1 必须由真实递推算出(因为 h0h_0h0 变了会影响 h1h_1h1),不能把两个 zzz 同时定死,否则时间依赖就断了。
text
==========================================================================
实操 B:δ 递推公式验算 ------ 对 z_t 逐个扰动,数值梯度应等于递推结果
==========================================================================
t=0: 数值 ∂L/∂z_t = [-0.243067 0.076948 0.264895]
递推 δ_t = [-0.243067 0.076948 0.264895]
最大误差 = 3.389e-11
t=1: 数值 ∂L/∂z_t = [ 0.379311 -0.030894 -0.338257]
递推 δ_t = [ 0.379311 -0.030894 -0.338257]
最大误差 = 2.162e-11
-> 最大误差 3.389e-11:递推公式与数值梯度一致,推导没有算错
误差在 10−1110^{-11}10−11 量级------这就是中心差分本身的浮点精度极限,可以认为完全一致 。你推导时哪一步乘错、哪一步漏了 ⊙\odot⊙、哪一步转置写反,这套体检立刻能抓出来。第五章的梯度校验就是它的"全面升级版"。
5.3 实操 C:把 WhhW_{hh}Whh 的"时间累加"拆开,逐项过目
公式说 ∂L/∂Whh=∑tht−1⊤δt\partial L/\partial W_{hh} = \sum_t h_{t-1}^{\top} \delta_t∂L/∂Whh=∑tht−1⊤δt,那就把两项贡献单独打印出来看:
text
==========================================================================
实操 C:为什么 ∂L/∂W_hh 是"时间轴累加和"?------ 逐项拆开看
==========================================================================
t=0 的贡献项: h_{-1}^T δ_0 = 0^T δ_0 = 全零矩阵(h_{-1}=0,起点没有记忆)
-> 全零? True
t=1 的贡献项: h_0^T δ_1 =
[[ 0.2292 -0.0187 -0.2044]
[-0.2292 0.0187 0.2044]
[ 0.0565 -0.0046 -0.0504]]
两项相加 ∂L/∂W_hh =
[[ 0.2292 -0.0187 -0.2044]
[-0.2292 0.0187 0.2044]
[ 0.0565 -0.0046 -0.0504]]
数值梯度 ∂L/∂W_hh =
[[ 0.2292 -0.0187 -0.2044]
[-0.2292 0.0187 0.2044]
[ 0.0565 -0.0046 -0.0504]]
最大误差 = 5.897e-11(累加和公式完全正确)
两个细节值得记住:
- t=0t=0t=0 的贡献是零矩阵 ------因为 h−1=h0=0h_{-1} = h_0 = \mathbf{0}h−1=h0=0(全零起点,第三章定的惯例)。起点没有"上一步记忆",自然不背锅;
- 累加和 vs 数值梯度误差 5.9×10−115.9\times10^{-11}5.9×10−11------"逐时间步相加"和"整体数值微分"是两条完全不同的计算路径,结果一致,说明累加不是"想当然",而是数学事实。
最后把五个参数一起过一遍体检(中心差分,ε=10−6\varepsilon = 10^{-6}ε=10−6):
text
三个权重 + 两个偏置的数值梯度体检(中心差分 eps=1e-6):
W_xh 解析 vs 数值 最大误差 = 1.280e-10 OK
W_hh 解析 vs 数值 最大误差 = 5.897e-11 OK
W_hy 解析 vs 数值 最大误差 = 4.474e-11 OK
b_h 解析 vs 数值 最大误差 = 1.110e-10 OK
b_y 解析 vs 数值 最大误差 = 9.377e-11 OK
五个全 OK。第三章的 cache、本章的递推公式、数值梯度三者互相印证------这条推导链是闭合的。
5.4 实操 D:谱范数实验,亲眼见证"指数"
把第四节的理论跑成数字(左列是"回传 k 步后梯度范数",右列同样)。先看纯连乘 (范数严格等于 ρk\rho^kρk):
text
实验A:纯连乘 ||δ_k|| = ρ^k(把最后一步误差当单位向量往回传)
回传步数 k | ρ=0.50 | ρ=0.90 | ρ=1.05 | ρ=1.50
--------------------------------------------------------------------------
k=0 | 1.0000e+00 | 1.0000e+00 | 1.0000e+00 | 1.0000e+00 |
k=3 | 1.2500e-01 | 7.2900e-01 | 1.1576e+00 | 3.3750e+00 |
k=7 | 7.8125e-03 | 4.7830e-01 | 1.4071e+00 | 1.7086e+01 |
k=11 | 4.8828e-04 | 3.1381e-01 | 1.7103e+00 | 8.6498e+01 |
k=15 | 3.0518e-05 | 2.0589e-01 | 2.0789e+00 | 4.3789e+02 |
k=19 | 1.9073e-06 | 1.3509e-01 | 2.5270e+00 | 2.2168e+03 |
k=23 | 1.1921e-07 | 8.8629e-02 | 3.0715e+00 | 1.1223e+04 |
k=29 | 1.8626e-09 | 4.7101e-02 | 4.1161e+00 | 1.2783e+05 |
ρ=0.5:回传 30 步,梯度缩水到 9.31e-10(十亿分之一!)
ρ=1.5:回传 30 步,梯度放大 191751 倍(19万倍!)
ρ=1.05:只比 1 大 5%,30 步后也放大 4.3 倍
再看真实 tanh RNN :同样的 WhhW_{hh}Whh 缩放,但 hhh 会饱和,(1−ht2)(1-h_t^2)(1−ht2) 因子持续"刮油"------
text
实验B:真实 tanh RNN(同样的 W_hh 缩放,h 会饱和)
回传步数 k | ρ=0.50 | ρ=0.90 | ρ=1.05 | ρ=1.50
--------------------------------------------------------------------------
k=0 | 1.0000e+00 | 1.0000e+00 | 1.0000e+00 | 1.0000e+00 |
k=3 | 2.1555e-02 | 1.2556e-01 | 1.9908e-01 | 5.7021e-01 |
k=7 | 3.7327e-04 | 2.2813e-02 | 6.6916e-02 | 7.5717e-01 |
k=11 | 3.3458e-06 | 2.1269e-03 | 1.1484e-02 | 5.0426e-01 |
k=15 | 3.1254e-08 | 2.0848e-04 | 2.0804e-03 | 3.6639e-01 |
k=19 | 3.0631e-10 | 2.1359e-05 | 3.9263e-04 | 2.8184e-01 |
k=23 | 3.0078e-12 | 2.1940e-06 | 7.4544e-05 | 2.2082e-01 |
k=29 | 2.9239e-15 | 7.2547e-08 | 6.1698e-06 | 1.3790e-01 |
对比:ρ=1.5 时纯连乘应放大 19 万倍,但真实 tanh 下只有 0.13 倍------
原因:h 被 tanh 压到饱和区后 (1-h^2)→0,把爆炸"拖慢"了;
而 ρ=0.5 时 tanh 让消失更彻底:缩水到 9.19e-16(比纯连乘还狠)。
结论一句话:连乘决定"趋势是指数",tanh 决定"消失比爆炸更常见"。 这也预告了后续章节的解药:LSTM 用"传送带"绕开连乘(第六章),GRU 用更少参数的近似(第九章),实在不行还有梯度裁剪兜底(第五章)。
六、一张图看懂 BPTT 推导全景
目录里本章的里程碑是"手写推导 ∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh 的链式法则全过程(含 δt\delta_tδt 的定义)"。我们用 Python 把推导全景画成一张印刷体大图------从损失出发,红色箭头走 WhyW_{hy}Why 直接路径,橙色箭头走 WhhW_{hh}Whh 时间路径,一路倒着推到三个梯度公式,右下角还挂着"连乘 → 梯度消失/爆炸"的根源公式:

读图顺序(从上到下):
① 损失 LLL(MSE)→ ② 输出层误差 δyt=yt−tt\delta_{y_t} = y_t - t_tδyt=yt−tt → ③ 最后一步 δzT\delta_{z_T}δzT(只有直接路径)→ ④ 核心递推 δt=(δt+1Whh⊤+δytWhy⊤)⊙(1−ht2)\delta_t = (\delta_{t+1} W_{hh}^{\top} + \delta_{y_t} W_{hy}^{\top}) \odot (1-h_t^2)δt=(δt+1Whh⊤+δytWhy⊤)⊙(1−ht2),沿时间轴倒着递推 T−1T-1T−1 次 → ⑤ 三个梯度公式各回各家、再对时间步求和(∑t\sum_t∑t)。
建议你对照这张图,在纸上把实操 A 的每一步自己写一遍------写一遍顶看十遍。
七、常见坑与 FAQ
🕳️ 常见坑
| # | 坑 | 现象 | 解法 |
|---|---|---|---|
| 1 | 反向遍历写成了正序 for t in range(S) |
算 δt\delta_tδt 时 δt+1\delta_{t+1}δt+1 还没算出来 | 必须倒序 for t in range(S-1, -1, -1),先算未来,再算过去 |
| 2 | 忘加"时间路径" δt+1Whh⊤\delta_{t+1} W_{hh}^{\top}δt+1Whh⊤ | 中间时间步的梯度算错(只有最后一步才对) | 除最后一步外,hth_tht 身后有两条路,缺一不可 |
| 3 | 忘乘 tanh\tanhtanh 导数 1−ht21-h_t^21−ht2 | 梯度整体偏大,方向还对但数值全错 | 递推式最后一步必须 ⊙(1−ht2)\odot (1-h_t^2)⊙(1−ht2),用的是缓存里的 hth_tht |
| 4 | 转置写反:Whhδt+1W_{hh} \delta_{t+1}Whhδt+1 还是 δt+1Whh⊤\delta_{t+1} W_{hh}^{\top}δt+1Whh⊤ | 维度对不上或梯度错误 | 行向量约定下,梯度右乘 转置:np.dot(delta_next, W_hh.T) |
| 5 | 梯度的"时间累加"少加了某一步 | 数值梯度体检立刻 FAIL | 每个时间步的贡献都要 +=,一个都不能少(t=0t=0t=0 贡献为零也要走一遍流程) |
| 6 | 把 δy\delta_yδy 和 δz\delta_zδz 混为一谈 | 公式张冠李戴 | 记清楚:δyt\delta_{y_t}δyt 是输出层误差(对 yty_tyt 求导),δt=δzt\delta_t = \delta_{z_t}δt=δzt 是隐藏层误差(对 ztz_tzt 求导),二者差一个 Why⊤W_{hy}^{\top}Why⊤ 和 tanh 闸门 |
❓ FAQ
Q1:为什么 WhyW_{hy}Why 的梯度不需要"时间累加"式的递推?
yty_tyt 只依赖 hth_tht(前向已经算好),不依赖未来的任何东西,所以 ∂L/∂Why\partial L/\partial W_{hy}∂L/∂Why 拿到 δyt\delta_{y_t}δyt 就能直接拼。但 WhhW_{hh}Whh 的每个"使用者" hth_tht 都影响了后面所有步,所以它的梯度必须把 δt\delta_tδt 一路递推出来、再跨时间步累加。
Q2:δt\delta_tδt 到底存的什么?为什么叫"误差信号"?
δt=∂L/∂zt\delta_t = \partial L/\partial z_tδt=∂L/∂zt:如果 ztz_tzt 的第 iii 个神经元微微变大一点,损失会变多少?这个"敏感度"就是误差信号。它同时携带"来自输出 yty_tyt 的账"和"来自未来所有时间步的账",是梯度公式的"通用货币"。
Q3:数值梯度那么准,为什么不用它代替解析梯度?
数值梯度对每个参数元素 都要跑两次前向,WhhW_{hh}Whh 有 H2H^2H2 个元素就要跑 2H22H^22H2 次------序列一长就彻底不可行。它的正确用途是体检工具:小模型上验证解析公式没错,然后放心用解析梯度(第五章就是这么干的)。
Q4:h−1h_{-1}h−1 全零,那 ∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh 在 t=0t=0t=0 的贡献为什么是零?
梯度公式里 t=0t=0t=0 的贡献是 h−1⊤δ0h_{-1}^{\top} \delta_0h−1⊤δ0,h−1=0h_{-1} = \mathbf{0}h−1=0 乘什么都为零。物理含义:WhhW_{hh}Whh 在第一步"没得记忆可回忆",自然不背第一步的锅。但 δ0\delta_0δ0 本身不为零 ------它还会通过 WxhW_{xh}Wxh 路径(陪 x0x_0x0 相乘)去更新 WxhW_{xh}Wxh。
Q5:梯度消失/爆炸是不是只在 RNN 里有?
不是。深层前馈网络也有(深度方向连乘),但 RNN 的"深度"是时间步 SSS,动辄几十上百,且权重共享导致连乘因子完全相同,所以问题格外严重。CNN 有残差连接、批归一化等一堆手段,RNN 这边最经典的解法就是第六章的 LSTM。
Q6:为什么 ρ>1\rho > 1ρ>1 会爆炸,实际训练却更多见到 NaN 而不是大梯度?
爆炸的梯度一旦进入梯度下降更新,权重一步飞出,下一步前向直接 NaN(梯度爆炸的"尸体")。而消失更隐蔽:不报错、不 NaN,只是梯度小得什么都不学,训练曲线趴地不动------这正是第六章里"普通 RNN 记不住长序列"的幕后黑手。
八、本章小结与下章预告
本章小结(一句话带走一个知识点)
| 知识点 | 一句话带走 |
|---|---|
| BPTT 的本质 | 就是沿时间轴展开的 BP:在展开图上逐层回传,权重共享导致梯度跨时间步累加 |
| 输出层误差 | δyt=∂L/∂yt=yt−tt\delta_{y_t} = \partial L/\partial y_t = y_t - t_tδyt=∂L/∂yt=yt−tt(MSE);交叉熵则是 pt−ttp_t - t_tpt−tt |
| 核心递推 | δt=(δt+1Whh⊤+δytWhy⊤)⊙(1−ht2)\delta_t = (\delta_{t+1} W_{hh}^{\top} + \delta_{y_t} W_{hy}^{\top}) \odot (1 - h_t^2)δt=(δt+1Whh⊤+δytWhy⊤)⊙(1−ht2),从 t=Tt=Tt=T 倒着递推到 t=0t=0t=0,边界 δT+1=0\delta_{T+1} = \mathbf{0}δT+1=0 |
| 三条路径 | 直接路径走 WhyW_{hy}Why,时间路径走 WhhW_{hh}Whh,汇合后过 tanh\tanhtanh 闸门 1−ht21-h_t^21−ht2 |
| 三个梯度 | ∂L/∂Why=∑tht⊤δyt\partial L/\partial W_{hy} = \sum_t h_t^{\top}\delta_{y_t}∂L/∂Why=∑tht⊤δyt;∂L/∂Whh=∑tht−1⊤δt\partial L/\partial W_{hh} = \sum_t h_{t-1}^{\top}\delta_t∂L/∂Whh=∑tht−1⊤δt;∂L/∂Wxh=∑txt⊤δt\partial L/\partial W_{xh} = \sum_t x_t^{\top}\delta_t∂L/∂Wxh=∑txt⊤δt |
| 为什么累加 | 权重被 SSS 个时间步共享,每个时间步都有一份"历史责任",缺一不可 |
| 梯度消失/爆炸 | ∂hT/∂h1=∏Whh⊤diag(1−ht2)\partial h_T/\partial h_1 = \prod W_{hh}^{\top}\mathrm{diag}(1-h_t^2)∂hT/∂h1=∏Whh⊤diag(1−ht2);谱范数 ρ<1\rho<1ρ<1 指数衰减(消失),ρ>1\rho>1ρ>1 指数放大(爆炸) |
| 验算结果 | 递推 vs 数值梯度:误差 10−1110^{-11}10−11 量级,五个参数全 OK ✅ |
下章预告:第 5 章《BPTT 算法(下)------代码实现与梯度裁剪》
公式推导完毕、验算通过------下一章把它写成能跑的完整 rnn_backward :拿着第三章的 cache,倒序遍历时间轴,把本章的递推式一行行翻译成 NumPy;再用数值梯度法 做全量体检(误差 10−610^{-6}10−6 以内算合格);最后给训练循环装上梯度裁剪,亲眼看看"裁剪前梯度范数冲上云霄、裁剪后老实趴线"的对比曲线。
另外补一句:本章的梯度消失分析,直接为第六章埋了引子------LSTM 是怎么用一条"传送带"绕过这个连乘的? 到时候见分晓。
BPTT(上),完结。
🧠 思考题与动手练习
思考题(先自己想,再看答案区,答案就在正文里):
- 为什么 WhhW_{hh}Whh 的梯度要"跨时间步累加",而 WhyW_{hy}Why 不用?(提示:谁被共享,谁就累加)
- 如果去掉递推式里的 (1−ht2)(1-h_t^2)(1−ht2) 因子,梯度会变大还是变小?误差会出现在哪里?(提示:tanh 闸门的作用)
- h−1=0h_{-1} = \mathbf{0}h−1=0 导致 ∂L/∂Whh\partial L/\partial W_{hh}∂L/∂Whh 在 t=0t=0t=0 的贡献为零,那 t=0t=0t=0 的误差信号 δ0\delta_0δ0 还重要吗?它去哪了?(提示:WxhW_{xh}Wxh 路径)
- 谱范数 ρ=0.9\rho = 0.9ρ=0.9 和 ρ=1.1\rho = 1.1ρ=1.1 只差 0.2,为什么训练效果天差地别?(提示:0.950≈0.0050.9^{50} \approx 0.0050.950≈0.005,1.150≈1171.1^{50} \approx 1171.150≈117)
- 为什么实际训练中"梯度消失"比"梯度爆炸"更常见?(提示:(1−ht2)≤1(1-h_t^2) \le 1(1−ht2)≤1 恒成立)
动手练习(改造本章实操代码):
- 把 MSE 换成二分类交叉熵 :yty_tyt 过 sigmoid,δyt=σ(yt)−tt\delta_{y_t} = \sigma(y_t) - t_tδyt=σ(yt)−tt,重新跑实操 A,验证梯度变化;
- 把 t=1t=1t=1 的贡献项单独放大 101010 倍再累加,看看数值梯度体检还过不过------体会"累加和"对每一项都敏感;
- 在实操 D 里把 ρ\rhoρ 取 0.990.990.99 和 1.011.011.01,对比回传 30 步的梯度范数,体会"分水岭"有多窄;
- 把 HHH 从 3 改成 5 重跑实操 A(权重随机),验证递推公式与数值梯度依然一致;
- 手写一份"推导清单":不查资料,从 δyt\delta_{y_t}δyt 开始把三个梯度公式推一遍,和实操 A 的输出逐位对照。
📌 下篇预告:第五章《BPTT 算法(下)------代码实现与梯度裁剪》------把本章的递推式写成完整的反向传播函数,用数值梯度给它做全身体检,再装上梯度裁剪驯服"爆炸"。我们下篇见!
本文为原创,遵循 CC 4.0 BY-SA 版权协议,转载需附原文链接。
🐍 附:实操代码运行
python
# -*- coding: utf-8 -*-
"""
本章不实现完整的 rnn_backward(那是第5章的事),只做四件事:
A. 小模型手算反向:用第3章实操B的模型(I=2,H=3,O=2,B=1,S=2),定义损失后手推
δ_y -> δ_z -> 三个梯度,打印每一步的数值;
B. δ 递推公式验算:把"反向递推算出的 δ_z" 与 "数值梯度(对 z_t 逐个扰动)" 对照,
证明递推公式正确;
C. W_hh 梯度的时间累加分解:逐时间步打印贡献项 h_{t-1}^T δ_t,证明
∂L/∂W_hh = Σ_t h_{t-1}^T δ_t 是"沿时间轴的累加和";
D. 梯度消失/爆炸数值实验:控制 W_hh 的谱范数 ρ,观察梯度沿时间轴回传时
范数随步数的变化(对数尺度),印证"连乘 -> 指数衰减/爆炸"。
超参数与第3章实操B完全一致,方便对照。
"""
import numpy as np
LINE = '=' * 74
# ---------------- 第3章实操B的同一套小模型 ----------------
w_xh = np.array([[0.5, -0.2, 0.3],
[-0.1, 0.4, 0.2]]) # (I=2, H=3)
w_hh = np.array([[0.8, 0.1, 0.0],
[0.2, 0.7, -0.1],
[0.0, 0.3, 0.6]]) # (H=3, H=3) 方阵
w_hy = np.array([[0.5, -0.4],
[0.2, 0.3],
[-0.1, 0.6]]) # (H=3, O=2)
bh = np.array([0.1, -0.1, 0.05])
by = np.array([0.0, 0.0])
x = np.array([[[1.0, -1.0], [0.5, 0.5]]]) # (B=1, S=2, I=2)
h0 = np.zeros((1, 3)) # h_{-1} 全零
# 损失的目标值(ground truth),随便挑一组方便心算的数
target = np.array([[[1.0, -1.0], [-0.5, 0.5]]]) # (B=1, S=2, O=2)
# ---------------- 前向(与第3章同款) ----------------
def rnn_forward(x, W_xh, W_hh, W_hy, b_h, b_y, h0):
S = x.shape[1]
h_prev = h0
h_seq, y_seq = [], []
cache = {'x': [], 'h_prev': [], 'z': [], 'h': []}
for t in range(S):
x_t = x[:, t, :]
z_t = np.dot(x_t, W_xh) + np.dot(h_prev, W_hh) + b_h
h_t = np.tanh(z_t)
y_t = np.dot(h_t, W_hy) + b_y
cache['x'].append(x_t)
cache['h_prev'].append(h_prev)
cache['z'].append(z_t)
cache['h'].append(h_t)
h_seq.append(h_t)
y_seq.append(y_t)
h_prev = h_t
h_seq = np.stack(h_seq, axis=0)
y_seq = np.stack(y_seq, axis=0)
for k in cache:
cache[k] = np.stack(cache[k], axis=0)
return h_seq, y_seq, cache
def mse_loss(y_seq, target):
"""L = 1/2 * Σ_t ||y_t - target_t||^2
y_seq 是 (S,B,O),target 是 (B,S,O),先把 target 转成 (S,B,O) 再相减。"""
tgt = target.transpose(1, 0, 2)
diff = y_seq - tgt
return 0.5 * np.sum(diff * diff)
# ---------------- 解析 BPTT(本章手推公式的代码形态) ----------------
def bptt_analytic(y_seq, target, cache, W_xh, W_hh, W_hy):
"""严格按正文推导的公式实现,返回各梯度与每一步的 delta_z / delta_y"""
S = y_seq.shape[0]
grad_W_xh = np.zeros_like(W_xh)
grad_W_hh = np.zeros_like(W_hh)
grad_W_hy = np.zeros_like(W_hy)
grad_b_h = np.zeros(cache['z'].shape[-1])
grad_b_y = np.zeros(y_seq.shape[-1])
delta_z = [None] * S
delta_y = [None] * S
for t in range(S - 1, -1, -1): # 倒序遍历时间轴
dy = y_seq[t] - target[:, t, :] # δ_{y_t} = ∂L/∂y_t = y_t - target_t
dh = np.dot(dy, W_hy.T) # 直接路径:经 W_hy 回传
if t + 1 < S: # 时间路径:来自后一刻的 δ_{z_{t+1}}
dh = dh + np.dot(delta_z[t + 1], W_hh.T)
dz = dh * (1.0 - cache['h'][t] ** 2) # δ_t = δ_h ⊙ (1 - h_t^2)
delta_z[t] = dz
delta_y[t] = dy
grad_W_hy += np.dot(cache['h'][t].T, dy) # Σ_t h_t^T δ_{y_t}
grad_W_hh += np.dot(cache['h_prev'][t].T, dz) # Σ_t h_{t-1}^T δ_t
grad_W_xh += np.dot(cache['x'][t].T, dz) # Σ_t x_t^T δ_t
grad_b_h += dz.sum(axis=0)
grad_b_y += dy.sum(axis=0)
return (grad_W_xh, grad_W_hh, grad_W_hy,
grad_b_h, grad_b_y, delta_z, delta_y)
# ---------------- 数值梯度(中心差分) ----------------
def numerical_grads(x, target, W_xh, W_hh, W_hy, bh, by, h0, eps=1e-6):
"""对每个参数矩阵逐个加微小扰动,用 (f(x+e)-f(x-e))/2e 估计梯度"""
def loss_with(Wx, Wh, Wy, bhh, byy):
_, ys, _ = rnn_forward(x, Wx, Wh, Wy, bhh, byy, h0)
return mse_loss(ys, target)
grads = {}
for name, P in [('W_xh', W_xh), ('W_hh', W_hh), ('W_hy', W_hy),
('b_h', bh), ('b_y', by)]:
G = np.zeros_like(P)
for idx in np.ndindex(P.shape):
orig = P[idx]
P[idx] = orig + eps
fp = loss_with(W_xh, W_hh, W_hy, bh, by)
P[idx] = orig - eps
fm = loss_with(W_xh, W_hh, W_hy, bh, by)
P[idx] = orig
G[idx] = (fp - fm) / (2 * eps)
grads[name] = G
return grads
# ================= 实操 A:小模型手算反向 =================
def demo_A():
print(LINE)
print('实操 A:小模型手算反向 ------ 从损失出发,把梯度一步步"追"出来')
print('模型与第3章实操B完全相同:I=2, H=3, O=2, B=1, S=2')
print(LINE)
h_seq, y_seq, cache = rnn_forward(x, w_xh, w_hh, w_hy, bh, by, h0)
L = mse_loss(y_seq, target)
print('前向结果(与第3章实操B逐位一致):')
for t in range(2):
print(' t=%d: z_t = %s' % (t, np.round(cache['z'][t, 0], 4)))
print(' h_t = %s' % np.round(h_seq[t, 0], 4))
print(' y_t = %s' % np.round(y_seq[t, 0], 4))
print('目标值 target: t=0: %s t=1: %s' % (target[0, 0], target[0, 1]))
print('损失 L = 1/2*Σ||y_t - target_t||^2 = %.6f' % L)
# 第 0 步:输出层误差 δ_y = ∂L/∂y = y - target
print('\n[第0步] 输出层误差 δ_{y_t} = y_t - target_t(MSE 的导数):')
for t in range(2):
dy = y_seq[t] - target[:, t, :]
print(' δ_y%d = y_%d - target_%d = %s' % (t, t, t, np.round(dy[0], 4)))
# 第 1 步:最后一步 δ_z(只有直接路径)
print('\n[第1步] 最后一步 t=1:δ_{z_1} = (δ_{y_1} @ W_hy^T) ⊙ (1 - h_1^2)')
dy1 = y_seq[1] - target[:, 1, :]
dh1 = np.dot(dy1, w_hy.T)
dz1 = dh1 * (1.0 - h_seq[1] ** 2)
print(' δ_{y_1} @ W_hy^T = %s' % np.round(dh1[0], 4))
print(' (1 - h_1^2) = %s' % np.round((1.0 - h_seq[1] ** 2)[0], 4))
print(' δ_{z_1} = %s' % np.round(dz1[0], 4))
# 第 2 步:t=0 的 δ_z(直接路径 + 时间路径)
print('\n[第2步] t=0:δ_{z_0} = (δ_{y_0} @ W_hy^T + δ_{z_1} @ W_hh^T) ⊙ (1 - h_0^2)')
dy0 = y_seq[0] - target[:, 0, :]
dh0_direct = np.dot(dy0, w_hy.T)
dh0_time = np.dot(dz1, w_hh.T)
dh0 = dh0_direct + dh0_time
dz0 = dh0 * (1.0 - h_seq[0] ** 2)
print(' 直接路径 δ_{y_0} @ W_hy^T = %s' % np.round(dh0_direct[0], 4))
print(' 时间路径 δ_{z_1} @ W_hh^T = %s' % np.round(dh0_time[0], 4))
print(' 两条路径相加,再 ⊙ (1 - h_0^2)')
print(' δ_{z_0} = %s' % np.round(dz0[0], 4))
# 第 3 步:三个梯度
print('\n[第3步] 用 δ 拼出三个梯度(这就是 BPTT 的终点站):')
gW_xh, gW_hh, gW_hy, gb_h, gb_y, _, _ = bptt_analytic(
y_seq, target, cache, w_xh, w_hh, w_hy)
print(' ∂L/∂W_hy = Σ_t h_t^T δ_{y_t}')
print(' = %s' % np.round(gW_hy, 4))
print(' ∂L/∂W_hh = Σ_t h_{t-1}^T δ_t')
print(' = %s' % np.round(gW_hh, 4))
print(' ∂L/∂W_xh = Σ_t x_t^T δ_t')
print(' = %s' % np.round(gW_xh, 4))
print(' ∂L/∂b_h = %s' % np.round(gb_h, 4))
print(' ∂L/∂b_y = %s' % np.round(gb_y, 4))
# ================= 实操 B:δ 递推 vs 数值梯度 =================
def demo_B():
print('\n' + LINE)
print('实操 B:δ 递推公式验算 ------ 对 z_t 逐个扰动,数值梯度应等于递推结果')
print(LINE)
h_seq, y_seq, cache = rnn_forward(x, w_xh, w_hh, w_hy, bh, by, h0)
_, _, _, _, _, delta_z, _ = bptt_analytic(
y_seq, target, cache, w_xh, w_hh, w_hy)
def loss_given_z(z0=None, z1=None):
"""注入某个 z_t(其余时间步走真实前向),重放完整网络计算损失。
注意:z_1 依赖 h_0 = tanh(z_0),所以扰动 z_0 时,z_1 必须由
真实递推算出(不能同时注入 z_1),否则时间依赖就断了。
"""
x0, x1 = x[0, 0, :], x[0, 1, :] # 两条输入 (I,)
if z0 is None:
z0 = np.dot(x0, w_xh) + np.dot(h0, w_hh) + bh
h0_t = np.tanh(z0)
y0 = np.dot(h0_t, w_hy) + by
loss = 0.5 * np.sum((y0 - target[0, 0, :]) ** 2)
if z1 is None:
z1 = np.dot(x1, w_xh) + np.dot(h0_t, w_hh) + bh
h1_t = np.tanh(z1)
y1 = np.dot(h1_t, w_hy) + by
loss += 0.5 * np.sum((y1 - target[0, 1, :]) ** 2)
return loss
eps = 1e-6
max_err = 0.0
z0_base, z1_base = cache['z'][0].copy(), cache['z'][1].copy()
for t in range(2):
G_num = np.zeros_like(cache['z'][t])
z_base = z0_base if t == 0 else z1_base
for idx in np.ndindex(z_base.shape):
orig = z_base[idx]
z_base[idx] = orig + eps
if t == 0:
fp = loss_given_z(z0=z_base, z1=None)
else:
fp = loss_given_z(z0=None, z1=z_base)
z_base[idx] = orig - eps
if t == 0:
fm = loss_given_z(z0=z_base, z1=None)
else:
fm = loss_given_z(z0=None, z1=z_base)
z_base[idx] = orig
G_num[idx] = (fp - fm) / (2 * eps)
err = np.abs(G_num - delta_z[t]).max()
max_err = max(max_err, err)
print('t=%d: 数值 ∂L/∂z_t = %s' % (t, np.round(G_num[0], 6)))
print(' 递推 δ_t = %s' % (np.round(delta_z[t][0], 6),))
print(' 最大误差 = %.3e' % err)
print('-> 最大误差 %.3e:递推公式与数值梯度一致,推导没有算错' % max_err)
# ================= 实操 C:W_hh 梯度 = 时间轴累加和 =================
def demo_C():
print('\n' + LINE)
print('实操 C:为什么 ∂L/∂W_hh 是"时间轴累加和"?------ 逐项拆开看')
print(LINE)
h_seq, y_seq, cache = rnn_forward(x, w_xh, w_hh, w_hy, bh, by, h0)
_, _, _, _, _, delta_z, _ = bptt_analytic(
y_seq, target, cache, w_xh, w_hh, w_hy)
print('t=0 的贡献项: h_{-1}^T δ_0 = 0^T δ_0 = 全零矩阵(h_{-1}=0,起点没有记忆)')
contrib0 = np.dot(cache['h_prev'][0].T, delta_z[0])
print(' -> 全零?', np.abs(contrib0).max() < 1e-15)
contrib1 = np.dot(cache['h_prev'][1].T, delta_z[1])
print('t=1 的贡献项: h_0^T δ_1 =')
print(np.round(contrib1, 4))
total = contrib0 + contrib1
print('\n两项相加 ∂L/∂W_hh =')
print(np.round(total, 4))
ng = numerical_grads(x, target, w_xh, w_hh, w_hy, bh, by, h0)
print('数值梯度 ∂L/∂W_hh =')
print(np.round(ng['W_hh'], 4))
err = np.abs(total - ng['W_hh']).max()
print('最大误差 = %.3e(累加和公式完全正确)' % err)
print('\n三个权重 + 两个偏置的数值梯度体检(中心差分 eps=1e-6):')
gW_xh, gW_hh, gW_hy, gb_h, gb_y, _, _ = bptt_analytic(
y_seq, target, cache, w_xh, w_hh, w_hy)
for name, ga, gn in [('W_xh', gW_xh, ng['W_xh']),
('W_hh', gW_hh, ng['W_hh']),
('W_hy', gW_hy, ng['W_hy']),
('b_h', gb_h, ng['b_h']),
('b_y', gb_y, ng['b_y'])]:
e = np.abs(ga - gn).max()
print(' %-5s 解析 vs 数值 最大误差 = %.3e %s'
% (name, e, 'OK' if e < 1e-4 else 'FAIL'))
# ================= 实操 D:梯度消失 / 爆炸数值实验 =================
def demo_D():
print('\n' + LINE)
print('实操 D:梯度消失 / 爆炸 ------ 谱范数 ρ 决定梯度是"衰减"还是"爆炸"')
print('实验A(纯连乘):去掉 tanh,梯度范数 = ρ^k,指数规律一目了然;')
print('实验B(真实tanh):同样缩放的 W_hh 跑真实 RNN,(1-h_t^2) 因子再插一脚。')
print(LINE)
np.random.seed(7)
H = 16
W = np.random.randn(H, H) / np.sqrt(H)
rho_vals = [0.5, 0.9, 1.05, 1.5]
S = 30
x_seq = np.random.randn(S, H) * 0.05
# ---- 实验A:线性展开,范数严格等于 ρ^k ----
print('\n实验A:纯连乘 ||δ_k|| = ρ^k(把最后一步误差当单位向量往回传)')
print('回传步数 k | ' + ' | '.join('ρ=%.2f' % r for r in rho_vals))
print('-' * (14 + 15 * len(rho_vals)))
for k in [0, 3, 7, 11, 15, 19, 23, 29]:
row = 'k=%-10d| ' % k
for rho in rho_vals:
row += '%.4e | ' % (rho ** k)
print(row)
print('\nρ=0.5:回传 30 步,梯度缩水到 %.2e(十亿分之一!)' % (0.5 ** 30))
print('ρ=1.5:回传 30 步,梯度放大 %.0f 倍(19万倍!)' % (1.5 ** 30))
print('ρ=1.05:只比 1 大 5%%,30 步后也放大 %.1f 倍' % (1.05 ** 30))
# ---- 实验B:真实 tanh RNN,观察 (1-h_t^2) 因子的影响 ----
print('\n实验B:真实 tanh RNN(同样的 W_hh 缩放,h 会饱和)')
print('回传步数 k | ' + ' | '.join('ρ=%.2f' % r for r in rho_vals))
print('-' * (14 + 15 * len(rho_vals)))
table = []
for rho in rho_vals:
W_hh = W * (rho / np.linalg.norm(W, 2))
h = np.random.randn(H) * 0.1
h_list = []
for t in range(S):
z = np.dot(h, W_hh) + x_seq[t]
h = np.tanh(z)
h_list.append(h)
dz = np.ones(H) / np.sqrt(H)
norms = [np.linalg.norm(dz)]
for t in range(S - 1, -1, -1):
dz = np.dot(dz, W_hh.T) * (1.0 - h_list[t] ** 2)
norms.append(np.linalg.norm(dz))
table.append((rho, norms))
for k in [0, 3, 7, 11, 15, 19, 23, 29]:
row = 'k=%-10d| ' % k
for _, norms in table:
row += '%.4e | ' % norms[k]
print(row)
print('\n对比:ρ=1.5 时纯连乘应放大 19 万倍,但真实 tanh 下只有 %.2f 倍------'
% (table[3][1][30] / table[3][1][0]))
print('原因:h 被 tanh 压到饱和区后 (1-h^2)→0,把爆炸"拖慢"了;')
print('而 ρ=0.5 时 tanh 让消失更彻底:缩水到 %.2e(比纯连乘还狠)。'
% (table[0][1][30] / table[0][1][0]))
def main():
demo_A()
demo_B()
demo_C()
demo_D()
print('\n' + LINE)
print('四个实操全部跑通!下一章(第5章)将把公式写成完整 rnn_backward,')
print('并用梯度裁剪驯服上面这种"爆炸"。')
if __name__ == '__main__':
main()