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

如何拆分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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 13:54:06