Keras绘制分类错误图像时遇TypeError与IndexError问题求助
解决分类错误图像绘制的两个问题
第一个错误:TypeError: Invalid shape () for image data
原因
你最初的代码里,test_images_filtered = [validation.classes[i] for i in indices_list]取的是类别标签(整数),而imshow()需要传入的是图像像素数组(形状为(X,Y,3)的RGB图像),传入整数会导致形状为空,触发错误。
解决方法
需要获取对应错误分类图像的文件路径,再读取图像像素数据:
- 将
test_images_filtered改为获取图像路径:test_images_filtered = [validation.filepaths[i] for i in indices_list] - 在
imshow()中用plt.imread()读取路径对应的图像:ax[i,j].imshow(plt.imread(test_images_filtered[l]))
第二个错误:IndexError: list index out of range
原因
random.randint(0, len(test_images_filtered))的取值范围是包含两端的,而列表的索引最大是len(test_images_filtered)-1,当随机数取到len(test_images_filtered)时,就会超出列表索引范围,触发错误。另外如果test_images_filtered是空列表(即没有错误分类的图像),也会触发该错误。
解决方法
- 调整随机数生成逻辑,避免越界:用
random.randrange(len(test_images_filtered))(上限不包含)或者random.randint(0, len(test_images_filtered)-1) - 添加空列表判断,避免无图像可绘制时执行后续代码
修正后的完整代码
import random import matplotlib.pyplot as plt import numpy as np predictions = model.predict(validation) # 预测概率向量 pred_labels = np.argmax(predictions, axis = 1) # 取概率最高的类别标签 def classification_evaluation(classification, predicted_labels, test_labels): if classification == "correct": indices_list = np.where(predicted_labels == test_labels)[0] else: indices_list = np.where(predicted_labels != test_labels)[0] # 获取对应图像的文件路径 test_images_filtered = [validation.filepaths[i] for i in indices_list] images_labels_original = [test_labels[i] for i in indices_list] images_labels_predicted = [predicted_labels[i] for i in indices_list] total = len(test_labels) count = len(test_images_filtered) print(f"{count} images were classified {classification.upper()} out of a total of {total} in the Validation dataset") unique, counts = np.unique(images_labels_original, return_counts=True) for cls, cnt in zip(unique, counts): print(f"For category {cls}, the number of {classification.upper()} classified images were: {cnt}") # 无对应图像时直接返回 if count == 0: print(f"No {classification} classified images to display.") return # 绘制样本图像 print("\n\n") fig, ax = plt.subplots(5, 2) fig.suptitle(f"Sample of {classification.upper()} Classified Images", fontsize=20) fig.set_size_inches(15, 15) for i in range(5): for j in range(2): # 生成合法索引 l = random.randrange(len(test_images_filtered)) img = plt.imread(test_images_filtered[l]) ax[i,j].imshow(img) ax[i,j].set_title(f"Predicted: {images_labels_predicted[l]}\nActual: {images_labels_original[l]}") ax[i,j].axis('off') # 隐藏坐标轴优化显示 plt.tight_layout() plt.show() # 调用函数,传入验证集真实标签 classification_evaluation("incorrect", pred_labels, validation.classes)
额外修正点
- 移除了原代码中错误的
test_labels_vector = np.argmax(validation.classes),validation.classes本身就是每个样本的类别标签数组,无需额外处理 - 简化了函数参数,去掉了冗余的
test_images参数 - 添加坐标轴隐藏逻辑,提升可视化效果
内容的提问来源于stack exchange,提问作者user979974
相关产品推荐
相关产品推荐

