循环神经网络(RNN)入门实战:从梯度消失问题到LSTM解决方案

1次阅读
没有评论

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

image.webp

为什么需要 RNN?

想象这样一个场景:我们要预测明天股票的收盘价。如果只用今天的价格作为输入,显然忽略了历史价格的变化趋势。这种 序列数据 的特点就是前后数据点之间存在依赖关系,而传统全连接网络无法捕捉这种时间维度上的模式。

循环神经网络 (RNN) 入门实战:从梯度消失问题到 LSTM 解决方案

这时候循环神经网络 (RNN) 就派上用场了——它通过引入隐藏状态 $h_t$ 来 ” 记住 ” 过去的信息:

$$
h_t = \tanh(W_{xh}x_t + W_{hh}h_{t-1} + b_h)
$$

但实际使用时会发现,当序列长度超过 20 步后,模型就很难学到长期依赖了。这是因为在反向传播时,梯度需要沿着时间步连续相乘(称为 BPTT 算法),导致梯度值指数级衰减——这就是著名的 梯度消失问题

LSTM 如何解决这个难题?

长短期记忆网络 (LSTM) 通过引入三个门控单元,实现了对信息流的精细控制:

  1. 遗忘门 决定丢弃哪些历史信息:
    $$
    f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)
    $$

  2. 输入门 控制新信息的写入:
    $$
    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)
    $$

  3. 输出门 决定当前隐藏状态的输出:
    $$
    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 维的隐藏状态降维后绘制动图,可以清晰看到随着输入序列的变化,模型内部状态的空间分布变化。

避坑指南

  1. 梯度裁剪:阈值通常设置在 0.25-1.0 之间,太小会阻碍学习,太大失去防护作用
  2. 序列 padding:一定要在 pack_padded_sequence 之前按长度降序排列,否则计算会浪费大量资源
  3. 双向 LSTM 优化 :可以使用batch_first=True 参数让批次维度在前,减少内存拷贝开销

拓展思考

  1. GRU 如何通过合并 LSTM 的门控结构来减少参数量?
  2. 当我们在 LSTM 的每个时间步加入注意力机制时,应该如何设计 score 函数?

(完整代码和训练曲线示例已上传 Github 仓库,文末链接可查看)

通过这次实践,我深刻体会到:理解 RNN 的核心缺陷比盲目调参更重要。LSTM 虽然结构复杂,但 PyTorch 已经提供了高度优化的实现,我们更应该关注如何根据任务特点选择合适的序列模型。

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