1104 字
6 分钟
RNN-BPTT,梯度消失与梯度爆炸以及对应的解决方案
2026-06-26
无标签

回顾符号定义,前向传播的过程#

参数符号含义
UU输入到隐藏层的权重矩阵
WW隐藏层到隐藏层的权重矩阵(循环权重)
VV隐藏层到输出层的权重矩阵
bb隐藏层的偏置向量
cc输出层的偏置向量
中间量含义
a(t)a^{(t)}tt 时刻隐藏层的净输入(激活前)
h(t)h^{(t)}tt 时刻隐藏层的输出(激活后)
o(t)o^{(t)}tt 时刻的输出层净输入
y^(t)\hat{y}^{(t)}tt 时刻的预测输出(经softmax后)
y(t)y^{(t)}tt 时刻的真实标签
L(t)L^{(t)}tt 时刻的损失值

前向传播公式:

路径:x(t)a(t)h(t)o(t)y^(t)x^{(t)} \to a^{(t)} \to h^{(t)} \to o^{(t)} \to \hat{y}^{(t)}

a(t)=b+Wh(t1)+Ux(t)a^{(t)} = b + Wh^{(t-1)} + Ux^{(t)}

h(t)=tanh(a(t))h^{(t)} = \tanh(a^{(t)})

o(t)=c+Vh(t)o^{(t)} = c + Vh^{(t)}

y^(t)=softmax(o(t))\hat{y}^{(t)} = \text{softmax}(o^{(t)})

BPTT (Backpropagation Through Time)过程#

路径:L(t)o(t)h(t)(a(t))h(t1)...h(1)L^{(t)} \to o^{(t)} \to h^{(t)} (\to a^{(t)}) \to h^{(t-1)}\to ... \to h^{(1)}

1. 对总损失求各时刻损失的梯度#

LL(t)=1\frac{\partial L}{\partial L^{(t)}} = 1

\nabla_{*}表示对 * 求梯度,在别的文献中也叫 δ\delta_{*}

2. 对输出层的净输入求梯度#

损失函数对输出层净输入的梯度 o(t)Ly^(t)y(t)\nabla_{o^{(t)}} L \triangleq \hat{y}^{(t)} - y^{(t)}

o(t)L=Lo(t)=LL(t)L(t)o(t)=y^(t)y(t)\nabla_{o^{(t)}} L = \frac{\partial L}{\partial o^{(t)}} = \frac{\partial L}{\partial L^{(t)}} \cdot \frac{\partial L^{(t)}}{\partial o^{(t)}}= \hat{y}^{(t)} - y^{(t)}

3. 对隐藏层的输出求梯度#

单个时间步内的损失函数对隐藏层输出的梯度 h(t)L(t)VTo(t)L\nabla_{h^{(t)}} L^{(t)}\triangleq V^T \nabla_{o^{(t)}} L ,是下面的方框项

整体损失函数对隐藏层输出的梯度 δ(t)Lh(t)\delta^{(t)} \triangleq \dfrac{\partial L}{\partial h^{(t)}} :

