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

基于sklearn的递归样本拆分方案定制(适配网格搜索)

自定义面板数据的递归递增交叉验证生成器

问题背景

我有一份面板数据,每个时间截面包含多个样本,示例数据构造代码如下:

import pandas as pd
import numpy as np

dates = ["2018-01-01", "2019-01-01", "2020-01-01", "2021-01-01", "2022-01-01"] * 2
dates.sort()
samples = [1, 2] * 5
df = pd.DataFrame(
    {
        "dates": dates,
        "samples": samples
    }
)

我需要创建一个交叉验证生成器,完成3次验证,规则是训练集递归递增,验证集固定为单个时间截面:

  • 第一次:训练集为["2018-01-01", "2019-01-01"]的样本,验证集为["2020-01-01"]的样本;
  • 第二次:训练集为["2018-01-01", "2019-01-01", "2020-01-01"]的样本,验证集为["2021-01-01"]的样本;
  • 第三次:训练集为["2018-01-01", "2019-01-01", "2020-01-01", "2021-01-01"]的样本,验证集为["2022-01-01"]的样本。

我尝试过sklearn.model_selection.PredefinedSplit(),但它无法满足需求:

  1. 每次拆分无法将所有样本纳入训练集或验证集;
  2. 像"2020-01-01"这样的样本,第一次是验证集,第二次要作为训练集,但PredefinedSplit的拆分规则无法动态调整这种归属。

请问如何定制符合要求的拆分方案?最好基于sklearn实现,能传入GridSearchCV()做网格搜索。


解决方案

你可以自定义一个继承自sklearn.model_selection.BaseCrossValidator的交叉验证器,既能贴合你的时间递归拆分需求,又能完全兼容sklearn的API(包括GridSearchCV)。

实现代码

import pandas as pd
from sklearn.model_selection import BaseCrossValidator

class RecursiveTimeSplit(BaseCrossValidator):
    def __init__(self, date_col, valid_start_idx=2, valid_end_idx=None):
        """
        递归时间交叉验证器:训练集从最早时间开始递增,验证集为单个后续时间截面
        
        参数:
            date_col: 数据框中存储时间的列名
            valid_start_idx: 第一个验证集对应的时间截面索引(从0开始计数)
            valid_end_idx: 最后一个验证集对应的时间截面索引,默认到最后一个
        """
        self.date_col = date_col
        self.valid_start_idx = valid_start_idx
        self.valid_end_idx = valid_end_idx

    def get_n_splits(self, X=None, y=None, groups=None):
        # 计算拆分次数
        unique_dates = sorted(X[self.date_col].unique())
        end_idx = self.valid_end_idx if self.valid_end_idx is not None else len(unique_dates)-1
        return end_idx - self.valid_start_idx + 1

    def split(self, X, y=None, groups=None):
        unique_dates = sorted(X[self.date_col].unique())
        end_idx = self.valid_end_idx if self.valid_end_idx is not None else len(unique_dates)-1
        
        for valid_idx in range(self.valid_start_idx, end_idx+1):
            # 训练集取valid_idx之前的所有时间截面
            train_dates = unique_dates[:valid_idx]
            # 验证集取当前valid_idx对应的时间截面
            valid_date = unique_dates[valid_idx]
            
            # 获取训练集和验证集的索引
            train_indices = X[X[self.date_col].isin(train_dates)].index
            valid_indices = X[X[self.date_col] == valid_date].index
            
            yield train_indices, valid_indices

使用示例

# 初始化自定义交叉验证器
# 这里unique_dates的索引是0:2018,1:2019,2:2020,3:2021,4:2022
# valid_start_idx=2对应第一个验证集是2020-01-01
tscv = RecursiveTimeSplit(date_col="dates", valid_start_idx=2)

# 测试拆分效果
for i, (train_idx, valid_idx) in enumerate(tscv.split(df)):
    print(f"第{i+1}次拆分:")
    print(f"训练集日期: {sorted(df.loc[train_idx, 'dates'].unique())}")
    print(f"验证集日期: {df.loc[valid_idx, 'dates'].unique()[0]}")
    print("---")

# 传入GridSearchCV使用
from sklearn.model_selection import GridSearchCV
from sklearn.linear_model import LinearRegression

# 假设我们有特征和目标变量(这里用samples作为示例特征,可替换为实际数据)
X = df[["samples"]]
y = df["samples"] * np.random.randn(len(df))  # 模拟目标变量

model = LinearRegression()
param_grid = {"fit_intercept": [True, False]}

grid_search = GridSearchCV(
    estimator=model,
    param_grid=param_grid,
    cv=tscv,
    scoring="neg_mean_squared_error"
)
grid_search.fit(X, y)

print("最佳参数:", grid_search.best_params_)

代码说明

  1. 继承BaseCrossValidator:这是sklearn自定义交叉验证器的强制要求,必须实现get_n_splits和split两个核心方法;
  2. split方法逻辑:
    • 先提取并排序所有唯一时间截面;
    • 遍历每个验证时间点,训练集取该时间点之前的所有样本,验证集取当前时间点的所有样本;
    • 返回训练集和验证集的索引,完全符合sklearn交叉验证器的输出格式;
  3. 兼容性:这个验证器可以直接传入GridSearchCV的cv参数,无缝对接网格搜索流程。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.06 03:11:28