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

按组随机拆分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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.20 07:21:49