PyTorch中如何获取双向LSTM的最后隐藏状态?
如何在PyTorch中提取双向LSTM的最后前向/反向隐藏/细胞状态
嘿,这个问题确实是用PyTorch做序列模型时的常见困惑,我来给你拆解清楚双向LSTM的输出结构,以及怎么拿到你需要的状态~
首先得明确PyTorch中双向LSTM的输出规则:当你运行output, (hn, cn) = bi_lstm(input, (h0, c0))时,各变量的形状和含义如下:
output:形状为(seq_len, batch_size, hidden_size * 2),每个时间步的输出是前向LSTM当前步状态 + 反向LSTM对应步状态的拼接hn:形状为(num_layers * 2, batch_size, hidden_size),存储了所有层的最后隐藏状态,前num_layers个是前向LSTM的各层最后状态,后num_layers个是反向LSTM的各层最后状态cn:和hn形状一致,存储的是各层的最后细胞状态,规则和hn完全相同
1. 从hn和cn提取目标状态
这是最直接的方式,因为PyTorch已经把最终状态打包好了:
假设你的双向LSTM有num_layers层,比如常见的1层:
- 前向LSTM的最后隐藏状态:取
hn的前num_layers层,代码示例:# 1层双向LSTM的情况 forward_last_h = hn[0, :, :] # 或者 hn[:1, :, :] 保留维度 # 多层的话,比如2层:forward_last_h = hn[:2, :, :] - 反向LSTM的最后隐藏状态(即处理完完整逆序序列后的状态):取
hn的后num_layers层,代码示例:# 1层双向LSTM的情况 backward_last_h = hn[1, :, :] # 或者 hn[1:, :, :] 保留维度 # 多层的话,比如2层:backward_last_h = hn[2:, :, :]
细胞状态cn的提取逻辑完全一样:
# 前向最后细胞状态 forward_last_c = cn[:num_layers, :, :] # 反向最后细胞状态 backward_last_c = cn[num_layers:, :, :]
2. 从output提取对应状态
如果你想从序列输出output中获取这些状态,也可以做到:
- 前向LSTM的最后一步状态:对应原序列的最后一个时间步,取
output的最后一行的前半段:hidden_size = 20 # 替换成你的LSTM hidden_size forward_last_h_from_output = output[-1, :, :hidden_size] - 反向LSTM的最后状态(处理完逆序序列后的状态):这里要注意,反向LSTM是从原序列的最后一个元素开始,逆序处理到第一个元素。它的最终状态是处理完原序列第一个元素后的状态,对应
output的第一行的后半段:
你可以验证一下,这个值和从backward_last_h_from_output = output[0, :, hidden_size:]hn中提取的backward_last_h是完全相等的~
举个完整的小例子验证一下:
import torch import torch.nn as nn # 初始化1层双向LSTM input_size = 10 hidden_size = 20 bi_lstm = nn.LSTM(input_size=input_size, hidden_size=hidden_size, num_layers=1, bidirectional=True) # 构造输入:seq_len=5, batch_size=3, input_size=10 input_seq = torch.randn(5, 3, input_size) # 初始化初始状态 h0 = torch.randn(2, 3, hidden_size) c0 = torch.randn(2, 3, hidden_size) # 前向传播 output, (hn, cn) = bi_lstm(input_seq, (h0, c0)) # 提取前向最后隐藏状态 forward_h = hn[0] forward_h_from_output = output[-1, :, :hidden_size] print(torch.allclose(forward_h, forward_h_from_output)) # 输出 True # 提取反向最后隐藏状态 backward_h = hn[1] backward_h_from_output = output[0, :, hidden_size:] print(torch.allclose(backward_h, backward_h_from_output)) # 输出 True
内容的提问来源于stack exchange,提问作者miditower
相关产品推荐
相关产品推荐

