深度解析塑造未来的技术文章。

K-Means大规模应用:为何失效与下一步

K-means看似简单,但面对十亿级数据点就会暴露内存、速度和初始化问题。看看现代算法如何解决这些难题。

无数彩色圆点环绕着夜空中几束明亮的光点

K-means是每个人在入门机器学习课上都会学到,然后职业生涯里却一直用错的算法。它看起来极其简单:选定K个聚类中心,把每个点分配给最近的中心,再把中心移动到其所属点的均值处,反复迭代。三行伪代码就够了。问题在于,这种简单性背后藏着一连串的陷阱,一旦把它用在真实的大规模数据上,这些陷阱就会接连爆发。

当数据只有1万个点时,k-means几毫秒就跑完了,没人会在意实现细节。数据增长到1000万个点时,初始化方法的选择能让运行时间相差100倍。到了10亿个点,标准算法根本放不进内存,你需要完全不同的思路。最近一篇关于Flash-KMeans的论文正是针对这一点,在大幅降低内存占用、加快收敛的同时,得到与精确k-means一致的结果。下面我们就来看看k-means在大规模下到底难在哪里,以及现代变体是如何解决的。

标准K-Means的工作原理

Lloyd算法——大家说“k-means”时指的通常就是它——每轮迭代包含两步。分配步骤:对每个数据点,计算它与所有K个质心的距离,并将其分配给最近的那个。更新步骤:把每个质心重新计算为分配给它的所有点的均值。

import numpy as np
def kmeans_lloyd(X, K, max_iter=100):
n, d = X.shape
# Random initialization (bad, but we'll fix this later)
centroids = X[np.random.choice(n, K, replace=False)]
for _ in range(max_iter):
# Assignment: O(n * K * d) — this is the bottleneck
distances = np.linalg.norm(X[:, None] - centroids[None, :], axis=2)
labels = np.argmin(distances, axis=1)
# Update: O(n * d)
new_centroids = np.array([
X[labels == k].mean(axis=0) for k in range(K)
])
if np.allclose(centroids, new_centroids):
break
centroids = new_centroids
return labels, centroids

分配步骤每轮的开销是O(n × K × d),其中n是点的数量,K是聚类数,d是维度。假设n = 10亿,K = 1000,d = 128(这是一个很现实的embedding聚类场景),那每轮就是128万亿次浮点运算。即便按一万亿次每秒的算力计算,每轮也要128秒,而k-means通常需要10到50轮才能收敛。

初始化问题

在k-means执行第一轮迭代之前,它必须先选定初始质心的位置。这个选择的重要性远超大多数人的想象——糟糕的初始化可能让算法收敛到一个比最优解差得离谱的结果。

随机初始化(从数据中随机挑选K个点作为初始质心)简单,但不可靠。如果两个初始质心恰好落在同一个簇里,就会有一个簇始终得不到代表。算法依然会收敛,只是收敛到了一个次优解。K = 100时,至少出现一次这种碰撞的概率很高。

K-means++(Arthur和Vassilvitskii,2007)是标准的解决方案。它先随机选第一个质心,之后每个新质心的选取概率与其到最近已有质心的距离平方成正比。远离现有质心的点更容易被选中,从而使初始质心分散在整个数据空间中。它提供了O(log K)的近似保证——解可以被证明与最优解相差不超过O(log K)倍。

def kmeans_plus_plus_init(X, K):
n, d = X.shape
centroids = [X[np.random.randint(n)]]
for _ in range(1, K):
# Distance from each point to nearest existing centroid
dists = np.min([
np.sum((X - c) ** 2, axis=1) for c in centroids
], axis=0)
# Sample proportional to distance squared
probs = dists / dists.sum()
next_idx = np.random.choice(n, p=probs)
centroids.append(X[next_idx])
return np.array(centroids)

问题在于,k-means++的初始化本身就是O(n × K × d)的开销,它需要对整个数据集扫描K遍。当K和n都很大时,初始化的耗时可能超过Lloyd算法的好几轮迭代。可扩展的变体,比如k-means||(Bahmani等人,2012),通过并行过采样候选点来减少遍历次数,然后再将这些候选点合并成K个质心。

内存墙

标准k-means需要同时将整个数据集装入内存。在分配步骤中,你需要访问每一个数据点,并将其与每一个质心进行比较。对于10亿个128维的float32向量,光数据本身就需要512 GB——这已经超过了大多数单机的内存容量。

一种直接的解决办法是mini-batch k-means:每轮迭代随机抽取一部分点,只根据这个样本更新质心。这种方法可行,但收敛更慢,而且得到的结果与精确k-means略有不同(通常质量略差)。对很多应用来说,质量损失可以接受;但对另一些场景,比如向量检索中的量化码本,质量差异就很关键了。

