Skip to content

均值漂移算法

[Mean Shift] 是一种强大的、[基于质心 (centroid-based)]的[聚类算法 (clustering algorithm)],用于[无监督学习 (unsupervised learning)]。与 [K-Means] 不同,[Mean Shift] 是一种[非参数算法 (non-parametric algorithm)],这意味着它无需事先对底层[数据分布 (data distribution)]或[聚类数量 (number of clusters)]进行假设。

[Mean Shift] 的核心思想是[迭代移动 (iteratively shift)]数据点朝其附近的[密度最高的区域 (densest region)]。本质上,点会爬向数据[密度景观 (density landscape)]中最近的[峰值 (mode)]。收敛到同一[峰值 (peak)]的点被分到同一个[聚类 (cluster)]中。

与 [K-Means] 的一个关键区别在于,[Mean Shift] 根据[数据结构 (data structure)]和一个称为[带宽 (bandwidth)]的参数自动确定[聚类数量 (number of clusters)],[带宽 (bandwidth)]定义了用于[密度估计 (density estimation)]的[邻域 (neighborhood)]大小。

[Mean Shift 聚类 (Mean Shift clustering)]过程可以概括如下:

  • 第 1 步:从每个数据点开始,将其作为初始[候选质心 (candidate centroid)]。
  • 第 2 步:对于每个[候选质心 (candidate centroid)],定义其周围由“[带宽 (bandwidth)]”参数确定的区域([窗口 (window)])。
  • 第 3 步:计算落在[窗口 (window)]内的数据点的均值。
  • 第 4 步:将[窗口 (window)]中心(即[候选质心 (candidate centroid)])移动到计算出的均值处。
  • 第 5 步:重复步骤 2-4,直到[质心 (centroids)][收敛 (converge)],即它们停止显著移动或达到[最大迭代次数 (maximum number of iterations)]。
  • 第 6 步:其[质心 (centroids)][收敛 (converge)]到附近位置([密度 (density)]的[众数 (modes)])的点被[分到同一个聚类 (grouped into the same cluster)]中。

“[带宽 (bandwidth)]”参数至关重要。较小的[带宽 (bandwidth)]可能导致更多的[聚类 (clusters)]([过度分割 (over-segmentation)]),而较大的[带宽 (bandwidth)]可能合并不同的[聚类 (clusters)]([分割不足 (under-segmentation)])。Scikit-learn 提供了 estimate_bandwidth 函数来帮助找到一个合理的起始值。

此示例演示了 [Mean Shift 聚类 (Mean Shift clustering)]。首先,我们使用 Scikit-learn 的 make_blobs 生成具有不同团块的 2D [数据集 (dataset)]。然后,我们应用 [Mean Shift 算法 (Mean Shift algorithm)]并可视化得到的[聚类 (clusters)]。

# 导入必要的库
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.cluster import MeanShift, estimate_bandwidth
from sklearn.datasets import make_blobs
# 生成示例数据
centers = [[1, 1], [5, 5], [3, 10]]
X, _ = make_blobs(n_samples=500, centers=centers, cluster_std=0.8, random_state=42)
# 可视化原始数据
plt.figure(figsize=(8, 6))
sns.scatterplot(x=X[:, 0], y=X[:, 1], s=50)
plt.title('Original Data Points')
plt.xlabel('Feature 1')
plt.ylabel('Feature 2')
plt.show()
# --- Mean Shift 聚类 ---
# 估计带宽(可选但推荐)
# quantile 控制用于估计的数据点比例
bandwidth = estimate_bandwidth(X, quantile=0.2, n_samples=500)
print(f"Estimated bandwidth: {bandwidth:.2f}")
# 如果带宽估计值偏低或偏高,可能需要手动调整
# 对于本例,让我们使用估计的带宽或略作调整的值
# bandwidth = 1.8 # 手动调整示例(如果需要)
# 初始化并拟合 MeanShift
# bandwidth: 核密度估计窗口的半径
# bin_seeding=True: 通过对点进行离散化来加速计算
ms = MeanShift(bandwidth=bandwidth, bin_seeding=True)
ms.fit(X)
# 获取聚类标签和中心
labels = ms.labels_
cluster_centers = ms.cluster_centers_
n_clusters_ = len(np.unique(labels))
print(f"\nCluster Centers Found:\n{cluster_centers}")
print(f"\nEstimated number of clusters: {n_clusters_}")
# 可视化结果
plt.figure(figsize=(8, 6))
sns.scatterplot(x=X[:, 0], y=X[:, 1], hue=labels, palette='viridis', s=50, legend='full')
# 绘制聚类中心
sns.scatterplot(x=cluster_centers[:, 0], y=cluster_centers[:, 1], marker='X', s=200, color='red', label='Cluster Centers')
plt.title(f'Mean Shift Clustering Results ({n_clusters_} clusters)')
plt.xlabel('Feature 1')
plt.ylabel('Feature 2')
plt.legend()
plt.show()

第一个图显示了原始的、未[聚类 (clustered)]的数据点。控制台输出将显示估计的[带宽 (bandwidth)]、找到的[聚类中心 (cluster centers)]的[坐标 (coordinates)]以及算法识别出的[聚类 (clusters)]总数。最后一个图显示根据其分配的[聚类 (cluster)]着色的原始数据点,并用标记(通常是“X”或类似标记)标注了识别出的[聚类中心 (cluster centers)]。理想情况下,[Mean Shift] 找到的[聚类 (clusters)]应该与原始数据中的[密集区域 (dense regions)]相对应。

[Mean Shift] 提供了几个优点:

  • 无需预先指定[聚类数量 (number of clusters)]。
  • 可以找到任意形状([非凸 (non-convex)])的[聚类 (clusters)]。
  • 对[离群点 (outliers)] [鲁棒 (robust)],因为它们不太可能形成[密集区域 (dense regions)]。
  • 只需一个主要参数:[带宽 (bandwidth)]。

然而,它也有局限性:

  • 性能对[带宽 (bandwidth)]的选择敏感。寻找[最优带宽 (optimal bandwidth)]可能具有挑战性。
  • 对于[大型数据集 (large datasets)],[计算开销 (computationally expensive)]大,因为它涉及[密度估计 (density estimation)]和潜在许多点的[迭代移动 (iterative shifts)](在[朴素实现 (naive implementations)]中为 O(N^2),尽管存在[优化 (optimizations)])。
  • 由于[维度灾难 (curse of dimensionality)]影响[密度估计 (density estimation)],可能难以处理[高维数据 (high-dimensional data)]。
  • 如果[带宽 (bandwidth)]太大,可能会合并附近的[聚类 (clusters)];如果太小,可能会创建[伪小的聚类 (spurious small clusters)]。

更多资源: