长短期记忆网络(LSTM)模型直观解析:从数学原理到PyTorch实战

1次阅读
没有评论

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

image.webp

为什么需要 LSTM?

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

长短期记忆网络 (LSTM) 模型直观解析:从数学原理到 PyTorch 实战

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 表示 ” 完全保留 ”。

输入门

输入门决定我们要在细胞状态中存储哪些新信息。它包含两部分:

  1. 决定更新哪些值的门:
    $$i_t = \sigma(W_i \cdot [h_{t-1}, x_t] + b_i)$$

  2. 候选值向量:
    $$\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_sequencepad_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

思考题

  1. 为什么 LSTM 比 GRU 参数更多但效果不一定更好?
  2. LSTM 有三个门 (遗忘门、输入门、输出门) 和细胞状态,而 GRU 只有两个门(重置门、更新门)
  3. 更多参数意味着更强的表达能力,但也更容易过拟合
  4. 在一些任务中,GRU 的简化结构反而能取得更好的效果

  5. 如何用 LSTM 的隐藏状态做注意力机制?

  6. 可以使用所有时间步的隐藏状态作为键和值
  7. 当前时间步的隐藏状态作为查询
  8. 计算注意力权重并加权求和

总结

LSTM 通过精巧的门控机制解决了 RNN 的长期依赖问题。理解其数学原理对于调参和问题诊断至关重要。在实践中,建议:

  • 从小规模模型开始,逐步增加复杂度
  • 监控梯度范数,适当使用梯度裁剪
  • 对于变长序列使用打包 / 解包操作
  • 尝试不同的初始化策略

希望这篇文章能帮助你直观理解 LSTM 的工作原理,并在实际项目中灵活应用。

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