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

如何展示Keras测试数据生成器输出的分类错误图片并解决维度报错问题

错误原因

你在预测前为了匹配模型输入要求,给单张图片添加了batch维度(np.expand_dims(img, axis=0)),最终存入images列表的图片张量形状为(1, 512, 512, 1)(4维),但array_to_img要求单张图片输入必须是(高度, 宽度, 通道数)的3维结构,因此触发维度不匹配报错。
另外你写的plt.plot()是绘制折线图的方法,展示图片需要用plt.imshow()。

修复方案

方案1:存储错误样本时直接去掉多余维度(推荐)

修改收集错误样本的循环逻辑,存储前去掉batch维度,同时只存储对应单张的标签而非整个batch的标签:

images = []
true_labels = []
pred_labels = []
for i in range(100):
    img_batch, label_batch = next(val_generator)
    img = img_to_array(img_batch[i])
    # 加batch维度用于预测
    img_input = np.expand_dims(img, axis=0)
    pred = model.predict_classes(img_input)[0]
    if pred != label_batch[i]:
        # 存储原图(3维),不要带batch维度的输入张量
        images.append(img)
        true_labels.append(label_batch[i])
        pred_labels.append(pred)

方案2:展示时临时处理维度

如果不想修改之前的存储逻辑,展示时直接用np.squeeze()去掉所有长度为1的多余维度即可:

# 单张图片展示示例
import matplotlib.pyplot as plt

# 去掉多余的batch维度
img0 = array_to_img(np.squeeze(images[0]))
# 展示图片,单通道灰度图需要指定cmap='gray'
plt.imshow(img0, cmap='gray')
# 可以添加标签对比信息
plt.title(f"真实标签: {labels[0][i]}, 预测标签: {model.predict_classes(images[0])[0]}")
plt.axis('off')
plt.show()

批量展示错误样本示例

如果需要批量展示所有分类错误的图片,可以用如下代码:

plt.figure(figsize=(12, 12))
# 最多展示9张,避免布局拥挤
show_num = min(9, len(images))
for idx in range(show_num):
    plt.subplot(3, 3, idx+1)
    img = array_to_img(np.squeeze(images[idx]))
    plt.imshow(img, cmap='gray')
    pred = model.predict_classes(images[idx])[0]
    # 如果用了方案1的存储逻辑,直接取true_labels和pred_labels即可
    # true = true_labels[idx]
    # pred = pred_labels[idx]
    plt.title(f"真实标签: {labels[idx][i]}\n预测标签: {pred}")
    plt.axis('off')
plt.tight_layout()
plt.show()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 16:39:04