TensorFlow查看Directory Iterator对象图像及形状报错解决
报错原因
flow_from_directory返回的目录迭代器调用next()时,输出的是**(批次图像张量, 批次标签张量)**的二元元组:
- 元组第一个元素是当前批次的所有图像,形状为
(batch_size, 图像高度, 图像宽度, 通道数),你当前批次共12张图,所以形状为(12, 224, 224, 3) plt.imshow()仅支持输入形状为(H,W,3)(彩色图)或(H,W)(灰度图)的单张图像数据,直接传入整个批次的4维张量就会触发形状不匹配的TypeError。
另外你使用了tf.keras.applications.vgg16.preprocess_input做预处理,该函数会将图像从RGB转为BGR通道顺序、逐通道减去ImageNet数据集的通道均值,即使正确取出单张图,直接显示也会出现颜色错乱、亮度异常的问题。
修复方案
按两步调整即可:
- 从批次张量中取出单张子图像,不要传入整个批次
- 逆向还原VGG16的预处理操作,将像素值转换为matplotlib可正常显示的0-255范围RGB格式
可直接运行的参考代码:
import matplotlib.pyplot as plt import tensorflow as tf import numpy as np # 数据集加载代码,新增shuffle=False可关闭乱序,方便按原顺序核对图像 imgs = tf.keras.preprocessing.image.ImageDataGenerator( preprocessing_function=tf.keras.applications.vgg16.preprocess_input ).flow_from_directory( r"C:\Users\Abhimanyu\Pictures\Camera Roll", target_size=(224, 224), classes=["Class 1", "Class 2"], shuffle=False ) # 获取一个批次的数据 batch_images, batch_labels = next(imgs) # 取出批次中第1张图像(索引从0开始,查看第N张就把索引改为N-1) single_img = batch_images[0] # 逆向还原VGG16预处理 # 加回之前减去的ImageNet通道均值 single_img[..., 0] += 103.939 single_img[..., 1] += 116.779 single_img[..., 2] += 123.68 # BGR通道顺序转回RGB single_img = single_img[..., ::-1] # 裁剪像素值到合法范围,转换为图像常用的uint8格式 single_img = np.clip(single_img, 0, 255).astype(np.uint8) # 显示图像 plt.imshow(single_img) plt.axis('off') plt.show()
补充说明
- 如果只是想查看原始图像、不需要保留VGG预处理效果,可以暂时去掉
ImageDataGenerator中的preprocessing_function参数,取出单张图后直接用.astype(np.uint8)转换格式即可显示,无需做逆向还原。 - 如果要查看整个批次的所有图像,可以循环遍历
batch_images的第一维,逐张处理后通过matplotlib子图批量绘制。
内容的提问来源于stack exchange,提问作者Abhimanyu Sharma
相关产品推荐
相关产品推荐

