RNN公式推导与BPTT反向传播:从结构设计到梯度消失爆炸的深度解析
2026/9/17 17:09:36 网站建设 项目流程

说实话,做深度学习最常遇到的一个现象就是:用 PyTorch 调 RNN 就像逛淘宝一样顺手的同学,你问他一句“反向传播时误差到底是怎么沿着时间步传回去的”,他大概率会愣住。循环神经网络、RNN 这些词在简历上写得滚瓜烂熟,但一落到公式推导,就立刻暴露出理解深度不够。这篇记录的初衷就是把 RNN 的公式推导从头到尾手写一遍,从网络结构怎么设计、前向传播每一步发生了什么,到 BPTT 反向传播的链式法则怎么展开,再到梯度消失和爆炸到底是怎么从公式里长出来的,都掰开揉碎讲清楚。如果你正处在深度学习的入门阶段,或者马上要准备面试、期末考,又或者已经调了很久 RNN 但总觉得心里没底,这篇内容应该能帮你把最后一块拼图补上。

1. RNN的结构设计与整体思路

1.1 RNN为什么需要循环结构

先用一句话概括:全连接网络和卷积神经网络都是针对“一个独立样本”设计的,输入输出的形状是固定的,但真实世界里有大量数据是序列——一段文本、一段语音、几天的股价、一个视频的多帧画面,都天然是“一串”而不是“一个”。处理这种数据时,我们希望模型能看到前后文,在决定 t 时刻的输出时,不仅依赖当前输入,还能利用前面几个时刻的信息。

RNN 的做法极其朴素:在当前时刻隐藏层的输入里,除了当前时刻的特征 $x_t$,再把上一时刻的隐藏状态 $h_{t-1}$ 一起拼接进来。这就形成了一个循环结构:$h_t$ 依赖于 $h_{t-1}$,而 $h_{t-1}$ 又依赖于 $h_{t-2}$,信息就这样沿着时间轴一级一级往后传。你甚至可以把它理解成同一个全连接网络被“复制”了一份用于每个时间步,且所有时间步共享同一套参数。这个参数共享是关键,因为序列长度变化很大,如果不共享参数,模型根本无法泛化到没见过的长度。

1.2 状态变量h_t到底在存什么

很多初学者把 $h_t$ 理解成“当前输出的一个中间变量”,这倒也没错,但不够本质。更准确的说法是:$h_t$ 是网络在 t 时刻维护的“记忆状态”,它编码了从序列起点到当前时刻所有输入信息的压缩摘要。这个摘要不是搞了一个像链表一样逐字存储的结构,而是把过去的信息压成了一个固定维度的向量。

因为压缩,它当然会丢信息,也不是所有历史信息都等权保留。这也解释了为什么普通 RNN 不太擅长捕捉特别长距离的依赖——信息经过多个时间步的“重复压缩”之后,早先的细节会越来越稀薄。LSTM 和 GRU 后来做的门控机制,本质上是给这个“记忆状态”加了一套读写控制,让信息可以更完整地穿越更长时间。但在理解 LSTM 之前,把标准 RNN 吃透是绕不开的一步,因为后面所有变体都是在这个循环框架上加结构的。

1.3 与全连接网络和CNN的对比

用一个表格可以很直观地看出 RNN 和另外两个家族的区别:

维度全连接网络CNNRNN
核心假设特征间相对独立,输入输出固定局部空间相关性,空间位置平移不变性序列时间相关性,前后依赖
输入形式向量网格结构(图像等)序列(长度可变)
参数共享方式不共享卷积核在空间上共享权重矩阵在所有时间步共享
典型任务表格分类、回归图像分类、目标检测文本分类、机器翻译、语音识别
梯度传播路径层间逐层传播空间上逐层传播空间上逐层 + 时间上逐时间步传播

