PyTorch中层数大于1的堆叠RNN(含GRU)实现规则问询
PyTorch 堆叠RNN(num_layers>1)实现逻辑解答
首先直接给出核心结论:你提到的官方文档描述中的堆叠逻辑,指的是第一层每个时间步的cell输出都会输入到下一层对应时间步的cell中,而非仅传递第一层最后一个cell的输出到下一层。
具体实现细节如下:
- PyTorch内置的RNN、GRU、LSTM层的堆叠逻辑完全统一:每一层会输出对应所有时间步的隐藏状态序列,该完整序列会直接作为下一层的输入序列,时间步严格一一对应:第
l层第t个时间步的输出,会作为第l+1层第t个时间步的输入。 - 以
num_layers=2的单向GRU为例,完整计算流程如下:- 第一层GRU接收原始输入序列
x = [x₁, x₂, ..., x_T](T为序列总长度),计算得到第一层的全时间步输出序列h¹ = [h¹₁, h¹₂, ..., h¹_T],以及第一层最后一个时间步的隐藏状态h¹_T - 第二层GRU的输入为第一层的完整输出序列
h¹,按对应时间步计算得到第二层的全时间步输出序列h² = [h²₁, h²₂, ..., h²_T],以及第二层最后一个时间步的隐藏状态h²_T - 最终GRU层对外返回的输出默认是最后一层的全时间步输出序列
h²,返回的隐藏状态为所有层最后一个时间步的隐藏状态拼接结果[h¹_T, h²_T]
- 第一层GRU接收原始输入序列
- 如果你需要实现「仅把上一层最后一个时间步的输出作为下一层输入」的堆叠逻辑,需要手动拆分RNN层逐个实现,无法通过内置的
num_layers参数直接实现。
官方文档相关描述参考:
循环层的数量。例如,设置num_layers=2意味着将两个GRU堆叠在一起形成堆叠GRU,第二个GRU接收第一个GRU的输出并计算最终结果。
你也可以通过输入输出维度快速验证该逻辑:当设置batch_first=True时,输入GRU的张量维度为(batch_size, seq_len, input_size),输出的序列张量维度为(batch_size, seq_len, hidden_size),刚好和下一层同隐藏维度RNN的输入维度要求匹配,也侧面说明传递的是完整的序列而非单个时间步的张量。
内容的提问来源于stack exchange,提问作者Justin
相关产品推荐
相关产品推荐

