TensorFlow MultiRNNCell中细胞顺序及状态索引对应关系问询
TensorFlow MultiRNNCell 细胞排列与状态元组索引说明
嘿,这个问题我当初刚用MultiRNNCell的时候也纠结过,确实官方文档这块说得不够直白!
直接给你结论:在MultiRNNCell中,状态元组的索引0对应的是最接近输入的底层网络,而索引最大的元素才是紧邻输出的顶层网络。
举个实际例子就好理解了:假设你定义了两层RNN细胞,按顺序传入MultiRNNCell:
from tensorflow.contrib.rnn import MultiRNNCell, BasicLSTMCell # 先定义靠近输入的第一层,再定义靠近输出的第二层 layer1_cell = BasicLSTMCell(units=64) layer2_cell = BasicLSTMCell(units=32) multi_cell = MultiRNNCell([layer1_cell, layer2_cell])
这时候,MultiRNNCell的状态元组(比如initial_state或者前向传播后的state)顺序和你传入的细胞列表完全一致:
state[0]是layer1_cell的状态,也就是输入数据首先经过的那一层;state[1]是layer2_cell的状态,也就是最后一层,输出结果直接来自这一层的输出。
简单来说,你把细胞按「输入到输出」的顺序放进列表,状态元组就会按同样的顺序存储每一层的状态,第一个元素对应最底层,最后一个对应最顶层。
内容的提问来源于stack exchange,提问作者Aswin Suresh
相关产品推荐
相关产品推荐

