构建可变数量输入图像的TensorFlow回归数据集报错求助
解决可变输入图像集合的TensorFlow数据集构建问题
错误原因
tf.keras.utils.load_img仅支持路径字符串或BytesIO对象,但tf.map_fn传递给read_image_tf的是Tensor类型的路径,无法被load_img直接识别,触发类型错误。
修改方案
核心是将Tensor路径转换为Python原生字符串,并用tf.py_function包装Python逻辑(规避TensorFlow图模式的限制),同时保留可变长度的图像集合结构。
完整修改代码
import tensorflow as tf def read_image_tf(path: tf.Tensor) -> tf.Tensor: # 将Tensor中的字节路径解码为Python字符串 path_str = path.numpy().decode('utf-8') image = tf.keras.utils.load_img(path_str) return tf.keras.utils.img_to_array(image) def read_image_list(x, y): # 用tf.py_function执行Python图像读取逻辑 images = tf.py_function( func=lambda paths: tf.stack([read_image_tf(p) for p in paths]), inp=[x], Tout=tf.float32 ) # 转回RaggedTensor以保留可变长度结构(按需选择) images = tf.RaggedTensor.from_tensor(images, ragged_rank=1) return images, y paths_list = [['image_1', 'image_2', 'image_3'], ['image_6'], ['image_4', 'image_5', 'image_8', 'image_19']] x = tf.ragged.constant(paths_list) y = tf.constant([1,2,3]) dataset = tf.data.Dataset.from_tensor_slices((x, y)) dataset = dataset.map(read_image_list) # 验证数据集输出 for images, label in dataset: print(f"标签: {label.numpy()}, 图像集合形状: {images.shape}")
关键修改说明
- 路径转换:通过
path.numpy().decode('utf-8')将Tensor中的字节型路径转为Python可识别的字符串格式 - Python逻辑包装:
tf.py_function允许在TensorFlow数据集的map操作中执行非图模式的Python代码,解决numpy操作的兼容性问题 - 保留可变结构:用
tf.RaggedTensor.from_tensor将堆叠后的张量转回可变长度结构,适配不同数量的输入图像需求
内容的提问来源于stack exchange,提问作者Alberto
相关产品推荐
相关产品推荐

