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

在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 21:24:04