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

