运行混淆矩阵代码时遇ValueError: multilabel-indicator is not supported求解决
解决ValueError: multilabel-indicator is not supported的方案
这个错误是因为你传入confusion_matrix的y_test或y_pred是多标签指示矩阵(也就是one-hot编码格式,比如二维数组,每一行是[0,1,0]这种形式表示类别),而confusion_matrix只支持一维的类别索引数组(比如[1,2,0]这种直接用数字表示类别的格式)。
给你两种针对性的解决思路:
1. 如果是单标签分类任务(每个样本只属于一个类别)
把one-hot格式的标签转成类别索引,用numpy.argmax()即可实现:
import numpy as np from sklearn.metrics import confusion_matrix import seaborn as sns # 将one-hot格式的标签转换为类别索引 y_test_idx = np.argmax(y_test, axis=1) y_pred_idx = np.argmax(y_pred, axis=1) # 重新绘制热力图 sns.heatmap(confusion_matrix(y_test_idx, y_pred_idx), annot=True)
2. 如果是多标签分类任务(每个样本可属于多个类别)
这种场景下普通confusion_matrix不适用,改用sklearn的multilabel_confusion_matrix,它会为每个类别生成单独的混淆矩阵:
from sklearn.metrics import multilabel_confusion_matrix import seaborn as sns import matplotlib.pyplot as plt # 生成多标签混淆矩阵 ml_cm = multilabel_confusion_matrix(y_test, y_pred) # 逐个绘制每个类别的混淆矩阵 for i, cm in enumerate(ml_cm): plt.figure() sns.heatmap(cm, annot=True, title=f"类别{i}的混淆矩阵")
内容的提问来源于stack exchange,提问作者Sumeet Yadav
相关产品推荐
相关产品推荐

