如何在MRI图像数据集上正确使用from_tensor_slices?
问题原因
tf.data.Dataset 流中的元素默认是 tf.Tensor 类型,你传入 from_tensor_slices 的路径会被封装为字符串Tensor,而 nib.load 是纯Python实现的工具函数,仅支持原生Python字符串、字节、路径对象作为输入,无法识别Tensor类型,因此触发类型错误。
修复方案
提供两种可用的实现方式,按需选择即可:
方案1:使用tf.py_function包裹加载函数
这是改动最小的方案,只需调整map阶段的调用逻辑,手动将Tensor转为Python字符串再传入加载函数:
def load_one_sample(image_path, label_path): # 新增:将Tensor转成utf-8编码的字符串路径 image_path = image_path.numpy().decode('utf-8') label_path = label_path.numpy().decode('utf-8') image = nib.load(image_path).get_fdata() image = tf.convert_to_tensor(image, dtype = 'float32') label = nib.load(label_path).get_fdata() label = tf.convert_to_tensor(label, dtype = 'uint8') return image, label # 用tf.py_function包裹函数,指定输出数据类型 all_data = dataset.map( lambda img_p, lab_p: tf.py_function( func=load_one_sample, inp=[img_p, lab_p], Tout=[tf.float32, tf.uint8] ) ) # 可选:手动声明张量形状,避免后续训练时形状不明确报错 # 形状需和你实际加载的MRI维度匹配,脑肿瘤任务数据维度一般为(240,240,155,4),标签为(240,240,155) all_data = all_data.map( lambda x, y: ( tf.ensure_shape(x, (240, 240, 155, 4)), tf.ensure_shape(y, (240, 240, 155)) ) )
方案2:从Python生成器构建数据集
如果不想和Tensor的类型转换打交道,可以直接用生成器遍历本地路径加载数据,逻辑更直观:
# 定义生成器,逐样本加载数据 def data_generator(): for img_path, lab_path in zip(image_paths, label_paths): image = nib.load(img_path).get_fdata().astype('float32') label = nib.load(lab_path).get_fdata().astype('uint8') yield image, label # 直接从生成器构建数据集,指定输出的类型和形状即可 all_data = tf.data.Dataset.from_generator( generator=data_generator, output_signature=( tf.TensorSpec(shape=(240, 240, 155, 4), dtype=tf.float32), tf.TensorSpec(shape=(240, 240, 155), dtype=tf.uint8) ) )
注意事项
- 上述代码中的张量形状需要根据你实际使用的数据集维度调整,避免形状不匹配报错
- 两种方式都支持后续链式调用
shuffle()、batch()、prefetch()等常用数据集操作
内容的提问来源于stack exchange,提问作者user15515518
相关产品推荐
相关产品推荐

