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

能否让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

关键改进点

  1. 添加中位数分数列:通过重写_post_process方法,在cv_results_中自动生成median_test_<scorer_name>字段,方便后续分析。
  2. 维护best_score_属性:在fit方法中手动计算并赋值最优模型的中位数分数。
  3. 修复score()方法的KeyError:直接使用指定的评分器名称获取scorer,绕过self.refit为函数导致的键查找失败问题。
  4. 手动训练最优模型:因为父类refit=False,所以在选中最优参数后手动调用fit训练模型,保证best_estimator_是训练好的状态。

内容的提问来源于stack exchange,提问作者roble

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 16:16:06