使用Keras ImageDataGenerator进行数据增强时matplotlib绘图报错:TypeError: Invalid shape (1, 256, 256, 3) for image data
修复ImageDataGenerator图像绘制的TypeError问题
这个报错的核心原因很明确:flow_from_directory返回的batch数据形状是(batch_size, height, width, channels),当你设置batch_size=1时,batch[0]的形状就变成了(1, 256, 256, 3)——多了一个长度为1的batch维度,但matplotlib的plt.imshow()只接受(height, width, channels)或者(height, width)的数组,直接传入就会触发形状不匹配的错误。
另外还有个隐藏的小问题:你设置了rescale=1./255,会把图像像素值缩放到0-1的浮点数范围,但直接用.astype('uint8')会把所有值截断为0(0-1的浮点数转整数就是0),最后画出来的图像会全黑,这个也得一起修复。
修复后的完整代码
import numpy as np import matplotlib.pyplot as plt from keras.preprocessing.image import ImageDataGenerator datagen = ImageDataGenerator( rescale=1./255, zoom_range=0.1, rotation_range=25, width_shift_range=0.1, height_shift_range=0.1, shear_range=0.1, horizontal_flip=True ) ite = datagen.flow_from_directory("Car Images", batch_size=1) for i in range(9): plt.subplot(330 + 1 + i) batch = ite.next() # 先把0-1的浮点数转回0-255的标准像素范围,再转成uint8类型 image = (batch[0] * 255).astype('uint8') # 移除多余的batch维度,把(1,256,256,3)转换成(256,256,3) image = np.squeeze(image) plt.imshow(image) plt.show()
关键修复点说明
- 移除多余维度:
np.squeeze()会自动去掉所有长度为1的维度,完美解决形状不匹配的问题。你也可以用索引image = batch[0][0]达到同样效果,但squeeze()更灵活,之后调整batch_size也不用修改这行代码。 - 恢复像素值范围:
(batch[0] * 255)把缩放后的0-1浮点数还原到图像标准的0-255像素范围,再转成uint8才能让matplotlib正确显示色彩。
这样修改后,你就能正常看到生成的增强图像了。
内容的提问来源于stack exchange,提问作者AJ2000
相关产品推荐
相关产品推荐

