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

多输出多类别数据集按比例拆分训练集的确定性方案需求

确定性拆分多输出多类别数据集的解决方案

核心思路

针对你的需求,我们通过先锁定每个目标列Class 1的样本选取范围,再采用交集优先、补选独有样本的确定性逻辑拆分数据集,完全规避随机循环的不确定性。

关键前提:先验证原数据集中每个目标列的Class 1样本数是否≥原总行数的15%,如果某列Class 1占比不足15%,需求无法满足,需提前提示。

步骤拆解

  1. 可行性校验:计算原数据集总行数,统计每个目标列Class 1的样本数,确保每列Class 1数量≥0.15×总行数。
  2. 确定每列Class 1选取数量:对每个目标列,选取的Class 1样本数取「原Class 1总数」和「0.3×总行数」的较小值,同时保证不低于0.15×总行数。
  3. 确定性选取Class 1样本:
    • 先提取所有目标列均为Class 1的共同样本,优先将这些样本纳入训练集(最大化样本复用)。
    • 对每个目标列,若共同样本不足以满足选取数量,从该列独有的Class 1样本(仅该列是1,其他列非1)中按索引顺序补选。
  4. 补充非Class 1样本:从非Class 1样本中按确定性规则(如索引顺序)选取任意数量(或全部)补充到训练集,不影响Class 1占比要求。
  5. 拆分数据集:训练集为选中的Class 1样本+补充的非Class 1样本,剩余为测试集。

示例代码与数据集样例

1. 构造样例数据集

import pandas as pd
import numpy as np

# 设置随机种子保证数据集可复现
np.random.seed(42)
N = 1000  # 原数据集总行数

# 构造特征列
features = pd.DataFrame({
    'f1': np.random.randn(N),
    'f2': np.random.randint(0, 10, size=N)
})

# 构造3个目标列(取值0/1/2)
targets = pd.DataFrame({
    'y1': np.random.choice([0,1,2], size=N, p=[0.5, 0.3, 0.2]),
    'y2': np.random.choice([0,1,2], size=N, p=[0.6, 0.25, 0.15]),
    'y3': np.random.choice([0,1,2], size=N, p=[0.4, 0.35, 0.25])
})

# 合并特征与目标
df = pd.concat([features, targets], axis=1)

2. 确定性拆分实现

def deterministic_split(df, target_cols, min_ratio=0.15, max_ratio=0.3):
    N = len(df)
    min_count = int(np.ceil(min_ratio * N))
    max_count = int(np.floor(max_ratio * N))
    
    # 1. 可行性检查
    for col in target_cols:
        class1_count = df[col].value_counts().get(1, 0)
        if class1_count < min_count:
            raise ValueError(f"目标列{col}的Class 1样本数({class1_count})低于最低要求({min_count}),无法满足需求")
    
    # 2. 确定每列需要选取的Class 1数量
    select_counts = {}
    for col in target_cols:
        class1_count = df[col].value_counts().get(1, 0)
        select_counts[col] = min(class1_count, max_count)
        # 兜底确保不低于最小值
        select_counts[col] = max(select_counts[col], min_count)
    
    # 3. 获取各目标列的Class 1样本索引
    class1_samples = {}
    for col in target_cols:
        class1_samples[col] = df[df[col] == 1].index.tolist()
    
    # 4. 提取所有目标列均为1的共同样本索引
    common_idx = set(class1_samples[target_cols[0]])
    for col in target_cols[1:]:
        common_idx.intersection_update(set(class1_samples[col]))
    common_idx = list(common_idx)
    
    # 5. 为每个目标列分配样本:先取共同样本,再补独有样本
    selected_idx = set()
    for col in target_cols:
        # 已选的该列样本数
        current_selected = len([idx for idx in selected_idx if idx in class1_samples[col]])
        need = select_counts[col] - current_selected
        
        if need <= 0:
            continue
        
        # 先从共同样本中取
        take_from_common = min(need, len(common_idx))
        if take_from_common > 0:
            add_idx = common_idx[:take_from_common]
            selected_idx.update(add_idx)
            # 移除已取的共同样本,避免重复分配
            common_idx = common_idx[take_from_common:]
            need -= take_from_common
        
        # 从该列独有样本中补选
        if need > 0:
            other_class1 = set()
            for other_col in target_cols:
                if other_col != col:
                    other_class1.update(class1_samples[other_col])
            unique_class1 = [idx for idx in class1_samples[col] if idx not in other_class1]
            add_idx = unique_class1[:need]
            selected_idx.update(add_idx)
    
    # 6. 补充非Class1样本(按索引顺序取前200个,可自定义数量)
    non_class1_idx = df.index.difference(selected_idx)
    add_non_class1 = non_class1_idx[:200]
    selected_idx.update(add_non_class1)
    
    # 7. 拆分训练集和测试集
    train_df = df.loc[selected_idx].copy()
    test_df = df.loc[df.index.difference(selected_idx)].copy()
    
    # 验证每个目标列的Class1占比
    print("训练集各目标列Class1占原总行数的比例:")
    for col in target_cols:
        cnt = train_df[col].value_counts().get(1, 0)
        ratio = cnt / N
        print(f"{col}: {ratio:.4f} (样本数: {cnt})")
    
    return train_df, test_df

# 执行拆分
target_cols = ['y1', 'y2', 'y3']
train_df, test_df = deterministic_split(df, target_cols)

代码说明

  • 所有样本选取均按索引顺序执行,完全确定性,无随机因素(若需随机选取非Class1样本,可设置固定随机种子)。
  • 输出会打印训练集每个目标列Class1占原总行数的比例,验证是否符合0.15-0.3的要求。
  • 若需调整非Class1样本的数量,修改add_non_class1 = non_class1_idx[:200]中的数字即可。

内容的提问来源于stack exchange,提问作者Caesar

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 08:14:54