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

PyTorch中RNN(nn.LSTM、nn.GRU等)的h_n输出结构及索引疑问

搞懂RNN输出h_n的维度索引逻辑

嘿,这个问题确实很容易绕晕,我来给你拆解清楚~

首先先明确官方文档的定义:h_n是包含所有层、所有方向最后一个时间步隐藏状态的张量,形状是(num_layers * num_directions, batch, hidden_size)。它的第一维度排列逻辑其实是有规律的,分两种情况来看:

1. 单向RNN(num_directions=1)

这时候第一维度的长度就是num_layers,元素是从底层到顶层依次排列的各层最后一步隐藏状态:

  • h_n[0]:最底层(第1层)最后一个时间步的隐藏状态
  • h_n[1]:第2层最后一个时间步的隐藏状态
  • ...
  • h_n[-1]:最上层(第num_layers层)最后一个时间步的隐藏状态

所以你问的「最上层的最后一个单元的隐藏状态」,在单向场景下就是h_n[-1]。

2. 双向RNN(num_directions=2)

这时候第一维度的长度是num_layers * 2,排列规则是先按层从下到上,每层内先放正向、再放反向的最后一步隐藏状态。举个例子,如果是2层双向RNN:

  • h_n[0]:第1层正向最后一个时间步的隐藏状态
  • h_n[1]:第1层反向最后一个时间步的隐藏状态
  • h_n[2]:第2层正向最后一个时间步的隐藏状态
  • h_n[3]:第2层反向最后一个时间步的隐藏状态

这种场景下,最上层的两个方向隐藏状态分别是:

  • 最上层正向:h_n[-2]
  • 最上层反向:h_n[-1]

用代码验证一下

拿PyTorch的LSTM举个实际例子:

import torch
import torch.nn as nn

# 定义2层双向LSTM,输入维度10,隐藏维度20
lstm = nn.LSTM(input_size=10, hidden_size=20, num_layers=2, bidirectional=True)
# 输入序列:(seq_len=5, batch=3, input_size=10)
input_seq = torch.randn(5, 3, 10)
output, (h_n, c_n) = lstm(input_seq)

print(h_n.shape)  # 输出: torch.Size([4, 3, 20]),对应2层*2方向=4个第一维度元素

这里的h_n[3]就是最上层反向的最后一步隐藏状态,也就是h_n[-1],完全符合上面的逻辑。

内容的提问来源于stack exchange,提问作者the-bass

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 08:06:52