共计 2260 个字符,预计需要花费 6 分钟才能阅读完成。
为什么需要 LSTM?
传统 RNN 在处理长序列时面临梯度消失(vanishing gradient)问题。简单来说,当误差反向传播时,梯度会随着时间步不断相乘。如果这个乘数小于 1,经过若干步后梯度就会趋近于零,导致网络无法学习长期依赖关系。

数学表达为:
$$
\frac{\partial E}{\partial W} = \sum_{t=1}^{T}\frac{\partial E}{\partial y_T}\frac{\partial y_T}{\partial h_t}\frac{\partial h_t}{\partial h_{t-1}}…\frac{\partial h_1}{\partial W}
$$
LSTM 通过引入门控机制(gate mechanism)和细胞状态(cell state)解决了这个问题。下面我们拆解它的核心组件。
LSTM 三大门控详解
1. 遗忘门(Forget Gate)
决定哪些信息从细胞状态中被丢弃,计算公式:
$$
f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)
$$
其中 $\sigma$ 是 sigmoid 函数,输出值在 0 到 1 之间,表示保留信息的比例。
2. 输入门(Input Gate)
控制新信息的加入,包含两个部分:
$$
\begin{aligned}
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{aligned}
$$
$i_t$ 决定更新程度,$\tilde{C}_t$ 是候选新值。
3. 细胞状态更新
结合遗忘门和输入门更新记忆:
$$
C_t = f_t \circ C_{t-1} + i_t \circ \tilde{C}_t
$$
$\circ$ 表示逐元素相乘。
4. 输出门(Output Gate)
控制最终输出:
$$
\begin{aligned}
o_t &= \sigma(W_o \cdot [h_{t-1}, x_t] + b_o) \
h_t &= o_t \circ \tanh(C_t)
\end{aligned}
$$
PyTorch 实现完整 LSTMCell
import torch
import torch.nn as nn
class LSTMCell(nn.Module):
def __init__(self, input_size, hidden_size):
super().__init__()
# 合并所有门的权重计算(PyTorch 风格)self.weight_ih = nn.Parameter(torch.randn(4 * hidden_size, input_size))
self.weight_hh = nn.Parameter(torch.randn(4 * hidden_size, hidden_size))
self.bias = nn.Parameter(torch.zeros(4 * hidden_size))
# 正交初始化提升训练稳定性
nn.init.orthogonal_(self.weight_hh)
self.hidden_size = hidden_size
def forward(self, x, state):
h_prev, c_prev = state
# 合并计算所有门(效率优化)gates = (x @ self.weight_ih.T + h_prev @ self.weight_hh.T + self.bias)
# 分割得到各门(按 hidden_size 切分)i, f, g, o = gates.chunk(4, dim=1)
# 计算门控信号
i = torch.sigmoid(i) # 输入门
f = torch.sigmoid(f) # 遗忘门
o = torch.sigmoid(o) # 输出门
g = torch.tanh(g) # 候选记忆
# 更新细胞状态
c_next = f * c_prev + i * g
# 计算隐藏状态
h_next = o * torch.tanh(c_next)
return h_next, c_next
RNN vs LSTM 实战对比
我们构造一个简单的序列复制任务测试:
- 普通 RNN:
- 序列长度 >10 时准确率急剧下降
-
损失函数波动剧烈
-
LSTM:
- 能稳定处理 100+ 长度的序列
- 训练曲线平滑收敛
生产环境注意事项
初始化技巧
- 使用正交初始化(orthogonal initialization)门控权重
- 偏置项特殊处理:遗忘门偏置初始化为 1(促进早期记忆)
训练稳定性
- 梯度裁剪(gradient clipping)阈值设为 1.0-5.0
- 监控门激活值分布:理想情况应接近 0 或 1
性能优化
- 启用 CuDNN 加速:
torch.backends.cudnn.enabled = True - 注意版本兼容性问题
可视化门控信号
from torch.utils.tensorboard import SummaryWriter
writer = SummaryWriter()
# 在训练循环中添加:if global_step % 100 == 0:
writer.add_histogram('forget_gate', f_t, global_step)
writer.add_histogram('input_gate', i_t, global_step)
writer.add_scalar('cell_state/mean', c_t.mean(), global_step)
思考题
如何验证遗忘门在长程依赖中的作用?建议实验方案:
- 构造需要记忆 50 步前信息的任务
- 固定遗忘门权重为 1(始终记忆)
- 固定遗忘门权重为 0(始终遗忘)
- 对比三种情况在验证集的表现差异
期待你在实践中发现更多有趣的现象!
