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
相关产品推荐
相关产品推荐

