如何拆分pandas DataFrame并保证两个子集均包含分类列全部标签
数据集分层拆分解决方案
你需要的是分层抽样能力,保证拆分后的两个子集都包含所有分类列的全量标签,以下两种方法都可以高效实现:
方法1:使用sklearn原生分层参数(推荐,效率最优)
train_test_split本身就支持stratify参数,指定分层依据的列后,会自动按类别比例拆分,保证每个子集都包含所有类别标签:
import pandas as pd from sklearn.model_selection import train_test_split # 构造示例数据 data = { 'Fruits': ['Banana','Grape','Apple','Papaya','Dragon','Mango','Banana','Grape','Apple','Papaya','Dragon','Mango'], 'Color': ['Yellow','Black','Red','Yellow','Pink','Yellow','Yellow','Black','Red','Yellow','Pink','Yellow'], 'Price': [60,100,200,50,150,400,75,106,190,60,120,390] } df = pd.DataFrame(data) # 单个分类列分层:按Fruits列分层,50:50拆分 df1, df2 = train_test_split(df, test_size=0.5, random_state=42, stratify=df['Fruits']) # 多分类列同时保证标签完整:生成联合分层键即可 df['stratify_key'] = df['Fruits'] + '_' + df['Color'] df1, df2 = train_test_split(df, test_size=0.5, random_state=42, stratify=df['stratify_key']) # 拆分后删除临时生成的分层键 df1.drop(columns=['stratify_key'], inplace=True) df2.drop(columns=['stratify_key'], inplace=True)
方法2:纯pandas分组采样(无额外依赖)
如果不想引入sklearn依赖,可以直接用pandas的分组采样实现:
# 按Fruits列分组,每组抽取50%样本作为df1,剩余作为df2 df1 = df.groupby('Fruits', group_keys=False).apply(lambda x: x.sample(frac=0.5, random_state=42)) df2 = df.drop(df1.index)
注意事项
- 以上方法生效的前提是,每个分类标签对应的样本数至少为2,否则无法实现两个子集都包含该标签,可根据业务场景对小样本类别做手动补充
- 可通过
df1['Fruits'].unique()和df2['Fruits'].unique()验证拆分后的标签完整性
内容的提问来源于stack exchange,提问作者swarna
相关产品推荐
相关产品推荐

