You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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图像),传入整数会导致形状为空,触发错误。

解决方法

需要获取对应错误分类图像的文件路径,再读取图像像素数据:

  1. 将test_images_filtered改为获取图像路径:test_images_filtered = [validation.filepaths[i] for i in indices_list]
  2. 在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是空列表(即没有错误分类的图像),也会触发该错误。

解决方法

  1. 调整随机数生成逻辑,避免越界:用random.randrange(len(test_images_filtered))(上限不包含)或者random.randint(0, len(test_images_filtered)-1)
  2. 添加空列表判断,避免无图像可绘制时执行后续代码

修正后的完整代码

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)

额外修正点

  1. 移除了原代码中错误的test_labels_vector = np.argmax(validation.classes),validation.classes本身就是每个样本的类别标签数组,无需额外处理
  2. 简化了函数参数,去掉了冗余的test_images参数
  3. 添加坐标轴隐藏逻辑,提升可视化效果

内容的提问来源于stack exchange,提问作者user979974

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.19 22:10:31