如何修改GridSearchCV的fit()方法以向自定义交叉验证器传递参数
GridSearchCV向自定义交叉验证器传参解决方案
不建议直接修改GridSearchCV的fit方法源码,会破坏scikit-learn原生API兼容性,后续版本升级会出现适配问题,你可以通过以下两种原生支持的方案实现需求:
方案一:初始化自定义CV时直接传入固定参数(最简便)
适合pred_times、eval_times在整个网格搜索过程中固定不变的场景,修改你自定义的CombPurgedKFoldCV类的初始化方法,将额外参数作为类属性存储,拆分时直接调用即可:
- 修改自定义CV类结构
class CombPurgedKFoldCV: # 原有参数保留,新增两个额外参数作为初始化入参 def __init__(self, n_splits=10, n_test_splits=2, embargo_td=None, pred_times=None, eval_times=None): self.n_splits = n_splits self.n_test_splits = n_test_splits self.embargo_td = embargo_td # 存储额外参数供拆分逻辑调用 self.pred_times = pred_times self.eval_times = eval_times def split(self, X, y=None, groups=None): # 原有拆分逻辑直接调用类属性即可,无需额外传参 # 你原来的split实现代码...
- 调整调用代码
import pandas as pd from sklearn.model_selection import GridSearchCV from sklearn.tree import DecisionTreeClassifier from sklearn.ensemble import BaggingClassifier # 初始化CV时直接传入额外参数 skf = CombPurgedKFoldCV( n_splits=10, n_test_splits=2, embargo_td=pd.Timedelta(minutes=100), pred_times=pred_times, eval_times=eval_times ) clf = DecisionTreeClassifier(criterion='entropy',max_features='auto',class_weight='balanced',min_weight_fraction_leaf=0.) classifier=BaggingClassifier(base_estimator=clf,n_estimators=1000,max_features=1., max_samples=avgU,oob_score=True,n_jobs=1) gs = GridSearchCV(estimator = classifier, param_grid = grid_param, scoring = 'f1',n_jobs = 1, cv=skf) # 正常调用fit即可,无需传入额外参数 gs.fit(X,y)
方案二:通过groups参数动态传参
适合需要动态调整传入参数、或者不想修改CV初始化逻辑的场景,GridSearchCV的fit方法自带groups参数,会原封不动传递给交叉验证器的split方法:
- 调整fit调用代码,将参数打包传入groups
# 把需要的两个参数打包成元组传给groups入参 gs.fit(X, y, groups=(pred_times, eval_times))
- 修改自定义CV的split方法,拆解参数
def split(self, X, y=None, groups=None): # 从groups中拆解出需要的两个参数 pred_times, eval_times = groups # 你原来的split实现代码...
注意:不要将给交叉验证器的参数放到
fit的额外关键字参数中,这部分参数会默认传递给基模型的fit方法,不会进入交叉验证器的逻辑,会触发参数不匹配报错。
内容的提问来源于stack exchange,提问作者Javier C Salaverri
相关产品推荐
相关产品推荐

