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

如何在GridSearchCV每次迭代后保存clf.cv_results_至文件?

保存GridSearchCV每次迭代的cv_results_到文件

我完全理解你的痛点——跑超参搜索经常要耗很久,中途崩溃前功尽弃真的太闹心了!下面给你两个可行的方案,能帮你在每次迭代完成后自动把cv_results_存到文件里:

方案1:使用Sklearn 1.1+的Callback API(推荐)

从Sklearn 1.1版本开始,官方引入了回调机制,你可以自定义一个回调类,在每个超参组合的交叉验证完成后自动保存结果:

import json
import numpy as np
from sklearn.model_selection import GridSearchCV
from sklearn.svm import SVC
from sklearn.datasets import load_iris

class SaveCVResultsCallback:
    def __init__(self, save_path="cv_results.json"):
        self.save_path = save_path

    def on_train_begin(self, estimator, **kwargs):
        # 训练开始时的初始化操作(可选)
        pass

    def on_iteration_end(self, estimator, **kwargs):
        # 每次迭代完成后保存结果,numpy数组要转成列表才能存JSON
        cv_results = {
            k: v.tolist() if isinstance(v, np.ndarray) else v 
            for k, v in estimator.cv_results_.items()
        }
        with open(self.save_path, "w") as f:
            json.dump(cv_results, f, indent=4)
        print(f"已保存当前迭代结果至 {self.save_path}")

# 示例用法
X, y = load_iris(return_X_y=True)
param_grid = {"C": [0.1, 1, 10], "gamma": [1, 0.1, 0.01]}
svc = SVC()

grid_search = GridSearchCV(
    svc, 
    param_grid, 
    cv=3,
    callbacks=[SaveCVResultsCallback(save_path="grid_search_results.json")]
)

grid_search.fit(X, y)

如果想要更高效的存储(比如结果里有复杂类型),可以把JSON换成pickle,修改回调里的保存逻辑就行:

import pickle

def on_iteration_end(self, estimator, **kwargs):
    with open(self.save_path, "wb") as f:
        pickle.dump(estimator.cv_results_, f)
    print(f"已保存当前迭代结果至 {self.save_path}")

方案2:手动重写GridSearchCV的_run_search方法(兼容旧版Sklearn)

要是你用的是Sklearn 1.1之前的版本,没有回调API,可以通过继承GridSearchCV并重写核心方法来实现:

import pickle
from sklearn.model_selection import GridSearchCV
from sklearn.svm import SVC
from sklearn.datasets import load_iris

class GridSearchCVWithSave(GridSearchCV):
    def __init__(self, estimator, param_grid, save_path="cv_results.pkl", **kwargs):
        super().__init__(estimator, param_grid, **kwargs)
        self.save_path = save_path

    def _run_search(self, evaluate_candidates):
        # 包装评估函数,每次候选参数跑完就保存结果
        def wrapped_evaluate(candidate_params):
            evaluate_candidates(candidate_params)
            with open(self.save_path, "wb") as f:
                pickle.dump(self.cv_results_, f)
            print(f"已保存当前迭代结果至 {self.save_path}")
        
        wrapped_evaluate(self.param_grid)

# 示例用法
X, y = load_iris(return_X_y=True)
param_grid = {"C": [0.1, 1, 10], "gamma": [1, 0.1, 0.01]}
svc = SVC()

grid_search = GridSearchCVWithSave(
    svc, 
    param_grid, 
    cv=3,
    save_path="grid_search_results.pkl"
)

grid_search.fit(X, y)

一些注意事项

  • 频繁写入文件会略微影响搜索速度,但和崩溃丢失所有数据比起来,这个代价完全值得
  • 用JSON保存的话,后续加载后可以把列表转回numpy数组:
    import json
    import numpy as np
    
    with open("grid_search_results.json", "r") as f:
        cv_results = json.load(f)
    cv_results = {k: np.array(v) if isinstance(v, list) else v for k, v in cv_results.items()}
    
  • 用pickle保存的结果只能在Python环境中读取,适合内部使用;JSON则更通用,方便跨语言查看

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 06:54:57