LSTM 为什么能缓解梯度消失?关键在哪里?
简化版
关键在于 LSTM 的细胞状态(cell state)用「加法」更新,而不是普通 RNN 那样用「乘法」反复变换。普通 RNN 的隐藏状态每步都被同一个权重矩阵乘一遍,梯度反向传播时连乘同一个 W 和 tanh 导数,导致指数消失。而 LSTM 的细胞状态更新是 C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t——旧记忆是「加」上来的,不经过反复的权重矩阵相乘。梯度沿细胞状态回传时,主要按遗忘门 f_t 传递(近似乘 f_t),当遗忘门接近 1(表示「一直记住这条信息」)时,直接细胞路径的梯度可以较少衰减地跨越多个时间步回传,形成一条「梯度高速公路(constant error carousel)」,从而缓解梯度消失、让 LSTM 学到长距离依赖。注意是「缓解」而非「彻底消除」。
详细版
普通 RNN 梯度消失的原因(对比):
h_t = tanh(W·h_{t-1} + ...) 隐藏状态每步「乘」W 变换
梯度回传:连乘 (∏ tanh' · Wᵀ) → 因子<1 指数消失
LSTM 的关键——加法更新的细胞状态:
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t
└── 加法:旧记忆直接加过来,不反复乘 W ──┘
梯度沿细胞状态回传(核心):
固定门值与其他间接依赖时,直接细胞路径的 ∂C_t/∂C_{t-1} = f_t
连乘变成 ∏ f_t,遗忘门≈1 时梯度不衰减 → 梯度高速公路
对比:
| 普通 RNN | LSTM | |
|---|---|---|
| 记忆更新 | 乘法(每步乘 W + tanh) | 加法(细胞状态相加) |
| 梯度回传 | 连乘 W,指数消失 | 按遗忘门传,f≈1 不衰减 |
| 长依赖 | 学不到 | 能学到(缓解,非根除) |
完整版教学
一、先回顾普通 RNN 为什么梯度消失
普通 RNN 的隐藏状态更新是乘法式的:
h_t = tanh(W · h_{t-1} + U · x_t + b)
每个时间步,历史记忆 h_{t-1} 都要乘一次权重矩阵 W、过一次 tanh。反向传播(BPTT)时,梯度从后往前跨时间步传播,要连乘很多个 tanh'(·) · Wᵀ 因子。这些因子的大小若普遍小于 1,连乘就指数级衰减 → 梯度消失,远处时间步的梯度传不回来、长依赖学不到(详见 RNN 梯度问题专题)。
根源是「乘法」:记忆反复经过同一个权重矩阵的乘性变换,连乘导致指数衰减。要缓解,就得改变这个「反复乘」的结构。
二、LSTM 的关键改变:细胞状态用加法更新
LSTM 引入了细胞状态 C_t,它的更新方式和普通 RNN 的隐藏状态截然不同——以加法为主:
C_t = f_t ⊙ C_{t-1} + i_t ⊙ C̃_t
└保留的旧记忆┘ └新增的信息┘
注意关键区别:旧记忆 C_{t-1} 是被「加」到新状态里的,而且被向量遗忘门 f_t(每维在 0~1)逐元素缩放,并没有经过「乘权重矩阵 W + tanh」的乘性变换。 这是 LSTM 缓解梯度消失的根本所在。
对比:
- 普通 RNN:
h_t = tanh(W · h_{t-1} + ...)——旧记忆乘 W 再过 tanh。 - LSTM 细胞状态:
C_t = f_t · C_{t-1} + ...——旧记忆只乘一个遗忘门(标量门控),加过来。
三、梯度为什么能沿主路径较少衰减——梯度高速公路
看梯度沿细胞状态的回传。相邻时间步细胞状态的梯度关系,主导项是:
直接细胞路径(暂不展开门值依赖):∂C_t / ∂C_{t-1} = f_t
于是梯度从时间步 T 传到 k,连乘的是一串遗忘门 ∏ f_t,而不是普通 RNN 那样连乘「同一个权重矩阵 W 和 tanh 导数」。
关键在于:
- 遗忘门 f_t 是可以学习、可以接近 1 的。当模型判断「这条信息需要一直记住」时,会让遗忘门接近 1。
- 遗忘门 ≈ 1 时,
∏ f_t ≈ 1——主路径梯度可以较少衰减地跨越多个时间步回传!
这就好比给梯度修了一条「高速公路(constant error carousel,恒定误差传送带)」:只要遗忘门保持打开(≈1),梯度和信息就能在细胞状态这条传送带上长距离、少衰减地流动。这正是 LSTM 能学到长距离依赖的机制。
对比普通 RNN 连乘 W(大小不受控、易 <1 衰减),LSTM 连乘的是可控的、可学习并可能接近 1 的遗忘门——这是本质区别。
四、直觉总结:从「反复改写」到「读写可控的记忆」
换个角度理解:
- 普通 RNN 每一步都强制改写整个记忆(乘 W 变换),旧信息被反复「揉搓」,很快面目全非、梯度也随之消散。
- LSTM 有一条专门的记忆通路(细胞状态),靠门控有选择地读、写、保持:需要记住的信息可以几乎原样保留(遗忘门≈1、输入门≈0,即「不忘、不覆盖」),一直传下去。信息保得住,梯度也就传得回来。
所以 LSTM 缓解梯度消失和它「能记住长依赖」是同一件事——加法更新 + 门控让重要信息和梯度都能长距离保存。
这种“保持”并非被动复制全部历史,而是每个维度由遗忘门和输入门共同决定保留或覆盖。若 f_t≈1 且 i_t≈0,该维度接近恒等传递;若两者都较大,状态幅值还可能累积,因此门控仍需通过数据学习。加性路径改善信用分配,但隐藏状态、门参数和输出路径仍包含非线性与乘法。
五、重要澄清:是「缓解」不是「根除」
要严谨:LSTM 缓解(mitigate)了梯度消失,但没有彻底消除它。
- 遗忘门不总是 1,序列极长时梯度仍会有所衰减。
- LSTM 对梯度爆炸帮助有限,通常仍需配合梯度裁剪。
- 对非常长的依赖,LSTM 依然吃力——这也是注意力机制/Transformer(任意位置直连、路径 O(1))最终胜出的原因(详见 RNN vs Transformer 专题)。
所以正确表述是「LSTM 显著缓解了梯度消失、能捕捉相当长的依赖」,而非「解决了」。
六、遗忘门连乘的数值边界
“遗忘门接近 1”仍要看序列有多长。若连续 100 步都取 f=0.99,主路径梯度比例是 0.99¹⁰⁰≈0.366;若 f=0.9,则 0.9¹⁰⁰≈2.66×10⁻⁵。因此门值只差 0.09,百步后会产生四个数量级以上的差异。
direct cell-path gradient ≈ ∏ f_t
0.99^100 ≈ 0.366
0.90^100 ≈ 0.0000266
这只是固定其他依赖后的直接路径分析;完整总导数还包含门值和候选状态通过 h_{t-1} 对早期状态的依赖。LSTM 提供一条更有利的加性路径,不是把整个计算图的雅可比恒定成 1。
记忆钩子:LSTM 把“不可控的矩阵连乘”改造成“可学习门控主路径”,但可学习不等于自动等于 1。
七、常见误区与追问
- 误区:LSTM 的细胞状态梯度恒等于 1。 主路径由遗忘门连乘决定,完整总导数还含其他依赖。
- 误区:只要使用 LSTM 就不需要梯度裁剪。 LSTM 主要缓解消失,爆炸仍可能发生。
- 追问:为什么说加法路径有帮助? 它允许旧状态不经新的权重矩阵和 tanh 强制重写而直接进入下一状态。
- 追问:遗忘门为什么不能总固定为 1? 模型还需要丢弃过期信息并腾出状态容量。
- 追问:门饱和会有什么问题? 门值稳定但 sigmoid 对门参数的导数可能很小,影响门控学习。
八、加强记忆
LSTM 缓解梯度消失的关键:细胞状态用「加法」更新 C_t = f_t⊙C_{t-1} + i_t⊙C̃_t,旧记忆经遗忘门逐元素缩放后进入加法更新,不经过普通 RNN 那种「乘 W + tanh」的乘性变换。分析直接细胞路径且暂不展开门值依赖时,局部因子是 ∂C_t/∂C_{t-1}=f_t,直接路径连乘的是一串可学习且可能接近 1 的遗忘门而非固定的 W——遗忘门长期接近 1 时主路径梯度衰减较慢,形成「梯度高速公路」。对比普通 RNN 连乘 W 指数衰减,LSTM 连乘可控门控是本质区别。直觉:LSTM 有条门控的记忆通路,重要信息在合适门控下能较少改写地跨越多个时间步、梯度随之传回。注意是「缓解」非「根除」:极长序列仍衰减、梯度爆炸仍需裁剪,在许多大规模任务上后来由注意力/Transformer 成为主流。