在 PyTorch 中,LSTM 模块位于 torch.nn 中,可以通过 torch.nn.LSTM 调用。LSTM 能够处理变长序列,并解决传统 RNN 的梯度消失问题。
LSTM 通过三个门(输入门、遗忘门、输出门)和一个细胞状态来控制信息流动,从而更好地捕捉长期依赖关系。
input_size: 输入特征的维度。
hidden_size: 隐藏状态的维度。
num_layers: LSTM 层数。
batch_first: 如果为 True,输入和输出的形状为 (batch_size, seq_len, input_size)。
bidirectional: 是否使用双向 LSTM。
以下是一个简单的 LSTM 模型示例:
import torch
import torch.nn as nn
# 定义 LSTM 模型
class LSTMModel(nn.Module):
def __init__(self, input_size, hidden_size, num_layers, output_size):
super(LSTMModel, self).__init__()
self.hidden_size = hidden_size
self.num_layers = num_layers
self.lstm = nn.LSTM(input_size, hidden_size, num_layers, batch_first=True)
self.fc = nn.Linear(hidden_size, output_size)
def forward(self, x):
h0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
c0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
out, _ = self.lstm(x, (h0, c0))
out = self.fc(out[:, -1, :])
return out
# 参数设置
input_size = 10
hidden_size = 20
num_layers = 2
output_size = 1
batch_size = 5
seq_len = 7
# 创建模型实例
model = LSTMModel(input_size, hidden_size, num_layers, output_size)
# 示例输入
x = torch.randn(batch_size, seq_len, input_size)
# 前向传播
output = model(x)
print(output.shape) # 输出形状: (batch_size, output_size)
代码说明
LSTMModel 类: 定义了 LSTM 模型,包含一个 LSTM 层和一个全连接层。
forward 方法: 实现前向传播,初始化隐藏状态和细胞状态,通过 LSTM 层和全连接层得到输出。
参数设置: 定义了输入、隐藏状态、层数和输出的维度。
示例输入: 创建了一个随机输入张量。
前向传播: 调用模型进行前向传播,输出形状为 (batch_size, output_size)。
总结
PyTorch 的 LSTM 模块简化了序列模型的构建和训练,适用于多种序列数据处理任务。