You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

构建可变数量输入图像的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.10 09:01:35