类R stratified实现:拆分训练测试集自动将单样本归入训练集
问题背景
数据集中的分类分层变量存在仅含1条观测值的类别时,不同工具的分层拆分行为存在差异:
- R语言
stratified函数按7:3比例拆分训练集、测试集时,单样本类别对应的观测会自动归入训练集,可正常完成拆分 - Python中调用
sklearn.model_selection.train_test_split传入相同分层参数时,会抛出ValueError,无法自动分配单样本行到训练集
R端可复现的正常运行代码
dataset = data.frame(target = c(100,200,300), Var1 = c("a","b","b")) split <- stratified(dataset,c("Var1"), 0.70, keep.rownames=TRUE, bothSets=TRUE) train <- split$SAMP1 train # rn target Var1 #1: 1 100 a #2: 3 300 b test <- split$SAMP2 test # rn target Var1 #1: 2 200 b
Python端对应实现及报错
import pandas as pd from sklearn.model_selection import train_test_split data = [[100, 'a'], [200, 'b'], [300, 'b']] dataset = pd.DataFrame(data, columns=['Target', 'Var1']) train, test = train_test_split(dataset, stratify = dataset[['Var1']], train_size = 0.7)
运行抛出如下错误:
ValueError: The least populated class in y has only 1 member, which is too few. The minimum number of groups for any class cannot be less than 2.
实现方法
提前将所有样本量小于2、无法跨集合拆分的稀有类样本全部划入训练集,剩余满足分层要求的样本再调用train_test_split按指定比例做分层拆分,即可完全对齐R中stratified函数的处理效果,实现代码如下:
import pandas as pd from sklearn.model_selection import train_test_split def stratified_split(df, stratify_cols, train_size=0.7, random_state=None): # 生成分层组合key,统计每组样本量 stratify_key = df[stratify_cols].astype(str).agg('-'.join, axis=1) group_count = stratify_key.value_counts() # 提取单样本稀有组,全部归入训练集 rare_group_keys = group_count[group_count < 2].index rare_mask = stratify_key.isin(rare_group_keys) train_part_rare = df[rare_mask].reset_index(drop=True) df_rest = df[~rare_mask].reset_index(drop=True) # 剩余可拆分样本执行常规分层拆分 if len(df_rest) > 0: train_part_rest, test_part = train_test_split( df_rest, stratify=df_rest[stratify_cols], train_size=train_size, random_state=random_state ) train_set = pd.concat([train_part_rare, train_part_rest], ignore_index=True) test_set = test_part.reset_index(drop=True) else: # 所有组均为单样本时,全部划入训练集,测试集为空 train_set = train_part_rare test_set = pd.DataFrame(columns=df.columns) return train_set, test_set # 调用测试 data = [[100, 'a'], [200, 'b'], [300, 'b']] dataset = pd.DataFrame(data, columns=['Target', 'Var1']) train, test = stratified_split(dataset, stratify_cols=['Var1'], train_size=0.7, random_state=42)
效果说明
- 单样本类别观测全部进入训练集,和R的
stratified默认处理规则一致 - 样本量≥2的类别严格按照指定比例分层抽样,训练、测试集的类别分布保持一致
- 支持传入多个分层列,兼容原生
train_test_split的分层逻辑
内容的提问来源于stack exchange,提问作者César Macieira
相关产品推荐
相关产品推荐

