必须自己写 split 方法当内置 stratifiedkfold、timeseriessplit 无法满足组约束、时间顺序或外部元数据划分需求时,需返回符合接口的 (train_index, test_index) 元组,类型为 int 型一维 numpy.ndarray,且不可重叠或为空。

什么时候必须自己写 split 方法?
当内置的 StratifiedKFold、TimeSeriesSplit 无法满足数据结构约束时,比如:样本间存在不可分割的组(如同一患者的多次测量)、时间依赖不能打乱顺序、或需按外部元数据(如实验批次)划分训练/验证集。这时 sklearn.model_selection.cross_val_score 或 GridSearchCV 的 cv 参数必须接收一个可迭代对象,每次 yield 一对 (train_index, test_index) 数组。
自定义分割器必须满足的接口要求
它得是可迭代的,且每次产出两个一维 numpy.ndarray,类型为 int,内容是原始数据索引。不能返回布尔掩码、DataFrame 索引或带重复/越界值的数组。常见错误包括:
- 忘记用
np.array包装索引列表,导致ValueError: Found array with dim 0. Expected >= 2 - yield 的
test_index为空或与train_index重叠 - 没实现
__len__,导致cross_val_score无法预估折数而报TypeError: object of type 'XXX' has no len()
推荐直接继承 sklearn.model_selection.KFold 并重写 _iter_test_indices,或更简单地定义一个生成器函数——只要它能被 list() 转成折列表即可。
一个按 patient_id 分组的示例
假设你有 X 和 y,还有一个 patient_id 列表,要求同一患者的所有样本只能出现在训练集或测试集之一:
from sklearn.model_selection import cross_val_score from sklearn.ensemble import RandomForestClassifier import numpy as np <p>def group_kfold_split(patient_ids, n_splits=3): unique_patients = np.unique(patient_ids) np.random.shuffle(unique_patients) # 可选,但建议打乱 fold_size = len(unique_patients) // n_splits for i in range(n_splits): val_patients = unique_patients[i<em>fold_size:(i+1)</em>fold_size] train_mask = ~np.isin(patient_ids, val_patients) test_mask = np.isin(patient_ids, val_patients) yield np.where(train_mask)[0], np.where(test_mask)[0]</p><h1>使用</h1><p>cv_splits = list(group_kfold_split(patient_ids=patient_ids, n_splits=3)) scores = cross_val_score(RandomForestClassifier(), X, y, cv=cv_splits) </p>
注意:这里用了 list(...) 强制展开生成器,因为 cross_val_score 内部会多次遍历 cv;若直接传生成器,第二轮遍历时将为空。
和 GroupKFold 的关键区别在哪?
GroupKFold 确实也按组划分,但它保证每组只出现在一个测试折中,且各折样本量尽量均衡——这在组大小差异大时会导致某些折的样本极少。而手写分割器可以加逻辑:比如跳过样本数 GroupKFold 不提供钩子。
真正容易被忽略的是:自定义分割器不参与 GridSearchCV 的参数缓存机制,每次调用都会重新计算分折;如果分折逻辑耗时,应提前算好并存为列表,别传生成器或类实例。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!











