如何在TensorFlow/Keras中调试CNN训练与评估时的误分类图像?
定位误分类图像的高效方法
因为你用tf.keras.utils.image_dataset_from_directory加载数据集,最优雅高效的方式是利用该API的return_filepaths参数,直接在数据集管道中完成预测与对比,精准定位误分类样本:
步骤1:重新加载数据集时开启文件路径返回
加载数据集时添加return_filepaths=True,这样每个batch会返回(图像数组, 真实标签, 文件路径)三元组,直接关联样本与原始文件:
import tensorflow as tf # 假设你的数据集目录结构是每个类别一个子文件夹 dataset = tf.keras.utils.image_dataset_from_directory( "你的数据集根目录", image_size=(你的图像尺寸), batch_size=32, # 保持你训练时的batch size即可 return_filepaths=True # 关键参数,开启文件路径返回 )
步骤2:批量遍历数据集,定位误分类样本
遍历每个batch,用训练好的模型做预测,对比真实标签与预测标签,记录误分类的文件路径、真实标签和预测标签:
misclassified = [] class_names = dataset.class_names # 获取类别名称列表 for images, true_labels, file_paths in dataset: # 对当前batch做预测 pred_probs = model.predict(images, verbose=0) pred_labels = tf.argmax(pred_probs, axis=1) # 找出当前batch中误分类的样本索引 mis_idx = tf.where(true_labels != pred_labels).numpy().flatten() # 记录每个误分类样本的信息 for idx in mis_idx: misclassified.append({ "file_path": file_paths[idx].numpy().decode("utf-8"), "true_class": class_names[true_labels[idx].numpy()], "pred_class": class_names[pred_labels[idx].numpy()] })
步骤3:可视化误分类样本
拿到误分类文件路径后,直接读取图像可视化,分析特征:
import matplotlib.pyplot as plt import cv2 # 展示前10个误分类样本 for i, item in enumerate(misclassified[:10]): img = cv2.imread(item["file_path"]) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) plt.subplot(2, 5, i+1) plt.imshow(img) plt.title(f"True: {item['true_class']}\nPred: {item['pred_class']}") plt.axis("off") plt.tight_layout() plt.show()
为什么这个方法更高效?
- 复用已有的数据集管道,无需重新加载或预处理数据,避免重复开销
- 直接通过文件路径关联原始图像,无需额外存储样本索引与路径的映射关系
- 批量处理与预测,符合TensorFlow的高效计算逻辑
内容的提问来源于stack exchange,提问作者Rodoleia
相关产品推荐
相关产品推荐

