使用相同BasicLSTM Cell构建MultiRNNCell报错,求问原因
为什么重复使用同一个BasicLSTMCell实例构建MultiRNNCell会报错?
两段代码的核心差异
- 第一段代码中,
MultiRNNCell的参数是同一个cell1实例被传入两次:multi = tf.contrib.rnn.MultiRNNCell([cell1, cell1]) - 第二段代码中,
MultiRNNCell的参数是两个独立的、结构完全一致的LSTM Cell实例:multi = tf.contrib.rnn.MultiRNNCell([cell1, cell2]),其中cell2是全新创建的、和cell1参数完全相同的BasicLSTMCell。
错误原因解析
TensorFlow中的RNN Cell实例是带有内部状态和变量的对象。当你把同一个cell1实例两次传入MultiRNNCell时,会触发以下问题:
- 第一次使用
cell1作为第一层RNN时,它会根据输入维度(这里是512维)初始化内部的权重矩阵(比如输入到隐藏层的变换矩阵)。 - 当第二次复用同一个
cell1实例作为第二层RNN时,它会尝试复用已经创建好的权重变量,但此时第二层的输入是第一层的输出(128维),和第一次初始化时的输入维度(512维)不匹配。 - 错误日志里的维度冲突就反映了这个问题:第二层输入拼接隐藏状态后的维度是
128(输入)+128(隐藏状态)=256,但权重矩阵是按照第一次的512+128=640维度创建的,矩阵乘法时自然维度不兼容。
而使用两个独立的实例cell1和cell2时,每个实例都会独立初始化自己的权重变量,分别适配第一层的512维输入和第二层的128维输入,因此可以正常运行。
错误信息
ValueError: Dimensions must be equal, but are 256 and 640 for 'multi_rnn_cell/cell_0/cell1/MatMul_1' (op: 'MatMul') with input shapes: [64,256], [640,512].
相关代码示例
失败的代码:
import tensorflow as tf cell1 = tf.contrib.rnn.BasicLSTMCell(128,reuse=False, name = "cell1") cell2 = tf.contrib.rnn.BasicLSTMCell(128,reuse=False,name = "cell2") multi = tf.contrib.rnn.MultiRNNCell([cell1, cell1] ) init = multi.zero_state(64, tf.float32) output,state = multi(tf.ones([64,512]),init)
正常运行的代码:
import tensorflow as tf cell1 = tf.contrib.rnn.BasicLSTMCell(128,reuse=False, name = "cell1") cell2 = tf.contrib.rnn.BasicLSTMCell(128,reuse=False,name = "cell2") multi = tf.contrib.rnn.MultiRNNCell([cell1, cell2] ) init = multi.zero_state(64, tf.float32) output,state = multi(tf.ones([64,512]),init)
内容的提问来源于stack exchange,提问作者Vinod Devarampati
相关产品推荐
相关产品推荐

