You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.13 07:53:32