δ(t)=Lh(t)=L(t)h(t)+L(t+1)h(t)+L(t+2)h(t)+=L(t)h(t)+h(t+1)h(t)L(t+1)h(t+1)+h(t+2)h(t+1)h(t+1)h(t)L(t+2)h(t+2)+=L(t)h(t)+k=1τth(t)L(t+k)\begin{aligned} \delta^{(t)}=\frac{\partial L}{\partial h^{(t)}} &= \boxed{\frac{\partial L^{(t)}}{\partial h^{(t)}}} + \frac{\partial L^{(t+1)}}{\partial h^{(t)}} + \frac{\partial L^{(t+2)}}{\partial h^{(t)}} + \ldots \\[1em] &= \boxed{\frac{\partial L^{(t)}}{\partial h^{(t)}}} + \frac{\partial h^{(t+1)}}{\partial h^{(t)}} \frac{\partial L^{(t+1)}}{\partial h^{(t+1)}} + \frac{\partial h^{(t+2)}}{\partial h^{(t+1)}}\frac{\partial h^{(t+1)}}{\partial h^{(t)}} \frac{\partial L^{(t+2)}}{\partial h^{(t+2)}} + \ldots\\ &=\boxed{\frac{\partial L^{(t)}}{\partial h^{(t)}}} + \sum_{k=1}^{\tau-t} \frac{\partial}{\partial h^{(t)}} L^{(t+k)} \end{aligned}h(t)L(t+k)=h(t+1)h(t)h(t+2)h(t+1)L(t+k)h(t+k)=(j=1kh(t+j)h(t+j1))L(t+k)h(t+k)\frac{\partial}{\partial h^{(t)}} L^{(t+k)} = \frac{\partial h^{(t+1)}}{\partial h^{(t)}} \cdot \frac{\partial h^{(t+2)}}{\partial h^{(t+1)}} \cdots \frac{\partial L^{(t+k)}}{\partial h^{(t+k)}}= \left(\prod_{j=1}^{k} \frac{\partial h^{(t+j)}}{\partial h^{(t+j-1)}}\right) \frac{\partial L^{(t+k)}}{\partial h^{(t+k)}}
  • 对于方框里面的部分,L(t)L^{(t)} 只依赖于 h(t)h^{(t)},所以可以直接求导;

  • 对于后续时刻的 L(t+1),L(t+2),L^{(t+1)}, L^{(t+2)}, \ldots,它们依赖于 h(t)h^{(t)} 是通过 h(t+1),h(t+2),h^{(t+1)}, h^{(t+2)}, \ldots 传递,所以需要用链式法则展开,但是展开依旧是一坨乘积,没法算

  • 距离越远,乘积越长,容易出现梯度消失或梯度爆炸问题,导致训练困难。

注意递推关系,把 tt 换成 t+1t+1

δ(t+1)=Lh(t+1)=L(t+1)h(t+1)+L(t+2)h(t+1)+L(t+3)h(t+1)+=L(t+1)h(t+1)+h(t+2)h(t+1)L(t+2)h(t+2)+h(t+3)h(t+2)h(t+2)h(t+1)L(t+3)h(t+3)+\begin{aligned} \delta^{(t+1)}=\frac{\partial L}{\partial h^{(t+1)}} &= \frac{\partial L^{(t+1)}}{\partial h^{(t+1)}} + \frac{\partial L^{(t+2)}}{\partial h^{(t+1)}} + \frac{\partial L^{(t+3)}}{\partial h^{(t+1)}} + \ldots \\[1em] &= \frac{\partial L^{(t+1)}}{\partial h^{(t+1)}} + \frac{\partial h^{(t+2)}}{\partial h^{(t+1)}} \frac{\partial L^{(t+2)}}{\partial h^{(t+2)}} + \frac{\partial h^{(t+3)}}{\partial h^{(t+2)}}\frac{\partial h^{(t+2)}}{\partial h^{(t+1)}} \frac{\partial L^{(t+3)}}{\partial h^{(t+3)}} + \ldots\\[1em] \end{aligned}

这个两边同时乘以 h(t+1)h(t)\dfrac{\partial h^{(t+1)}}{\partial h^{(t)}},就可以得到 δ(t)\delta^{(t)} 非方框 中的部分!

重大发现:δ(t)\delta^{(t)} 可以递推计算:

δ(t)=Lh(t)=L(t)h(t)+(h(t+1)h(t))Tδ(t+1)=(o(t)h(t))TL(t)o(t)+(h(t+1)h(t))TLh(t+1)=VTo(t)L+WTdiag(1(h(t+1))2)h(t+1)L\begin{aligned} \delta^{(t)}=\dfrac{\partial L}{\partial h^{(t)}}&= \frac{\partial L^{(t)}}{\partial h^{(t)}} + \left(\frac{\partial h^{(t+1)}}{\partial h^{(t)}}\right)^T \delta^{(t+1)}\\[5pt] &= \left(\frac{\partial o^{(t)}}{\partial h^{(t)}}\right)^T \frac{\partial L^{(t)}}{\partial o^{(t)}} + \left(\frac{\partial h^{(t+1)}}{\partial h^{(t)}}\right)^T \frac{\partial L}{\partial h^{(t+1)}} \\[1em] &= V^T \nabla_{o^{(t)}} L + W^T\text{diag}(1 - (h^{(t+1)})^2) \nabla_{h^{(t+1)}} L \end{aligned}

