如何展示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
相关产品推荐
相关产品推荐

