如何基于独立标签数组按比例拆分数据集且不打乱数据
按标签比例拆分数据集(不打乱顺序)的解决方案
刚好我之前也遇到过类似的需求,要按每个标签类别以7:3比例拆分训练/验证集,同时保留原始数据的顺序。下面给你几个用numpy、pandas实现的方案,也会说明scikit-learn工具的局限性:
方法1:Numpy手动实现(最灵活可控)
这个方法直接操作数组索引,能精准控制每个类别的拆分比例,完全保留原始顺序:
import numpy as np # 模拟你的数据(替换成你实际的dataset和labels) dataset = np.random.rand(18, 6, 10) # 对应你提到的(128,6,-1)形状 labels = np.array([0,0,0,0,1,1,1,1,2,2,2,2,2,2,2,2,2,2]) # 初始化训练/验证索引列表 train_indices = [] eval_indices = [] # 遍历每个唯一标签类别 for label in np.unique(labels): # 获取当前标签的所有样本索引 label_indices = np.where(labels == label)[0] total_samples = len(label_indices) # 计算7:3的拆分点(用ceil保证训练集占比不低于70%,和你的示例匹配) split_idx = int(np.ceil(total_samples * 0.7)) # 前split_idx个样本归训练集,剩下的归验证集 train_indices.extend(label_indices[:split_idx]) eval_indices.extend(label_indices[split_idx:]) # 提取最终的训练/验证数据和标签 train_dataset = dataset[train_indices] train_labels = labels[train_indices] eval_dataset = dataset[eval_indices] eval_labels = labels[eval_indices] # 验证结果(和你的示例完全一致) print("训练标签:", train_labels) print("验证标签:", eval_labels)
运行后输出的结果正好符合你的预期:
训练标签: [0 0 0 1 1 1 2 2 2 2 2 2 2]
验证标签: [0 1 2 2 2]
方法2:Pandas实现(更直观的结构化处理)
如果习惯用DataFrame处理数据,这个方法可读性更强,同样能保留原始顺序:
import pandas as pd import numpy as np # 模拟数据 dataset = np.random.rand(18, 6, 10) labels = np.array([0,0,0,0,1,1,1,1,2,2,2,2,2,2,2,2,2,2]) # 构造DataFrame,将每个样本的数组存入列中 df = pd.DataFrame({ 'label': labels, 'data': list(dataset) # 把3D数组的每个样本转成列表元素 }) # 按标签分组,对每个组拆分7:3 train_df = df.groupby('label').apply( lambda group: group.iloc[:int(np.ceil(len(group)*0.7))] ).reset_index(drop=True) eval_df = df.groupby('label').apply( lambda group: group.iloc[int(np.ceil(len(group)*0.7)):] ).reset_index(drop=True) # 提取数据和标签 train_dataset = np.array(train_df['data'].tolist()) train_labels = train_df['label'].values eval_dataset = np.array(eval_df['data'].tolist()) eval_labels = eval_df['label'].values print("训练标签:", train_labels) print("验证标签:", eval_labels)
关于Scikit-learn工具的说明
你提到的train_test_split虽然好用,但它的问题是:
- 即使设置
shuffle=False,它也是全局按比例拆分,不是按每个标签类别拆分,无法保证每个类别都符合7:3的比例 - 而
StratifiedShuffleSplit虽然能按类别分层,但它默认会打乱数据,不符合你“不打乱”的需求
所以如果必须用sklearn工具,其实不如上面的手动方案直接可控,毕竟逻辑本身很简单,不需要依赖复杂的库函数。
内容的提问来源于stack exchange,提问作者Michael24
相关产品推荐
相关产品推荐

