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

使用image_dataset_from_directory加载数据集时如何计算混淆矩阵

用image_dataset_from_directory加载数据集时计算混淆矩阵的方法

当使用tf.keras.preprocessing.image_dataset_from_directory()加载的tf.data.Dataset类型数据集时,确实没有单独的x_train/x_val、y_train/y_val变量,但可以通过以下步骤计算混淆矩阵:

步骤1:提取真实标签

遍历数据集,收集所有样本的真实标签并转换为NumPy数组:

import numpy as np

# 以验证集为例,训练集操作完全相同
true_labels = []
for _, labels in val_ds:
    true_labels.extend(labels.numpy())  # 提取批次中的标签并加入列表
true_labels = np.array(true_labels)

步骤2:获取模型预测标签

使用训练好的模型对数据集进行预测,通过argmax得到每个样本的预测类别:

# 对数据集进行预测,得到每个样本的类别概率
predictions = model.predict(val_ds)
# 取概率最大的索引作为预测标签
pred_labels = np.argmax(predictions, axis=1)

步骤3:计算并可视化混淆矩阵

借助sklearn.metrics中的confusion_matrix计算混淆矩阵,还可以用Seaborn可视化:

from sklearn.metrics import confusion_matrix
import seaborn as sns
import matplotlib.pyplot as plt

# 计算混淆矩阵
cm = confusion_matrix(true_labels, pred_labels)

# 获取数据集自动生成的类别名称
class_names = val_ds.class_names

# 可视化混淆矩阵
plt.figure(figsize=(10, 8))
sns.heatmap(cm, annot=True, fmt="d", cmap="Blues", 
            xticklabels=class_names, yticklabels=class_names)
plt.xlabel("Predicted Label")
plt.ylabel("True Label")
plt.title("Confusion Matrix")
plt.show()

注意事项

  • 确保遍历数据集收集标签的顺序与model.predict()的预测顺序一致,验证集建议关闭shuffle(加载时设置shuffle=False),避免标签和预测结果不匹配。
  • 若数据集过大,内存不足以一次性存储所有标签和预测结果,可以分批处理后再合并,上述方法在大多数场景下都适用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.19 10:37:40