PyTorch中GRU的output返回值是否就是隐藏状态?
PyTorch GRU 返回值
output 与 h_n 的区别 首先直接回答你的第一个疑问:
- 对单层、单向的GRU来说,
output就是每个时间步计算出的原始隐藏状态的按序集合,没有经过任何额外的线性变换或激活处理,这时候你说的“从output最后一个元素提取h_n”是完全成立的,两者数值完全一致,可以用下面的代码验证:
import torch import torch.nn as nn # 单层单向GRU测试 gru = nn.GRU(input_size=10, hidden_size=20, num_layers=1, bidirectional=False) # 默认输入维度顺序: (序列长度, batch大小, 输入特征维度) x = torch.randn(5, 3, 10) output, h_n = gru(x) # 验证output最后一个时间步的隐状态和h_n相等 print(torch.allclose(output[-1], h_n[0])) # 输出True
至于为什么框架要同时返回两个值,核心是覆盖不同网络配置、不同任务场景的需求,避免用户手动做冗余的切片、拼接操作:
- 当使用多层GRU时,
output仅保留最后一层所有时间步的隐藏状态,而h_n会存储每一层在最后一个时间步的隐藏状态。比如2层GRU的h_n形状为(2, batch_size, hidden_size),分别对应第一层、第二层的最终隐状态,这部分非最后层的最终隐状态你是没法从output里直接拿到的。 - 当使用双向GRU时,
output会把每个时间步的正向、反向隐状态做拼接,每个时间步的特征维度为2*hidden_size;而h_n会分开存储所有层的正向最终隐状态、反向最终隐状态,比如单层双向GRU的h_n形状为(2, batch_size, hidden_size),第一个元素是正向传播到序列末尾的隐状态,第二个是反向传播到序列开头的隐状态,这部分拆分好的结果如果从output里提取需要手动做切片,很容易出错。 - 从编码便利性来说,两类最常见的序列任务刚好对应两个返回值的使用场景:做序列标注、序列到序列生成这类需要每个时间步隐状态的任务,直接用
output接下游输出层即可;做文本分类这类只需要序列最终隐状态的任务,直接取h_n即可,不需要手动索引output的最后一位。
注意:如果你初始化GRU时传入了
batch_first=True参数,output的维度顺序会变为(batch_size, 序列长度, hidden_size * 方向数),此时单层单向场景下和h_n相等的是output[:, -1, :],数值依然完全一致。
内容的提问来源于stack exchange,提问作者DiveIntoML
相关产品推荐
相关产品推荐

