TensorFlow RNN初始化报错:用张量设置batch_size的问题排查
解决TensorFlow中RNN初始化batch_size时的ValueError问题
你的问题出在cell.zero_state()对batch_size参数的要求上——它需要一个能确定具体数值的输入,要么是Python标量,要么是运行时可获取实际值的Tensor,但你传入的x_rnn.shape[0]是静态形状,在图构建阶段它的值是None(因为你把x_rnn的第一维设为None来支持可变batch size),导致TensorFlow无法创建形状包含?的状态张量,所以抛出了这个错误。
下面给你两种可行的解决方案:
方案1:使用动态形状获取batch size
直接把调用RNN时的batch_size=x_rnn.shape[0]替换成batch_size=tf.shape(x_rnn)[0]。tf.shape()是TensorFlow的动态形状操作,会在运行时根据输入数据的实际情况获取batch大小,完全适配可变batch的场景。
修改后的调用代码:
y_rnn = RNN(x_rnn, weights_rnn, biases_rnn, n_inputs, n_steps, n_hidden_units, batch_size=tf.shape(x_rnn)[0])
方案2:显式定义batch_size占位符
如果你希望更清晰地管控batch size,可以单独定义一个占位符来传递这个参数:
# 定义batch_size的占位符,类型为int32 batch_size = tf.placeholder(tf.int32, []) # 调用RNN时传入这个占位符 y_rnn = RNN(x_rnn, weights_rnn, biases_rnn, n_inputs, n_steps, n_hidden_units, batch_size=batch_size)
之后在运行会话喂数据时,记得同时feed这个占位符的实际值(比如你每次喂入的batch大小,比如128)。
额外提示
我注意到你的RNN函数里漏写了tf.nn.dynamic_rnn的调用代码——你在处理outputs的unstack之前,应该先获取动态RNN的输出:
# 补上这一行,获取RNN的输出和最终状态 outputs, final_state = tf.nn.dynamic_rnn(cell, X_in, initial_state=init_state, dtype=tf.float32)
如果没这行代码,后续的outputs会报未定义的错误,记得补上哦。
内容的提问来源于stack exchange,提问作者rosefun
相关产品推荐
相关产品推荐

