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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 03:52:18