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

为Scikit-learn的GridSearchCV注入进度条的更优方案

更优的Scikit-learn GridSearchCV进度条实现方案

针对你提到的现有进度条方案存在的「需自定义评分器、不支持多进程、实现粗糙」等问题,这里提供两种更实用的解决方案,均无需修改Scikit-learn源码,且支持多进程并行:

方案一:基于Joblib回调+Alive-progress实现

利用GridSearchCV底层依赖Joblib做并行计算的特性,通过Joblib的回调机制跟踪任务进度,配合Alive-progress生成直观的进度条。

代码示例

from alive_progress import alive_bar
from sklearn.model_selection import GridSearchCV
from sklearn.svm import SVC
from sklearn.datasets import load_iris
from joblib import parallel_backend

# 自定义Joblib进度回调类
class ProgressCallback:
    def __init__(self, total_tasks):
        self.total = total_tasks
        self.completed = 0
        self.bar = alive_bar(total_tasks)
    
    def __call__(self, _):
        self.completed += 1
        self.bar()
        if self.completed == self.total:
            self.bar.close()

# 加载测试数据
X, y = load_iris(return_X_y=True)

# 定义模型与参数网格
model = SVC()
param_grid = {'C': [0.1, 1, 10, 100], 'kernel': ['linear', 'rbf']}
cv_folds = 5
# 计算总任务数:参数组合数 × CV折数
total_tasks = len(param_grid['C']) * len(param_grid['kernel']) * cv_folds

# 初始化进度回调
progress_cb = ProgressCallback(total_tasks)

# 启用多进程并行并绑定进度回调
with parallel_backend('loky', n_jobs=-1, callback=progress_cb):
    grid_search = GridSearchCV(model, param_grid, cv=cv_folds, verbose=0)
    grid_search.fit(X, y)

print("最优参数组合:", grid_search.best_params_)

方案优势

  • 无需自定义评分器,完全兼容Scikit-learn内置的所有评估指标
  • 原生支持多进程(n_jobs>1),无序列化问题
  • 实现简洁,不依赖修改Scikit-learn内部逻辑

方案二:子类化GridSearchCV集成进度条

通过子类化GridSearchCV并重写_run_search方法,直接在类内部集成进度条逻辑,更贴合Scikit-learn的架构设计。

代码示例

from alive_progress import alive_bar
from sklearn.model_selection import GridSearchCV
from sklearn.svm import SVC
from sklearn.datasets import load_iris

class ProgressGridSearchCV(GridSearchCV):
    def _run_search(self, evaluate_candidates):
        # 计算总任务数
        total_tasks = len(self.param_grid) * self.cv
        with alive_bar(total_tasks) as bar:
            # 包装评估函数,每次完成任务更新进度条
            def progress_wrapper(candidate_params, *args, **kwargs):
                bar()
                return evaluate_candidates(candidate_params, *args, **kwargs)
            
            super()._run_search(progress_wrapper)

# 测试使用
X, y = load_iris(return_X_y=True)
model = SVC()
param_grid = {'C': [0.1, 1, 10, 100], 'kernel': ['linear', 'rbf']}

grid_search = ProgressGridSearchCV(model, param_grid, cv=5, n_jobs=-1, verbose=0)
grid_search.fit(X, y)

print("最优参数组合:", grid_search.best_params_)

方案优势

  • 封装性强,直接替换原有GridSearchCV即可使用
  • 兼容多进程并行,无序列化障碍
  • 无需额外依赖Joblib的回调配置,逻辑更紧凑

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 21:55:44