使用混淆矩阵遇NameError:predicted_class未定义及标签显示问题求助
问题修复:混淆矩阵代码的两处问题解决
1. 解决NameError: name 'predicted_class' is not defined错误
你代码最后一行的print(predicted_class)里,predicted_class从未被定义。你已经通过classes_x=np.argmax(predict_x,axis=1)得到了预测类别,两种修复方式:
- 直接打印已生成的
classes_x - 将
classes_x赋值给predicted_class后再打印
修改后的代码片段:
predict_x=model.predict(X) classes_x=np.argmax(predict_x,axis=1) # 方案1:直接打印预测结果 print(classes_x) # 方案2:先赋值再打印 predicted_class = classes_x print(predicted_class)
2. 修复混淆矩阵标签显示问题
你的混淆矩阵函数已设置刻度标签,但可以通过优化DataFrame的创建逻辑,让标签显示更精准:
修改后的完整可运行代码:
import numpy as np import pandas as pd import seaborn as sn from sklearn.metrics import confusion_matrix import matplotlib.pyplot as plt def print_confusion_matrix(y_true, y_pred): cm = confusion_matrix(y_true, y_pred) print('True positive = ', cm[0][0]) print('False positive = ', cm[0][1]) print('False negative = ', cm[1][0]) print('True negative = ', cm[1][1]) print('\n') # 直接指定DataFrame的行列标签,对应真实和预测类别 df_cm = pd.DataFrame(cm, index=['Fake', 'Real'], # 实际类别标签 columns=['Fake', 'Real']) # 预测类别标签 sn.set(font_scale=1.4) sn.heatmap(df_cm, annot=True, annot_kws={"size": 16}, fmt='d') plt.ylabel('Actual label', size=20) plt.xlabel('Predicted label', size=20) # 调整刻度位置,让标签居中显示在热图格子上 plt.xticks(np.arange(2)+0.5, ['Fake', 'Real'], size=16) plt.yticks(np.arange(2)+0.5, ['Fake', 'Real'], size=16, rotation=0) plt.ylim([2, 0]) plt.show() # 生成预测类别 predict_x=model.predict(X) classes_x=np.argmax(predict_x,axis=1) # 调用混淆矩阵函数,传入真实标签和预测标签 print_confusion_matrix(Y_val_org, classes_x)
关键优化点:
- 创建
df_cm时直接定义行列标签,避免依赖手动刻度设置 - 调整刻度位置为
np.arange(2)+0.5,让标签与热图格子对齐 - 给
yticks添加rotation=0,防止标签旋转导致显示混乱 - 调用函数时传入正确的真实标签
Y_val_org和预测标签classes_x
内容的提问来源于stack exchange,提问作者Sanjiro
相关产品推荐
相关产品推荐

