【问题标题】:Selecting a Subset of Features in scikit-learn for Training在 scikit-learn 中选择特征子集进行训练
【发布时间】:2019-10-22 06:25:51
【问题描述】:

假设我有一个包含 5 个特征的数据集,并且我想使用特征 1、2 和 5 进行训练(跳过特征 3 和 4)。我不想更改数据集,因为我希望在预测期间将相同的 5 个特征提供给模型。我只希望预处理管道的第一步删除功能 3 和 4。

此外,我希望能够在训练结束时根据要加载和运行的任何其他对象或代码来腌制/joblib 管道对象,而无需腌制对象。因此,我不想使用FunctionTransformer,因为我必须编写一个自定义函数(传递给这个转换器),然后腌制并与腌制的模型对象一起发送。

在 scikit-learn 中有什么好的方法吗?

【问题讨论】:

    标签: python scikit-learn


    【解决方案1】:

    您可以创建自己的转换器对象来为您执行列选择。当您将要提取的列放入管道时,您会将其作为参数传递。通过在您的管道中,它将与您的其余步骤一起被腌制。

    为了包含这个自定义转换器,您的类需要从两个基本 sklearn 类继承:TransformerMixinBaseEstimator。只要您自己定义fittransform,从TransformerMixin 继承即可为您提供fit_transform 方法。从BaseEstimator 继承提供get_paramsset_params。由于 fit 方法除了返回对象本身不需要做任何事情,您真正需要做的就是定义 transform 方法。

    这是一个示例,假设您的数据 (X) 是 pandas DataFrame,您可以传入要提取的列名列表。

    from sklearn.base import BaseEstimator, TransformerMixin
    
    
    class FeatureSelector(BaseEstimator, TransformerMixin):
    
        def __init__(self, feature_names):
            self._feature_names = feature_names 
    
        def fit(self, X, y = None):
            return self 
    
        def transform(self, X, y = None):
            return X[self._feature_names]
    

    现在您已经有了转换器,您可以将它包含在您的管道中,它可以按照您的要求进行腌制。

    至于您不使用FunctionTransformer 的要求,我假设您看到了示例here,他们在其中全局定义了all_but_first_column。使用上面定义的FeatureSelector 类,您总是可以将all_but_first_column 之类的东西作为另一种方法移动到该类中。

    【讨论】:

    • 感谢您的解决方案。然而问题是我需要避免编写任何自定义 Python 代码并完全使用 scikit-learn 的库来实现这一点(因为如果你腌制管道对象,joblib 或 pickle 将无法正确腌制自定义代码。请参阅:github.com/scikit-learn/scikit-learn/issues/12903)。我也认为 scikit-learn 绝对应该在他们的原生库中添加一个像你这样的转换器(根据他们的名字或他们的位置选择特征)。
    • 我认为 [issue](github.com/scikit-learn/scikit-learn/issues/12903) 中提出的根本问题可以通过使用dill 腌制您的 sklearn 管道来解决。我亲自完成了,dill 作者甚至有一个 SO 帖子,概述了 here 的优缺点。如果您给我一些时间,我可以使用 iris 之类的玩具数据集通过端到端示例更新我的答案。
    【解决方案2】:

    为了将来参考,有一个解决方法可以使用包mlxtend 中的feature_selection.ColumnSelector 来执行此任务。它采用以下要选择的列的索引:

    from mlxtend.feature_selection import ColumnSelector
    
    ...
    
    pipeline = Pipeline(steps=[
        ('selector', ColumnSelector([1,2,3])),
        ('kmeans', KMeans()), 
    ])
    
    ...
    

    有关更多信息,请参阅docs

    【讨论】:

      猜你喜欢
      • 2015-08-10
      • 2014-11-05
      • 2018-02-24
      • 2021-03-26
      • 2018-06-01
      • 1970-01-01
      • 2014-05-22
      • 1970-01-01
      相关资源
      最近更新 更多