这里最值得注意的就是最后一行。RNN 的反向传播路径额外多了一条时间维度,这会带来完全不同的梯度问题。你训练 CNN 时很少遇到“梯度爆炸导致直接 NaN”,但在 RNN 里这几乎是家常便饭。理解了这一点,你就能明白为什么我们在训练 RNN 时有那么多“小技巧”——那些技巧本质上都是在和这条时间传播路径上的连乘效应作斗争。

2. 前向传播公式逐步拆解

2.1 先把符号约定写清楚

公式推导最怕符号混乱,我在下面统一约定,后续所有推导都基于这套符号,建议你手推的时候也固定一套自己的记号。

  • 输入序列:$x_1, x_2, \ldots, x_T$,其中每个时间步的输入 $x_t \in \mathbb{R}^n$,n 是输入特征维度。
  • 隐藏状态:$h_t \in \mathbb{R}^m$,m 是隐藏层神经元数量,它承担记忆存储和传递的角色。
  • 输出:$o_t \in \mathbb{R}^K$,K 是输出维度。如果是分类任务,K 就是类别数;如果是回归或语言模型,K 对应相应输出维度。
  • 权重矩阵:$W_{xh} \in \mathbb{R}^{m \times n}$ 负责把输入投影到隐藏空间;$W_{hh} \in \mathbb{R}^{m \times m}$ 负责隐藏到隐藏的时间递归;$W_{hy} \in \mathbb{R}^{K \times m}$ 负责从隐藏映射到输出。
  • 偏置:$b_h \in \mathbb{R}^m$、$b_y \in \mathbb{R}^K$。

有些资料也会用 $U$、$W$、$V$ 来分别表示输入到隐藏、隐藏到隐藏、隐藏到输出的权重,比如 $h_t = \tanh(U x_t + W h_{t-1} + b_h)$。为了简便,我在下文混用这两套符号时,会以笔画的清晰为先,但核心原则不变:三个权重矩阵分别对应三个连接关系,推导时看清下标就行。

2.2 隐藏状态更新公式

标准 RNN 的前向传播核心就两个式子。第一个是从输入到隐藏状态的更新:

$$h_t = \tanh(W_{xh} x_t + W_{hh} h_{t-1} + b_h)$$

从计算图的角度看,这个公式做了三件事:先把当前输入 $x_t$ 线性变换到隐藏空间;再把上一时刻的隐藏状态 $h_{t-1}$ 也线性变换到同一个隐藏空间;把两者相加后加偏置,最后用逐元素的 tanh 激活函数做一个压缩和“非线性化”。

这里很多教材直接给出公式就过了,但有个问题值得停下来想清楚:为什么激活函数选 tanh,而不是 ReLU 或者 Sigmoid?

在实际的工程经验里,RNN 的隐藏激活函数几乎默认就是 tanh。原因有两层。第一,tanh 的输出范围是 $[-1, 1]$,均值接近 0,这让隐藏状态在不同时间步之间传递时不容易产生系统性偏移;而 Sigmoid 输出是非负的,一个非负的隐藏状态经过反复迭代,很容易在时间维度上不断累积正的偏移量,数值稳定性更差。第二,tanh 的导数最大值是 1(在 0 附近),虽然这个性质不足以根治梯度消失,但至少比 Sigmoid 的 0.25 上限要好四个倍,梯度在时间维度上能多撑一段时间。

那 ReLU 呢?ReLU 的正区间导数恒为 1,理论上最有利于缓解梯度消失,但在标准 RNN 里它有一个很头疼的问题:隐藏状态 $h_t$ 本身会进入下一时刻的输入,如果某个神经元一直处于正区间,它的值可能一路朝正方向增长没有约束,导致隐藏状态发散,数值直接溢出。所以我个人在搭建标准 RNN 时,tanh 永远是默认第一选择,这不算创新,只是前人踩过太多坑后的共识。

2.3 输出层与损失函数

第二个前向公式是从隐藏状态到输出的变换:

$$o_t = W_{hy} h_t + b_y$$

如果做分类任务,通常在 $o_t$ 后面接一个 softmax,把它变成类别概率分布:

