【问题标题】:How to use a user defined metric for nearest neighbors in scikit-learn?如何在 scikit-learn 中使用用户定义的最近邻度量?
【发布时间】:2016-06-01 19:11:03
【问题描述】:

我正在使用 scikit-learn 0.18.dev0。我知道在here 之前也有人问过同样的问题。我尝试了那里提供的答案,我收到以下错误

>>> def mydist(x, y):
...     return np.sum((x-y)**2)
...
>>> X = np.array([[-1, -1], [-2, -1], [-3, -2], [1, 1], [2, 1], [3,   2]])

>>> nbrs = NearestNeighbors(n_neighbors=4, algorithm='ball_tree',
...            metric='pyfunc', func=mydist)

错误信息 _init_params() got an unexpected keyword argument 'func'

此选项似乎已被删除。如何在sklearn.neighbors 中使用用户定义的矩阵?

【问题讨论】:

    标签: scikit-learn distance metrics nearest-neighbor


    【解决方案1】:

    正确的关键字是metric:

    import numpy as np
    from sklearn.neighbors import NearestNeighbors
    
    def mydist(x, y):
        return np.sum((x-y)**2)
    
    nn = NearestNeighbors(n_neighbors=4, algorithm='ball_tree', metric=myfunc)
    
    X = np.array([[-1, -1], [-2, -1], [-3, -2], [1, 1], [2, 1], [3,   2]])
    nn.fit(X)
    

    这在开发版的docstring中也有提到:https://github.com/scikit-learn/scikit-learn/blob/86b1ba72771718acbd1e07fbdc5caaf65ae65440/sklearn/neighbors/unsupervised.py#L48

    【讨论】:

      猜你喜欢
      • 2020-10-31
      • 2018-09-06
      • 2021-11-04
      • 2016-09-26
      • 1970-01-01
      • 1970-01-01
      • 2019-09-10
      • 2015-06-07
      • 2011-12-06
      相关资源
      最近更新 更多