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

TensorFlow图像分类教程适配:非X_test/y_test格式下搭建混淆矩阵

基于TensorFlow官方图像加载教程的混淆矩阵实现方案

你参考的官方教程使用tf.data.Dataset批次格式存储验证集(通常命名为val_ds,对应传统方案的测试集分组),不需要手动拆分出X_test/y_test,直接按以下方法实现即可:

实现代码

import numpy as np
from sklearn.metrics import classification_report, confusion_matrix

# 方式1:逐批次遍历提取(逻辑直观,适合小数据集)
y_true = []
y_pred = []

for images, labels in val_ds:
    batch_pred = np.argmax(model.predict(images, verbose=0), axis=1)
    y_true.extend(labels.numpy())
    y_pred.extend(batch_pred)

# 方式2:批量处理(效率更高,适合大数据集,可替换上面的遍历逻辑)
# y_pred = np.argmax(model.predict(val_ds, verbose=0), axis=1)
# y_true = np.concatenate([labels.numpy() for _, labels in val_ds])

# 输出混淆矩阵和分类报告
print('Confusion Matrix')
print(confusion_matrix(y_true, y_pred))
print('Classification Report')
print(classification_report(y_true, y_pred))

注意事项

  • 代码中的val_ds就是教程中用image_dataset_from_directory接口加载得到的验证集对象,不需要额外修改格式
  • sklearn评估接口第一个参数为真实标签,第二个为预测标签,不要写反避免结果错误
  • 即使验证集开启了shuffle、prefetch等性能优化配置,也不会影响标签和样本的对应关系

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 06:54:00