关于tf.Tensor.set_shape()的澄清

jks*_*hin 29 tensorflow

我的图像是478 x 717 x 3 = 1028178像素,等级为1.我通过调用tf.shape和tf.rank验证了它.

当我调用image.set_shape([478,717,3])时,它会抛出以下错误.

"Shapes %s and %s must have the same rank" % (self, other)) 
ValueError: Shapes (?,) and (478, 717, 3) must have the same rank
Run Code Online (Sandbox Code Playgroud)

我通过首次测试再次测试到1028178,但错误仍然存​​在.

ValueError: Shapes (1028178,) and (478, 717, 3) must have the same rank
Run Code Online (Sandbox Code Playgroud)

嗯,这确实有意义,因为一个是等级1而另一个是等级3.但是,为什么有必要抛出一个错误,因为像素的总数仍然匹配.

我当然可以使用tf.reshape并且它有效,但我认为这不是最佳的.

正如TensorFlow常见问题解答中所述

x.set_shape()和x = tf.reshape(x)之间有什么区别?

tf.Tensor.set_shape()方法更新Tensor对象的静态形状,并且通常用于在无法直接推断时提供其他形状信息.它不会改变张量的动态形状.

tf.reshape()操作创建一个具有不同动态形状的新张量.

创建新的张量涉及内存分配,并且当涉及更多培训示例时可能会更昂贵.这是设计,还是我在这里遗漏了什么?

mrr*_*rry 61

据我所知(并且我编写了该代码),没有错误Tensor.set_shape().我认为这种误解源于该方法令人困惑的名称.

要详细说明您引用的FAQ条目,Tensor.set_shape()是一个纯Python函数,它可以改善给定tf.Tensor对象的形状信息.通过"改进",我的意思是"更具体".

因此,当你有一个具有形状的Tensor物体t时(?,),这是一个未知长度的一维张量.你可以打电话t.set_shape((1028178,)),然后在你打电话时t会有形状.这不会影响底层存储,或者实际上后端上的任何内容:它仅仅意味着后续使用的形状推断可以依赖于它是长度为1028178的向量的断言.(1028178,)t.get_shape()t

如果t有形状(?,),调用t.set_shape((478, 717, 3))将失败,因为TensorFlow已经知道它t是一个向量,所以它不能有形状(478, 717, 3).如果你想从内容中制作一个具有该形状的新Tensor t,你可以使用reshaped_t = tf.reshape(t, (478, 717, 3)).这tf.Tensor在Python中创建了一个新对象; 实际的实现tf.reshape()使用张量缓冲区的浅拷贝来做到这一点,所以在实践中它很便宜.

有一个比喻Tensor.set_shape()就像是像Java这样的面向对象语言的运行时强制转换.例如,如果你有一个指向a的指针Object但实际上知道它是a String,你可以进行转换(String) obj以传递obj给需要String参数的方法.但是,如果你有一个String s并尝试将它强制转换为a java.util.Vector,编译器会给你一个错误,因为这两种类型是无关的.