TensorFlow中decode_jpeg解码后如何获取图像真实形状?
纯图模式下获取tf.data中解码后图像的真实形状
在图模式的tf.data.Dataset流水线里,直接访问张量的.shape属性得到的是静态形状——而tf.image.decode_jpeg在图构建阶段无法确定图像的具体尺寸,所以静态形状会显示(None, None, 3)。想要获取真实的动态形状,你可以用tf.shape()函数,它会在图运行时返回张量的实际尺寸,完全符合纯图模式的要求,性能也比tf.py_function更好。
具体实现代码
把你的映射函数改成这样:
def load_and_get_shape(file_path): img = tf.io.read_file(file_path) img = tf.image.decode_jpeg(img, channels=3) # 用tf.shape()获取动态的高和宽 img_height = tf.shape(img)[0] img_width = tf.shape(img)[1] return img, img_height, img_width # 构建数据集 dataset = tf.data.Dataset.from_tensor_slices(file_paths) dataset = dataset.map(load_and_get_shape)
为什么这个方法有效?
.shape返回的是张量的静态形状信息,这是在图构建时就能确定的固定值,但解码JPEG时图像尺寸是动态的,所以静态形状只能是None。tf.shape()是一个图操作,它会在图执行阶段(也就是实际读取并解码图像时)计算出张量的真实尺寸,返回的是一个标量张量,完全兼容图模式的流水线,不会有tf.py_function带来的性能开销。
额外提示
如果你的数据集里所有图像尺寸都是固定的,也可以在解码后手动设置静态形状来优化性能:
img = tf.image.decode_jpeg(img, channels=3) # 假设所有图像都是512x512 img.set_shape((512, 512, 3)) # 此时img.shape[0]和img.shape[1]就能返回512了
但如果图像尺寸不固定,还是优先用tf.shape()的方案。
内容的提问来源于stack exchange,提问作者Diego Palacios
相关产品推荐
相关产品推荐

