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

类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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 17:57:35