如何绘制混淆矩阵:非零值保留1位小数,零值显示为0而非0.0?
实现混淆矩阵的自定义数值格式化(0显示为0,非零保留1位小数)
问题背景
需要绘制混淆矩阵,要求:
- 精确为0的数值显示为
0而非0.0 - 非零数值保留1位小数(比如
4.0保留为4.0)
直接向ConfusionMatrixDisplay.plot()的values_format参数传入自定义格式化函数会触发TypeError,因为该参数仅支持格式字符串,不接受函数。
解决方案
方法1:绘制后遍历修改文本元素
先按常规方式绘制混淆矩阵(用.1f保证非零值格式),再逐个修改单元格的文本内容:
import matplotlib.pyplot as plt from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay y_true = [0, 0, 0, 0, 1, 1, 1, 1] y_pred = [0, 0, 0, 0, 1, 1, 1, 1] cm = confusion_matrix(y_true, y_pred) disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=[0, 1]) # 先绘制,非零值按.1f格式显示 disp.plot(include_values=True, cmap='Blues', xticks_rotation='vertical', values_format='.1f', ax=None, colorbar=True) # 遍历所有单元格文本,将0.0替换为0 for text in disp.text_.ravel(): value = float(text.get_text()) if value == 0: text.set_text('0') plt.show()
方法2:预先生成格式化字符串矩阵
关闭默认的数值显示,手动生成格式化后的字符串矩阵并添加到单元格中:
import matplotlib.pyplot as plt import numpy as np from sklearn.metrics import confusion_matrix, ConfusionMatrixDisplay y_true = [0, 0, 0, 0, 1, 1, 1, 1] y_pred = [0, 0, 0, 0, 1, 1, 1, 1] cm = confusion_matrix(y_true, y_pred) disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=[0, 1]) # 绘制时不显示默认数值 disp.plot(include_values=False, cmap='Blues', xticks_rotation='vertical', ax=None, colorbar=True) # 定义格式化逻辑,向量化处理整个混淆矩阵 def custom_format(val): return '0' if val == 0 else f'{val:.1f}' formatted_cm = np.vectorize(custom_format)(cm) # 逐个单元格添加格式化后的文本 for i in range(cm.shape[0]): for j in range(cm.shape[1]): disp.ax_.text(j, i, formatted_cm[i, j], horizontalalignment='center', verticalalignment='center') plt.show()
关键说明
ConfusionMatrixDisplay.plot()的values_format参数仅接受标准的Python格式字符串(如.1f、d等),无法直接传入自定义函数,因此需要通过上述两种方式实现自定义格式化逻辑。
内容的提问来源于stack exchange,提问作者scriptgirl_3000
相关产品推荐
相关产品推荐

