You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

K-Means++是否适用于大规模数据?及代码优化问询

嘿,这个问题问得太关键了!先直接给你结论:K-Means++完全可以适配大规模数据集,但你的实现细节拖慢了速度——咱们一步步拆解来看:

K-Means++对大规模数据的适用性

K-Means++的核心是通过概率化方式选择初始中心,避免普通K-Means随机初始化的糟糕情况,它的时间复杂度是O(n*k)(n是样本数,k是聚类数),属于线性复杂度,理论上天生适合大规模数据。你遇到的耗时问题,根源在代码的实现效率上,而非算法本身。

你的代码为什么慢?

看你写的_dist_from_centers方法:

D2 = np.array([min([np.linalg.norm(x-c)**2 for c in cent]) for x in X])

这里有两个致命的低效点:

  • 双重Python循环:外层遍历40万样本,内层遍历已选中心,相当于几十万次的循环嵌套——Python的for循环在处理百万级运算时速度会断崖式下跌,完全没用到numpy的向量化优势。
  • 单次距离计算:np.linalg.norm在循环里逐个计算样本与中心的距离,没有利用numpy的批量运算能力,额外开销极大。
优化你的初始化代码

核心思路是用numpy向量化操作替代Python循环,把距离计算改成批量运算。下面是优化后的关键实现:

import numpy as np
from sklearn.cluster import KMeans
import random

class KPlusPlus(KMeans):
    def _dist_from_centers(self):
        # 向量化批量计算所有样本到每个已选中心的距离平方
        # X形状:(n_samples, 2),cent形状:(k_selected, 2)
        distances = np.sum((self.X[:, np.newaxis, :] - self.mu)**2, axis=2)
        # 取每个样本到最近中心的距离平方
        self.D2 = np.min(distances, axis=1)
    
    def _choose_next_center(self):
        self.probs = self.D2 / self.D2.sum()
        self.cumprobs = self.probs.cumsum()
        r = random.random()
        # 用argmax找第一个满足条件的索引,比where更高效
        ind = np.argmax(self.cumprobs >= r)
        return self.X[ind]
    
    def init_centroids(self, k):
        # 初始化第一个随机中心
        self.mu = [self.X[random.randint(0, len(self.X)-1)]]
        for _ in range(k-1):
            self._dist_from_centers()
            next_center = self._choose_next_center()
            self.mu.append(next_center)
        self.mu = np.array(self.mu)

优化点说明:

  • 用np.sum((X[:, np.newaxis, :] - mu)**2, axis=2)一次性完成所有样本与所有中心的距离计算,彻底摆脱Python循环,速度能提升几个数量级。
  • 用np.argmax(self.cumprobs >= r)替代np.where(...)[0][0],argmax在布尔数组上会直接返回第一个True的索引,效率更高。
更进阶的提速方案

如果数据量再往上走(比如千万级),还可以试试:

  • 采样初始化:先从全量数据中随机采样一小部分(比如1万条),用K-Means++在采样数据上选初始中心,再用这些中心跑全量数据的K-Means——精度损失极小,但速度会大幅提升。
  • 直接用sklearn的原生实现:sklearn的KMeans类本身支持init='k-means++',它的底层是Cython实现的,比纯Python代码快得多,用法也简单:
from sklearn.cluster import KMeans
# n_init=1是关键,因为K-Means++只需要初始化一次,默认n_init=10会重复计算浪费时间
kmeans = KMeans(n_clusters=10, init='k-means++', n_init=1, random_state=42)
kmeans.fit(your_400k_data)

内容的提问来源于stack exchange,提问作者user9562553

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.26 10:40:29