LSTM

长短期记忆网络(Long Short-Term Memory,LSTM)由 Hochreiter 与 Schmidhuber 于 1997 年提出,是循环神经网络(Recurrent Neural Network,RNN)家族中最具代表性的门控变体之一。它通过引入细胞状态(cell state)与三类门控(gate),在理论上缓解 vanilla RNN 的梯度消失/爆炸问题,使模型能够跨越数十到数百个时间步保留有效信息。

段末注释LSTM 的核心不是”记住一切”,而是有选择地写入、保留与读出信息;遗忘门决定”忘什么”,输入门决定”记什么”,输出门决定”此刻输出什么”。

读前说明:配图位于同名目录 2104.神经网络-循环神经网络-LSTM/。本文独立成篇,不依赖系列内其他 RNN 入门文档的阅读顺序;若需对比基础 RNN,可参阅 2104.神经网络-循环神经网络-RNN


1. 从何处来:演进脉络

图 1 序列模型演进:从 Elman RNN 到 LSTM(科普示意)

1.1 前驱:Elman 网络与 Jordan 网络

1980 年代末,Elman(1990)与 Jordan(1986)分别提出带上下文单元(context unit)的递归结构:隐藏状态 $\mathbf{h}_{t-1}$ 与当前输入 $\mathbf{x}_t$ 共同决定下一时刻输出。这类模型已具备”记忆”雏形,但仍是浅层无门控的循环计算。

1.2 标准 RNN 及其瓶颈

将 Elman 结构参数化并堆叠后,得到现代意义上的 vanilla RNN

$$
\mathbf{h}t = \phi!\left(\mathbf{W}{xh},\mathbf{x}t + \mathbf{W}{hh},\mathbf{h}_{t-1} + \mathbf{b}_h\right), \qquad
\mathbf{y}t = g!\left(\mathbf{W}{hy},\mathbf{h}_t + \mathbf{b}_y\right)
$$

其中 $\phi$ 通常为 $\tanh$ 或 ReLU,$g$ 为 softmax(分类)或恒等映射(回归)。

vanilla RNN 要解决的:对序列 $\mathbf{x}_1, \ldots, \mathbf{x}T$ 建模条件分布 $P(y_t \mid \mathbf{x}{\le t})$,捕获短期上下文(相邻词、相邻帧)。

vanilla RNN 解决不了的——也是 LSTM 诞生的直接动机:

问题 数学表现 直觉
梯度消失 沿时间反向传播时 $\prod_{k} \frac{\partial \mathbf{h}k}{\partial \mathbf{h}{k-1}}$ 的谱半径 $< 1$,远端梯度指数衰减 无法学习”100 步之前”的依赖
梯度爆炸 谱半径 $> 1$,梯度指数增长 训练不稳定,需梯度裁剪
信息覆盖 $\mathbf{h}_t$ 同时承担”存储”与”计算”,新输入不断改写旧记忆 长序列中早期信息被冲刷

Bengio 等(1994)从理论上指出:长期依赖(long-term dependency)在 vanilla RNN 中难以通过梯度学习;Hochreiter(1991)的 diploma thesis 已分析该现象,为 LSTM 奠基。

1.3 LSTM 的提出与后续变体

年份 里程碑 要点
1997 Hochreiter & Schmidhuber,LSTM 细胞状态 $\mathbf{C}_t$ + 输入/输出/遗忘门(早期版本)
1999 遗忘门(forget gate)正式加入 可控地擦除 $\mathbf{C}_{t-1}$ 中的信息
2000s 双向 LSTM、深层堆叠 机器翻译、手写识别等任务
2014 Cho 等提出 GRU 合并细胞状态与隐藏状态,门数更少
2015–2017 seq2seq + attention LSTM 编码器仍主流,但 attention 减轻”仅靠固定向量压缩整句”的压力
2017+ Transformer 自注意力替代循环结构成为 NLP 默认,LSTM 在部分时序/边缘场景仍常用

段末注释GRU(Gated Recurrent Unit)可视为 LSTM 的简化版:将 $\mathbf{C}_t$ 与 $\mathbf{h}_t$ 合并,仅保留更新门与重置门,参数量约为 LSTM 的 75%。


