梯度消失、爆炸与 LSTM/GRU 门控
朴素 RNN 的梯度沿时间连乘 W_h 的幂,谱半径偏离 1 即指数消失或爆炸;LSTM/GRU 用输入门、遗忘门、输出门构建「加法式」的细胞状态通路,让梯度可沿门控的加法路径跨越数百步存活。
一句话定义
朴素 RNN 的梯度沿时间连乘 W_h 的幂,谱半径偏离 1 即指数消失或爆炸;LSTM/GRU 用输入门、遗忘门、输出门构建「加法式」的细胞状态通路,让梯度可沿门控的加法路径跨越数百步存活。
直觉
朴素 RNN 的记忆像「复印机复印的复印件」:每过一天再复印一次,几十天后字迹全无(消失)或墨点成灾(爆炸);LSTM 的细胞状态像「保管在保险柜的原件」,遗忘门决定扔掉哪几页、输入门决定存入哪几页——原件本身不经反复复印。
为什么重要
梯度消失/爆炸是 kp-007 方差链分析在时间维度的重演,也是全库第二次正面攻坚这一主题:空间维度用残差(kp-020)解决,时间维度用门控解决——两条路线共享「为梯度修高速公路」的公理。LSTM 从 1997 年存活到 2017 年 Transformer 前夜,是序列建模的二十年代言人;读懂它的门控方程,才算真正读懂 kp-007 的乘性结构有多致命。
前置知识
- kp-022(RNN 的时间展开)、kp-012(方差连乘的数值后果)。
核心概念
- 梯度消失/爆炸(时间维度):∂L/∂h_t ≈ Π W_hᵀ·diag(φ') 的幂连乘趋于 0 或 ∞。
- 细胞状态 c_t:LSTM 的长期记忆主线,默认线性直通、仅受门控微调。
- 三个门:遗忘门 f_t(旧记忆留多少)、输入门 i_t(新信息存多少)、输出门 o_t(记忆透露多少)。
- GRU:两门简化版(更新门+重置门),性能接近 LSTM、参数少约 25%。
- 正交初始化(orthogonal init):让 W_h 初值为正交矩阵,谱半径恰为 1,缓解起步期消失/爆炸(kp-012 的 RNN 特化)。
- 梯度裁剪:爆炸的治疗手段(norm 超阈值即缩放),工程标配(kp-025 详述)。
原理与机制
病根定量:朴素 RNN 展开 T 步后 ∂L/∂h_0 ≈ Π_{t} W_hᵀ diag(φ'(z_t))。设 W_h 最大奇异值为 σ_max,若 (σ_max·mean|φ'|) < 1,连乘 T 次指数趋零——100 步后 0.9^100 ≈ 2.7e−5,「50 步之前」的信息梯度已低于浮点噪声,学习信号到不了远方;若 >1 则梯度数值爆炸、训练 NaN。致命处在于这是乘性结构:只要连乘因子均值偏离 1,深度(时间)就是敌人。LSTM 的解法是改乘为加。三门皆由 sigmoid 输出 (0,1) 作为「阀门开度」:f_t = σ(W_f[h_{t-1},x_t]+b_f),i_t = σ(W_i[·]+b_i),候选记忆 g_t = tanh(W_g[·]+b_g),细胞更新 c_t = f_t⊙c_{t-1} + i_t⊙g_t,输出 o_t = σ(W_o[·]+b_o),h_t = o_t⊙tanh(c_t)。关键在 c_t 这一行:细胞状态的梯度路径是 ∂c_t/∂c_{t-1} = f_t——一个 (0,1) 之间的逐元素因子,且网络可以学着把 f_t≈1(遗忘门全开);于是梯度沿 c 链是「近似加法直通」,如同 kp-020 残差的「1 地基」:∂L/∂c_{t-1} = f_t⊙∂L/∂c_t + (经门的旁路项),首项为线性保底,不再有 W_h 幂连乘的指数衰减。对比表:朴素 RNN 的传递因子是「矩阵×激活导数」的双重随机积,LSTM 的记忆主线的传递因子是学出来的标量门——从「祈祷乘积≈1」进化到「主动学习乘积=1」。GRU 把 f 与 i 合并为更新门 z_t:h_t = (1−z_t)⊙h_{t-1} + z_t⊙ĥ_t(加法结构保留),参数更少、小数据上常更快收敛。C 地址:Hochreiter 1997 引入的让梯度「流过恒等误差桥」的原初表述与本节的门控分析完全一致。
公式或模型
LSTM 全套方程(对照朴素 RNN 逐行阅读):
f_t = σ( W_f [h_{t-1}; x_t] + b_f ) 遗忘门:旧记忆保留比例
i_t = σ( W_i [h_{t-1}; x_t] + b_i ) 输入门:新信息写入比例
g_t = tanh( W_g [h_{t-1}; x_t] + b_g ) 候选记忆
c_t = f_t ⊙ c_{t-1} + i_t ⊙ g_t ★加法更新:梯度主线
o_t = σ( W_o [h_{t-1}; x_t] + b_o ) 输出门:读取比例
h_t = o_t ⊙ tanh( c_t )
关键导数: ∂c_t/∂c_{t-1} = f_t (逐元素、可学、可≈1 → 长程梯度存活)变量说明:[a;b] 为向量拼接;⊙ 为逐元素乘;所有门为逐时刻独立计算。对照 kp-020:c 链之于时间,正如残差之于深度。
图示
朴素 RNN 记忆通路: h_{t-1} ──[×W_h ×φ']──► h_t ──[×W_h ×φ']──► h_t+1 …
(乘性链:每步随机衰减,50 步后≈0)
LSTM 细胞状态通路: c_{t-1} ──[×f_t≈1]──┐
+ ──► c_t ──[×f≈1]──► c_{t+1}
┌─ i_t⊙g_t ┘
(加法主干:f_t 是学出来的「阀门」,梯度线性存活)直观类比
三门像图书管理员的三问:旧档案扔不扔(f)、新书收不收(i)、今天透露多少给访客(o)——档案原件(c)从不外借,只给「复印摘要」(h)。
实例或案例
数值实验对比记忆寿命:构造「复制第 1 个字符」的合成任务(序列长 100,末位输出首位字符),朴素 RNN 训练准确率停在 20% 以下(梯度到不了首步),LSTM 数百步内冲到 99%——用 30 行代码即可复现门控的威力。工程配套三件套:W_h/c 链用正交初始化(谱半径 1 起步,kp-012 的 RNN 特化)、遗忘门 bias 初始化为 1~2(默认「多记住」)、梯度裁剪(kp-025)防残余爆炸。kp-025 将把本节模型接入完整的截断 BPTT 训练循环。
常见误区
- 「LSTM 消灭了梯度消失」:它只修复了细胞状态主线的消失,h 链与各门自身仍会衰减;爆炸仍需裁剪。
- 「门是 0/1 开关」:sigmoid 输出连续值,门是「阀门开度」而非布尔闸。
- 「GRU 一定不如 LSTM」:多数任务两者接近,GRU 参数少、更不易过拟合小数据。
- 忽略遗忘门 bias 初始化:默认 0 时模型起步倾向「全遗忘」,长序列任务收敛显著变慢(经验偏置而非玄学)。
与其他知识点的关系
本节是 kp-007 乘性结构分析的时间维重演、kp-012 正交初始化的应用现场、kp-020 加法结构的序列孪生;kp-025 提供其训练工程;kp-026 的注意力从「绕过长程检索」这一角度最终取代了门控的绝大多数场景。
自测题
- 推导朴素 RNN 时间梯度连乘式,说明谱半径如何决定消失/爆炸。
答案要点:∂L/∂h_t≈Π W_hᵀdiag(φ');σ_max·E|φ'|<1 指数消失、>1 指数爆炸。
- 写出 LSTM 细胞更新式,指出 ∂c_t/∂c_{t-1} 等于什么、为何这是关键。
答案要点:c_t=f⊙c_{t-1}+i⊙g;∂c_t/∂c_{t-1}=f_t,可学习的逐元素因子≈1 时长程梯度线性存活。
- GRU 如何用两个门实现 LSTM 的功能?
答案要点:更新门 z 合并遗忘/输入(h_t=(1−z)h_{t-1}+z·ĥ_t),重置门 r 控制候选记忆对历史的依赖。
延伸阅读
- Hochreiter、Schmidhuber《Long Short-Term Memory» Neural Computation 1997。
- Pascanu、Mikolov、Bengio ICML 2013(梯度爆炸/裁剪的系统分析,与 kp-025 共读)。
学习状态
状态保存在浏览器本地,用于首页与路径页的进度统计。