$$\hat{y}t = \operatorname{softmax}(o_t) = \frac{\exp(o_t)}{\sum_k \exp(o{t,k})}$$

这里的 $\hat{y}_t$ 是模型预测的 K 类概率分布,而 $y_t$ 表示真实标签的 one-hot 向量。对单个时间步,我们定义交叉熵损失:

$$L_t = -\sum_{k=1}^{K} y_{t,k} \log \hat{y}_{t,k}$$

整个序列的损失通常取所有时间步之和(也可以取平均,看任务风格,但数学上没有本质差别):

$$L = \sum_{t=1}^{T} L_t$$

这里有一个在公式推导中功劳最大、但教材往往一笔带过的结论:当输出层使用 softmax + 交叉熵损失时,损失对输出 $o_t$ 的导数非常简洁:

$$\frac{\partial L}{\partial o_t} = \hat{y}_t - y_t$$

这个结论值得单独拿出来说。很多人第一次看到时觉得是魔法,其实推导也不复杂:交叉熵对第 i 类 logit 的偏导,结合 softmax 的雅可比展开,中间那一大坨跨项求和最后恰好全部抵消,剩下一项。这个小结论如果你能自己手推一次,不仅后续 BPTT 流畅很多,而且你在写代码手写梯度检查时也会方便不少。

3. 反向传播BPTT核心推导

3.1 误差项是绕不开的枢纽

反向传播的核心思路永远一条:用链式法则把损失对参数的偏导拆成可计算的中间项。RNN 的反向传播叫 BPTT,Backpropagation Through Time,翻译过来是“随时间反向传播”,意思是它不仅要像普通网络那样从输出层往输入层回传,还需要在时间轴上从最后一个时间步往第一个时间步回传。

为了推导简洁,先定义两个误差项。

第一个是输出误差项,记作:

$$\delta_t^o = \frac{\partial L}{\partial o_t}$$

第二个是隐藏状态误差项,记作:

$$\delta_t^h = \frac{\partial L}{\partial h_t}$$

几乎所有参数的梯度最终都能用这两个误差项表达出来,所以推导的第一个任务就是想办法把 $\delta_t^h$ 算出来。

3.2 误差项沿时间轴递归传播

$h_t$ 在计算图中出现在两条路径上:一条是它的输出侧,直接影响 $o_t$,然后对 $L$ 产生贡献;另一条是它的时间侧,它被当作输入送到下一步计算 $z_{t+1}$,从而间接影响后续所有损失。因此链式法则需要同时考虑这两条路径:

$$\delta_t^h = \frac{\partial L}{\partial o_t} \frac{\partial o_t}{\partial h_t} + \frac{\partial L}{\partial h_{t+1}} \frac{\partial h_{t+1}}{\partial h_t}$$

第一项比较好算,由 $o_t = W_{hy} h_t + b_y$ 可以得到:

$$\frac{\partial o_t}{\partial h_t} = W_{hy}^T$$

所以第一项就是 $W_{hy}^T \delta_t^o$。

第二项稍微绕一下。$h_{t+1}$ 是作用在 $z_{t+1} = W_{hh} h_t + \dots$ 上的 tanh 函数,所以:

$$\frac{\partial h_{t+1}}{\partial h_t} = \frac{\partial h_{t+1}}{\partial z_{t+1}} \frac{\partial z_{t+1}}{\partial h_t}$$

其中:

$$\frac{\partial h_{t+1}}{\partial z_{t+1}} = \operatorname{diag}(1 - h_{t+1}^2)$$

这里要特别强调“diag”这个记号。$h_{t+1}$ 是一个 m 维向量,tanh 是逐元素作用在每个分量上的,所以它对 $z_{t+1}$ 的导数是一个对角矩阵,第 i 个对角元素是 $1 - h_{t+1,i}^2$。初学者常常在这里把维度弄错,写成一个向量,然后后面维度怎么都对不上。

再说:

$$\frac{\partial z_{t+1}}{\partial h_t} = W_{hh}^T$$

这两项一乘,第二项就是:

$$W_{hh}^T \operatorname{diag}(1 - h_{t+1}^2) \delta_{t+1}^h$$

最终得到隐藏状态误差项的递归表达式:

$$\delta_t^h = W_{hy}^T \delta_t^o + W_{hh}^T \operatorname{diag}(1 - h_{t+1}^2) \delta_{t+1}^h$$

边界条件也清晰:最后一个时间步 T 后面没有 $h_{T+1}$,所以:

$$\delta_T^h = W_{hy}^T \delta_T^o$$

整个计算过程是从后往前算的。先算最后一个时间步的 $\delta_T^h$,然后按 $t = T-1, T-2, \ldots, 1$ 的顺序一路回推。这个递归式就是 BPTT 的发动机,也是你面试时最值得在纸上画出来的式子。

3.3 W、U、V 三个权重的梯度

有了 $\delta_t^o$ 和 $\delta_t^h$,梯度求解就变成了“对号入座”。

对输出权重 $W_{hy}$,因为 $L$ 是所有时间步损失之和,而每一步的输出只依赖该步的 $h_t$,所以梯度是每个时间步贡献之和:

$$\frac{\partial L}{\partial W_{hy}} = \sum_{t=1}^T \delta_t^o \otimes h_t = \sum_{t=1}^T \delta_t^o h_t^T$$

加上输出偏置:

$$\frac{\partial L}{\partial b_y} = \sum_{t=1}^T \delta_t^o$$

对隐藏到隐藏的权重 $W_{hh}$,这一步要注意对 $h_t$ 的梯度先乘上 tanh 的逐元素导数。因为 $h_t = \tanh(z_t)$,而 $z_t = W_{hh} h_{t-1} + W_{xh} x_t + b_h$,所以:

$$\frac{\partial L}{\partial W_{hh}} = \sum_{t=1}^T \left( \delta_t^h \odot (1 - h_t^2) \right) h_{t-1}^T$$

其中 $\odot$ 表示逐元素乘法。(有的同学会发现我的式子里这里写的是 $\delta_t^h$ 而不是 $\delta_t^h$ 与另一个量相乘的复杂形式,其实就是把 diag 矩阵与 $\delta_t^h$ 相乘等价写成了逐元素乘,更符合书写习惯。)

对输入到隐藏的权重 $W_{xh}$:

$$\frac{\partial L}{\partial W_{xh}} = \sum_{t=1}^T \left( \delta_t^h \odot (1 - h_t^2) \right) x_t^T$$

隐藏偏置:

$$\frac{\partial L}{\partial b_h} = \sum_{t=1}^T \delta_t^h \odot (1 - h_t^2)$$

这里还有一个细节值得提一下:$h_0$ 是初始隐藏状态,一般初始化为零向量。因为它没有参与任何计算图的生成,所以 $h_0$ 本身没有梯度,不需要更新。

我自己在做维度校验时有个习惯:写完每个梯度公式,先看左右维度是否一致。例如 $W_{hh}$ 的维度是 $m \times m$,右边 $\delta_t^h \odot (1 - h_t^2)$ 是 m 维列向量,$h_{t-1}^T$ 是 $1 \times m$ 行向量,外积是 $m \times m$,齐了。这一招在看 Pytorch 自动求导报错时尤其救命。

4. 梯度消失与爆炸的原因分析

4.1 从公式看指数效应从哪来

前面推导的递归式是分析梯度问题的钥匙。把 $\delta_t^h$ 展开,可以看到误差从时间步 T 传回时间步 t,中间需要连乘一系列因子:

$$\delta_t^h = \prod_{k=t}^{T-1} W_{hh}^T \operatorname{diag}(1 - h_{k+1}^2) \cdot (\text{来自输出层的项})$$