2. 要解决什么问题

2.1 长期依赖的形式化

给定序列长度 $T$,我们希望模型有效利用 $t - \tau$ 时刻的信息来预测 $t$ 时刻输出,其中 $\tau$ 可能远大于 1。例如:

  • 语言:”我出生在中国,……(中间 50 词)…… 因此我的母语是 ____”
  • 时序:股票在 $t-200$ 的异常波动影响 $t$ 的波动率
  • 生物序列:远端启动子区域影响当前位点的表达

理想情况下,损失 $\mathcal{L}$ 对 $\mathbf{h}{t-\tau}$ 的梯度 $\partial \mathcal{L}/\partial \mathbf{h}{t-\tau}$ 不应在 $\tau$ 增大时迅速趋零。

2.2 LSTM 的设计思路

LSTM 将”记忆载体”与”计算状态”解耦

  • 细胞状态 $\mathbf{C}_t \in \mathbb{R}^{d_c}$:慢变化的记忆通道,类似传送带;
  • 隐藏状态 $\mathbf{h}_t \in \mathbb{R}^{d_h}$:对外暴露的、经门控过滤后的读出;
  • 门控 $\in (0,1)^{d}$:逐元素控制写入、保留、读出比例。

关键机制:加法更新 $\mathbf{C}_t = \mathbf{f}t \odot \mathbf{C}{t-1} + \mathbf{i}_t \odot \tilde{\mathbf{C}}_t$,使误差可在 $\mathbf{C}$ 上线性传递(在门接近 1 时),缓解连乘导致的梯度消失。


3. 模型架构与原理

图 2 LSTM 单元结构:细胞状态与三门控(科普示意)

3.1 符号约定

符号 含义 典型维度
$\mathbf{x}_t$ 时刻 $t$ 的输入 $\mathbb{R}^{d_x}$
$\mathbf{h}_{t-1}$ 上一时刻隐藏状态 $\mathbb{R}^{d_h}$
$\mathbf{C}_{t-1}, \mathbf{C}_t$ 细胞状态 $\mathbb{R}^{d_c}$(常取 $d_c = d_h$)
$\mathbf{i}_t, \mathbf{f}_t, \mathbf{o}_t$ 输入门、遗忘门、输出门 $\mathbb{R}^{d_c}$
$\tilde{\mathbf{C}}_t$ 候选记忆 $\mathbb{R}^{d_c}$
$\sigma$ sigmoid:$\sigma(z) = 1/(1+e^{-z})$ 逐元素
$\odot$ Hadamard 积(逐元素乘)

拼接向量记为 $[\mathbf{h}_{t-1}; \mathbf{x}_t] \in \mathbb{R}^{d_h + d_x}$。

3.2 前向传播(标准 LSTM)

Step 1 — 门控与候选记忆

$$
\begin{aligned}
\mathbf{i}_t &= \sigma!\left(\mathbf{W}i,[\mathbf{h}{t-1}; \mathbf{x}_t] + \mathbf{b}_i\right) & \text{(输入门:写入多少新信息)} \
\mathbf{f}_t &= \sigma!\left(\mathbf{W}f,[\mathbf{h}{t-1}; \mathbf{x}_t] + \mathbf{b}_f\right) & \text{(遗忘门:保留多少旧记忆)} \
\mathbf{o}_t &= \sigma!\left(\mathbf{W}o,[\mathbf{h}{t-1}; \mathbf{x}_t] + \mathbf{b}_o\right) & \text{(输出门:读出多少到 } \mathbf{h}_t\text{)} \
\tilde{\mathbf{C}}_t &= \tanh!\left(\mathbf{W}C,[\mathbf{h}{t-1}; \mathbf{x}_t] + \mathbf{b}_C\right) & \text{(候选记忆内容)}
\end{aligned}
$$

Step 2 — 更新细胞状态(核心)

$$
\mathbf{C}_t = \mathbf{f}t \odot \mathbf{C}{t-1} + \mathbf{i}_t \odot \tilde{\mathbf{C}}_t
$$

