为什么tf.data.Dataset.map包裹tf.numpy_function会改变图像掩码形状
tf.numpy_function调用新增冗余维度的原因及修复方案
核心原因
tf.numpy_function本身不会修改Python侧函数的返回值形状,问题由静态形状丢失、参数配置错误两类原因导致:
- 静态形状丢失:
tf.numpy_function将Python函数包装为TF算子时,默认不会继承Python侧返回值的静态形状信息,仅能保证运行时动态形状和原输出一致。而tf.data.Dataset的batch、对齐逻辑依赖静态形状做推断,当掩码、标签的静态形状未知时,TF会错误隐式插入大小为1的维度匹配已知静态形状的图像张量,最终batch后就会出现多余的1维度。 - 参数配置错误:如果调用
tf.numpy_function时,Tout参数配置和实际返回值结构不匹配,或是将返回的单个张量额外套了一层列表/元组包装,TF会将返回值识别为长度为1的序列,转换为张量时自动在最前面插入大小为1的维度。
你之前用from_tensor_slices构造数据集时形状正常,是因为构造阶段所有数据的静态形状已经被TF识别并固定,后续batch阶段不会触发错误的维度插入逻辑。
修复方法
在tf.numpy_function调用后,手动为输出张量设置静态形状即可解决问题,参考代码如下:
def dataset_map_func(img_path): # 读取图像,本身是TF原生张量,自带正确静态形状 img = readTFImage(img_path) # 调用numpy_function读取掩码和标签 mask, label = tf.numpy_function( func=lambda path: (getMask(path), getClassificationLabel(path)), inp=[img_path], # Tout需要和返回值的数量、类型一一对应,不要多写也不要少写 Tout=[tf.uint8, tf.float32] ) # 手动设置静态形状,和getMask、getClassificationLabel的实际返回形状完全匹配 mask.set_shape(OUTPUT_SHAPE) label.set_shape((class_number,)) return img, mask, label
内容的提问来源于stack exchange,提问作者Phan Nguyen
相关产品推荐
相关产品推荐

