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

如何在scikit-learn的cross_validate中输出各折测试集样本/索引?

如何获取cross_validate中每折的测试集样本/索引?

sklearn.model_selection.cross_validate本身不会直接返回交叉验证拆分的测试集索引,但可以通过以下两种简单方法实现:

方法一:预存拆分索引后传入cross_validate

  1. 先初始化你要用的交叉验证拆分器(比如KFold、StratifiedKFold等),遍历拆分器的split()方法,提前记录每折的测试集索引。
  2. 将同一个拆分器实例传入cross_validate的cv参数,确保拆分逻辑完全一致。

示例代码:

import numpy as np
from sklearn.model_selection import KFold, cross_validate
from sklearn.linear_model import LogisticRegression

# 生成示例数据
X = np.random.rand(10, 2)
y = np.random.randint(0, 2, size=10)

# 初始化拆分器
kf = KFold(n_splits=5, shuffle=True, random_state=42)

# 存储每折的测试集索引
test_indices_list = []
for _, test_idx in kf.split(X):
    test_indices_list.append(test_idx)

# 运行cross_validate,使用同一个拆分器
model = LogisticRegression()
cv_results = cross_validate(model, X, y, cv=kf)

# 查看结果与对应索引
for fold, (test_idx, score) in enumerate(zip(test_indices_list, cv_results['test_score'])):
    print(f"第{fold+1}折测试集索引: {test_idx}, 测试分数: {score:.4f}")

方法二:自定义拆分器记录索引

如果需要更灵活的记录方式,可以自定义一个拆分器的包装类,在每次拆分时自动保存测试索引:

class IndexRecordingKFold(KFold):
    def __init__(self, *args, **kwargs):
        super().__init__(*args, **kwargs)
        self.test_indices = []
    
    def split(self, X, y=None, groups=None):
        for train_idx, test_idx in super().split(X, y, groups):
            self.test_indices.append(test_idx)
            yield train_idx, test_idx

# 使用自定义拆分器
recording_kf = IndexRecordingKFold(n_splits=5, shuffle=True, random_state=42)
cv_results = cross_validate(model, X, y, cv=recording_kf)

# 直接获取所有测试索引
print("所有折的测试集索引:", recording_kf.test_indices)

两种方法都能保证你拿到cross_validate使用的每折测试集索引,第一种更直观易维护,适合大多数场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 20:25:20