Step 3 — 输出隐藏状态

$$
\mathbf{h}_t = \mathbf{o}_t \odot \tanh(\mathbf{C}_t)
$$

Step 4 — 任务头(依任务而定)

$$
\hat{\mathbf{y}}t = \mathbf{W}{hy},\mathbf{h}_t + \mathbf{b}y \quad\text{(回归)}; \qquad
P(y_t \mid \cdot) = \mathrm{softmax}(\mathbf{W}
{hy},\mathbf{h}_t + \mathbf{b}_y) \quad\text{(分类)}
$$

3.3 门控的直觉

  • $\mathbf{f}_t \approx 0$:几乎清空 $\mathbf{C}_{t-1}$ 对应维度(”忘记”);
  • $\mathbf{i}_t \approx 0$:不写入 $\tilde{\mathbf{C}}_t$(”不记新内容”);
  • $\mathbf{f}_t \approx 1, \mathbf{i}_t \approx 0$:记忆原样保留(”长期持有”)—— 这正是跨越长距离依赖的关键模式;
  • $\mathbf{o}_t$:控制 $\mathbf{h}_t$ 暴露给下游多少信息,不影响 $\mathbf{C}_t$ 本身(Peephole 变体除外)。

3.4 多层与双向扩展

堆叠 LSTM:第 $l$ 层在时刻 $t$ 的输入为第 $l-1$ 层的 $\mathbf{h}_t^{(l-1)}$,顶层 $\mathbf{h}_t^{(L)}$ 接输出头。

双向 LSTM(BiLSTM):前向 LSTM 得 $\overrightarrow{\mathbf{h}}_t$,后向 LSTM 得 $\overleftarrow{\mathbf{h}}_t$,拼接 $[\overrightarrow{\mathbf{h}}_t; \overleftarrow{\mathbf{h}}_t]$ 作为下游特征。适合整句编码(命名实体识别、情感分析),不适合纯自回归生成(生成时无法看到未来)。

3.5 与 vanilla RNN 的对比

维度 Vanilla RNN LSTM
状态 仅 $\mathbf{h}_t$ $\mathbf{h}_t$ + $\mathbf{C}_t$
更新方式 替换式:$\mathbf{h}_t = \phi(\cdots)$ 门控加法式:$\mathbf{C}_t = \mathbf{f}t \odot \mathbf{C}{t-1} + \cdots$
长期梯度 易消失/爆炸 细胞路径上可近似线性传播
参数量 $O((d_h+d_x)d_h)$ 约 $4\times$ 上述(四门) + 输出头

4. 训练哪些参数

4.1 单个 LSTM 单元的可学习参数

每个 LSTM 层包含 4 组 仿射变换(对应四个 $\mathbf{W}$)及偏置:

参数 形状(常见设定) 作用
$\mathbf{W}_i, \mathbf{b}_i$ $\mathbf{W}_i \in \mathbb{R}^{d_c \times (d_h + d_x)}$ 输入门
$\mathbf{W}_f, \mathbf{b}_f$ 同上 遗忘门
$\mathbf{W}_o, \mathbf{b}_o$ 同上 输出门
$\mathbf{W}_C, \mathbf{b}_C$ 同上 候选记忆
$\mathbf{W}_{hy}, \mathbf{b}_y$ $\mathbb{R}^{d_y \times d_h}$ 输出层(若每层独立头)

参数量估算(单层、单方向):

$$
|\Theta|_{\mathrm{LSTM}} \approx 4 \cdot d_c ,(d_h + d_x + 1) + d_y(d_h + 1)
$$

若 $d_h = d_c = d_x = 256$,仅 LSTM 核心约 $4 \times 256 \times 513 \approx 525\text{K}$,尚未计嵌入层与输出层。

权重共享:同一 LSTM 层在所有时间步 $t = 1,\ldots,T$ 共享 $\mathbf{W}_i, \mathbf{W}_f, \mathbf{W}_o, \mathbf{W}_C$——这是 RNN 族”循环”的含义。

