TensorFlow中tf.Print意外修复形状错误问题求助
移除tf.Print后LSTM动态RNN报错的原因及解决办法
嘿,这个问题我之前也碰到过类似的,太有意思了——加个tf.Print就正常运行,删掉直接炸锅,简直像薛定谔的代码😂。咱们来拆解一下:
问题根源:形状推断的“延迟性” + dynamic_rnn的输入格式要求
先看你代码里的关键操作:
new_x = tf.unstack(x, 90, 1) output, state = tf.nn.dynamic_rnn(cell, new_x, dtype=tf.float32)
你把输入x(形状[batch, 90, 4])unstack成了90个形状为[?, 4]的张量组成的列表。而tf.nn.dynamic_rnn对输入格式有两种预期:
- 默认
time_major=False:输入是[batch_size, max_time, features]的三维张量 - 如果传的是张量列表,它默认认为是
time_major=True的格式,也就是列表里每个元素是[batch_size, features],整体对应[max_time, batch_size, features]
但这里的问题是,没有tf.Print时,TensorFlow的形状推断没完全跟上:unstack后的张量形状里的?(batch维度)没被明确,导致dynamic_rnn内部做transpose操作时,误以为张量维度不对,抛出了Dimension must be 2 but is 3的错误。
而加了tf.Print(new_x, [tf.shape(new_x)])之后,这个操作会强制触发TensorFlow去计算new_x的形状,相当于提前告诉框架“这90个张量每个都是[batch,4]”,dynamic_rnn就能正确识别输入格式,自然不会报错。
两种靠谱的解决办法
办法1:直接用三维张量输入(推荐)
dynamic_rnn本身就支持直接处理[batch, time, features]的三维张量,完全不需要unstack,代码更简洁,也不会有形状推断的问题:
# 删掉unstack和对应的tf.Print # new_x = tf.unstack(x, 90, 1) # new_x = tf.Print(new_x, [tf.shape(new_x)], message='newx is: ') # 直接把x传给dynamic_rnn output, state = tf.nn.dynamic_rnn(cell, x, dtype=tf.float32) # 此时output形状是[batch, 90, 200],取最后一个时间步的输出 logits = tf.matmul(output[:, -1, :], w_out) + b_out
办法2:明确设置time_major=True
如果你一定要保留unstack的操作,记得给dynamic_rnn加上time_major=True参数,明确告诉框架输入是时间步在前的格式:
new_x = tf.unstack(x, 90, 1) # 加上time_major=True output, state = tf.nn.dynamic_rnn(cell, new_x, dtype=tf.float32, time_major=True) # 此时output形状是[90, batch, 200],取最后一个时间步的输出就没问题 logits = tf.matmul(output[-1], w_out) + b_out
这样不管加不加tf.Print,代码都能正常运行啦~
内容的提问来源于stack exchange,提问作者Lzkatz
相关产品推荐
相关产品推荐

