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

如何基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 15:06:23