TensorFlow Dataset API加载图片报错'无法迭代标量张量'问题求助
错误原因
- 最后绘图循环遍历的是未经过处理的原始
dataset,该数据集的每个元素是单个图片文件路径的标量张量,既不是处理好的(图像,标签)对,也没有批量维度支持i索引操作,是触发cannot iterate over a scalar tensor报错的核心原因 - 图像解码步骤缺失:
tf.io.read_file读取的是编码后的二进制数据,未解码直接调用resize会产生格式错误 - 变量和语法错误:
class_names未定义直接使用,etiqueta[i].numpy缺少调用括号,应为etiqueta.numpy() - 循环逻辑错误:外层遍历数据集单样本,内层循环9次尝试从单样本中取多个元素,本身不符合数据结构逻辑
修复方案
import tensorflow as tf import numpy as np import matplotlib.pyplot as plt import os dataset = tf.data.Dataset.list_files('ai minecraft ores dataset/*/*', shuffle=False) def etiquetar(ubicacion_archivo): return tf.strings.split(ubicacion_archivo, os.path.sep)[-2] def procesar_imagen(ubicacion_archivo): etiqueta = etiquetar(ubicacion_archivo) imagen = tf.io.read_file(ubicacion_archivo) # 补充图像解码步骤,png格式可替换为decode_png imagen = tf.image.decode_jpeg(imagen, channels=3) imagen = tf.image.resize(imagen, [128,128]) # 提前转为uint8类型,适配绘图要求 imagen = tf.cast(imagen, tf.uint8) return imagen, etiqueta # 处理后的有效数据集 dataset1 = dataset.map(procesar_imagen) # 绘图逻辑修正 plt.figure(figsize=(10, 10)) # 直接取9个样本遍历绘制,无需双重循环 for i, (imagen, etiqueta) in enumerate(dataset1.take(9)): ax = plt.subplot(3, 3, i + 1) plt.imshow(imagen.numpy()) # 标签解码为字符串直接作为标题 plt.title(etiqueta.numpy().decode()) plt.axis("off") plt.show()
修改后直接运行即可正常输出3*3的样本图像网格,无需额外调整。
内容的提问来源于stack exchange,提问作者Martin Arriola
相关产品推荐
相关产品推荐