Flash-KMeans采用了另一种思路。它不是把整个数据集都放进内存处理,而是通过带有智能缓存的流式遍历来完成计算。关键洞察在于:大多数点在相邻迭代之间并不会改变所属簇。如果一个点牢牢地属于第7个簇(离质心7的距离远小于离其他任何质心的距离),那么逐一计算它到全部K个质心的距离就是白费力气。通过维护距离的上下界,Flash-KMeans可以在大多数迭代中跳过大部分点的完整距离计算。

三角不等式优化

Elkan算法(2003)利用三角不等式来跳过不必要的距离计算。三角不等式是说:点P到质心A的距离,不会超过P到质心B的距离加上B到A的距离。如果你知道P当前被分配给质心B,又知道质心A和B之间的距离,那么有时就能在不计算实际距离的情况下证明A离P太远。

在实践中,经过最初几轮迭代后,大多数点已经靠近各自正确的质心,这种方法能消除80%到95%的距离计算。剩下需要计算的,主要是位于簇边界附近的点——它们才有可能真正改变所属簇。

代价是:Elkan算法需要额外的O(n × K)内存来存储距离界限(每个点到每个质心的下界,以及到所属质心的上界)。当K很大时,这部分内存开销可能相当可观。Hamerly算法则只为每个点维护一个下界,将内存降到O(n),代价是能剪枝掉的计算变少了。

GPU加速

k-means的分配步骤天然适合并行:每个点的距离计算彼此独立,因此非常适合用GPU加速。NVIDIA的cuML库和Facebook的FAISS都内置了GPU版的k-means实现,在大规模数据集上相比CPU实现可以获得10到50倍的加速。

麻烦在于GPU显存有限。一块A100有80 GB显存,大约只能容纳1.5亿个128维向量。更大的数据集需要多GPU分片,或者采用流式方案:分块加载数据,在GPU上处理,再在CPU上汇总结果。

# Using FAISS for GPU-accelerated k-means
import faiss
import numpy as np
# 10 million 128-dimensional vectors
n, d, K = 10_000_000, 128, 1000
X = np.random.randn(n, d).astype('float32')
# CPU k-means (for comparison)
kmeans_cpu = faiss.Kmeans(d, K, niter=20, verbose=True)
kmeans_cpu.train(X)  # ~120 seconds
# GPU k-means (single GPU)
kmeans_gpu = faiss.Kmeans(d, K, niter=20, verbose=True, gpu=True)
kmeans_gpu.train(X)  # ~8 seconds — 15x faster

什么时候K-Means不是正确的选择

在优化你的k-means实现之前,先考虑一下k-means是否适合你的问题。它做了一些很强的假设,而这些假设对很多真实数据集并不成立。

  • 球形簇。 K-means假设簇大致呈球形且大小相近。如果你的簇是细长的、形状不规则的,或者大小差异悬殊,k-means会把大簇切开、把小簇合并。高斯混合模型(GMM)能处理椭圆形的簇,DBSCAN则可以处理任意形状。
  • 需要预先指定K。 你必须事先确定聚类数量。如果不知道K,就需要用不同的K值多次运行k-means,并借助某个指标(轮廓系数、肘部法则)来挑选最佳值。这会让总计算量乘以你尝试的K值个数。
  • 欧氏距离。 K-means默认使用欧氏距离。对于文本embedding,余弦相似度通常更合适。你可以通过对向量做L2归一化来变通(这样欧氏距离就等价于余弦距离),但这一步很容易被忘掉。
  • 对离群值敏感。 一个远离所有簇的离群点会把它所属的质心拉向自己。像k-medoids这样的鲁棒变体(用中位数代替均值)能更好地处理离群值,但计算成本更高。

实用建议

我在从数千到数十亿个点的数据集上都用过k-means,以下是我觉得最重要的几点经验:

  1. 始终使用k-means++初始化。 随机初始化从来不值得冒这个风险。最终簇质量的差距往往能达到10%到30%,而对于中小规模的K,k-means++带来的额外开销几乎可以忽略。
  2. 多次运行。 K-means找到的是局部最优而非全局最优。用5到10个不同的随机种子运行,然后选取最优结果(簇内总距离最小的那个)。这是防范糟糕运行结果的低成本保险。
  3. 对特征进行归一化。 如果一个特征的取值范围是[0, 1000000],另一个是[0, 1],那么取值范围大的特征会主导距离计算。聚类之前先做标准化(减去均值、除以标准差)或最小-最大归一化。
  4. 大规模聚类使用FAISS。 如果你的数据超过一百万个点,scikit-learn的k-means会比较慢。FAISS的实现经过大量优化,而且开箱即用地支持GPU加速。
  5. 探索性工作可以考虑近似方法。 Mini-batch k-means比精确k-means快约10倍,对于探索性分析来说结果通常已经足够好。对于质量至关重要的生产环境聚类,则应使用精确k-means。

k-means是一种用起来容易、用好却很难,但又非常值得深入理解的算法。朴素实现与优化实现之间的差距——无论是速度还是结果质量——都非常大。在小规模下,这些都无关紧要;而在真正需要它们的规模上,对初始化、内存管理、距离剪枝和GPU加速的理解,决定了聚类是花几小时还是几分钟。