使用Scikit-learn包的ClusterMixin
【摘要】 ClusterMixin 是 Scikit-learn 中一个基础的混入类 (Mixin class),它为所有聚类估计器(cluster estimators)提供统一的标准接口。 核心作用它的主要目的是确保 Scikit-learn 中所有的聚类算法(如 K-Means、DBSCAN 等)都遵循相同的 API 规范,方便用户使用和记忆。 主要功能ClusterMixin 主要提供了以下两...
ClusterMixin 是 Scikit-learn 中一个基础的混入类 (Mixin class),它为所有聚类估计器(cluster estimators)提供统一的标准接口。
核心作用
它的主要目的是确保 Scikit-learn 中所有的聚类算法(如 K-Means、DBSCAN 等)都遵循相同的 API 规范,方便用户使用和记忆。
主要功能
ClusterMixin 主要提供了以下两个标准化的功能:
_estimator_type类属性:将该估计器的类型标识为"clusterer"。fit_predict方法:一个便捷方法,用于一次性执行聚类并返回聚类标签,相当于先调用fit()再调用predict()。
核心方法:fit_predict
这是 ClusterMixin 提供的最主要的方法,其定义如下:
fit_predict(X, y=None, **kwargs)
- 参数:
X:形状为(n_samples, n_features)的输入数据。y:被忽略,仅为保持API一致性而存在。**kwargs:要传递给fit方法的额外参数(New in version 1.4)。
- 返回值:
labels:形状为(n_samples,)的 ndarray,即每个样本的聚类标签。
如何使用
ClusterMixin 通常与 BaseEstimator 结合使用,来创建自定义的聚类器。
示例:创建一个简单的自定义聚类器
下面的代码展示了如何创建一个将所有样本都归为同一类的自定义聚类器:
import numpy as np
from sklearn.base import BaseEstimator, ClusterMixin
class MyClusterer(ClusterMixin, BaseEstimator):
def fit(self, X, y=None):
# 这里实现你的聚类逻辑,并将结果存储在 self.labels_ 中
self.labels_ = np.ones(shape=(len(X),), dtype=np.int64)
return self
# 使用自定义聚类器
X = [[1, 2], [2, 3], [3, 4]]
labels = MyClusterer().fit_predict(X)
print(labels) # 输出: array([1, 1, 1])
一个实际案例:InductiveClusterer
在 Scikit-learn 的官方示例中,ClusterMixin 被用于创建一个名为 InductiveClusterer 的元估计器。它先使用一个聚类器(如 AgglomerativeClustering)对数据进行聚类,然后用一个分类器(如 RandomForestClassifier)来学习这些聚类标签,从而实现对新数据的快速分类,避免了每次都对全量数据重新聚类的昂贵开销。
from sklearn.base import BaseEstimator, ClusterMixin
from sklearn.cluster import AgglomerativeClustering
from sklearn.ensemble import RandomForestClassifier
class InductiveClusterer(ClusterMixin, BaseEstimator):
def __init__(self, clusterer, classifier):
self.clusterer = clusterer
self.classifier = classifier
def fit(self, X, y=None):
# 1. 先进行聚类
self.clusterer_ = clone(self.clusterer)
y = self.clusterer_.fit_predict(X)
# 2. 再用分类器学习聚类结果
self.classifier_ = clone(self.classifier)
self.classifier_.fit(X, y)
self.labels_ = y
return self
# ... 还需要实现 predict 等方法
注意:以上是示例的核心部分,完整实现可参考 [Inductive Clustering 示例]
相关的其他 Mixin 类
Scikit-learn 中还有类似的 Mixin 类,用于定义其他类型估计器的标准接口:
ClassifierMixin:用于所有分类器。RegressorMixin:用于所有回归器。TransformerMixin:用于所有数据转换器。
【版权声明】本文为华为云社区用户原创内容,未经允许不得转载,如需转载请自行联系原作者进行授权。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱:
cloudbbs@huaweicloud.com
- 点赞
- 收藏
- 关注作者
评论(0)