适配Scikit-learn新版本:替换plot_confusion_matrix为ConfusionMatrixDisplay
适配Scikit-learn新API:用ConfusionMatrixDisplay替代已弃用的plot_confusion_matrix
方案1:重构自定义函数,兼容原调用逻辑
原自定义函数的核心效果包括指定颜色映射、自定义字体/刻度样式、红色单元格文本、保留小数格式等。基于ConfusionMatrixDisplay重构后,可保留原函数的参数接口,无需修改调用代码:
import matplotlib.pyplot as plt from sklearn.metrics import ConfusionMatrixDisplay import numpy as np from matplotlib import rc def plot_confusMatrix(cm, classes, title='Confusion matrix', cmap=plt.cm.Blues): # 初始化混淆矩阵显示对象 disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=classes) # 绘制基础混淆矩阵,指定数值格式 disp.plot(cmap=cmap, values_format='.1f') # 应用原自定义样式 plt.rcParams.update({'font.size': 19}) disp.ax_.set_title(title, fontdict={'size':'16'}) # 设置刻度标签样式 disp.ax_.set_xticklabels(classes, rotation=45, fontsize=12, color="blue") disp.ax_.set_yticklabels(classes, fontsize=12, color="blue") # 设置轴标签样式 disp.ax_.set_ylabel('True label', fontdict={'size':'16'}) disp.ax_.set_xlabel('Predicted label', fontdict={'size':'16'}) # 设置字体加粗 rc('font', weight='bold') # 修改单元格文本颜色为红色 for text in disp.text_.flatten(): text.set_color("red") plt.tight_layout() # 原调用代码完全复用 plot_confusMatrix(confusion_matrix(y_test, y_pred=y_pred), classes=['Non Fraud','Fraud'], title='Confusion matrix')
方案2:直接使用ConfusionMatrixDisplay替代原plot_confusion_matrix调用
如果不需要保留自定义函数,可直接用ConfusionMatrixDisplay的API实现相同效果:
from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay import matplotlib.pyplot as plt from matplotlib import rc # 计算混淆矩阵 cm = confusion_matrix(y_test, y_pred=y_pred) classes = ['Non Fraud','Fraud'] # 初始化并绘制混淆矩阵 disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=classes) disp.plot(cmap=plt.cm.Blues, values_format='.1f') # 应用自定义样式 plt.rcParams.update({'font.size': 19}) disp.ax_.set_title('Confusion matrix', fontdict={'size':'16'}) disp.ax_.set_xticklabels(classes, rotation=45, fontsize=12, color="blue") disp.ax_.set_yticklabels(classes, fontsize=12, color="blue") disp.ax_.set_ylabel('True label', fontdict={'size':'16'}) disp.ax_.set_xlabel('Predicted label', fontdict={'size':'16'}) rc('font', weight='bold') for text in disp.text_.flatten(): text.set_color("red") plt.tight_layout() plt.show()
关键适配点说明
ConfusionMatrixDisplay需先传入混淆矩阵和类别标签,通过plot()方法生成图像- 原函数中的
fmt='.1f'对应values_format='.1f'参数 - 原函数中手动调整的样式,均可通过
disp.ax_(获取绘图的Axes对象)修改 - 单元格文本颜色可通过遍历
disp.text_集合统一设置
内容的提问来源于stack exchange,提问作者Xray25
相关产品推荐
相关产品推荐

