TensorFlow输入管道报错:无法批量处理形状不同的张量
解决TensorFlow Dataset Batch时的形状不匹配问题
这个报错完全在情理之中——你的generator返回的样本里,第一个元素的形状不一致(一会儿是[2,48,48,3],一会儿是[27,48,48,3]),而batch()方法要求同一个batch里的所有样本必须有完全相同的形状,自然会抛出这个错误。
下面分两种情况给你解决方案:
情况1:你需要保留这个可变维度(比如是序列长度、多帧数据等)
这种情况下不要用batch(),改用padded_batch(),它会自动把不同长度的维度填充到统一大小,让batch可以正常创建。
修改你的代码如下:
dataset = tf.data.Dataset.from_generator(generator, (tf.float32, tf.int64, tf.int64, tf.float32, tf.int64, tf.float32)) # 使用padded_batch替代batch,并指定需要填充的形状 dataset = dataset.padded_batch( max_buffer_size, # padded_shapes对应generator返回的每个元素的形状,用None标记可变维度 padded_shapes=([None, 48, 48, 3], [], [], [], [], []) )
这里的padded_shapes参数需要和generator返回的每个输出元素一一对应:
- 第一个元素的第一个维度是可变的(2、27这类),所以用
None表示这个维度会被自动填充到当前batch的最大长度; - 后面的元素都是标量(形状为空),所以用
[]即可。
如果需要指定填充值(默认是0),可以加上padding_values参数,注意要和每个元素的数据类型匹配:
dataset = dataset.padded_batch( max_buffer_size, padded_shapes=([None, 48, 48, 3], [], [], [], [], []), padding_values=(tf.constant(0.0, dtype=tf.float32), 0, 0, 0.0, 0, 0.0) )
情况2:这个可变维度是意外产生的
如果你的generator本来应该返回固定形状的样本,那问题出在generator的逻辑里——你需要检查生成第一个元素的代码,确保每次输出的第一个维度都是相同的数值(比如统一为2或者27),修正后再用batch()就不会报错了。
内容的提问来源于stack exchange,提问作者Derk
相关产品推荐
相关产品推荐

