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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 13:54:21