如何用sklearn.model_selection.GroupShuffleSplit实现按组划分训练测试集
用GroupShuffleSplit实现按产品/类型分组的训练测试集拆分
嘿,这个需求我之前处理过类似的——要保证同一产品的所有数据不跨训练/测试集,GroupShuffleSplit绝对是最适合的工具,它专门解决这种分组级别的数据拆分问题,完美匹配你的场景。
完整步骤&代码示例
假设你已经把CSV读入了Pandas DataFrame,且IMAGE_ID是索引列,咱们一步步来:
- 先导入需要的库
import pandas as pd from sklearn.model_selection import GroupShuffleSplit
- 初始化分组拆分器
我们先设置好测试集比例(比如20%),还有拆分次数(如果只需要一组训练测试集,n_splits设为1就行):
# 初始化拆分器:测试集占20%,固定随机种子保证结果可复现 gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42)
- 执行核心拆分操作
这里的关键是把PRODUCT_ID作为分组依据传入groups参数——这样就能保证同一个PRODUCT_ID对应的所有行,要么全进训练集,要么全进测试集:
# 获取训练集和测试集的行索引(因为split返回迭代器,用next取第一次拆分结果) train_idx, test_idx = next(gss.split(df, groups=df['PRODUCT_ID'])) # 根据索引生成训练集和测试集 train_df = df.iloc[train_idx] test_df = df.iloc[test_idx]
- 验证拆分是否符合要求
为了确保没有产品同时出现在两个集合里,咱们可以做个简单的检查:
train_unique_products = set(train_df['PRODUCT_ID'].unique()) test_unique_products = set(test_df['PRODUCT_ID'].unique()) print(f"训练集包含 {len(train_unique_products)} 个唯一产品") print(f"测试集包含 {len(test_unique_products)} 个唯一产品") print(f"两个集合是否有重叠产品:{'是' if len(train_unique_products & test_unique_products) > 0 else '否'}")
正常情况下最后一行应该输出否,说明拆分完全符合你的要求。
扩展:按其他列(比如PRODUCT_TYPE)分组拆分
如果你的需求是按PRODUCT_TYPE来分组(即同一类型的所有产品要么全在训练集,要么全在测试集),只需要把groups参数换成df['PRODUCT_TYPE']就行:
train_idx, test_idx = next(gss.split(df, groups=df['PRODUCT_TYPE']))
原理完全一致,只是分组的粒度从单个产品变成了产品类型。
几个小提醒
- 虽然你把
IMAGE_ID设为了索引,但拆分时我们用的是DataFrame的原始行位置索引(iloc),这完全没问题,因为分组的依据是PRODUCT_ID,和IMAGE_ID的索引无关。 - 如果需要多次拆分(比如做交叉验证),只需要把
n_splits设成你需要的次数,然后循环迭代拆分结果就行。
内容的提问来源于stack exchange,提问作者Aravind Chamakura
相关产品推荐
相关产品推荐

