You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.24 23:17:26