如何在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
相关产品推荐
相关产品推荐