4.2 完整模型通常还包括

  • 输入嵌入 $\mathbf{E} \in \mathbb{R}^{V \times d_x}$($V$ 为词表大小);
  • Dropout(非参数,但影响训练);
  • LayerNorm / BatchNorm(可选,有独立仿射参数);
  • Peephole 连接(可选变体):门还依赖 $\mathbf{C}{t-1}$,增加 $\mathbf{W}{*c}$。

5. 如何训练

图 3 LSTM 训练:前向展开、损失汇总与 BPTT 反向传播(科普示意)

5.1 数据与展开

给定序列 $\mathcal{D} = {(\mathbf{x}^{(n)}{1:T_n}, \mathbf{y}^{(n)})}{n=1}^N$:

  1. 截断/填充 至 batch 内最大长度(或固定窗口 $T$);
  2. 按时间步前向 计算 $\mathbf{h}_t, \mathbf{C}_t$;
  3. 初始状态常置 $\mathbf{h}_0 = \mathbf{0}, \mathbf{C}_0 = \mathbf{0}$,或作为可学习参数。

5.2 损失函数

序列标注(每步有标签):交叉熵求和或平均

$$
\mathcal{L} = -\frac{1}{T}\sum_{t=1}^{T} \sum_{c} y_{t,c},\log \hat{y}_{t,c}
$$

序列级分类(仅最后一步):$\mathcal{L} = \ell(\mathbf{y}, \hat{\mathbf{y}}_T)$。

回归(时序预测):MSE 或 MAE,$\mathcal{L} = \frac{1}{T}\sum_t |\mathbf{y}_t - \hat{\mathbf{y}}_t|^2$。

语言建模:next-token prediction,$\mathcal{L} = -\sum_t \log P(x_{t+1} \mid \mathbf{x}_{\le t})$。

5.3 反向传播:BPTT

随时间反向传播(Backpropagation Through Time,BPTT)将展开图视为深层前馈网络,从 $t=T$ 向 $t=1$ 回传梯度:

$$
\frac{\partial \mathcal{L}}{\partial \mathbf{W}*} = \sum{t=1}^{T} \frac{\partial \mathcal{L}t}{\partial \mathbf{W}*}
$$

对 $\mathbf{C}_t$ 的递推梯度(略去门细节)体现”加法路径”:

$$
\frac{\partial \mathcal{L}}{\partial \mathbf{C}_{t-1}} = \frac{\partial \mathcal{L}}{\partial \mathbf{C}_t} \odot \mathbf{f}_t + \cdots
$$

当 $\mathbf{f}t \approx 1$ 时,梯度可较完整地传回 $\mathbf{C}{t-\tau}$。

Truncated BPTT(TBPTT):每 $K$ 步截断一次反向传播,降低显存与计算;牺牲部分超长依赖的学习能力。

5.4 优化实践

技巧 说明
梯度裁剪 $\mathbf{g} \leftarrow \mathbf{g} \cdot \min(1, \theta / |\mathbf{g}|)$,防爆炸
Adam / AdamW 默认优化器;学习率 $10^{-3}$ 量级起步
Teacher forcing 训练 seq2seq 时用真值前缀作解码输入,加速收敛
Scheduled sampling 逐步用模型预测替代真值,缓解 exposure bias
Dropout 常加在 $\mathbf{h}_t$ 或层间;变体 Gal & Ghahramani(2016)在同一 mask 上沿时间重复
LayerNorm 稳定深层 LSTM 训练

5.5 PyTorch 最小示例

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
import torch
import torch.nn as nn

class LSTMClassifier(nn.Module):
"""序列分类:输入 (batch, seq_len, d_x),输出 (batch, num_classes)"""

def __init__(self, d_x: int, d_h: int, num_classes: int, num_layers: int = 1):
super().__init__()
self.lstm = nn.LSTM(
input_size=d_x,
hidden_size=d_h,
num_layers=num_layers,
batch_first=True, # 输入形状 (N, T, d_x)
)
self.head = nn.Linear(d_h, num_classes)

def forward(self, x: torch.Tensor) -> torch.Tensor:
# output: (N, T, d_h); h_n: (num_layers, N, d_h)
output, (h_n, c_n) = self.lstm(x)
return self.head(h_n[-1]) # 取最后一层最后时刻的隐藏状态

