You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用sklearn.model_selection.GroupShuffleSplit实现按组划分训练测试集

用GroupShuffleSplit实现按产品/类型分组的训练测试集拆分

嘿,这个需求我之前处理过类似的——要保证同一产品的所有数据不跨训练/测试集,GroupShuffleSplit绝对是最适合的工具,它专门解决这种分组级别的数据拆分问题,完美匹配你的场景。

完整步骤&代码示例

假设你已经把CSV读入了Pandas DataFrame,且IMAGE_ID是索引列,咱们一步步来:

  1. 先导入需要的库
import pandas as pd
from sklearn.model_selection import GroupShuffleSplit
  1. 初始化分组拆分器
    我们先设置好测试集比例(比如20%),还有拆分次数(如果只需要一组训练测试集,n_splits设为1就行):
# 初始化拆分器:测试集占20%,固定随机种子保证结果可复现
gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42)
  1. 执行核心拆分操作
    这里的关键是把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]
  1. 验证拆分是否符合要求
    为了确保没有产品同时出现在两个集合里,咱们可以做个简单的检查:
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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.15 08:35:17