【发布时间】:2018-03-28 09:47:56
【问题描述】:
我的输入管道中有以下错误:
tensorflow.python.framework.errors_impl.InvalidArgumentError: 不能 组件 0 中具有不同形状的批量张量。第一个元素有 形状为 [2,48,48,3],元素 1 的形状为 [27,48,48,3]。
使用此代码
dataset = tf.data.Dataset.from_generator(generator,
(tf.float32, tf.int64, tf.int64, tf.float32, tf.int64, tf.float32))
dataset = dataset.batch(max_buffer_size)
这是完全合乎逻辑的,因为批处理方法试图创建一个 (batch_size, ?, 48, 48, 3) 张量。但是我希望它为这种情况创建一个 [29,48,48,3] 张量。所以连接而不是堆栈。 tf.data 可以吗?
我可以在生成器函数中用 Python 进行连接,但我想知道这是否也可以通过 tf.data 管道实现
【问题讨论】:
-
所以一个实例(数据点)的形状是(48、48、3)?为什么您的生成器首先会产生大量实例?
-
因为它们通过消息总线进入消息群。另一种方法确实是在生成器中产生 (48,48,3) 个实例。但是,我需要一种方法来使批量大小可变,因为我需要再次将实例块一起发送。
-
我明白了。所以团块的大小是可变的,但是你想要连接的团块的数量是固定的?那么我可能有一个解决方案。我会尽快将其发布为答案。
-
抱歉,我想的解决方案没有成功。
-
是的,我想用更大的批次通过网络进行前向传递。但不要打破团块
标签: python tensorflow tensorflow-datasets