【问题标题】:Matlab split into train/valid/test set and keep proportionMatlab分成训练/有效/测试集并保持比例
【发布时间】:2016-04-17 22:17:33
【问题描述】:

我有 12 列 + 1 个目标(二进制)和大约 4000 行的数据集。我需要将它分成训练集(70%)、验证集(20%)和测试集(10%)。

数据集的采样率很低(0 类的 95% 到 1 类的 5%),所以我需要保持每个样本中目标的比例。

我能够以某种方式拆分数据集,但我不知道如何保持比率。

我正在处理子集葡萄酒质量数据here

【问题讨论】:

  • 查看内置crossvalind函数的文档。您可以(例如)生成 10 个折叠,并将其中 7 个的并集用于训练,2 个用于验证,1 个用于测试。
  • 我看过那个。确实是个好主意。但是crossvalind 是否考虑了目标列中的比例?
  • 我 99% 确定您可以给它这个选项。看看文档..

标签: matlab


【解决方案1】:

如果您可以访问 Matlab 的统计处理工具箱,您可以使用 cvpartition 功能。

来自 cvpartition 上的 matlab 帮助 -:

c = cvpartition(group,'HoldOut',p) 使用组中的类信息将观察随机划分为训练集和测试集,并分层;也就是说,训练集和测试集的类比例与组中的大致相同。

您可以应用该函数两次以获得三个分区。该函数保留了原始的类分布。

【讨论】:

    【解决方案2】:

    到目前为止,我想出了这个,如果有人知道更好的解决方案,请告诉我。 我按目标列拆分数据集,然后将这两个拆分中的每一个进一步拆分为前 70%、下 20% 和最后 10% 的数据,然后合并在一起。 之后,我拆分特征和目标。

    %split in 0/1 samples
    winedataset_0 = winedataset(winedataset(:, 13) == 0, :);
    winedataset_1 = winedataset(winedataset(:, 13) == 1, :);
    
    %train
    split_tr_0 = round(length(winedataset_0)*0.7);
    split_tr_1 = round(length(winedataset_1)*0.7);
    train_0 = winedataset_0(1:split_tr_0,:);
    train_1 = winedataset_1(1:split_tr_1,:);
    train_set = vertcat(train_0, train_1);
    train_set = train_set(randperm(length(train_set)),:);
    
    %valid
    split_valid_0 = split_tr_0 + round(length(winedataset_0)*0.2);
    split_valid_1 = split_tr_1 + round(length(winedataset_1)*0.2);
    valid_0 = winedataset_0(split_tr_0+1:split_valid_0,:);
    valid_1 = winedataset_1(split_tr_1+1:split_valid_1,:);
    valid_set = vertcat(valid_0, valid_1);
    valid_set = valid_set(randperm(length(valid_set)),:);
    
    %test
    test_0 = winedataset_0(split_valid_0+1:end,:);
    test_1 = winedataset_1(split_valid_1+1:end,:);
    test_set = vertcat(test_0, test_1);
    test_set = test_set(randperm(length(test_set)),:);
    
    
    %Split into X and y
    X_train = train_set(:,1:12);
    y_train = train_set(:,13);
    
    X_valid = valid_set(:,1:12);
    y_valid = valid_set(:,13);
    
    X_test = test_set(:,1:12);
    y_test = test_set(:,13);
    

    【讨论】:

      猜你喜欢
      • 2019-03-07
      • 2021-01-28
      • 2017-11-01
      • 1970-01-01
      • 1970-01-01
      • 2018-01-21
      • 2021-03-17
      • 2022-12-24
      • 2016-04-04
      相关资源
      最近更新 更多