如何在scikit-learn交叉验证的Logistic Regression管道中获取各折Confusion Matrix?
好问题!cross_val_score确实能快速拿到交叉验证的分数,但如果想获取每折的混淆矩阵,咱们得换个思路——要么借助cross_validate拿到每折训练好的模型,要么手动遍历每一轮交叉验证的拆分。下面给你两种实用的方案:
方法1:用cross_validate获取每折模型并计算混淆矩阵
cross_validate比cross_val_score更灵活,我们可以通过return_estimator=True让它返回每折训练完成的Pipeline,同时用return_indices=True拿到每折的测试集索引,之后就能用这些信息计算混淆矩阵了:
from sklearn.pipeline import make_pipeline from sklearn.preprocessing import MinMaxScaler from sklearn.linear_model import LogisticRegression from sklearn.model_selection import cross_validate, KFold from sklearn.metrics import confusion_matrix # 初始化你的Pipeline和交叉验证策略 clf = make_pipeline(MinMaxScaler(), LogisticRegression()) cv = KFold(n_splits=3, shuffle=True, random_state=42) # shuffle可选,根据你的需求调整 # 执行交叉验证,返回模型和测试集索引 results = cross_validate(clf, X_train, y_train, cv=cv, return_estimator=True, return_indices=True) # 遍历每折,计算并保存混淆矩阵 confusion_matrices = [] for fold_num, (estimator, test_idx) in enumerate(zip(results['estimator'], results['test_indices']), 1): # 根据测试集索引获取对应的数据 # 如果X_train/y_train是numpy数组,直接用X_train[test_idx]即可,不用iloc y_true = y_train.iloc[test_idx] y_pred = estimator.predict(X_train.iloc[test_idx]) # 计算混淆矩阵 cm = confusion_matrix(y_true, y_pred) confusion_matrices.append(cm) # 打印每折结果 print(f"第{fold_num}折混淆矩阵:") print(cm) print("---") # 后续可以对confusion_matrices做分析,比如求平均等
方法2:手动遍历交叉验证折
这种方式更直观,完全由你控制每一轮的训练和测试流程,适合需要自定义更多步骤的场景:
from sklearn.model_selection import KFold cv = KFold(n_splits=3, shuffle=True, random_state=42) confusion_matrices = [] for fold_num, (train_idx, test_idx) in enumerate(cv.split(X_train), 1): # 拆分当前折的训练和测试数据 X_fold_train, X_fold_test = X_train.iloc[train_idx], X_train.iloc[test_idx] y_fold_train, y_fold_test = y_train.iloc[train_idx], y_train.iloc[test_idx] # 训练Pipeline clf.fit(X_fold_train, y_fold_train) # 预测并计算混淆矩阵 y_pred = clf.predict(X_fold_test) cm = confusion_matrix(y_fold_test, y_pred) confusion_matrices.append(cm) print(f"第{fold_num}折混淆矩阵:") print(cm) print("---")
小提示
- 如果你的数据集是numpy数组而非Pandas DataFrame/Series,直接用
X_train[train_idx]这种索引方式就行,不用iloc。 - 可以给
confusion_matrix传入labels参数(比如labels=[0,1]),固定标签的顺序,避免因为类别顺序不一致导致混淆矩阵结果混乱。 - 如果需要同时计算其他指标(比如精确率、召回率),可以在每折里调用对应的评估函数(如
precision_score)一起计算。
内容的提问来源于stack exchange,提问作者Nick
相关产品推荐
相关产品推荐

