共计 1809 个字符,预计需要花费 5 分钟才能阅读完成。
为什么需要 RNN?
想象这样一个场景:我们要预测明天股票的收盘价。如果只用今天的价格作为输入,显然忽略了历史价格的变化趋势。这种 序列数据 的特点就是前后数据点之间存在依赖关系,而传统全连接网络无法捕捉这种时间维度上的模式。

这时候循环神经网络 (RNN) 就派上用场了——它通过引入隐藏状态 $h_t$ 来 ” 记住 ” 过去的信息:
$$
h_t = \tanh(W_{xh}x_t + W_{hh}h_{t-1} + b_h)
$$
但实际使用时会发现,当序列长度超过 20 步后,模型就很难学到长期依赖了。这是因为在反向传播时,梯度需要沿着时间步连续相乘(称为 BPTT 算法),导致梯度值指数级衰减——这就是著名的 梯度消失问题。
LSTM 如何解决这个难题?
长短期记忆网络 (LSTM) 通过引入三个门控单元,实现了对信息流的精细控制:
-
遗忘门 决定丢弃哪些历史信息:
$$
f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)
$$ -
输入门 控制新信息的写入:
$$
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)
$$ -
输出门 决定当前隐藏状态的输出:
$$
o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o)
$$
最终记忆单元和隐藏状态的更新公式为:
$$
C_t = f_t \odot C_{t-1} + i_t \odot \tilde{C}_t
$$
$$
h_t = o_t \odot \tanh(C_t)
$$
通过这种设计,LSTM 的参数量比普通 RNN 多 4 倍,但实际计算时可以通过合并矩阵运算来优化。
PyTorch 实战演练
1. 手动实现 RNNCell
class MyRNNCell(nn.Module):
def __init__(self, input_size, hidden_size):
super().__init__()
self.Wxh = nn.Linear(input_size, hidden_size) # 输入到隐藏层的权重
self.Whh = nn.Linear(hidden_size, hidden_size) # 隐藏层到隐藏层
def forward(self, x, h_prev):
"""
x: (batch_size, input_size)
h_prev: (batch_size, hidden_size)
"""
h_next = torch.tanh(self.Wxh(x) + self.Whh(h_prev))
return h_next
2. 完整 LSTM 训练示例
# 构建一个情感分类模型
model = nn.LSTM(input_size=300, # 词向量维度
hidden_size=128,
num_layers=2,
bidirectional=True)
# 训练循环关键代码
for epoch in range(10):
for batch in dataloader:
text, lengths, labels = batch
packed = nn.utils.rnn.pack_padded_sequence(text, lengths)
output, (hn, cn) = model(packed)
loss = criterion(output, labels)
# 梯度裁剪防止爆炸
nn.utils.clip_grad_norm_(model.parameters(), max_norm=0.25)
optimizer.step()
3. 可视化隐藏状态
使用 PCA 将 128 维的隐藏状态降维后绘制动图,可以清晰看到随着输入序列的变化,模型内部状态的空间分布变化。
避坑指南
- 梯度裁剪:阈值通常设置在 0.25-1.0 之间,太小会阻碍学习,太大失去防护作用
- 序列 padding:一定要在 pack_padded_sequence 之前按长度降序排列,否则计算会浪费大量资源
- 双向 LSTM 优化 :可以使用
batch_first=True参数让批次维度在前,减少内存拷贝开销
拓展思考
- GRU 如何通过合并 LSTM 的门控结构来减少参数量?
- 当我们在 LSTM 的每个时间步加入注意力机制时,应该如何设计 score 函数?
(完整代码和训练曲线示例已上传 Github 仓库,文末链接可查看)
通过这次实践,我深刻体会到:理解 RNN 的核心缺陷比盲目调参更重要。LSTM 虽然结构复杂,但 PyTorch 已经提供了高度优化的实现,我们更应该关注如何根据任务特点选择合适的序列模型。
