在TensorFlow中如何获取图像分类任务的混淆矩阵及评估指标
测试代码逻辑修正
- 你现在的代码存在冗余加载问题:已经通过
tf.keras.utils.image_dataset_from_directory生成了标准化的测试集test_ds,不需要再手动遍历图片文件做单张预测,单张预测的效率远低于批量推理。 - 存在类别错位风险:
image_dataset_from_directory默认会按文件夹名称的字典序生成类别编号,需要确保测试集的类别顺序和训练集的train_class_names完全一致,否则会出现标签对应错误的问题。
优化后推理与指标计算代码
首先加载测试集时直接指定训练集的类别顺序,彻底规避标签错位问题:
train_class_names = train_ds.class_names test_data_dir = pathlib.Path('test_data/') test_ds = tf.keras.utils.image_dataset_from_directory( test_data_dir, image_size=(img_height, img_width), batch_size=batch_size, class_names=train_class_names # 固定类别顺序和训练集完全一致 )
批量推理收集所有真实标签和预测标签,效率远高于单张循环:
import numpy as np from sklearn.metrics import confusion_matrix, classification_report, accuracy_score import seaborn as sns import matplotlib.pyplot as plt true_labels = [] pred_labels = [] # 批量遍历测试集推理 for images, labels in test_ds: predictions = model.predict(images, verbose=0) pred_batch = np.argmax(predictions, axis=1) true_labels.extend(labels.numpy()) pred_labels.extend(pred_batch)
生成混淆矩阵并可视化:
# 生成数值型混淆矩阵 cm = confusion_matrix(true_labels, pred_labels) # 可视化混淆矩阵 plt.figure(figsize=(6,6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=train_class_names, yticklabels=train_class_names) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.title('测试集混淆矩阵') plt.show()
计算每个类别的精确率、召回率以及整体准确率:
# 直接输出各类别的精确率、召回率、F1值 print(classification_report(true_labels, pred_labels, target_names=train_class_names)) # 计算测试集整体准确率 overall_acc = accuracy_score(true_labels, pred_labels) print(f"测试集整体准确率:{overall_acc:.4f}")
如果你需要保留原来的单张预测逻辑做逐张结果校验,只需要在你原有循环中,把预测类别的编号和真实类别的编号分别存入两个列表,后续同样可以用上面的指标计算代码生成混淆矩阵和各项指标,只是运行速度会更慢。
内容的提问来源于stack exchange,提问作者lilyming
相关产品推荐
相关产品推荐

