能否让sklearn GridSearchCV基于中位数而非均值评估候选模型?
基于中位数选择GridSearchCV最优模型的解决方案
针对小数据集或LeaveOneOut这类CV场景下,均值评估易受极端值干扰的问题,我们可以通过自定义GridSearchCV子类来实现基于中位数的模型选择,同时解决你当前方案中的两个核心问题:
问题根源
你当前的自定义refit函数仅返回最优模型的索引,但GridSearchCV本身不会自动:
- 在
cv_results_中添加中位数分数列 - 为
best_score_赋值 - 处理
score()方法中因refit为函数导致的KeyError(因为scorer_字典的键是评分器名称,而非函数)
完整解决方案代码
import numpy as np from sklearn.model_selection import GridSearchCV class MedianGridSearchCV(GridSearchCV): def __init__(self, estimator, param_grid, refit_scorer_name, **kwargs): # 保存要用于中位数评估的评分器名称 self.refit_scorer_name = refit_scorer_name # 调用父类构造,refit暂时设为False,后续手动处理 super().__init__(estimator, param_grid, refit=False, **kwargs) def _post_process(self): # 先执行父类的后处理逻辑,生成基础的cv_results_ super()._post_process() # 遍历所有测试折的分数,提取目标评分器的结果 split_scores = [] for key in self.cv_results_: if key.startswith('split') and f'test_{self.refit_scorer_name}' in key: split_scores.append(self.cv_results_[key]) # 计算每个候选模型的中位数分数 median_scores = np.median(split_scores, axis=0) # 将中位数分数添加到cv_results_中 self.cv_results_[f'median_test_{self.refit_scorer_name}'] = median_scores def fit(self, X, y=None, **fit_params): # 执行父类的fit逻辑 super().fit(X, y, **fit_params) # 基于中位数分数选择最优模型索引 median_key = f'median_test_{self.refit_scorer_name}' best_idx = self.cv_results_[median_key].argmax() # 设置最优模型相关属性 self.best_index_ = best_idx self.best_params_ = self.cv_results_['params'][best_idx] self.best_estimator_ = self.estimator.set_params(**self.best_params_) # 重新训练最优模型(因为父类refit=False,所以需要手动fit) self.best_estimator_.fit(X, y, **fit_params) # 赋值best_score_为选中模型的中位数分数 self.best_score_ = self.cv_results_[median_key][best_idx] return self def score(self, X, y=None): # 直接使用指定的评分器计算分数,避免KeyError if isinstance(self.scorer_, dict): scorer = self.scorer_[self.refit_scorer_name] else: scorer = self.scorer_ return scorer(self.best_estimator_, X, y)
使用示例
# 假设你的pipe、params、scorer已定义 grid = MedianGridSearchCV( pipe, params, refit_scorer_name='your_scorer_name', # 替换为你实际使用的评分器名称 scoring=scorer, cv=10, verbose=1, n_jobs=-1 ) grid.fit(X_train, y_train) # 现在可以正常访问以下属性: print(grid.best_score_) # 最优模型的中位数分数 print(grid.cv_results_['median_test_your_scorer_name']) # 所有候选模型的中位数分数 print(grid.score(X_test, y_test)) # 不再触发KeyError
关键改进点
- 添加中位数分数列:通过重写
_post_process方法,在cv_results_中自动生成median_test_<scorer_name>字段,方便后续分析。 - 维护best_score_属性:在
fit方法中手动计算并赋值最优模型的中位数分数。 - 修复score()方法的KeyError:直接使用指定的评分器名称获取scorer,绕过
self.refit为函数导致的键查找失败问题。 - 手动训练最优模型:因为父类
refit=False,所以在选中最优参数后手动调用fit训练模型,保证best_estimator_是训练好的状态。
内容的提问来源于stack exchange,提问作者roble
相关产品推荐
相关产品推荐

