使用tf.concat()拼接张量未得到预期结果,是什么原因?
问题根因
你输出的l_image、l_augim、u_image都被额外的圆括号包裹,说明从dataset中迭代取出的每个元素都是长度为1的元组,不是直接的张量对象。你将三个元组直接传入tf.concat拼接时,TensorFlow会自动将元组转为张量,相当于每个输入多了一层长度为1的维度,三个张量沿axis=0拼接后就会得到(3, 4, 32, 32, 3)的不符合预期的shape。
解决方案
只需要在拼接前,从每个元组中取出第一个元素拿到实际的图像张量即可。
修正后代码
dataset = tf.data.Dataset.zip((l_image, l_augim, u_image)).batch(4) for x, (l_image, l_augim, u_image) in enumerate(dataset): # 提取元组内的实际张量 l_image = l_image[0] l_augim = l_augim[0] u_image = u_image[0] concat_tensor = tf.concat([l_image, l_augim, u_image], axis = 0) print(concat_tensor) print(l_image) print(l_augim) print(u_image) break
运行后即可得到期望的shape=(12, 32, 32, 3)的拼接张量。
内容的提问来源于stack exchange,提问作者IdeaKing
相关产品推荐
相关产品推荐