其中 h(t+1)h(t)\dfrac{\partial h^{(t+1)}}{\partial h^{(t)}} 是通过链式法则展开得到的:

h(t+1)=tanh(a(t+1))a(t+1)=b+Wh(t)+Ux(t+1)h(t+1)h(t)=h(t+1)a(t+1)a(t+1)h(t)=WTdiag(1(h(t+1))2)\begin{aligned} h^{(t+1)} &= \tanh(a^{(t+1)})\\ a^{(t+1)} &= b + W\boxed{h^{(t)}} + Ux^{(t+1)}\\[2pt] \Rightarrow \frac{\partial h^{(t+1)}}{\partial h^{(t)}} &= \frac{\partial h^{(t+1)}}{\partial a^{(t+1)}} \cdot \frac{\partial a^{(t+1)}}{\partial h^{(t)}} = W^T\text{diag}(1 - (h^{(t+1)})^2) \end{aligned}

对参数求梯度#

注意每个参数在所有时刻共享,所以要对所有时刻求和。

UL=tdiag(1(h(t))2)h(t)L(t)(x(t))TWL=tdiag(1(h(t))2)h(t)L(t)(h(t1))TVL=to(t)L(h(t))TbL=tdiag(1(h(t))2)h(t)LcL=to(t)L\begin{aligned} \nabla_{U} L &= \sum_t \text{diag}(1 - (h^{(t)})^2) \nabla_{h^{(t)}} L^{(t)} \cdot (x^{(t)})^T\\ \nabla_{W} L &= \sum_t \text{diag}(1 - (h^{(t)})^2) \cdot \nabla_{h^{(t)}} L^{(t)} \cdot (h^{(t-1)})^T\\ \nabla_{V} L &= \sum_t \nabla_{o^{(t)}} L \cdot (h^{(t)})^T\\ \nabla_{b} L &= \sum_t \text{diag}(1 - (h^{(t)})^2) \cdot \nabla_{h^{(t)}} L\\ \nabla_{c} L &= \sum_t \nabla_{o^{(t)}} L \end{aligned}

长期依赖#

当间隔时间步很长时,δ(t)\delta^{(t)} 中的乘积项 j=1kh(t+j)h(t+j1)\prod_{j=1}^{k} \dfrac{\partial h^{(t+j)}}{\partial h^{(t+j-1)}} 会导致梯度消失或梯度爆炸问题,训练困难。

在消失的情况下,tt 时刻的梯度对 t+kt+k 时刻以及之前的损失几乎没有影响,导致模型无法学习长期依赖关系。

梯度消失与梯度爆炸的解决方案#

1. 截断梯度#

实际应用中,两种方式性能表现类似

  • 方式1:在参数更新之前,逐元素地截断Mini-batch 产生的参数梯度
  • 方式2:在参数更新之前,整体约束参数梯度大小 (不改变梯度方向)

2. 时间维度的跳跃连接#

直接构造从 tt 时刻单元到 t+dt + d 时刻单元的连接

3. 渗漏单元#

对于隐藏层之间的连接:

从开始的 h(t)=tanh(b+Wh(t1)+Ux(t))h^{(t)} = \tanh(b + Wh^{(t-1)} + Ux^{(t)}) 改结构成:

h(t)=αh(t1)+(1α)g(x(t),h(t1))h^{(t)} = \alpha h^{(t-1)} + (1 - \alpha) g(x^{(t)}, h^{(t-1)})
  • gg 可以是任意的非线性函数,比如 g(x(t),h(t1))=tanh(b+Wh(t1)+Ux(t))g(x^{(t)}, h^{(t-1)}) = \tanh(b + Wh^{(t-1)} + Ux^{(t)}),也可以是其他的函数。
  • α\alpha 是一个小于 1 的常数,表示“渗漏”比例。它可以让梯度在时间维度上有一个“泄漏”,从而缓解梯度消失问题。
  • α1\alpha\approx 1 容易饱和。
  • α0\alpha\approx 0 退化为普通RNN。
RNN-BPTT,梯度消失与梯度爆炸以及对应的解决方案
https://biscuit0613.github.io/posts/ml/rnn-bptt/
作者
Biscuit
发布于
2026-06-26
许可协议
CC BY-NC-SA 4.0
视觉先验-在神经网络结构中的体现
RNN-双向RNN