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

Scikit-Learn cross_validate无法使用非连续索引自定义折的问题求助

解决Scikit-Learn cross_validate自定义交叉验证折的索引越界问题

问题场景

向Scikit-Learn的cross_validate传入自定义交叉验证折时,触发IndexError: indices are out-of-bounds错误,相关代码示例如下:

import pandas as pd
import numpy as np
from sklearn.model_selection import cross_validate
from sklearn.linear_model import LinearRegression

# 生成示例数据
df2 = pd.DataFrame(np.random.rand(8, 3), columns=['feature_1', 'feature_2', 'feature_3'])

# 拆分训练集(基于原DataFrame的标签索引)
train_index_list = [0,1,2,5,6,7]
X_train = df2.loc[train_index_list].drop(columns='feature_3').copy()
y_train = df2.loc[train_index_list]['feature_3'].copy()

# 自定义2折交叉验证(使用原DataFrame的标签索引)
cv_folds = [ ([0,1,2], [5,6,7]), ([5,6,7], [0,1,2])]

# 初始化模型
model = LinearRegression()

# 触发错误的调用
cv_output = cross_validate(model, X_train, y_train, scoring=['neg_mean_squared_error'], cv=cv_folds)

原因分析

Scikit-Learn的cross_validate处理自定义交叉验证折时,采用位置索引(即iloc逻辑),而非DataFrame的标签索引:

  • X_train和y_train是原DataFrame的子集,行标签保留原数据的[0,1,2,5,6,7],但实际只有6行,位置索引范围为0~5。
  • 你定义的cv_folds使用了原DataFrame的标签值(如5,6,7),这些值超出了X_train的位置索引范围,因此触发越界错误。

解决方案

方法1:重置训练集索引(最简单直接)

将X_train和y_train的索引重置为连续的位置索引,让标签与位置索引完全匹配:

# 重置索引,drop=True丢弃原标签索引
X_train = X_train.reset_index(drop=True)
y_train = y_train.reset_index(drop=True)

# 基于新的连续索引定义交叉验证折
cv_folds = [ ([0,1,2], [3,4,5]), ([3,4,5], [0,1,2])]

# 重新调用cross_validate
cv_output = cross_validate(model, X_train, y_train, scoring=['neg_mean_squared_error'], cv=cv_folds)

方法2:将原标签索引转换为训练集的位置索引

如果不想重置索引,可以手动将cv_folds中的原标签转换为X_train对应的位置索引:

# 构建原标签到训练集位置索引的映射
label_to_pos = {label: idx for idx, label in enumerate(X_train.index)}

# 转换cv_folds中的标签为位置索引
cv_folds_pos = [
    ([label_to_pos[label] for label in train_labels], [label_to_pos[label] for label in test_labels])
    for train_labels, test_labels in cv_folds
]

# 使用转换后的交叉验证折调用cross_validate
cv_output = cross_validate(model, X_train, y_train, scoring=['neg_mean_squared_error'], cv=cv_folds_pos)

补充说明

  • 若基于哈希值选择拆分,可以直接针对X_train的位置索引生成交叉验证折,避免后续转换步骤。
  • 测试X_train.loc[train_index_list]正常是因为loc使用标签索引,而cross_validate内部对自定义cv折的处理逻辑是按位置索引取数。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 23:25:31