Sklearn GroupShuffleSplit拆分后验证测试集出现组重叠问题
问题诱因
- 核心错误是混淆了两类索引:
GroupShuffleSplit.split()返回的val_inds、test_inds是传入数据集的整数位置下标(从0开始计数,对应传入df的第N行),不是DataFrame自带的index标签值。 - 之前功能运行正常,是因为当时使用的df保留了pandas默认的从0开始连续递增的RangeIndex,位置下标和标签索引值完全相等,用
df.index.isin()匹配刚好能选对行;一旦df经过行过滤、排序、拼接、去重等操作,导致索引不再连续/不再从0起始,这个匹配逻辑就会选到错误的行,拆分结果自然出现cl_uid分组跨集合重叠的问题。 - 代码中
n_splits=2的参数设置不会导致分组泄漏,该参数仅控制拆分器生成多少组独立的拆分结果,调用next()取第一组结果的写法本身没有问题。sklearn 0.24.1版本的GroupShuffleSplit实现逻辑无原生bug,只要行匹配正确,拆分出的两个集合不会出现同组样本重叠。
修复方案
直接使用iloc按位置下标选行,不要用标签索引做匹配,建议同时固定随机种子保证结果可复现:
val_inds, test_inds = next(GroupShuffleSplit( test_size=0.5, n_splits=2, random_state=42 ).split(df, groups=df['cl_uid'].values)) # 按位置下标选行,严格匹配split的划分结果 df_val = df.iloc[val_inds] df_test = df.iloc[test_inds] # 此时两个集合的cl_uid交集长度必然为0 len(set(df_val.cl_uid).intersection(set(df_test.cl_uid)))
验证提示:修复后可以先打印两个集合的cl_uid交集确认结果,只要groups参数和传入split的df行顺序一一对应,就不会出现分组泄漏问题。
内容的提问来源于stack exchange,提问作者Giovanni Rescia
相关产品推荐
相关产品推荐

