跳到正文
北京时间
原文
MarkTechPost(RSS)· Asif Razzaq·· 2026-06-15精选AI 评分73

Flash-KMeans:IO感知的精确K-Means,在GPU上比FAISS快200倍以上

Meet Flash-KMeans: An IO-Aware, Exact K-Means That Runs Over 200× Faster Than FAISS on GPUs

AI 导读

UC Berkeley与UT Austin团队开源Flash-KMeans(Apache 2.0,`pip install flash-kmeans`),精确实现标准Lloyd's k-Means,通过重构GPU数据流而非改变数学或近似来提速。在NVIDIA H200上,端到端速度比最佳基线快17.9×,比cuML快33×,比FAISS快200×以上。其FlashAssign核避免物化完整N×K距离矩阵,将IO复杂度从O(NK)降至O(Nd+Kd),单核加速最高21.2×;Sort-Inverse Update核通过排序聚类ID减少原子争用,单核加速最高6.3×。支持out-of-core处理,在1B数据点、K=32768时单次迭代仅41.4s。适用于向量搜索索引、稀疏注意力路由、KV缓存压缩等在线场景。

推荐理由

Flash-KMeans 把 k-means 从离线预处理拉进了在线循环,200 倍加速不是纸面数字,而是让向量索引重建、稀疏注意力路由这些场景突然可行了。做大规模聚类的可以立刻换掉 FAISS。

正文 · AI 翻译

k-means 作为离线工具已有数十年历史。你运行它一次来预处理数据,然后就不再管它。来自 UC Berkeley 和 UT Austin 的一个研究团队发布了 Flash-KMeans,这是一个面向不同场景的新开源库。现代 AI 流水线如今会在训练和推理循环内部调用 k-means。在这种调用频率下,每次调用的延迟比理论 FLOPs 更重要。

Flash-KMeans 是标准 Lloyd's k-means 的一种 IO 感知实现。它不改变数学原理,也不做近似。它只是重构了算法在 GPU 上移动数据的方式。在一块 NVIDIA H200 上,研究团队报告称相比最佳基线实现了最高 17.9× 的端到端加速。相比 NVIDIA cuML,他们报告为 33×。相比 FAISS,他们报告超过 200×。

什么是 Flash-KMeans

Flash-KMeans 是一个用 Triton GPU kernel 编写的批量 k-means 库。它以 Apache 2.0 许可发布,可通过 pip install flash-kmeans 安装。

其输出在数学上与标准 Lloyd's k-means 完全一致。加速来自 kernel 级别的数据流,而非跳过计算。这使它区别于三角不等式剪枝或 coreset 采样等算法类方法。

一次标准的 Lloyd 迭代分为两个阶段。分配阶段计算每个点到每个质心的距离,然后选出最近的那个。更新阶段对每个簇内的点取平均,形成新的质心。这两个阶段都是简单的算术运算。在 GPU 上,两者的瓶颈都在内存,而非计算。

它攻克的兩大瓶頸

第一个瓶颈是分配阶段。标准代码会在高带宽显存(HBM)中构建一个形状为 N×K 的完整距离矩阵 D。它先写入该矩阵,然后再读回来执行 argmin。对于 N=65536、K=1024、d=128、B=32 的情况,距离计算耗时 2.6ms。写入和读取 D 大约耗时 23ms。真正的开销在于这个矩阵,而非算术运算。

Flash-KMeans 用 FlashAssign 取而代之。该设计借鉴了 FlashAttention。FlashAssign 将点和质心的分块从 HBM 流式传入片上 SRAM。它将距离计算与在线 argmin 融合在一起。完整的 N×K 矩阵从不实际生成。这将主导性的 IO 复杂度从 O(NK) 降至 O(Nd + Kd)。在内核层面,FlashAssign 最高可达 21.2×。在一个案例中,它将分配阶段从 122.5ms 缩短至 5.8ms。

第二个瓶颈是质心更新阶段。标准代码使用散射式原子加法。每个线程将其点按簇 id 为键累加到一个共享的求和缓冲区中。许多线程会同时命中同一个“热点”簇。这会导致原子竞争和硬件串行化。研究团队在 H200 上测得此处的有效带宽仅为 50 GB/s。

