公司动态

黎曼流形上的K-Means聚类

📅 2026/8/8 11:29:05
黎曼流形上的K-Means聚类
文章目录黎曼流形上的K-Means聚类数据生成聚类黎曼流形上的K-Means聚类KMeans最非常简单的一种聚类方法其目的是将样本分割为k个簇其基本流程是选择k个点作为k个簇的初始质心计算样本到这k个质心(簇)的距离并将其划入距离最近的簇中计算每个簇的均值并使用该均值更新簇的质心重复上述2-3的操作直到质心区域稳定或者达到最大迭代次数。常规的KMeans算法默认在欧式空间中进行也就是说在计算距离的时候采取平直的欧式距离。但在某些特殊情况下比如下面这种球面上的点从直觉来说更应该采取测地线作为计算距离的依据。Python的【geomstats】库可以轻松做到这一点。本文是geomstats的官方教程。数据生成我们要做的是在球面上随机生成两簇数据并且将这两簇数据随机旋转一个角度代码如下。importnumpyasnpimportgeomstats.backendasgs np.random.seed(1)gs.random.seed(1000)fromgeomstats.geometry.hypersphereimportHyperspherefromgeomstats.geometry.special_orthogonalimportSpecialOrthogonal sphereHypersphere(dim2,equipFalse)clustersphere.random_von_mises_fisher(kappa20,n_samples140)SO3SpecialOrthogonal(3,equipFalse)rotation1SO3.random_uniform()rotation2SO3.random_uniform()c1cluster rotation1 c2cluster rotation2 datags.concatenate((c1,c2),axis0)【Hypersphere】用于生成一个嵌入到n 1 n1n1维空间中的n nn维单位超球面。在上述代码中生成了一个三维空间中的二维球面。【random_von_mises_fisher】可在超球面上生成符合von Mises-Fisher (vMF) 分布的随机样本vMF可以理解为超球面上的高斯分布其kappa为浓度参数对标高斯分布中的1 σ 2 \frac{1}{\sigma^2}σ21​n_samples为点数。【SpecialOrthogonal】生成一个特殊正交群S O ( n ) SO(n)SO(n)表示n nn维空间中的旋转群。上述代码中用random_uniform方法生成了两个随机的三维旋转矩阵。在生成数据之后可通过visualization函数对数据进行可视化代码如下。importmatplotlib.pyplotasplt plt.rcParams[font.sans-serif]Times New Romanimportgeomstats.visualizationasvisualization figplt.figure(figsize(15,15))axvisualization.plot(c1,spaceS2,colorred,alpha0.7,labelData points 1 )axvisualization.plot(c2,spaceS2,axax,colorblue,alpha0.7,labelData points 2)ax.auto_scale_xyz([-1,1],[-1,1],[-1,1])ax.legend()plt.show()聚类【RiemannianKMeans】是geomstats中用于在黎曼流形上执行 K-Means 聚类的算法类。其不可缺省的参数有而分别是【metric】定义流形几何结构的度量对象算法依赖它来计算测地线距离和Fréchet均值。【n_clusters】: 聚类的簇数简单来说我们需要在生成数据的超球面中定义一个黎曼K-Means算法类。然后调用其中的fit方法对前面创建的数据集data进行分类。代码如下fromgeomstats.learning.kmeansimportRiemannianKMeans manifoldHypersphere(dim2)kmeansRiemannianKMeans(manifold,2,tol1e-3)kmeans.fit(data)labelskmeans.labels_ censkmeans.cluster_centers_ figplt.figure(figsize(15,15))colors[red,blue]axvisualization.plot(data,spaceS2,marker.,colorgrey)foriinrange(2):axvisualization.plot(pointsdata[labelsi],axax,spaceS2,marker.,colorcolors[i])fori,cinenumerate(cens):axvisualization.plot(c,axax,spaceS2,marker*,s2000,colorcolors[i])ax.auto_scale_xyz([-1,1],[-1,1],[-1,1])plt.show()聚类效果如下