使用Scikit-learn包的Birch
【摘要】 下面是一个使用 scikit-learn 中 Birch 进行聚类的完整示例和说明。 1. 基本用法import numpy as npimport matplotlib.pyplot as pltfrom sklearn.datasets import make_blobsfrom sklearn.cluster import Birchfrom sklearn.preprocessing...
下面是一个使用 scikit-learn 中 Birch 进行聚类的完整示例和说明。
1. 基本用法
import numpy as np
import matplotlib.pyplot as plt
from sklearn.datasets import make_blobs
from sklearn.cluster import Birch
from sklearn.preprocessing import StandardScaler
from sklearn.metrics import silhouette_score
# 生成示例数据:3000 个样本,5 个簇
X, y_true = make_blobs(
n_samples=3000,
centers=5,
cluster_std=0.7,
random_state=42
)
# BIRCH 对特征尺度敏感,建议先标准化
X = StandardScaler().fit_transform(X)
# 创建 BIRCH 模型
birch = Birch(
threshold=0.5, # 子簇半径阈值
branching_factor=50, # CF 树节点最大分支数
n_clusters=5, # 最终聚类数;None 表示只生成子簇
compute_labels=True,
copy=True
)
# 拟合并预测
labels = birch.fit_predict(X)
# 查看结果
print("最终簇标签:", np.unique(labels))
print("子簇中心数量:", len(birch.subcluster_centers_))
print("轮廓系数:", silhouette_score(X, labels))
# 可视化
plt.figure(figsize=(8, 6))
plt.scatter(X[:, 0], X[:, 1], c=labels, cmap="viridis", s=10)
# 注意:subcluster_centers_ 是子簇中心,不是最终簇中心
plt.scatter(
birch.subcluster_centers_[:, 0],
birch.subcluster_centers_[:, 1],
c="red",
marker="x",
s=100,
label="子簇中心"
)
plt.title("BIRCH 聚类结果")
plt.legend()
plt.show()
2. 预测新数据
X_new = np.array([
[0.1, 0.2],
[-1.0, 0.5],
[1.5, -0.3]
])
pred = birch.predict(X_new)
print("新样本预测簇:", pred)
3. 关键参数说明
| 参数 | 说明 |
|---|---|
threshold |
子簇合并的半径阈值,默认 0.5。越大,子簇越少,聚类数可能越少;越小,子簇越多。 |
branching_factor |
CF 树节点最大分支数,默认 50。越大,内存占用越高,但树更浅。 |
n_clusters |
最终聚类数。None 表示只生成子簇;整数表示最后用层次聚类合并成指定簇数。 |
compute_labels |
是否计算样本标签,默认 True。 |
copy |
是否复制输入数据,默认 True。 |
常用属性:
birch.labels_:训练数据的最终簇标签birch.subcluster_centers_:子簇中心birch.subcluster_labels_:子簇对应的最终簇标签
4. 使用建议
-
先标准化数据
BIRCH 的threshold对特征尺度敏感,通常先用StandardScaler。 -
适合大规模、球状簇数据
BIRCH 通过 CF 树压缩数据,适合样本量较大的场景,但对非凸簇、密度差异很大的数据效果可能不好。 -
如何选择
n_clusters
如果不知道簇数,可以先设n_clusters=None,观察子簇数量,再结合轮廓系数、肘部法等确定最终簇数。 -
增量学习
BIRCH 支持partial_fit,可用于增量构建 CF 树,适合流式或大数据场景。
简单总结:
Birch(threshold=..., branching_factor=..., n_clusters=...) → fit_predict(X) → 得到聚类标签。实际使用时重点调节 threshold 和 n_clusters。
【版权声明】本文为华为云社区用户原创内容,未经允许不得转载,如需转载请自行联系原作者进行授权。如果您发现本社区中有涉嫌抄袭的内容,欢迎发送邮件进行举报,并提供相关证据,一经查实,本社区将立刻删除涉嫌侵权内容,举报邮箱:
cloudbbs@huaweicloud.com
- 点赞
- 收藏
- 关注作者
评论(0)