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

使用scikit-learn cross_val_predict做k折预测及模型对比是否合理?

你的预测及模型对比方式不合理,以下是具体分析和正确流程

核心问题:误用cross_val_predict在测试集上的操作完全违背了训练/测试集的拆分原则

cross_val_predict的设计逻辑是:对输入的数据集做k折拆分,每次用k-1折数据训练模型,再预测剩下的1折数据。如果把测试集传入该方法,相当于让模型在部分测试数据上完成了训练——这直接破坏了测试集作为「未见过的真实场景数据」的定位,得到的评估结果会严重高估模型的泛化能力,完全失去了参考价值。

正确的流程(针对XGBoost缓解过拟合+模型对比)

1. 固定训练/测试集拆分,测试集全程"锁死"

先把整个数据集拆分为(x_train, y_train)和(x_test, y_test),测试集只留到最后一步做最终评估,中间绝对不能参与任何训练或交叉验证环节。

2. 用k折交叉验证做两件关键事(缓解XGBoost过拟合的核心)

  • 评估模型在训练集内的泛化能力:用cross_val_score看模型在训练集拆分后的验证折上的表现,判断是否过拟合
  • 超参数调优:XGBoost默认参数很容易过拟合,结合k折交叉验证搜索最优正则化参数(比如max_depth、learning_rate、subsample等),这才是缓解过拟合的有效手段

示例代码(调参+交叉验证):

from sklearn.model_selection import GridSearchCV, StratifiedKFold
from xgboost import XGBClassifier

# 定义XGBoost的参数搜索网格(正则化相关参数为主)
param_grid = {
    'max_depth': [3, 5],          # 限制树深,防止过拟合
    'learning_rate': [0.01, 0.1], # 减小步长,降低模型对噪声的敏感度
    'subsample': [0.8],           # 随机采样训练样本
    'colsample_bytree': [0.8]     # 随机采样特征
}

kfold = StratifiedKFold(n_splits=10, shuffle=True, random_state=42)
model = XGBClassifier()

# 网格搜索+交叉验证,自动找到最优参数的模型
grid_search = GridSearchCV(estimator=model, param_grid=param_grid, cv=kfold, scoring='accuracy')
grid_search.fit(x_train, y_train)

# 得到经过交叉验证调优后的最优XGBoost模型
best_xgb = grid_search.best_estimator_

3. 用全训练集训练模型,再在测试集上做最终预测和评估

经过交叉验证调参后,用完整的训练集重新训练最优模型(GridSearchCV默认已经帮你完成了这一步),然后用这个模型预测测试集,得到真实的泛化性能指标:

from sklearn import metrics

# XGBoost测试集预测与评估
y_pred_xgb = best_xgb.predict(x_test)
accuracy_xgb = metrics.accuracy_score(y_test, y_pred_xgb)
cm_xgb = metrics.confusion_matrix(y_test, y_pred_xgb)

# Logistic Regression遵循同样流程(可选交叉验证调参)
from sklearn.linear_model import LogisticRegression

lr = LogisticRegression()
lr_param_grid = {'C': [0.1, 1, 10]} # 调正则化参数C
lr_grid = GridSearchCV(lr, lr_param_grid, cv=kfold, scoring='accuracy')
lr_grid.fit(x_train, y_train)
best_lr = lr_grid.best_estimator_

y_pred_lr = best_lr.predict(x_test)
accuracy_lr = metrics.accuracy_score(y_test, y_pred_lr)
cm_lr = metrics.confusion_matrix(y_test, y_pred_lr)

4. 模型对比的正确姿势

基于同一个未被污染的测试集上的预测结果,对比两个模型的准确率、混淆矩阵等指标——这样的对比才是公平且有参考意义的。

补充:cross_val_predict的正确用法

它应该用在训练集内部,比如生成训练样本的交叉验证预测结果,用来做误差分析、特征堆叠等:

y_train_pred = cross_val_predict(model, x_train, y_train, cv=kfold)
# 分析训练集内的交叉验证误差,判断模型是否过拟合
train_cv_accuracy = metrics.accuracy_score(y_train, y_train_pred)

内容的提问来源于stack exchange,提问作者Jaime Andrés Castañeda

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 20:01:33