Logistic回归模型:训练/测试集预测与ROC AUC计算及可视化
关于statsmodels逻辑回归的预测结果与模型评估解答
一、获取训练集与测试集的预测结果
你当前使用result.predict()的方式是正确的,statsmodels的Logistic回归模型predict()方法默认返回正类的预测概率(取值0到1之间)。如果需要得到分类标签(0或1),可以设定一个阈值(通常取0.5)将概率转换为类别:
# 获取训练集、测试集的预测概率 in_sample_pred = result.predict(x_train_data) out_sample_pred = result.predict(x_test_data) # 将概率转换为类别标签(阈值设为0.5) in_sample_pred_class = (in_sample_pred >= 0.5).astype(int) out_sample_pred_class = (out_sample_pred >= 0.5).astype(int)
二、常用分类评估指标
逻辑回归是分类任务,需要用到分类评估指标。从sklearn.metrics导入相关指标即可,常用的包括准确率、精确率、召回率、F1分数、混淆矩阵:
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix # 训练集指标计算 print("训练集准确率:", accuracy_score(y_train_data, in_sample_pred_class)) print("训练集精确率:", precision_score(y_train_data, in_sample_pred_class)) print("训练集召回率:", recall_score(y_train_data, in_sample_pred_class)) print("训练集F1分数:", f1_score(y_train_data, in_sample_pred_class)) print("训练集混淆矩阵:\n", confusion_matrix(y_train_data, in_sample_pred_class)) # 测试集指标计算 print("\n测试集准确率:", accuracy_score(y_test_data, out_sample_pred_class)) print("测试集精确率:", precision_score(y_test_data, out_sample_pred_class)) print("测试集召回率:", recall_score(y_test_data, out_sample_pred_class)) print("测试集F1分数:", f1_score(y_test_data, out_sample_pred_class)) print("测试集混淆矩阵:\n", confusion_matrix(y_test_data, out_sample_pred_class))
三、用scikit-learn计算ROC AUC并绘制曲线
ROC AUC是评估二分类模型性能的常用指标,结合sklearn.metrics和matplotlib即可完成计算与绘图:
from sklearn.metrics import roc_auc_score, roc_curve import matplotlib.pyplot as plt # 计算ROC AUC分数(直接用预测概率计算) train_roc_auc = roc_auc_score(y_train_data, in_sample_pred) test_roc_auc = roc_auc_score(y_test_data, out_sample_pred) print("训练集ROC AUC:", train_roc_auc) print("测试集ROC AUC:", test_roc_auc) # 绘制ROC曲线 plt.figure(figsize=(8, 6)) # 绘制训练集ROC曲线 fpr_train, tpr_train, _ = roc_curve(y_train_data, in_sample_pred) plt.plot(fpr_train, tpr_train, label=f'Train ROC AUC = {train_roc_auc:.2f}') # 绘制测试集ROC曲线 fpr_test, tpr_test, _ = roc_curve(y_test_data, out_sample_pred) plt.plot(fpr_test, tpr_test, label=f'Test ROC AUC = {test_roc_auc:.2f}') # 绘制随机模型的对角线 plt.plot([0, 1], [0, 1], 'k--') plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('ROC Curve') plt.legend() plt.show()
内容的提问来源于stack exchange,提问作者user3003605
相关产品推荐
相关产品推荐