训练循环与标准监督学习相同:loss.backward()optimizer.step()nn.LSTM 内部已实现 BPTT。


6. 应用与局限

图 4 LSTM 的四类典型局限(科普示意)

6.1 典型应用场景(仍具价值)

场景 说明
中小规模 NLP 资源受限时的文本分类、序列标注
时序预测 金融、传感器、工业信号(中等长度窗口)
语音/手写 与 CTC 等结合的历史基线
边缘部署 参数量可控、无 attention 二次复杂度
生物序列 蛋白质/ DNA 局部到中等长度 motif 建模

6.2 主要局限性

(1)超长依赖仍有限

门控缓解但未消除长程衰减;依赖跨度 $\tau \gg 500$ 时,Transformer 或结构化记忆通常更优。细胞状态维度 $d_c$ 有限,信息瓶颈客观存在。

(2)无法时间步并行

时刻 $t$ 依赖 $\mathbf{h}_{t-1}$,训练/推理必须顺序执行;GPU 利用率低于 Transformer 的矩阵并行。长序列训练慢是工程常态。

(3)参数量与过拟合

四门结构参数量约为 vanilla RNN 的 4 倍;小数据集上 BiLSTM 易过拟合,需 Dropout、正则、数据增广或预训练。

(4)固定上下文压缩

经典 seq2seq 编码器将整句压入单个 $\mathbf{h}_T$,信息损失大;需 attention 或分层结构补救。LSTM + attention 曾是 2014–2016 机器翻译 SOTA,现已被 Transformer 取代。

(5)可解释性弱

门值 $\mathbf{f}_t, \mathbf{i}_t$ 可视作”软”解释,但高维下难以对应人类语义;不如注意力权重直观。

(6)变体选择成本

LSTM / GRU / Peephole / LayerNorm-LSTM 等需调参对比;GRU 在多数任务上与 LSTM 打平且更快,成为常见默认。

6.3 与 Transformer 的定位

维度 LSTM Transformer
复杂度 每步 $O(d^2)$,序列 $O(Td^2)$ 自注意力 $O(T^2 d)$
并行 时间维串行 时间维并行
长依赖 门控 + 截断 BPTT,中等 直接连接任意 $(i,j)$,更长
归纳偏置 顺序、局部递归 弱顺序偏置,需位置编码

结论:LSTM 并未”过时”,但在大规模预训练 NLP 中已让位于 Transformer;在资源受限、序列中等长度、需递归归纳偏置的场景,LSTM/GRU 仍是合理选择。


7. 小结

主题 要点
由来 源于 Elman RNN → vanilla RNN 的长期依赖与梯度问题 → 1997 LSTM
解决什么 有选择地保留/写入/读出序列信息,缓解梯度消失
架构 $\mathbf{C}_t$ 加法更新 + 输入/遗忘/输出三门
训练参数 $\mathbf{W}_i, \mathbf{W}_f, \mathbf{W}_o, \mathbf{W}_C$ 及偏置,时间步共享
训练方式 BPTT + 交叉熵/MSE + 梯度裁剪 + Adam 等
局限 超长依赖、串行慢、参数多、固定向量瓶颈、可解释性弱

参考文献

  1. Hochreiter, S., & Schmidhuber, J. (1997). Long Short-Term Memory. Neural Computation, 9(8), 1735–1780.
  2. Bengio, Y., Simard, P., & Frasconi, P. (1994). Learning long-term dependencies with gradient descent is difficult. IEEE TNN.
  3. Gers, F. A., Schmidhuber, J., & Cummins, F. (2000). Learning to forget: Continual prediction with LSTM. Neural Computation.
  4. Cho, K., et al. (2014). Learning phrase representations using RNN encoder-decoder. EMNLP (GRU).
  5. Graves, A. (2012). Supervised Sequence Labelling with Recurrent Neural Networks. Springer.
  6. Goodfellow, I., Bengio, Y., & Courville, A. (2016). Deep Learning. MIT Press, Chapter 10.
-------------本文结束感谢您的阅读-------------