PyCaret中自定义sklearn兼容交叉验证无法生效问题求解
问题根因
自定义交叉验证类未完全遵循scikit-learn分割器接口规范,PyCaret内部调用CV逻辑时会严格匹配sklearn接口约定,不兼容的实现会触发内部异常被静默捕获,最终表现为compare_models()执行无输出。
现有实现的不符合项:
split方法签名不满足要求:sklearn规定split方法必须支持X, y=None, groups=None三个入参,当前实现仅接收X一个参数,调用时传参不匹配会直接报错- 缺少必须实现的
get_n_splits方法:sklearn兼容的CV分割器必须实现该方法返回拆分折数,PyCaret初始化流程时会调用该方法读取折数配置 - 索引逻辑不统一:训练集索引基于
reset_index(drop=True)后的新索引计算,测试集索引基于原始DataFrame索引计算,基准不一致会导致索引越界、样本匹配错误 - 实例化代码缩进错误:
custom_CV = custom_cv(...)写在了类定义的缩进块内,属于类属性而非全局可用实例,外部调用时无法拿到正确的CV对象
兼容PyCaret/sklearn的修正实现
import numpy as np class CustomTimeSeriesCV: def __init__(self, train_end, test_size, n_splits): self.train_end = train_end self.test_size = test_size self.n_splits = n_splits def split(self, X, y=None, groups=None): # 统一重置为连续整数索引,避免原始索引乱序/非连续导致的匹配错误 X_ = X.reset_index(drop=True) for i in range(self.n_splits, 0, -1): tr_thresh = self.train_end - self.test_size * i te_thresh = tr_thresh + self.test_size tr_idx = np.where(X_['N_month'] <= tr_thresh)[0] te_idx = np.where((X_['N_month'] > tr_thresh) & (X_['N_month'] <= te_thresh))[0] yield tr_idx, te_idx def get_n_splits(self, X=None, y=None, groups=None): return self.n_splits # 实例化CV对象,注意该行需顶格写,不要缩进在类定义内部 custom_cv = CustomTimeSeriesCV(train_end=365, test_size=28, n_splits=5)
使用注意事项
- 调用PyCaret的
setup函数时,传入参数fold=custom_cv即可,不需要提前手动调用split方法传入分割好的索引 - 如果不想写类实现,也可以直接把你之前函数式写法返回的「(训练集索引列表, 验证集索引列表)元组组成的列表」传给
fold参数,PyCaret同样支持这种格式 - 若仍出现无输出的情况,可以在
setup时加参数verbose=True, html=False,就能看到被静默捕获的具体报错信息,方便定位问题
内容的提问来源于stack exchange,提问作者User
相关产品推荐
相关产品推荐

