torch.nn.LSTM 是 PyTorch 中实现长短期记忆网络(Long Short-Term Memory)的核心模块。它是一个多层、双向的 LSTM 实现,广泛应用于序列建模任务,如自然语言处理、时间序列预测和语音识别。
在PyTorch中,torch.nn.LSTM 是一个用于构建和训练LSTM网络的模块。它是 torch.nn 中的一个重要层(Layer ),支持多层堆叠、双向LSTM等特性。

1. LSTM基本原理

LSTM 通过引入门控机制细胞状态来解决标准 RNN 的长期依赖与梯度消失问题。一个 LSTM 单元包含:

  • 输入门 iti_tit:控制新信息写入细胞状态。
  • 遗忘门 ftf_tft:控制旧细胞状态被丢弃的程度。
  • 输出门 oto_tot:控制细胞状态输出到隐藏状态的比例。
  • 细胞状态 ctc_tct:贯穿时间的“记忆传送带”。
  • 隐藏状态 hth_tht:当前时刻的输出。

在这里插入图片描述
前向计算公式如下(以单层为例):
ft=σ(Wifxt+bif+Whfht−1+bhf)(遗忘门)it=σ(Wiixt+bii+Whiht−1+bhi)(输入门)gt=tanh(Wigxt+big+Whght−1+bhg)=C^t(候补单元状态)ot=σ(Wioxt+bio+Whoht−1+bho)(输出门)ct=ft⊗ct−1+it⊗gt(更新单元状态)ht=σt⊗tanh(ct)(输出) f_t=\sigma(W_{if}x_t+b_{if}+W_{hf}h_{t-1}+b_{hf})(遗忘门)\\ i_t=\sigma(W_{ii}x_t+b_{ii}+W_{hi}h_{t-1}+b_{hi})(输入门)\\ g_t=tanh(W_{ig}x_t+b_{ig}+W_{hg}h_{t-1}+b_{hg})=\hat C_t(候补单元状态)\\ o_t=\sigma(W_{io}x_t+b_{io}+W_{ho}h_{t-1}+b_{ho})(输出门)\\ c_t=f_t \otimes c_{t-1}+i_t \otimes g_t(更新单元状态)\\ h_t = \sigma_t \otimes tanh(c_t)(输出) ft=σ(Wifxt+bif+Whfht1+bhf)(遗忘门)it=σ(Wiixt+bii+Whiht1+bhi)(输入门)gt=tanh(Wigxt+big+Whght1+bhg)=C^t(候补单元状态)ot=σ(Wioxt+bio+Whoht1+bho)(输出门)ct=ftct1+itgt(更新单元状态)ht=σttanh(ct)(输出)

torch.nn.LSTM 将上述计算封装为一个可堆叠多层的模块,并支持双向传播。

2. 初始化参数详解

torch.nn.LSTM(
    input_size,      # 输入特征维度(每个时间步的 x_t 维度)
    hidden_size,     # 隐藏状态(和细胞状态)的特征维度
    num_layers=1,    # LSTM 层的堆叠数量,默认1
    bias=True,       # 是否使用偏置项 b_ih 和 b_hh,默认 True
    batch_first=False,  # 输入张量形状是否将 batch 放在第一维,默认 False(seq_len, batch, input_size)
    dropout=0,       # 若 num_layers > 1,在层间(除最后一层外)添加 dropout,默认0
    bidirectional=False, # 是否使用双向 LSTM,默认 False
    proj_size=0,     # 如果 >0,将对隐藏状态进行投影降维(PyTorch 1.8+)
    device=None,
    dtype=None
)

关键参数说明:

参数 类型 作用
input_size int 每个时间步输入向量的维度。
hidden_size int 隐藏状态hth_tht的维度。LSTM 单元内部所有门控的计算均依赖此维度。
num_layers int 堆叠的 LSTM 层数。例如 num_layers=2 时,第一层的输出作为第二层的输入。
bidirectional bool 若为 True,则变为双向 LSTM,此时输出维度变为 2 * hidden_size
batch_first bool 默认为 False,输入形状为 (seq_len, batch, input_size);设为 True 时形状为 (batch, seq_len, input_size),更符合常规习惯。
dropout float 在除最后一层外的层之间添加 Dropout,用于正则化。仅在 num_layers > 1 时生效。
proj_size int 若 >0,则在输出hth_tht之前插入一个线性投影层,将 hidden_size 压缩至 proj_size(类似于 LSTMP)。

