【问题标题】:scikit learn: custom classifier compatible with GridSearchCVscikit learn:与 GridSearchCV 兼容的自定义分类器
【发布时间】:2018-06-21 01:11:57
【问题描述】:

我已经实现了自己的分类器,现在我想对其进行网格搜索,但出现以下错误:estimator.fit(X_train, y_train, **fit_params) TypeError: fit() takes 2 positional arguments but 3 were given

我关注了this tutorial,并使用了scikit's official documentation提供的this template。我的类定义如下:

class MyClassifier(BaseEstimator, ClassifierMixin):
    def __init__(self, lr=0.1):
        self.lr=lr

    def fit(self, X, y):
        # Some code
        return self
    def predict(self, X):
        # Some code
        return y_pred
    def get_params(self, deep=True)
        return {'lr'=self.lr}
    def set_params(self, **parameters):
        for parameter, value in parameters.items():
            setattr(self, parameter, value)
        return self

我正在尝试网格搜索,如下所示:

params = {
    'lr': [0.1, 0.5, 0.7]
}
gs = GridSearchCV(MyClassifier(), param_grid=params, cv=4)

编辑我

我是这样称呼它的: gs.fit(['hello world', 'trying','hello world', 'trying', 'hello world', 'trying', 'hello world', 'trying'], ['I', 'Z', 'I', 'Z', 'I', 'Z', 'I', 'Z'])

结束编辑我

错误是由文件python3.5/site-packages/sklearn/model_selection/_validation.py中的_fit_and_score方法产生的

它用 3 个参数调用 estimator.fit(X_train, y_train, **fit_params),但我的估算器只有两个,所以这个错误对我来说是有意义的,但我不知道如何解决它......我还尝试向 @ 添加一些虚拟参数987654330@ 方法,但是没有用。

编辑二

完整的错误输出:

Traceback (most recent call last):
  File "/home/rodrigo/no_version/text_classifier/MyClassifier.py", line 355, in <module>
    ['I', 'Z', 'I', 'Z', 'I', 'Z', 'I', 'Z'])
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/model_selection/_search.py", line 639, in fit
    cv.split(X, y, groups)))
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/parallel.py", line 779, in __call__
    while self.dispatch_one_batch(iterator):
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/parallel.py", line 625, in dispatch_one_batch
    self._dispatch(tasks)
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/parallel.py", line 588, in _dispatch
    job = self._backend.apply_async(batch, callback=cb)
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/_parallel_backends.py", line 111, in apply_async
    result = ImmediateResult(func)
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/_parallel_backends.py", line 332, in __init__
    self.results = batch()
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/parallel.py", line 131, in __call__
    return [func(*args, **kwargs) for func, args, kwargs in self.items]
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/externals/joblib/parallel.py", line 131, in <listcomp>
    return [func(*args, **kwargs) for func, args, kwargs in self.items]
  File "/home/rodrigo/no_version/text_classifier/.env/lib/python3.5/site-packages/sklearn/model_selection/_validation.py", line 458, in _fit_and_score
    estimator.fit(X_train, y_train, **fit_params)
TypeError: fit() takes 2 positional arguments but 3 were given

结束编辑 II

已解决 谢谢大家,我犯了一个愚蠢的错误:有两个不同的函数具有相同的名称(fit),(我用不同的参数实现了另一个用于自定义目的,只要我重命名我的'custom fit',它就可以正常工作。)

谢谢你,对不起

【问题讨论】:

  • 如何调用网格搜索拟合方法?还要确保从 fit 方法返回 self。您的示例不会重现错误。
  • 是的,我正在返回适合自己的方法
  • 您对gs.fit的确切呼叫是什么?
  • 你也可以包括确切的错误输出
  • 这里:完成错误输出:

标签: python machine-learning scikit-learn


【解决方案1】:

以下代码适用于我:

class MyClassifier(BaseEstimator, ClassifierMixin):
     def __init__(self, lr=0.1):
         self.lr = lr
         # Some code
         pass
     def fit(self, X, y):
         # Some code
         pass
     def predict(self, X):
         # Some code
         return X % 3

params = {
    'lr': [0.1, 0.5, 0.7]
}
gs = GridSearchCV(MyClassifier(), param_grid=params, cv=4)

x = np.arange(30)
y = np.concatenate((np.zeros(10), np.ones(10), np.ones(10) * 2))
gs.fit(x, y)

我能想到的最好的结果是您将某些内容传递给 gs.fit 方法,超出了 x 和 y 或您的 MyClassifier.fit 方法缺少 self 参数。

只有在将 kwarg 传递给 gs.fit 方法时才应填充 fit_params kwargs,否则它是一个空字典 ({}) 并且 **fit_params 不会引发参数错误。要对此进行测试,请创建一个分类器实例并传递**{}。例如:

clf = MyClassifier()
clf.fit(x, y, **{})

这不会引发位置参数错误。

因此,除非有东西被传递给gs.fit,例如gs.fit(x, y, some_arg=123) 在我看来,您缺少MyClassifier.fit 定义中的位置参数之一。您包含的错误消息似乎支持这一假设,因为它指出fit() takes 2 positional arguments but 3 were given。如果您将 fit 定义如下,则需要 3 个位置参数:

def fit(self, X, y): ...

【讨论】:

  • 感谢您对 fit_params 的解释!我已经解决了我的问题,这是一个愚蠢的错误,对不起(我也编辑了我的原始帖子)再次感谢您!
【解决方案2】:

看起来像是一些自定义参数的直通。只需为您的 fit-Method 添加一个包罗万象的关键字参数:

def fit(self, X, y, **_k):
    ...

【讨论】:

    猜你喜欢
    • 2018-05-14
    • 2023-04-09
    • 2019-08-18
    • 2017-03-14
    • 2014-06-27
    • 2017-01-21
    • 1970-01-01
    • 2014-08-09
    相关资源
    最近更新 更多