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

PyCaret中如何按分组列划分训练/测试集 保证同组数据不拆分

PyCaret 按分组完整划分训练集/测试集实现方案

核心要求:同一group_id对应的所有样本必须完整归入训练集或测试集,禁止同组数据拆分到两个子集。
PyCaret默认的随机拆分逻辑是按行抽样,无法满足组级不拆分的要求,我们可以先手动完成组级拆分,再将拆分结果传入PyCaret即可,全程保证同组数据不泄露。

实现步骤

  • 提取数据集所有唯一的分组ID
  • 使用专用的组抽样工具按比例拆分分组ID,保证同组不跨子集
  • 匹配分组ID得到全量数据对应的训练、测试行索引
  • 将预生成的索引传入PyCaret初始化配置,同时将交叉验证也设置为组级拆分,全流程避免数据泄露

可直接运行的代码示例

import pandas as pd
from sklearn.model_selection import GroupShuffleSplit
# 回归任务导入pycaret.regression的setup,分类任务导入pycaret.classification的setup
from pycaret.regression import setup, compare_models
# 1. 读取/构造数据集,此处用问题给出的样例数据演示
df = pd.DataFrame({
    'group_id': [1,1,1,2,2,3,3,3,4,4],
    'measure1': [3455,6455,6444,23,623,3455,6155,6434,93,693],
    'measure2': [3425,825,225,34,22,3425,525,325,345,222],
    'measure3': [345,945,145,233,888,345,645,845,233,808]
})
# 2. 按组完成8:2拆分,固定随机种子保证结果可复现
gss = GroupShuffleSplit(n_splits=1, test_size=0.2, random_state=42)
train_idx, test_idx = next(gss.split(df, groups=df['group_id']))
# 3. 验证拆分结果(可选,用于确认没有同组拆分问题)
train_groups = set(df.iloc[train_idx]['group_id'].unique())
test_groups = set(df.iloc[test_idx]['group_id'].unique())
print(f"训练集分组ID: {sorted(train_groups)}")
print(f"测试集分组ID: {sorted(test_groups)}")
print(f"是否存在同组拆分: {len(train_groups & test_groups) > 0}")
# 4. 传入PyCaret完成初始化
exp_setup = setup(
    data = df,
    target = 'measure3', # 替换为实际任务的目标列名
    train_size = 0.8,
    data_split_shuffle = False, # 关闭默认行级打乱,避免覆盖自定义拆分
    # 交叉验证环节也配置为组级拆分,全流程避免数据泄露
    fold_strategy = 'groupkfold',
    fold_groups = 'group_id',
    index = True,
    train_indices = list(train_idx),
    test_indices = list(test_idx),
    verbose = False
)
# 后续正常调用PyCaret的建模API即可,比如compare_models()

关键说明

  • 不建议手写unique分组后随机抽样的逻辑,GroupShuffleSplit是经过验证的成熟工具,天然保证同组样本不跨子集,能避免抽样比例偏差、索引匹配错误等手写逻辑容易出现的问题
  • 如果是分类任务需要保证训练/测试集的目标标签分布一致,可以将GroupShuffleSplit替换为StratifiedGroupShuffleSplit,在保证组不拆分的前提下实现分层抽样
  • 必须配置fold_strategy='groupkfold',否则交叉验证阶段还是会按行拆分,出现同组数据同时出现在训练和验证折的情况,导致模型效果评估虚高

样例数据拆分效果

运行上述代码对提供的10行样例数据拆分后,结果完全符合预期:

  • 训练集包含分组ID:1、3、4,对应共8行样本
  • 测试集包含分组ID:2,对应共2行样本
  • 无任何同组ID被拆分到两个子集的情况

内容的提问来源于stack exchange,提问作者Forge

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 14:31:00