返回值:
LSTM返回一个元组:

output, (h_n, c_n) = lstm(input, (h_0, c_0))

output:输出张量,形状为

  • seq_len, batch_size, num * hidden_size)(batch_first=False)
  • batch_size, seq_len, num * hidden_size)(batch_first=True)
  • 如果bidirectional=Truenum=2,否则为1。

h_n:最后一个时间步的隐藏状态,形状为
(numlayers×num,batch_size,hidden_size) (num_layers × num, batch\_size, hidden\_ size) (numlayers×num,batch_size,hidden_size)
c_n:最后一个时间步的细胞状态,形状同h_n

3. 基本使用

3.1 单层LSTM

import torch
import torch.nn as nn

input_size = 10		# 每个时间步的输入维度
hidden_size = 20		# 隐藏层维度
seq_len = 5				# 序列长度
batch_size = 3		# 批次大小

lstm = nn.LSTM(input_size, hidden_size, num_layers=1, batch_first=True)

# 生成随机输入数据
x = torch.randn(batch_size, seq_len, input_size)

# 初始化隐藏状态和细胞状态
h0 = torch.zeros(1, batch_size, hidden_size)  # (num_layers * num, batch_size, hidden_size)
c0 = torch.zeros(1, batch_size, hidden_size)

# 前向传播
output, (hn, cn) = lstm(x, (h0, c0))

print("输出形状:", output.shape)  	  # (batch_size, seq_len, hidden_size)
print("最后的隐藏状态形状:", hn.shape)  # (num_layers, batch_size, hidden_size)
print("最后的细胞状态形状:", cn.shape)  # (num_layers, batch_size, hidden_size)

(h0, c0)如果不提供,默认是0。

3.2 堆叠多层 LSTM

import torch
import torch.nn as nn

input_size = 10		# 每个时间步的输入维度
hidden_size = 20		# 隐藏层维度
seq_len = 5				# 序列长度
batch_size = 3		# 批次大小

lstm = nn.LSTM(input_size, hidden_size, num_layers=3, batch_first=True)

# 生成随机输入数据
x = torch.randn(batch_size, seq_len, input_size)

# 初始化多层隐藏状态和细胞状态
h0 = torch.zeros(3, batch_size, hidden_size)
c0 = torch.zeros(3, batch_size, hidden_size)

output, (hn, cn) = lstm(x, (h0, c0))

print("多层 LSTM 的隐藏状态形状:", hn.shape)  # (num_layers, batch_size, hidden_size)
print("多层 LSTM 的细胞状态形状:", cn.shape)

3.3 双向 LSTM

import torch
import torch.nn as nn

input_size = 10		# 每个时间步的输入维度
hidden_size = 20		# 隐藏层维度
seq_len = 5				# 序列长度
batch_size = 3		# 批次大小

lstm = nn.LSTM(input_size, hidden_size, num_layers=1, batch_first=True, bidirectional=True)

# 生成随机输入数据
x = torch.randn(batch_size, seq_len, input_size)

# 初始化双向 LSTM
h0 = torch.zeros(2, batch_size, hidden_size)  # 2 = num_directions
c0 = torch.zeros(2, batch_size, hidden_size)

output, (hn, cn) = lstm(x, (h0, c0))

print("双向 LSTM 的输出形状:", output.shape)  # (batch_size, seq_len, 2 * hidden_size)
print("双向 LSTM 的隐藏状态形状:", hn.shape)  # (num_layers * 2, batch_size, hidden_size)

3.4 带投影的 LSTM

import torch
import torch.nn as nn

input_size = 10		# 每个时间步的输入维度
hidden_size = 20		# 隐藏层维度
seq_len = 5				# 序列长度
batch_size = 3		# 批次大小

lstm = nn.LSTM(input_size, hidden_size, num_layers=1, batch_first=True, proj_size=8)

# 生成随机输入数据
x = torch.randn(batch_size, seq_len, input_size)

# 初始化隐藏和细胞状态
h0 = torch.zeros(1, batch_size, 8)  # proj_size 维度
c0 = torch.zeros(1, batch_size, hidden_size)

output, (hn, cn) = lstm(x, (h0, c0))

print("带投影的 LSTM 的输出形状:", output.shape)  	# (batch_size, seq_len, proj_size)
print("带投影的隐藏状态形状:", hn.shape)  				# (num_layers, batch_size, proj_size)
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