LSTM 的门结构与长程依赖问题
LSTM(Long Short-Term Memory)通过引入细胞状态(Cell State)和三个门控机制来解决标准 RNN 的梯度消失问题。核心设计思想是:为梯度提供一条贯穿所有时间步的加法通路,避免矩阵连乘导致的梯度衰减。
整体结构
C_{t-1} ────────────●──────────────── Cₜ ← 细胞状态(信息高速公路) ↑ h_{t-1} ──┬──→ [遗忘门 fₜ] ──┐ │ │ xₜ ────┤──→ [输入门 iₜ] ──→●(加)──→ Cₜ │ ↑ │ │ [候选 C̃ₜ] │ │ │ └──→ [输出门 oₜ] ──→ tanh(Cₜ) ⊙ oₜ ──→ hₜ三个门的详细作用
1. 遗忘门(Forget Gate)—— 决定丢弃什么
fₜ = σ(W_f · [h_{t-1}, xₜ] + b_f)- 输入:上一步隐藏状态
h_{t-1}+ 当前输入xₜ - 输出:0 到 1 之间的向量,逐元素控制
- 作用:决定细胞状态
C_{t-1}中哪些信息需要遗忘,哪些需要保留 fₜ → 0:完全遗忘对应维度的历史信息fₜ → 1:完全保留对应维度的历史信息
直觉:在语言模型中,处理新主语时,遗忘门会关闭以丢弃旧主语的性别信息。
2. 输入门(Input Gate)—— 决定写入什么
iₜ = σ(W_i · [h_{t-1}, xₜ] + b_i) C̃ₜ = tanh(W_C · [h_{t-1}, xₜ] + b_C)- 输入门
iₜ:0 到 1 的向量,控制每个维度允许写入多少新信息 - 候选状态
C̃ₜ:tanh 输出(-1 到 1),表示候选的新信息内容 - 作用:两者逐元素相乘
iₜ ⊙ C̃ₜ,决定将多少新信息写入细胞状态
直觉:遇到新主语 “Alice” 时,输入门打开,将 “女性” 信息写入细胞状态。
3. 输出门(Output Gate)—— 决定输出什么
oₜ = σ(W_o · [h_{t-1}, xₜ] + b_o) hₜ = oₜ ⊙ tanh(Cₜ)- 作用:决定细胞状态
Cₜ中哪些信息作为当前时间步的隐藏状态hₜ输出 - 细胞状态先经 tanh 压缩到 [-1, 1],再与输出门逐元素相乘
直觉:细胞状态中可能存储了主语性别信息,但当前词是动词,不需要输出性别,输出门关闭。
细胞状态更新:加法通路的关键
三个门协同工作,细胞状态的更新公式:
Cₜ = fₜ ⊙ C_{t-1} + iₜ ⊙ C̃ₜ ↑ ↑ 遗忘门控制保留 输入门控制写入这是 LSTM 解决梯度消失的核心所在。
为什么能解决梯度消失?
对比标准 RNN 和 LSTM 的梯度传播路径:
标准 RNN:
hₜ = tanh(W_hh · h_{t-1} + W_xh · xₜ) ∂hₜ/∂h_{t-1} = W_hh · diag(tanh'(hₜ)) ← 矩阵乘法,连乘 T 次 → (W_hh)^T → 消失/爆炸LSTM:
Cₜ = fₜ ⊙ C_{t-1} + iₜ ⊙ C̃ₜ ∂Cₜ/∂C_{t-1} = fₜ ← 逐元素乘法,不是矩阵乘法关键区别:
| 标准 RNN | LSTM | |
|---|---|---|
| 梯度传播方式 | 矩阵连乘(W_hh)^T | 逐元素乘法∏ fₜ |
| 连乘问题 | 特征值连乘 → 指数衰减/增长 | 无矩阵连乘,无特征值问题 |
| 梯度控制 | 无 | 遗忘门fₜ可学习为接近 1 |
当遗忘门fₜ ≈ 1时:
∂L/∂C₁ = ∂L/∂C_T · ∏(t=2→T) fₜ ≈ ∂L/∂C_T · 1梯度几乎无损地从最后一步传回第一步,长距离依赖得以保留。
各门协作的完整示例
以语言模型处理句子"The cat, which already ate fish, was full."为例:
| 时间步 | 当前词 | 遗忘门 | 输入门 | 输出门 | 细胞状态变化 |
|---|---|---|---|---|---|
| t=1 | The | 打开 | 写入"定冠词" | 输出 | C: [定冠词] |
| t=2 | cat | 保留 | 写入"猫、单数" | 输出 | C: [定冠词, 猫, 单数] |
| t=3-7 | which…fish | 保留主语信息 | 写入从句信息 | 从句输出 | C: [猫, 单数, …从句] |
| t=8 | was | 保留"单数" | 写入"过去时" | 输出 | C: [猫,单数, 过去时] |
| t=9 | full | 保留 | — | 输出 | C: [猫, 单数, 过去时] |
was需要回溯到cat(距离 6 步)获取"单数"信息 → 遗忘门保持打开,细胞状态中的"单数"信息一路保留- 标准 RNN 在 6 步后梯度已严重衰减,无法将
cat的单数信息传递到was
总结
| 门 | 公式 | 作用 | 类比 |
|---|---|---|---|
| 遗忘门 fₜ | σ(W_f·[h_{t-1},xₜ]) | 控制历史信息的保留/丢弃 | 橡皮擦 |
| 输入门 iₜ | σ(W_i·[h_{t-1},xₜ]) | 控制新信息的写入量 | 笔 |
| 候选状态 C̃ₜ | tanh(W_C·[h_{t-1},xₜ]) | 生成候选新信息 | 墨水 |
| 输出门 oₜ | σ(W_o·[h_{t-1},xₜ]) | 控制细胞状态到输出的量 | 阅读窗口 |
核心结论:LSTM 解决长程依赖的本质不是"门"本身,而是细胞状态更新公式
Cₜ = fₜ ⊙ C_{t-1} + iₜ ⊙ C̃ₜ中的加法结构。加法使得梯度传播变为逐元素乘法(由可学习的遗忘门控制),而非矩阵连乘,从而让遗忘门可以学会保持接近 1,使梯度长距离无损传播。三个门的作用是让模型自适应地控制信息的保留、写入和输出,使加法通路携带的是有用信息而非噪声。