Flash-KMeans 用排序-逆更新(Sort-Inverse Update)取而代之。它使用 argsort 按簇 id 对一维分配向量进行排序。相同的簇 id 随后形成连续的分段。每个线程块在片上对一个分段进行归约,然后每个分段发出一次原子加法。庞大的点矩阵从不进行物理置换。原子操作从 (O((K+NBN)d)) 下降。该内核最高可达 6.3×。

基准测试

研究团队在配备 CUDA 12.8、FP16 数据、d=128 的 H200 上进行测试。他们遍历了 N、K 和批量大小 B。他们与四个经过优化的基线进行了对比:fast_pytorch_kmeans、fastkmeans、cuML 和 FAISS。

对比报告的加速比工作负载背景
端到端 vs 最佳基线最高 17.9×N=8M,K=1024(大 N,小 K)
vs NVIDIA cuML33×行业库
vs FAISS超过 200×行业库
FlashAssign kernel最高 21.2×N=1M,K=8192(分配)
排序-逆更新 kernel最高 6.3×N=33M,K=4096(更新)
核外、大规模最高 10.5×N=400M,K=16384 对比 fastkmeans

有一种失败模式对上下文至关重要。标准 PyTorch 实现在大 K 场景下会耗尽内存。它们无法将 N×K 矩阵实际物化。FAISS 是许多生产级向量搜索系统所依赖的行业标准库。

该库同样支持核外运行。在十亿个点(K=32768,d=128)上,它完成一次迭代仅需 41.4s,而基线需要 261.8s。它采用分块流重叠,将 PCIe 传输隐藏在计算之后。一种缓存感知的编译启发式方法还能将调优开销最多降低 175×,且与调优后速度的差距在 0.3% 以内。

MTP 交互式讲解

Marktechpost · 交互式讲解

Flash-KMeans:围绕 GPU 内存重建的精确 k-means

与标准 k-means 相同的 Lloyd's 数学——仅因数据流而更快。实时运行聚类,观察更新瓶颈,并衡量它所消除的 IO。

17.9×

端到端 vs 最佳基线

33×

vs NVIDIA cuML

200×+

vs FAISS

1B

个点,核外计算

1 · 实时聚类

2 · 更新争用

3 · IO 计算器

迭代

0

质心偏移

—

状态

空闲

这段代码在你的浏览器中对二维点运行真正的 Lloyd's k-means。该算法与 Flash-KMeans 所加速的算法完全相同——只是 GPU 数据流不同。每一步 = 一次分配 + 一次质心更新。

按下播放。当多个块写入同一个“热”质心时,标准 scatter-update 会串行化(红色表示停顿)。Sort-Inverse Update 先对聚类 ID 排序,因此每个块通过一次原子加即可合并连续段——无冲突。

标准原子操作

O(N·d)

Sort-Inverse 原子操作

O((K+N/B)·d)

实测标准带宽

50 GB/s

内核加速比

6.3×

标准更新对每个 token执行一次原子加。许多线程同时命中同一个质心,导致争用。按聚类 ID 排序可将 scatter 转化为片上内存中的段级归约。

标准

— 物化 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 让你可以在数据变化时重新索引,而不必花一整夜重建。
  • 稀疏注意力路由:Routing Transformer 和 Tactic 会对 token 进行聚类以路由注意力。毫秒级的 k-means 让这在推理循环内变得可行。
  • KV-cache 压缩: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

该 kernel 根据形状和 dtype 自动分派。小 D 路径处理 d≤512。拆分 D 路径处理更大的 d,且无需物化距离矩阵。对于存放在 CPU 内存中的大 N 数据,会自动触发多 GPU 运行。

核心要点

  • Flash-KMeans 是精确的,而非近似的——相同的 Lloyd's 数学,纯粹靠 GPU 数据流加速。
  • FlashAssign 融合了距离计算与在线 argmin,将分配阶段的 IO 从 O(NK) 降至 O(Nd+Kd)——最高达 21.2×。
  • 排序-逆更新 将聚类 ID 排序为分段,取代 scatter 原子操作——最高达 6.3×。
  • 报告显示,端到端最高可达 17.9×,相比 cuML 最高可达 33×,在 H200 上相比 FAISS 超过 200×。
  • 可 在核外扩展至十亿个点,并将调优开销最多降低 175×。

来源:MarkTechPost(RSS) · marktechpost.com