使用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
相关产品推荐
相关产品推荐

