ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

RNN公式推导与BPTT反向传播:从结构设计到梯度消失爆炸的深度解析

RNN公式推导与BPTT反向传播:从结构设计到梯度消失爆炸的深度解析 说实话做深度学习最常遇到的一个现象就是用 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_{k1}^{K} y_{t,k} \log \hat{y}_{t,k}$$整个序列的损失通常取所有时间步之和也可以取平均看任务风格但数学上没有本质差别$$L \sum_{t1}^{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 的反向传播叫 BPTTBackpropagation 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_{t1}$从而间接影响后续所有损失。因此链式法则需要同时考虑这两条路径$$\delta_t^h \frac{\partial L}{\partial o_t} \frac{\partial o_t}{\partial h_t} \frac{\partial L}{\partial h_{t1}} \frac{\partial h_{t1}}{\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_{t1}$ 是作用在 $z_{t1} W_{hh} h_t \dots$ 上的 tanh 函数所以$$\frac{\partial h_{t1}}{\partial h_t} \frac{\partial h_{t1}}{\partial z_{t1}} \frac{\partial z_{t1}}{\partial h_t}$$其中$$\frac{\partial h_{t1}}{\partial z_{t1}} \operatorname{diag}(1 - h_{t1}^2)$$这里要特别强调“diag”这个记号。$h_{t1}$ 是一个 m 维向量tanh 是逐元素作用在每个分量上的所以它对 $z_{t1}$ 的导数是一个对角矩阵第 i 个对角元素是 $1 - h_{t1,i}^2$。初学者常常在这里把维度弄错写成一个向量然后后面维度怎么都对不上。再说$$\frac{\partial z_{t1}}{\partial h_t} W_{hh}^T$$这两项一乘第二项就是$$W_{hh}^T \operatorname{diag}(1 - h_{t1}^2) \delta_{t1}^h$$最终得到隐藏状态误差项的递归表达式$$\delta_t^h W_{hy}^T \delta_t^o W_{hh}^T \operatorname{diag}(1 - h_{t1}^2) \delta_{t1}^h$$边界条件也清晰最后一个时间步 T 后面没有 $h_{T1}$所以$$\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_{t1}^T \delta_t^o \otimes h_t \sum_{t1}^T \delta_t^o h_t^T$$加上输出偏置$$\frac{\partial L}{\partial b_y} \sum_{t1}^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_{t1}^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_{t1}^T \left( \delta_t^h \odot (1 - h_t^2) \right) x_t^T$$隐藏偏置$$\frac{\partial L}{\partial b_h} \sum_{t1}^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_{kt}^{T-1} W_{hh}^T \operatorname{diag}(1 - h_{k1}^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 里的注意力计算都会觉得这只是同一个问题的不同解法而不是一座座孤岛。
返回列表