问题就出在这个连乘上。这个式子类似于把同一个矩阵 $W_{hh}^T$ 反复作用在误差信号上,连乘了 $T - t$ 次。如果 $W_{hh}$ 的谱半径(简单理解就是最大特征值的绝对值)小于 1,每乘一次信号都会缩小;经过几十个时间步之后,梯度会小到什么程度?会小到对参数更新几乎没有贡献,这就是梯度消失。如果谱半径大于 1,梯度就会像滚雪球一样指数放大,很快变成 NaN,这就是梯度爆炸。

这件事可以用复利来类比。你往一个年利率 5% 的账户存钱,50 年后翻了 11.5 倍;但利率变成 105% 时,同样 50 年账户就变成了天文数字。RNN 的时间步就是那 50 年的复利周期,$W_{hh}$ 就是利率。普通网络也有梯度连乘问题,但 RNN 因为时间步可以长达几百,而且所有步共享同一个 $W_{hh}$,所以连乘效应尤其极端。

4.2 工程上的应对手段

梯度爆炸虽然在数学上看起来吓人,但工程处理反而最简单直接:梯度截断。

梯度截断的做法是,在每次参数更新前,计算整个梯度的二范数:

$$|g|_2$$

如果它超过预设阈值 $\tau$,就把梯度整体缩放到以 $\tau$ 为范数:

$$g \leftarrow g \cdot \frac{\tau}{|g|_2}$$

阈值我自己随手用过 5,也有不少人固定在 1 左右。这不是什么精密的超参,但能有效防止一次更新把模型参数踢到不可恢复的角落。你在训练 RNN 时如果发现 loss 在某一 step 突然变成 NaN,先查梯度是不是爆了,然后用 clip 基本能止住。

梯度消失就麻烦得多,截断解决不了,因为问题不是“梯度过大”而是“梯度太小,前面时间步啥都学不到”。有几个行之有效的实践方向:

  • 权重初始化:对 $W_{hh}$ 使用正交初始化,能让初始的谱半径接近 1,从源头延缓梯度消失。我自己实验下来,正交初始化比随机高斯初始化在标准 RNN 上稳定很多。
  • 使用截断 BPTT:在长序列上反向传播时,只反传到最近 K 个时间步,而不是一路回传到序列开头,这其实是在“承认模型记不了那么远”。
  • 换结构:把标准 RNN 换成 LSTM 或 GRU。这两种结构通过门控机制引入了一条从 $c_{t-1}$ 到 $c_t$ 的线性传递路径,误差信号在这条路径上不需要经过 tanh 压缩,可以直接无损地传播,因此能有效支持长距离依赖。

很多人是在这里第一次体会到“结构和梯度是绑定在一起的”这个道理。你的网络结构本质上决定了误差信号是否有一条“高速公路”可以安全通行,这也是为什么后来 Transformer 里的残差连接、LayerNorm 都在做同一件事——保证梯度信号在深层网络里可以顺畅回流。

5. 常见问题与排查技巧实录

5.1 训练RNN时最容易踩的坑

我在日常训练和帮别人排查问题的时候,积累了一些很典型的错误,把它们整理成一张速查表,遇到问题可以直接对照:

现象可能原因处理方式
训练刚开始 loss 就出现 NaN梯度爆炸加梯度截断,初始学习率调小到 1e-3 以下
loss 一直不下降,像条平线梯度消失或学习率过小检查 W 初始化,换正交初始化;调大学习率;尝试换 GRU
训练到一半 loss 突然跳高学习率过大导致参数震荡使用学习率衰减或自适应优化器
测试集效果差,但训练集正常模型容量不足或序列长度截断不合理增加隐藏层维度;检查是否忘记 masking padding 部分
不同 batch 的结果波动巨大未合理初始化 $h_0$每个新 batch 重置 $h_0$,或用可学习的初始状态
长序列任务效果很差模型记不住太久以前的信息换 LSTM/GRU;增大隐藏维度;考虑注意力机制

5.2 我的一些训练实战经验

