如何基于Keras prefetch数据集为多标签分类生成混淆矩阵
报错根因
报错是因为传入sklearn.metrics.confusion_matrix的两个标签格式不匹配:
predicted_categories经过tf.argmax处理后是一维类别索引数组,属于普通多分类标签格式,每个值代表对应样本预测的类别编号(0~7)true_categories是直接从prefetch测试集拼接得到的one-hot编码矩阵,属于多标签指示器格式,形状为(样本总数, 8),每个位置用0/1标记样本是否属于对应类别confusion_matrix要求输入的真实标签和预测标签必须是同一种编码格式,因此触发类型不匹配错误。
修复方法
如果你的任务是单标签8分类(每个样本只属于一个类别,你当前代码监控categorical_accuracy也符合这个场景,只是标签存为了one-hot格式),只需要把真实标签也通过tf.argmax转为类别索引,和预测标签格式对齐即可。
修正后的核心代码如下:
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt import tensorflow as tf # 获取模型预测结果 predictions = shallow_mlp_model.predict(test_dataset) # 预测结果转为类别索引 predicted_categories = tf.argmax(predictions, axis=1) # 拼接测试集真实标签,同时将one-hot标签转为类别索引 true_categories = tf.argmax( tf.concat([y for x, y in test_dataset], axis=0), axis=1 ) # 生成混淆矩阵 cm = confusion_matrix(y_true=true_categories, y_pred=predicted_categories) # 可选:可视化混淆矩阵 disp = ConfusionMatrixDisplay( confusion_matrix=cm, display_labels=[f"类别{i}" for i in range(8)] # 替换成你自己的类别名称 ) disp.plot() plt.show()
注意事项
- 拼接测试集标签前要确认测试集没有开启shuffle,否则数据集迭代时样本顺序随机,拼接出的标签和预测结果的样本顺序无法对应,会导致混淆矩阵计算结果完全错误。如果测试集加了shuffle,需要先去掉shuffle操作再做预测和标签拼接。
- 如果你实际做的是多标签分类(一个样本可以同时属于多个类别),不能直接使用普通的多分类混淆矩阵,需要针对每个类别单独计算二分类混淆矩阵,普通混淆矩阵仅支持单标签分类场景。
- 单标签多分类任务建议将损失函数从
binary_crossentropy替换为categorical_crossentropy,更适配单标签one-hot标签的训练场景,binary_crossentropy更适合多标签分类的损失计算。
内容的提问来源于stack exchange,提问作者ForeverLearner
相关产品推荐
相关产品推荐

