LSTM长短期记忆网络:从1997年诞生到解决RNN长期记忆问题的技术演进

1次阅读
没有评论

共计 2015 个字符,预计需要花费 6 分钟才能阅读完成。

image.webp

背景痛点: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 到现代变体

  1. 1997 年原始 LSTM(Hochreiter & Schmidhuber)
  2. 引入细胞状态 $C_t$ 作为 ” 记忆通道 ”
  3. 三个门控单元:输入门、输出门、遗忘门
  4. 初始版本未包含 peephole 连接

  5. 2000 年改进 LSTM(Gers et al.)

  6. 添加遗忘门偏置(初始值通常设为 1)
  7. 引入 peephole 连接,允许门控查看细胞状态

  8. 2014 年 GRU 简化版(Cho et al.)

  9. 合并遗忘门与输入门为更新门
  10. 耦合细胞状态与隐藏状态
  11. 参数减少 33% 但性能相近

核心架构:门控机制详解

LSTM 长短期记忆网络:从 1997 年诞生到解决 RNN 长期记忆问题的技术演进

遗忘门 决定保留多少旧记忆:
$$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

训练避坑指南

  1. 梯度爆炸
  2. 使用梯度裁剪:torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5)
  3. 初始化遗忘门偏置为 1(促进初始阶段记忆保留)

  4. 超参数敏感

  5. 学习率:从 3e- 4 开始尝试,配合学习率调度器
  6. 隐藏层大小:根据 GPU 内存选择(通常 128-512)
  7. Dropout 率:层间 dropout 建议 0.2-0.5(需设置 lstm.dropout 参数)

  8. 序列填充干扰

  9. 使用 pack_padded_sequence 处理变长序列
  10. 设置 enforce_sorted=False 允许乱序输入

延伸思考:Transformer 时代的 LSTM

虽然 Transformer 在长序列建模中表现突出,但 LSTM 仍有其优势场景:

  • 低资源环境:LSTM 参数更少,训练成本低
  • 在线学习:LSTM 可逐时间步处理,而 Transformer 需要完整序列
  • 小规模数据:LSTM 不易过拟合

实际选择建议:
– 序列长度 <100 且数据量小时优先考虑 LSTM
– 需要处理文档级依赖时选用 Transformer
– 可尝试 LSTM+Attention 的混合架构

正文完
 0
评论(没有评论)