代码之家  ›  专栏  ›  技术社区  ›  davidrpugh

如何使用SciKit Learn选择最佳集群数?[副本]

  •  1
  • davidrpugh  · 技术社区  · 7 年前

    是否可以在不进行交叉验证的情况下使用GridSearchCV?我试图通过网格搜索来优化kmeans集群中的集群数量,因此我不需要或不需要交叉验证。

    这个 documentation 这也让我困惑,因为在fit()方法下,它有一个无监督学习的选项(即对无监督学习使用none)。但是如果你想进行无监督的学习,你需要在没有交叉验证的情况下进行,而且似乎没有摆脱交叉验证的选择。

    0 回复  |  直到 9 年前
        1
  •  16
  •   DataMan    7 年前

    经过多次搜索,我终于找到了 this thread . 如果使用以下选项,则可以在GridSearchCv中消除交叉验证:

    cv=[(slice(None), slice(None))]

    我已经在没有交叉验证的情况下对我自己的网格搜索编码版本进行了测试,并且从这两种方法中得到了相同的结果。我把这个答案贴在我自己的问题上,以防其他人也有同样的问题。

    编辑:要在注释中回答JJRR的问题,下面是一个示例用例:

    from sklearn.metrics import silhouette_score as sc
    
    def cv_silhouette_scorer(estimator, X):
        estimator.fit(X)
        cluster_labels = estimator.labels_
        num_labels = len(set(cluster_labels))
        num_samples = len(X.index)
        if num_labels == 1 or num_labels == num_samples:
            return -1
        else:
            return sc(X, cluster_labels)
    
    cv = [(slice(None), slice(None))]
    gs = GridSearchCV(estimator=sklearn.cluster.MeanShift(), param_grid=param_dict, 
                      scoring=cv_silhouette_scorer, cv=cv, n_jobs=-1)
    gs.fit(df[cols_of_interest])
    
        2
  •  6
  •   Scratch'N'Purr    9 年前

    我要回答你的问题,因为它似乎还没有回答。使用并行方法 for 循环,您可以使用 multiprocessing 模块。

    from multiprocessing.dummy import Pool
    from sklearn.cluster import KMeans
    import functools
    
    kmeans = KMeans()
    
    # define your custom function for passing into each thread
    def find_cluster(n_clusters, kmeans, X):
        from sklearn.metrics import silhouette_score  # you want to import in the scorer in your function
    
        kmeans.set_params(n_clusters=n_clusters)  # set n_cluster
        labels = kmeans.fit_predict(X)  # fit & predict
        score = silhouette_score(X, labels)  # get the score
    
        return score
    
    # Now's the parallel implementation
    clusters = [3, 4, 5]
    pool = Pool()
    results = pool.map(functools.partial(find_cluster, kmeans=kmeans, X=X), clusters)
    pool.close()
    pool.join()
    
    # print the results
    print(results)  # will print a list of scores that corresponds to the clusters list
    
        3
  •  2
  •   PavelLes MrD    7 年前

    我最近推出了以下自定义交叉验证器,基于 this answer . 我把它传给 GridSearchCV 它正确地禁用了我的交叉验证:

    import numpy as np
    
    class DisabledCV:
        def __init__(self):
            self.n_splits = 1
    
        def split(self, X, y, groups=None):
            yield (np.arange(len(X)), np.arange(len(y)))
    
        def get_n_splits(self, X, y, groups=None):
            return self.n_splits
    

    我希望能有所帮助。

        4
  •  1
  •   ihebiheb    8 年前

    我认为使用cv=shufflesplit(test_size=0.20,n_splits=1)和n_splits=1是一个更好的解决方案 post 建议

    推荐文章