k-means 几十年来一直是一种离线工具。你运行一次来预处理数据,然后就继续下一步。来自加州大学伯克利分校和德克萨斯大学奥斯汀分校的一个研究团队发布了 Flash-KMeans,这是一个新的开源库,针对的是不同的场景。现代 AI 流水线现在会在训练和推理循环内部调用 k-means。在这种频率下,每次调用的延迟比理论上的 FLOPs 更为重要。
Flash-KMeans 是标准 Lloyd k-means 的一种 IO 感知实现。它没有改变数学原理,也没有进行近似计算。它只是重构了算法在 GPU 上移动数据的方式。在 NVIDIA H200 上,研究团队报告称,相比最佳基线实现了高达 17.9 倍的端到端加速。相比 NVIDIA cuML,他们报告了 33 倍的加速。相比 FAISS,他们报告了超过 200 倍的加速。
什么是 Flash-KMeans
Flash-KMeans 是一个用 Triton GPU 内核编写的批处理 k-means 库。它采用 Apache 2.0 许可发布,并通过 `pip install flash-kmeans` 安装。
其输出在数学上与标准 Lloyd k-means 完全相同。加速来自于内核级别的数据流,而不是跳过计算。这使其与三角不等式剪枝或核心集采样等算法方法区分开来。
一个标准的 Lloyd 迭代包含两个阶段。分配阶段计算每个点到每个质心的距离,然后选择最近的质心。更新阶段对每个簇中的点取平均,以形成新的质心。这两个阶段都是简单的算术运算。在 GPU 上,这两个阶段都受限于内存,而非计算。
它解决的两个瓶颈
第一个瓶颈是分配阶段。标准代码在高带宽内存(HBM)中构建一个形状为 N×K 的完整距离矩阵 D。它写入该矩阵,然后读回它以运行 argmin。对于 N=65536、K=1024、d=128、B=32 的情况,距离计算需要 2.6 毫秒。写入和读取 D 大约需要 23 毫秒。矩阵本身是成本所在,而非算术运算。
Flash-KMeans 用 FlashAssign 替代了这一过程。其设计借鉴了 FlashAttention。FlashAssign 将数据点和质心的数据块从 HBM 流式传输到片上 SRAM。它将距离计算与在线 argmin 操作融合在一起。完整的 N×K 矩阵从未被实例化。这便将主要的 IO 复杂度从 O(NK) 降低到了 O(Nd + Kd)。在核函数层面,FlashAssign 的加速比最高可达 21.2 倍。在一个案例中,它将分配时间从 122.5 毫秒缩短到了 5.8 毫秒。
第二个瓶颈是质心更新阶段。标准代码使用的是分散式原子加法。每个线程将其数据点添加到一个由聚类 ID 索引的共享求和缓冲区中。许多线程会同时命中同一个“热点”聚类。这会导致原子操作争用和硬件串行化。研究团队测量到,在 H200 上,此处的有效带宽仅为 50 GB/s。
Flash-KMeans 用排序-逆更新(Sort-Inverse Update)替代了这一过程。它使用 argsort 按聚类 ID 对一维分配向量进行排序。相同的聚类 ID 随后形成连续的段。每个线程块在片上对一个段进行归约,然后每个段只发起一次原子加法操作。庞大的数据点矩阵从未被实际重排。原子操作次数大幅减少。该核函数的加速比最高可达 6.3 倍。
基准测试
研究团队在配备 CUDA 12.8 的 H200 上进行了测试,使用 FP16 数据,维度 d=128。他们遍历了 N、K 和批次大小 B 的不同取值。并与四个优化后的基线方法进行了对比:fast_pytorch_kmeans、fastkmeans、cuML 和 FAISS。
| 对比 | 报告加速比 | 工作负载场景 |
|---|---|---|
| 端到端 vs 最佳基线 | 最高 17.9 倍 | N=800万,K=1024(大 N,小 K) |
| 对比 NVIDIA cuML | 33 倍 | 行业库 |
| 对比 FAISS | 超过 200 倍 | 行业库 |
| FlashAssign 核函数 | 最高 21.2 倍 | N=100万,K=8192(分配阶段) |
| 排序-逆更新核函数 | 最高 6.3 倍 | N=3300万,K=4096(更新阶段) |
| 外存计算,大规模 | 最高 10.5 倍 | N=4亿,K=16384 对比 fastkmeans |
有一种失败模式值得注意。标准的 PyTorch 实现在大 K 场景下会耗尽内存。它们无法实例化 N×K 矩阵。FAISS 是许多生产级向量搜索系统背后的行业标准库。
该库还支持外存计算。在十亿个数据点(K=32768,d=128)上,它完成一次迭代仅需 41.4 秒,而基线方法需要 261.8 秒。它采用分块流重叠技术,将 PCIe 传输隐藏在计算背后。一种缓存感知的编译启发式方法还将调优开销降低了高达 175 倍,速度仅比调优后慢 0.3%。
MTP 交互式解释器
Marktechpost · 交互式解释器
Flash-KMeans:围绕 GPU 内存重构的精确 k-means 算法
与标准 k-means 使用相同的 Lloyd 算法数学原理——其速度提升完全源于数据流的优化。运行实时聚类,观察更新瓶颈,并衡量它消除了多少 IO 开销。
端到端对比 vs 最佳基线
对比 NVIDIA cuML
对比 FAISS
10 亿
数据点,外存计算
1 · 实时聚类
2 · 更新争用
3 · IO 计算器
迭代
质心偏移
状态
空闲
此程序在您的浏览器中对二维数据点运行真正的 Lloyd k-means 算法。其算法与 Flash-KMeans 加速的算法完全相同——区别仅在于 GPU 的数据流。每一步 = 一次分配 + 一次质心更新。
按下播放键。当多个线程块写入同一个“热门”质心(红色表示停滞)时,标准的散点更新会发生串行化。排序-逆序更新首先对聚类 ID 进行排序,这样每个线程块就能通过一次原子加法合并连续的数据段——从而避免冲突。
标准原子操作
O(N·d)
排序-逆序原子操作
O((K+N/B)·d)
实测标准带宽
50 GB/s
内核加速比
标准更新为每个模型 token 发起一次原子加法。许多线程同时访问同一个质心,导致争用。按聚类 ID 排序后,散点操作转变为在片上内存中进行的段级归约。
标准方法
—— 实例化 N×K 矩阵,复杂度 O(NK)
FlashAssign
—— 流式输入,复杂度 O(Nd+Kd)
分配步骤的 HBM 流量更少(理论值)
标准 k-means 会在 HBM 中写入并读取一个完整的 N×K 距离矩阵。FlashAssign 从不构建该矩阵——它只读取一次 X 和 C,并写入一次分配结果。柱状图显示的是相对 HBM 往返次数,采用 FP16 格式。
加速比数据来源:Flash-KMeans 论文(arXiv:2603.09229),NVIDIA H200。演示在浏览器中运行,仅供示意说明 ·
github.com/svg-project/flash-kmeans
应用场景
更快的精确 k-means 算法不仅改变了离线场景,也改变了在线场景的可能性。
- 向量搜索索引:FAISS 使用 k-means 构建其搜索索引。更快的 k-means 允许您在数据发生变化时重新构建索引,而无需等待整夜重建。
- 稀疏注意力路由:路由Transformer与策略簇token实现注意力路由。毫秒级k-means使得该操作在推理循环内部成为可行方案。
- KV缓存压缩:ClusterKV在语义空间中对token进行聚类以压缩缓存。更低成本的聚类使得逐层、逐步压缩具备实用性。
- 低位KV量化:近期方法将KV条目反复聚类至码本中。更快的聚类可缩减该预处理成本。
- 扩散Transformer:Sparse VideoGen2在前向传播过程中调用批量k-means。它依据语义相似度对token进行重排,以利用稀疏性。
使用方法
该API与faiss及sklearn类似。以下调用将对一个批量化(B, N, d)张量进行聚类。
import torch
from flash_kmeans import batch_kmeans_Euclid
x = torch.randn(32, 75600, 128, device="cuda", dtype=torch.float16)
cluster_ids, centers, _ = batch_kmeans_Euclid(
x, n_clusters=1000, tol=1e-4, verbose=True
) 同时提供scikit-learn风格的接口。
from flash_kmeans import FlashKMeans
km = FlashKMeans(d=128, k=8192, niter=100)
labels = km.fit_predict(large_cpu_tensor) # device=None uses all visible GPUs 内核根据形状和数据类型自动分发。小维度路径处理d≤512的情况。拆分维度路径处理更大维度,无需实例化距离矩阵。对于CPU内存中存储的大规模N数据,将自动触发多GPU运行。
核心要点
- Flash-KMeans是精确算法,而非近似算法——采用相同的Lloyd数学原理,仅通过GPU数据流实现加速。
- FlashAssign融合了距离计算与在线argmin操作,将分配IO从O(NK)降低至O(Nd+Kd)——最高可达21.2倍。
- 排序-逆序更新将簇ID排序为分段,替代了分散原子操作——最高可达6.3倍。
- 在H200上报告了最高17.9倍的端到端加速、相较于cuML最高33倍加速、以及相较于FAISS最高200倍以上加速。
- 可扩展至十亿数据点的外存处理,并将调优开销降低最高175倍。