【问题标题】:Downsampling for more than 2 classes2 类以上的下采样
【发布时间】:2019-03-12 10:48:48
【问题描述】:

我正在创建一个简单的代码,当您的目标变量具有 2 个以上的类时,它允许对数据帧进行下采样。

df 是我们的任意数据集,'TARGET_VAR' 是一个具有 2 个以上类别的分类变量。

import pandas as pd
label='TARGET_VAR' #define the target variable

num_class=df[label].value_counts() #creates list with the count of each class value
temp=pd.DataFrame() #create empty dataframe to be filled up

for cl in num_class.index: #loop through classes
    #iteratively downsample every class according to the smallest
    #class 'min(num_class)' and append it to the dataframe.
    temp=temp.append(df[df[label]==cl].sample(min(num_class)))

df=temp #redefine initial dataframe as the subsample one

del temp, num_class #delete temporary dataframe

现在我想知道,有没有办法以更精致的方式做到这一点?例如无需创建临时数据集? 我试图找出一种方法来“矢量化”多个类的操作,但没有得到任何结果。以下是我的想法,可以轻松实现 2 个类,但我不知道如何将其扩展到多个类的情况。

如果您有 2 个类,这将非常有效

 df= pd.concat([df[df[label]==num_class.idxmin()],\
 df[df[label]!=num_class.idxmin()].sample(min(num_class))])

这允许您为其他类选择正确数量的观察值,但这些类不一定具有同等的代表性。

 df1= pd.concat([df[df[label]==num_class.idxmin()],\
 df[df[label]!=num_class.idxmin()].sample(min(num_class)*(len(num_class)-1))])

【问题讨论】:

    标签: python pandas downsampling


    【解决方案1】:

    您可以尝试类似的方法:

    label='TARGET_VAR'
    
    g = df.groupby(label, group_keys=False)
    balanced_df = pd.DataFrame(g.apply(lambda x: x.sample(g.size().min()))).reset_index(drop=True)
    

    我相信这会产生您想要的结果,请随时提出任何其他问题。

    编辑

    根据 OP 的建议修复了代码。

    【讨论】:

    • 您好 Gustavo,感谢您的回答!我不知道 groupby 方法。它真的很强大!但是,这样我仍然在创建一个临时对象,对它的属性有什么想法吗?如果原始数据框df 变得非常大,您知道是否会出现问题(主要是内存方面)?
    • 我相信它的性能与您之前的方法相似,只是更简洁。但是请随意模拟一个大数据框并对这两种方法进行基准测试。选择更适合您的。
    【解决方案2】:

    此代码用于少数类的oversampling 实例或多数类的undersampling 实例。它应该只在训练集上使用。注意:activity 是标签

    balanced_df=Pdf_train.groupby('activity',as_index = False,group_keys=False).apply(lambda s: s.sample(100,replace=True))
    

    【讨论】:

      【解决方案3】:

      Gustavo 的答案是正确的,但有一个小问题(出于某种原因,我无法编辑他的答案)。

      label='TARGET_VAR'
      
      g = df.groupby(label, group_keys=False)
      balanced_df = pd.DataFrame(g.apply(lambda x: 
      x.sample(g.size().min()).
      reset_index(drop=True)))
      

      此处将为每个组重置索引,最终数据帧将具有重复的行索引。如果我们定义少数类的元素个数为n

      idx, data 
      0,   ...
      1,   ...
      .,   ...
      .,   ...
      .,   ...
      n,   ...
      0,   ...
      1,   ...
      .,   ...
      .,   ...
      .,   ...
      n,   ...
      

      以下调整将解决问题

      g = df.groupby(label, group_keys=False)
      balanced_df = pd.DataFrame(g.apply(lambda x: 
      x.sample(g.size().min()))).reset_index(drop=True)
      

      如果我们现在将balanced_df 的元素总数定义为N=n*kk 是不同类的数量。索引将如下所示:

      idx, data 
      0,   ...
      1,   ...
      .,   ...
      .,   ...
      .,   ...
      N,   ...
      

      【讨论】:

        猜你喜欢
        • 1970-01-01
        • 1970-01-01
        • 2015-12-27
        • 1970-01-01
        • 1970-01-01
        • 2020-08-25
        • 2014-05-04
        • 2016-11-04
        • 2016-02-14
        相关资源
        最近更新 更多