python神经网络编程入门(十七)——RNN BPTT算法(上):梯度沿时间轴传播的链式法则

引言:前向传播把路铺好了,这一章让梯度"倒着走"

先花 30 秒回顾一下上一篇(第十六篇)。我们把 RNN 的前向传播真正"跑"了起来:

  1. 函数落地 ------rnn_forward(x, W_xh, W_hh, W_hy, b_h, b_y, h0) 用一行 for t in range(S) 统治了所有时间步,返回 (h_seq, y_seq, cache)
  2. 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,⋅) 的"口粮",就是为了今天;
  3. 维度验证 ------输出 (10,32,256)(10,32,256)(10,32,256) 与第二章推演表逐位一致,前向没写错;
  4. 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),差在哪

学完本章,你要能独立完成三件事:

🎯 本章三大目标

  1. 在纸上完整推导出三个梯度公式:∂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;
  2. 用自己的话讲清楚:为什么 WhhW_{hh}Whh 的梯度是"沿时间轴的累加和",中间变量 δt\delta_tδt 是怎么递推出来的;
  3. 用代码把推导"验算"一遍:递推公式 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,再算 −log⁡pt-\log p_t−logpt。它的输出层误差长得更漂亮:

L=−∑tlog⁡pt,δ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−tanh⁡21 - \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−tanh⁡2(z)1 - \tanh^2(z)1−tanh2(z)。而我们缓存里存的是 ht=tanh⁡(zt)h_t = \tanh(z_t)ht=tanh(zt) 本身,所以:

dhtdzt=1−tanh⁡2(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−tanh⁡2(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.41140.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.41140.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(累加和公式完全正确)

两个细节值得记住:

  1. t=0t=0t=0 的贡献是零矩阵 ------因为 h−1=h0=0h_{-1} = h_0 = \mathbf{0}h−1=h0=0(全零起点,第三章定的惯例)。起点没有"上一步记忆",自然不背锅;
  2. 累加和 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(上),完结。


🧠 思考题与动手练习

思考题(先自己想,再看答案区,答案就在正文里):

  1. 为什么 WhhW_{hh}Whh 的梯度要"跨时间步累加",而 WhyW_{hy}Why 不用?(提示:谁被共享,谁就累加)
  2. 如果去掉递推式里的 (1−ht2)(1-h_t^2)(1−ht2) 因子,梯度会变大还是变小?误差会出现在哪里?(提示:tanh 闸门的作用)
  3. 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 路径)
  4. 谱范数 ρ=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)
  5. 为什么实际训练中"梯度消失"比"梯度爆炸"更常见?(提示:(1−ht2)≤1(1-h_t^2) \le 1(1−ht2)≤1 恒成立)

动手练习(改造本章实操代码):

  1. 把 MSE 换成二分类交叉熵 :yty_tyt 过 sigmoid,δyt=σ(yt)−tt\delta_{y_t} = \sigma(y_t) - t_tδyt=σ(yt)−tt,重新跑实操 A,验证梯度变化;
  2. 把 t=1t=1t=1 的贡献项单独放大 101010 倍再累加,看看数值梯度体检还过不过------体会"累加和"对每一项都敏感;
  3. 在实操 D 里把 ρ\rhoρ 取 0.990.990.99 和 1.011.011.01,对比回传 30 步的梯度范数,体会"分水岭"有多窄;
  4. 把 HHH 从 3 改成 5 重跑实操 A(权重随机),验证递推公式与数值梯度依然一致;
  5. 手写一份"推导清单":不查资料,从 δ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()
相关推荐
千禧皓月1 小时前
【深度学习】参数量和GFLOPs的计算
深度学习·参数量·gflops·thop·torchinfo
韩师傅2 小时前
重生之我成为模型 番外 · 饲养员厨房下——大厨古法餐
深度学习·机器学习·计算机视觉
幻影123!12 小时前
从零训练一个会下五子棋的AI
python·深度学习·神经网络·强化学习·五子棋·alpha zero·mokugo
mingo_敏13 小时前
DeepAgents : 检索(Retrieval)
人工智能·深度学习·langchain
LaughingZhu15 小时前
Product Hunt 每日热榜 | 2026-08-01
人工智能·深度学习·神经网络·搜索引擎·百度
卡梅德生物科技小能手18 小时前
卡梅德生物科普 TNFSF4(肿瘤坏死因子超家族成员 4)
经验分享·深度学习·生活
逻辑君19 小时前
ANNA 认知引擎 · Humanoid 机器人训练白皮书
人工智能·深度学习·机器学习·机器人
hans汉斯19 小时前
计算机科学与应用|改进MeanShift算法在智能监控视频中的应用研究
图像处理·人工智能·功能测试·深度学习·算法·音视频
余俊晖1 天前
多模态大模型细粒度视觉理解:Vision-OPD在线策略自蒸馏技术方案概述
人工智能·深度学习·算法·多模态·opd