TensorFlow写入TFRecord后解析重塑浮点列表遇形状不兼容错误求助
解决TFRecord解析时的形状不兼容问题
这个错误的核心原因是你在读取TFRecord时,对Feature的定义和写入时的格式不匹配。让我拆解一下问题:
问题分析
在写入TFRecord时,你把每个3D数据(比如data_rgb、data_decalib)扁平化后作为一个浮点数组存入FloatList。但在读取时,你用了tf.FixedLenFeature([], tf.float32)——这个定义表示每个Feature是一个单个float标量,而不是数组。这就导致解析后得到的张量形状是()(标量),但你后续试图把它重塑成(8,)或者更高维的形状,自然会触发形状不兼容的错误。
解决方案
你需要在读取时,给每个FixedLenFeature指定对应的数组长度,也就是你写入时扁平化后的元素个数。具体步骤如下:
计算每个Feature的扁平化长度:
data_rgb的长度是IMAGE_HEIGHT * IMAGE_WIDTH * 3(对应INPUT_RGB_SHAPE的元素总数)data_depth的长度是IMAGE_HEIGHT * IMAGE_WIDTH * 1(对应INPUT_DEPTH_SHAPE的元素总数)data_decalib的长度是8(对应LABEL_CALIB_SHAPE的元素总数)
修改
input_parser中的Feature定义,把每个FixedLenFeature的shape参数从[]改成对应的长度值。
修正后的代码
下面是调整后的input_parser函数和相关部分:
def input_parser(example_proto): # 计算每个feature的扁平化长度 rgb_flat_len = np.prod(INPUT_RGB_SHAPE) depth_flat_len = np.prod(INPUT_DEPTH_SHAPE) calib_flat_len = np.prod(LABEL_CALIB_SHAPE) # 修正Feature定义:指定每个feature的数组长度 features = { 'data_rgb': tf.FixedLenFeature([rgb_flat_len], tf.float32), 'data_depth': tf.FixedLenFeature([depth_flat_len], tf.float32), 'data_decalib': tf.FixedLenFeature([calib_flat_len], tf.float32) } parsed_features = tf.parse_single_example(example_proto, features) # 直接重塑为目标形状,不需要额外set_shape(因为FixedLenFeature已经指定了长度) img_rgb = tf.reshape(parsed_features['data_rgb'], INPUT_RGB_SHAPE) img_depth = tf.reshape(parsed_features['data_depth'], INPUT_DEPTH_SHAPE) # data_decalib本身就是一维数组,不需要reshape,直接返回即可 data_decalib = parsed_features['data_decalib'] return img_rgb, img_depth, data_decalib
额外说明
- 如果你还是想保留
set_shape,其实也可以,但在指定了FixedLenFeature的长度后,张量已经自带正确的形状了,所以set_shape不是必须的。 - 确保写入时的扁平化操作和读取时的长度计算完全一致,比如如果写入时用了
value.flatten(),读取时的长度必须是原形状的元素总数,否则依然会出现形状错误。
内容的提问来源于stack exchange,提问作者Mark Rofail
相关产品推荐
相关产品推荐

