PyTorch LSTM / GRU
循环神经网络(RNN)在处理序列数据时面临梯度消失问题,导致难以学习长距离依赖关系。
长短期记忆网络(Long Short-Term Memory,LSTM)和门控循环单元(Gated Recurrent Unit,GRU)通过引入门控机制解决了这一问题,是处理时间序列、自然语言、语音等序列任务的核心模型。
1. RNN 的局限与门控机制
标准 RNN 在每个时间步将当前输入与上一步的隐藏状态合并计算新的隐藏状态:
这一结构存在两个核心问题:
梯度消失:反向传播时梯度需要逐步乘以权重矩阵,序列较长时梯度指数级衰减,导致早期时间步的参数几乎不更新,模型无法学习长距离依赖。
梯度爆炸:当权重矩阵的最大特征值大于 1 时,梯度反向传播过程中指数级增大,训练不稳定(通常用梯度裁剪缓解)。
门控机制的核心思想是引入可学习的"开关",让网络自主决定:在当前时间步,哪些信息应该被记住,哪些应该被遗忘,哪些新信息应该写入记忆。
LSTM 使用三个门(遗忘门、输入门、输出门)加上一个单独的细胞状态;GRU 将结构简化为两个门(重置门、更新门),参数更少,训练更快。
2. LSTM 原理
2.1 核心结构与三个门
LSTM 维护两个状态向量在时间步之间传递:
细胞状态(Cell State):长期记忆的载体,信息可以在其中几乎无损地流动 隐藏状态(Hidden State):短期记忆,也是当前时间步的输出
三个门均为 Sigmoid 激活的线性变换,输出值在 0~1 之间,起到"阀门"的作用:
python
遗忘门(Forget Gate):决定从细胞状态中丢弃哪些信息
输入门(Input Gate):决定将哪些新信息写入细胞状态
输出门(Output Gate):决定基于细胞状态输出什么2.2 前向计算公式
计算逻辑解读:
- 重置门
r_t接近 0 时,候选状态h̃_t几乎不依赖历史,相当于重新开始 - 更新门
z_t接近 1 时,新状态更多采用候选值;接近 0 时,更多保留历史状态 - GRU 没有独立的细胞状态,参数量约为 LSTM 的 75%
4. PyTorch 中的 LSTM
本节详细介绍 nn.LSTM 的参数、输入输出形状以及隐藏状态初始化方法。
4.1 nn.LSTM 参数详解
实例
python
import torch
import torch.nn as nn
lstm = nn.LSTM(
input_size=64, # 每个时间步输入向量的维度
hidden_size=128, # 隐藏状态(以及细胞状态)的维度
num_layers=2, # 堆叠层数,默认为 1
bias=True, # 是否使用偏置项,默认 True
batch_first=False, # 输入/输出 shape 中 batch 是否在第一维,默认 False
dropout=0.0, # 层间 dropout 概率(仅在 num_layers > 1 时生效)
bidirectional=False, # 是否使用双向 LSTM,默认 False
proj_size=0, # 投影层维度(LSTM with projection),默认 0 表示不使用
)
# 查看参数量
total_params = sum(p.numel() for p in lstm.parameters())
print(f"LSTM 参数量: {total_params:,}")
# input_size=64, hidden_size=128, num_layers=2 时约为 197,632参数量估算公式(单层单向):