如何创建可绘制混淆矩阵的类?解决TypeError报错问题
正确封装混淆矩阵绘制类的解决方案
错误原因分析
- 类的实例方法必须接收
self作为第一个参数,你的printPlot()方法未定义该参数,导致调用时Python自动传入的实例对象被判定为多余参数。 - 调用
printPlot时额外传入了cm3,但类初始化时已将混淆矩阵保存为实例属性self.cm,无需重复传递。 printPlot方法内直接使用cm变量,而非实例属性self.cm,会触发变量未定义错误。
修正后的完整代码
import matplotlib.pyplot as plt import seaborn as sns from sklearn.metrics import confusion_matrix class ConfusionPlot(): def __init__(self, cm): self.cm = cm def printPlot(self): dx = plt.subplot() # 使用实例属性self.cm而非全局变量 sns.heatmap(self.cm, annot=True, fmt='g', ax=dx) # 设置标签、标题和刻度 dx.set_xlabel('Predicted labels') dx.set_ylabel('True labels') dx.set_title('Confusion Matrix') dx.xaxis.set_ticklabels(['Died', 'Survived']) dx.yaxis.set_ticklabels(['Died', 'Survived']) # 显示图形 plt.show()
正确调用方式
# 生成混淆矩阵 cm3 = confusion_matrix(y_train, pred_train2) # 初始化实例 cplot3 = ConfusionPlot(cm3) # 调用绘图方法,无需额外传参 cplot3.printPlot()
可选优化:让类更灵活
当前类的刻度标签为硬编码,可将标签作为参数传入,适配更多分类场景:
class ConfusionPlot(): def __init__(self, cm, class_labels=None): self.cm = cm # 默认使用['Died', 'Survived'],传入自定义标签则替换 self.class_labels = class_labels or ['Died', 'Survived'] def printPlot(self, title='Confusion Matrix'): dx = plt.subplot() sns.heatmap(self.cm, annot=True, fmt='g', ax=dx) dx.set_xlabel('Predicted labels') dx.set_ylabel('True labels') dx.set_title(title) dx.xaxis.set_ticklabels(self.class_labels) dx.yaxis.set_ticklabels(self.class_labels) plt.show() # 自定义标签调用示例 cm3 = confusion_matrix(y_train, pred_train2) cplot3 = ConfusionPlot(cm3, class_labels=['Negative', 'Positive']) cplot3.printPlot(title='Training Set Confusion Matrix')
内容的提问来源于stack exchange,提问作者Jordi Pacreu Antunez
相关产品推荐
相关产品推荐

