共计 2015 个字符,预计需要花费 6 分钟才能阅读完成。
背景痛点:RNN 的长期依赖困境
传统 RNN 通过循环连接处理时序数据,其隐藏状态 $h_t$ 的计算可表示为:
$$h_t = \sigma(W_h h_{t-1} + W_x x_t + b)$$
但在反向传播时需计算损失函数对参数的梯度,当时间步 $T$ 较大时会出现连乘效应:
$$\frac{\partial h_T}{\partial h_1} = \prod_{t=2}^T \frac{\partial h_t}{\partial h_{t-1}}$$
- 梯度消失问题:当导数 $\frac{\partial h_t}{\partial h_{t-1}} < 1$ 时,梯度呈指数衰减,导致早期时间步的参数几乎不更新
- 梯度爆炸问题:当导数 $\frac{\partial h_t}{\partial h_{t-1}} > 1$ 时,梯度呈指数增长,引发数值不稳定
- 传统解决方案局限:截断 BPTT 仅缓解但未根治问题,手工设计时间窗口难以适应不同序列长度
技术演进:从原始 LSTM 到现代变体
- 1997 年原始 LSTM(Hochreiter & Schmidhuber)
- 引入细胞状态 $C_t$ 作为 ” 记忆通道 ”
- 三个门控单元:输入门、输出门、遗忘门
-
初始版本未包含 peephole 连接
-
2000 年改进 LSTM(Gers et al.)
- 添加遗忘门偏置(初始值通常设为 1)
-
引入 peephole 连接,允许门控查看细胞状态
-
2014 年 GRU 简化版(Cho et al.)
- 合并遗忘门与输入门为更新门
- 耦合细胞状态与隐藏状态
- 参数减少 33% 但性能相近
核心架构:门控机制详解

遗忘门 决定保留多少旧记忆:
$$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$$
输入门 控制新信息流入:
$$
\begin{cases}
i_t = \sigma(W_i \cdot [h_{t-1}, x_t] + b_i) \
\tilde{C}t = \tanh(W_C \cdot [h, x_t] + b_C)
\end{cases}
$$
细胞状态更新:
$$C_t = f_t \odot C_{t-1} + i_t \odot \tilde{C}_t$$
输出门 调控最终输出:
$$
\begin{cases}
o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o) \
h_t = o_t \odot \tanh(C_t)
\end{cases}
$$
PyTorch 实现示例
import torch
import torch.nn as nn
class LSTMModel(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_size, num_layers):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
# hidden_size: 每层 LSTM 单元的数量
# num_layers: 堆叠的 LSTM 层数
self.lstm = nn.LSTM(embed_dim, hidden_size, num_layers, batch_first=True)
self.fc = nn.Linear(hidden_size, vocab_size)
def forward(self, x, hidden):
x = self.embedding(x) # (batch, seq_len) -> (batch, seq_len, embed_dim)
out, hidden = self.lstm(x, hidden)
out = self.fc(out[:, -1, :]) # 只取最后一个时间步
return out, hidden
性能对比实验
在 Penn Treebank 语言建模任务上的表现:
| 模型 | 困惑度(Perplexity) | 参数量 |
|---|---|---|
| RNN | 120.3 | 4.2M |
| LSTM | 78.4 | 4.5M |
| GRU | 80.1 | 4.3M |
训练避坑指南
- 梯度爆炸
- 使用梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5) -
初始化遗忘门偏置为 1(促进初始阶段记忆保留)
-
超参数敏感
- 学习率:从 3e- 4 开始尝试,配合学习率调度器
- 隐藏层大小:根据 GPU 内存选择(通常 128-512)
-
Dropout 率:层间 dropout 建议 0.2-0.5(需设置
lstm.dropout参数) -
序列填充干扰
- 使用
pack_padded_sequence处理变长序列 - 设置
enforce_sorted=False允许乱序输入
延伸思考:Transformer 时代的 LSTM
虽然 Transformer 在长序列建模中表现突出,但 LSTM 仍有其优势场景:
- 低资源环境:LSTM 参数更少,训练成本低
- 在线学习:LSTM 可逐时间步处理,而 Transformer 需要完整序列
- 小规模数据:LSTM 不易过拟合
实际选择建议:
– 序列长度 <100 且数据量小时优先考虑 LSTM
– 需要处理文档级依赖时选用 Transformer
– 可尝试 LSTM+Attention 的混合架构
