按组随机拆分DataFrame:避免训练测试集共享group_id的优雅方案
按Group ID拆分训练/测试集的优雅方案
这个问题我之前处理带分组的数据集时也踩过坑——直接按行随机拆分确实会出现同一group_id跨集的情况,事后修正不仅麻烦还容易出错。其实最优雅的思路是先对group_id本身做随机拆分,再根据拆分后的组筛选原DataFrame的行,从根源上杜绝跨集的group_id问题。
具体实现方法
这里提供两种常用的实现方式,你可以根据自己的环境选择:
方法1:用sklearn的train_test_split(推荐)
sklearn的拆分工具可以直接对唯一的group_id进行拆分,代码简洁还支持可复现:
import pandas as pd from sklearn.model_selection import train_test_split # 假设你的数据集是df,分组列名为'group_id' unique_groups = df['group_id'].unique() # 拆分group_id,测试集占比可根据需求调整,random_state保证结果可复现 train_groups, test_groups = train_test_split(unique_groups, test_size=0.2, random_state=42) # 根据拆分后的group_id筛选得到训练集和测试集 train_df = df[df['group_id'].isin(train_groups)] test_df = df[df['group_id'].isin(test_groups)]
方法2:纯numpy实现(无需额外依赖)
如果你的环境没有安装sklearn,用numpy也能轻松实现:
import pandas as pd import numpy as np unique_groups = df['group_id'].unique() np.random.seed(42) # 设置随机种子保证可复现 # 打乱group_id顺序后拆分 shuffled_groups = np.random.permutation(unique_groups) split_point = int(len(shuffled_groups) * 0.8) # 80%作为训练集 train_groups = shuffled_groups[:split_point] test_groups = shuffled_groups[split_point:] train_df = df[df['group_id'].isin(train_groups)] test_df = df[df['group_id'].isin(test_groups)]
为什么这个方法更优?
- 逻辑清晰:从分组层面拆分,天然保证同一group_id不会同时出现在训练集和测试集里,不需要事后校验修正
- 效率更高:尤其是当数据集很大、分组很多时,直接操作group_id的效率远高于逐行处理再修正
- 可复现性强:通过设置随机种子,每次拆分的结果都一致,方便后续调试和验证
内容的提问来源于stack exchange,提问作者dozyaustin
相关产品推荐
相关产品推荐

