共计 4072 个字符,预计需要花费 11 分钟才能阅读完成。
为什么需要 LSTM?
在传统 RNN 处理长序列时,我们经常遇到梯度消失问题。简单来说,当误差反向传播时,梯度会随着时间步不断乘积(特别是当权重矩阵的特征值小于 1 时),导致距离当前时刻较远的时刻几乎无法更新参数。这就像试图记住 100 个单词前的第一个单词,但记忆在传递过程中逐渐模糊。

LSTM 通过引入门控机制和细胞状态(cell state)解决了这个问题。与普通 RNN 相比,LSTM 多了三个控制门:
- 遗忘门:决定丢弃哪些信息
- 输入门:决定存储哪些新信息
- 输出门:决定输出什么信息
这种结构让信息可以像高速公路一样在细胞状态中流动,减少了梯度消失的影响。
门控机制详解
遗忘门
遗忘门控制着哪些信息应该被丢弃。数学表达式为:
$$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$$
其中:
– $\sigma$ 是 sigmoid 函数,输出在 0 到 1 之间
– $W_f$ 是权重矩阵
– $[h_{t-1}, x_t]$ 表示将上一时刻的隐藏状态和当前输入拼接起来
这个门输出一个 0 到 1 之间的值,0 表示 ” 完全忘记 ”,1 表示 ” 完全保留 ”。
输入门
输入门决定我们要在细胞状态中存储哪些新信息。它包含两部分:
-
决定更新哪些值的门:
$$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)$$
然后将这两部分结合起来更新细胞状态:
$$C_t = f_t \ast C_{t-1} + i_t \ast \tilde{C}_t$$
这里的 $\ast$ 表示逐元素相乘。
输出门
输出门决定下一个隐藏状态应该是什么。首先我们计算输出门的值:
$$o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o)$$
然后结合当前细胞状态计算新的隐藏状态:
$$h_t = o_t \ast \tanh(C_t)$$
PyTorch 实现
手工实现 LSTM 单元
import torch
import torch.nn as nn
class LSTMCell(nn.Module):
def __init__(self, input_size, hidden_size):
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
# 遗忘门参数
self.W_f = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_f = nn.Parameter(torch.Tensor(hidden_size))
# 输入门参数
self.W_i = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_i = nn.Parameter(torch.Tensor(hidden_size))
# 候选值参数
self.W_c = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_c = nn.Parameter(torch.Tensor(hidden_size))
# 输出门参数
self.W_o = nn.Parameter(torch.Tensor(hidden_size, hidden_size + input_size))
self.b_o = nn.Parameter(torch.Tensor(hidden_size))
self.reset_parameters()
def reset_parameters(self):
"""参数初始化"""
# 使用正交初始化
nn.init.orthogonal_(self.W_f)
nn.init.orthogonal_(self.W_i)
nn.init.orthogonal_(self.W_c)
nn.init.orthogonal_(self.W_o)
# 偏置初始化
nn.init.constant_(self.b_f, 0.)
nn.init.constant_(self.b_i, 0.)
nn.init.constant_(self.b_c, 0.)
# 输出门偏置初始化为 1,有助于训练初期保留更多信息
nn.init.constant_(self.b_o, 1.)
def forward(self, x, h_prev, c_prev):
"""
参数:
x: 当前输入 (batch_size, input_size)
h_prev: 前一个隐藏状态 (batch_size, hidden_size)
c_prev: 前一个细胞状态 (batch_size, hidden_size)
"""
# 拼接输入和前一隐藏状态
combined = torch.cat((h_prev, x), dim=1) # (batch_size, hidden_size + input_size)
# 计算遗忘门
f = torch.sigmoid(combined @ self.W_f.t() + self.b_f)
# 计算输入门
i = torch.sigmoid(combined @ self.W_i.t() + self.b_i)
# 计算候选值
c_tilde = torch.tanh(combined @ self.W_c.t() + self.b_c)
# 更新细胞状态
c = f * c_prev + i * c_tilde
# 计算输出门
o = torch.sigmoid(combined @ self.W_o.t() + self.b_o)
# 计算新隐藏状态
h = o * torch.tanh(c)
return h, c
使用 PyTorch 内置 LSTM
import torch.nn as nn
# 定义 LSTM 模型
lstm = nn.LSTM(
input_size=100, # 输入特征维度
hidden_size=256, # 隐藏层维度
num_layers=2, # LSTM 层数
batch_first=True, # 输入格式为(batch, seq_len, feature)
dropout=0.2, # 层间 dropout
bidirectional=False # 是否双向
)
# 前向传播示例
batch_size = 32
seq_len = 50
input_dim = 100
# 输入数据 (batch_size, seq_len, input_dim)
inputs = torch.randn(batch_size, seq_len, input_dim)
# 初始隐藏状态和细胞状态
h0 = torch.zeros(2, batch_size, 256) # (num_layers, batch_size, hidden_size)
c0 = torch.zeros(2, batch_size, 256)
# 前向传播
output, (hn, cn) = lstm(inputs, (h0, c0))
print(output.shape) # torch.Size([32, 50, 256])
print(hn.shape) # torch.Size([2, 32, 256])
print(cn.shape) # torch.Size([2, 32, 256])
工程实践技巧
1. 参数初始化
- 权重矩阵使用正交初始化
- 遗忘门偏置初始化为 1(有助于保留长期依赖)
2. 梯度裁剪
在训练过程中添加梯度裁剪:
optimizer.zero_grad()
loss.backward()
# 裁剪梯度范数不超过 5
nn.utils.clip_grad_norm_(model.parameters(), max_norm=5)
optimizer.step()
3. 处理变长序列
使用 pack_padded_sequence 和pad_packed_sequence:
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence
# 假设我们有一个 batch 的序列,长度分别为[10,8,6,4]
lengths = torch.tensor([10,8,6,4])
# 按照长度从长到短排序
lengths, perm_idx = lengths.sort(0, descending=True)
inputs = inputs[perm_idx]
# 打包序列
packed_input = pack_padded_sequence(inputs, lengths, batch_first=True)
# 通过 LSTM
packed_output, (hn, cn) = lstm(packed_input)
# 解包序列
output, _ = pad_packed_sequence(packed_output, batch_first=True)
# 恢复原始顺序
_, unperm_idx = perm_idx.sort(0)
output = output[unperm_idx]
hn = hn[:, unperm_idx]
cn = cn[:, unperm_idx]
内存占用分析
对于一个输入维度为 100,隐藏层维度为 256 的单层 LSTM:
- 参数数量:4 × (100 + 256) × 256 = 364,544
- 每个时间步的激活内存:(256 × 4) × batch_size
思考题
- 为什么 LSTM 比 GRU 参数更多但效果不一定更好?
- LSTM 有三个门 (遗忘门、输入门、输出门) 和细胞状态,而 GRU 只有两个门(重置门、更新门)
- 更多参数意味着更强的表达能力,但也更容易过拟合
-
在一些任务中,GRU 的简化结构反而能取得更好的效果
-
如何用 LSTM 的隐藏状态做注意力机制?
- 可以使用所有时间步的隐藏状态作为键和值
- 当前时间步的隐藏状态作为查询
- 计算注意力权重并加权求和
总结
LSTM 通过精巧的门控机制解决了 RNN 的长期依赖问题。理解其数学原理对于调参和问题诊断至关重要。在实践中,建议:
- 从小规模模型开始,逐步增加复杂度
- 监控梯度范数,适当使用梯度裁剪
- 对于变长序列使用打包 / 解包操作
- 尝试不同的初始化策略
希望这篇文章能帮助你直观理解 LSTM 的工作原理,并在实际项目中灵活应用。