第一个建议是:手推公式和实际调试一定要结合。第一次自己实现 RNN 的手写反向传播时,我把前向、反向所有公式都写在纸上,然后对着 PyTorch 的自动求导结果做数值对比,用 torch.autograd.gradcheck 逐参数核对。那条路虽然慢,但走一遍之后你对 BPTT 的记忆深度远超别人。

第二个建议关乎调试顺序。如果你手写 BPTT 出问题,别上来就调参,先做“小测试”:把时间步长度设为 1,这时 RNN 退化成普通的单隐层网络,梯度公式也应该退化,这是最简单直接的合理性检查;然后再把时间步长度设为 2,手动展开计算图,正向反向都拿笔算一遍,再和代码结果对齐。这样两步走,几乎所有维度错误和公式错误都能暴露出来。

第三个建议是关于 batch 和 padding。RNN 处理变长序列时需要 padding,但 padding 区域没有真实信息,反向传播时如果不对这些位置加 mask,模型会把这些无关位置当成有效信息去学习。实现上很机械:在损失计算和梯度累加时把 padding 位置对应的 loss 置零即可,但漏掉的人非常多。

最后说一个我自己在初始化上的偏好。$W_{hh}$ 用正交初始化,$W_{xh}$ 用 Xavier 初始化,隐藏偏置 $b_h$ 初始化为零。这样做的好处是,在训练早期梯度信号更稳定,尤其是序列比较长的时候。你要是去翻一些经典开源实现,会发现里面也是这么搞的——这不是偶然,而是大家用大量实验“交过学费”之后沉淀下来的默认配置。

6. 从标准RNN到LSTM的一个自然延伸

写到这里,我想再补充一点理解 LSTM 和 GRU 时的视角,因为这个点被问到的频率极高,而且和标准 RNN 的公式推导直接相关。

LSTM 在结构上最核心的改动是引入了一条“细胞状态” $c_t$ 的线性传送带。在标准 RNN 里,$h_t$ 是唯一的信息载体,每次更新都要经过 tanh 压缩,必然磨损历史信息。而 LSTM 里 $h_t$ 和 $c_t$ 是分家的:$c_t$ 可以在“遗忘门”和“输入门”的控制下,以近乎线性的方式更新:

$$c_t = f_t \odot c_{t-1} + i_t \odot \tilde{c}_t$$

$f_t$ 接近 1 的时候,$c_t \approx c_{t-1}$,这条路径对梯度的传递系数接近 1,误差信号就可以顺着这条传送带安全地穿过几十甚至上百个时间步。这正是标准 RNN 的 BPTT 连乘项把梯度啃光之后,LSTM 能“续命”的原因。

你从公式推导的角度来看,这其实就是在解答一个设计问题:既然连乘 $W_{hh}^T \operatorname{diag}(1 - h^2)$ 会导致梯度消失,那就设计一条“不需要经过激活函数压缩”的捷径,让梯度能在时间轴上直接穿过去。理解了这一点,你再看 LSTM 的各种门控公式,就不会觉得是天上掉下来的杰作,而是一个有明确目标导向的工程设计。

GRU 则是 LSTM 的一个精简版本,把遗忘门和输入门合成了更新门,把细胞状态和隐藏状态合并成一个向量,参数更少,训练开销更低,在不少任务上和 LSTM 效果相当。如果在项目里不确定用哪个,我通常会先试 GRU,数据量不大时它的效率和效果往往最均衡。

Transformer 这类基于自注意力的架构虽然在长序列上替代了 RNN,但理解 RNN 的这条学习曲线依旧值得走完。因为序列模型里反复出现的核心问题——如何高效传递信息、如何处理变长序列、如何在长距离依赖和计算效率之间做取舍——在 RNN 的公式里已经全部以最朴素的形式出现过一次。你把标准 RNN 的推导弄清楚,之后无论是看 LSTM、GRU,还是去啃 Transformer 里的注意力计算,都会觉得这只是同一个问题的不同解法,而不是一座座孤岛。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询