如何在scikit-learn的cross_validate中输出各折测试集样本/索引?
如何获取cross_validate中每折的测试集样本/索引?
sklearn.model_selection.cross_validate本身不会直接返回交叉验证拆分的测试集索引,但可以通过以下两种简单方法实现:
方法一:预存拆分索引后传入cross_validate
- 先初始化你要用的交叉验证拆分器(比如
KFold、StratifiedKFold等),遍历拆分器的split()方法,提前记录每折的测试集索引。 - 将同一个拆分器实例传入
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
相关产品推荐
相关产品推荐

