把 FlashAttention 那套搬到老 K-Means 上,他们说比 FAISS 快 200 倍
这两天在 GitHub 上刷到一个仓库,叫 Flash-KMeans,名字一看就有点眼熟。
Flash 这个前缀,玩过大模型的人多少有点条件反射。FlashAttention 嘛,这几年但凡推理跑得快一点的框架,底下大概率都压着它。现在有人把这个 Flash 的味道,端到了 K-Means 头上。
K-Means。聚类算法里最老的那批,老到你大一上数据挖掘课就学过。Lloyd 那篇奠基的东西是 1957 年在贝尔实验室攒出来的,比我爸都大。这么一个被人讲烂了的算法,2026 年了居然还能被人拿出来重写一遍,而且重写完,报道里写的是在 GPU 上比 FAISS 快 200 倍以上。
我看到 200 倍这个数的时候,第一反应不是兴奋,是警惕。
干这行久了会有点职业病,凡是看到「快 N 倍」「提升 N 倍」这种标题,手会先去摸 baseline 在哪。因为太多所谓的提速,是拿一个没人会那么用的朴素实现当陪练,赢了也是欺负老实人。所以这篇我想跟你唠唠的,不只是这玩意有多快,而是它快在哪、跟谁比、那个 200 倍到底是不是标题党。
先把链接放这,省得你以为我编的。项目在 github.com/svg-project/flash-kmeans,Apache 2.0,一行 pip install flash-kmeans 装得上。论文是 arXiv:2603.09229,标题就叫《Flash-KMeans: Fast and Memory-Efficient Exact K-Means》,作者一长串,Shuo Yang、Haocheng Xi 这些人,挂的是 UC Berkeley、UT Austin 那一挂。项目主页在 svg-project.github.io。
好了,进正题。
先说一个容易被略过、但其实是整件事最关键的前提,这个 K-Means 不是凭空造出来的玩具。
它是从一个叫 Sparse VideoGen2 的项目里长出来的,论文 arXiv:2505.18875,SVG2,做的是视频生成加速。视频生成你知道的,attention 矩阵大到离谱,他们要在里面挑出真正重要的那部分 token 来算,省下的算力是实打实的钱。挑哪些 token,就得对 token 做聚类,于是 K-Means 被推到了一个非常尴尬的位置,它成了整条加速流水线里的一个热点环节。
这就有意思了。平时我们用 K-Means,跑一次几秒钟,慢就慢了,喝口水的事。可一旦它被塞进一个每一帧、每一步都要调用的推理管线里,它的速度就不再是无所谓的事,它直接卡在你的关键路径上。原本的瓶颈消除工具,自己成了新瓶颈。
所以他们干脆把这块单拎出来重写,然后开源。这个出身很重要,它意味着 Flash-KMeans 不是实验室里为了刷 benchmark 攒的,是真有人在生产场景里被它慢到了,逼出来的。被 deadline 怼出来的代码,和为了发论文写的代码,质感是不一样的,这个你做工程的应该懂。
那它到底快在哪。
这里得先讲个老道理,可能有些做后端的朋友天天打交道但没往这上面想过。在 GPU 上,很多时候你的瓶颈根本不是算力。
我换个大白话。GPU 的算力,就是那个动辄几百 TFLOPS 的数字,强得离谱,强到大部分时候它在闲着。真正卡你的,是数据搬运。芯片上有一小块特别快的内存叫 SRAM,巴掌大,金子做的;旁边有一大坨慢一些的显存叫 HBM,就是你买显卡时看的那个 24G、80G。算一个东西,数据得从 HBM 搬进 SRAM,算完再搬回去。这一进一出,才是真正吃时间的地方。
FlashAttention 当年炸场子,核心就这一句话,它没有改 attention 的数学,一个公式都没动,结果还是那个结果,它只是想办法让数据少在 HBM 和 SRAM 之间来回搬。IO-aware,知道 IO 在哪、绕着 IO 设计,就这么个朴素到有点欠揍的思路,把 attention 干快了好几倍。
Flash-KMeans 是把这套原样搬到了 K-Means 上。
K-Means 每一轮要干啥?要算每个数据点到每个质心的距离,然后把点分给最近的质心。假设你有 N 个点、K 个质心、batch 是 B,朴素的写法会怎么做?它会老老实实把一个 (B, N, K) 的距离矩阵整个算出来、整个堆到显存里,再去找每行最小值。
这个矩阵有多吓人,你代个数。N 七万五千个点,K 一千个质心,好家伙,光这一个矩阵就好几个 G。点再多一点,直接 OOM 给你看。显存爆了,程序崩了,凌晨三点你盯着那行 CUDA out of memory,那种感觉,做过的都懂,真的就是一声叹息。
Flash-KMeans 的做法叫 K-streaming。它一次只加载一小块数据进 SRAM,对着质心一块一块地流着算距离,算完一块这块的中间结果直接累加进寄存器,那个完整的 (B, N, K) 距离矩阵压根不往显存里落。它根本不存在。
你品一下这个设计的妙处。距离矩阵不物化,省下的不只是那几个 G 的显存,更是省掉了把这几个 G 数据写回 HBM 再读出来的那一整趟搬运。又快又省,快和省在这里是同一件事的两面,这就是 IO-aware 最爽的地方。
说真的,我看到 K-streaming 这个词的时候愣了一下。倒不是设计多复杂,是那种朴素到让你拍大腿的感觉。这思路 FlashAttention 已经验证过了,K-Means 这么个天天有人用的算法,居然到 2026 年才被人认认真真按这个思路在 GPU 上重写一遍。你说前面那么多年大家都在干嘛。
当然,它具体落地还是有工程细节的,不是嘴上说说。它的 assign kernel 用 Triton 写,分了两条路自动切换。数据维度小的时候(D≤512)走一条精打细算的小核,维度大了或者 SRAM 装不下了,自动切到一条 split-D 的路子,把 D 这个维度也拆开来分块算,但 K-streaming 那个核心性质,距离矩阵不落显存,一直保着。它还按显卡型号 H200、H100、A100 这些手调了一堆分块参数,遇到没调过的卡,就退回一套保守配置,保证你至少能跑起来,虽然不是最优。
这一手分发设计有点子牛逼,因为它把性能和兼容性这对老冤家给捏到一块了。你拿一张冷门卡也能跑,跑得慢但不会直接报错摆烂。
好,现在回到那个 200 倍。
我去翻了它 PyPI 页面上挂的官方 benchmark。这里我得替你把话说清楚,免得你被标题带歪。
它官方自己列的对比对象,是这么几个,fast_pytorch_kmeans(一个 PyTorch 实现的 K-Means),AnswerDotAI 家的 fastkmeans(另一个 Triton 实现,外加它的 PyTorch 兜底版),还有一个完全朴素、压根没考虑 OOM 的 batched torch kmeans。测试是在 NVIDIA H200、FP16、128 维数据上跑的,变着 K、变着 N、变着 batch size 扫了一圈。结论是它的 Triton 实现确实带来了显著提升,有些配置下另一个 Triton 实现在 K 比较大时甚至直接报错跑不出来。
但你注意到了吗,这份官方 benchmark 里,我没看见 FAISS。
「比 FAISS 快 200 倍以上」,是 MarkTechPost 那篇报道的标题说法,原文在这,你自己点进去看。我不是说这个数一定是假的,FAISS 的 K-Means 在某些场景下确实有它的短板,被拉开很大差距不是不可能。我只是想说,项目方自己挂出来的、能复现的对比,baseline 是那几个 PyTorch 和 Triton 的实现,不是 FAISS。
这俩是两码事。一个是「我跑了图,你可以复现」,一个是「报道帮我总结了一个更炸的数字」。
我对后者一向保留三分。不是觉得人家骗人,是「比业界标杆快 200 倍」这种话,传一传就容易变成「Flash-KMeans 让 FAISS 原地退役」,再传一传就成了「FAISS 已死」。每次都是这么死的。真要较真,得看在同样的数据规模、同样的精度、同样的硬件上,两边各自最优配置跑出来的数。这个活儿,等哪天我自己手头有 H 卡了再说,现在我没有,我也不装。
但抛开那个 200 倍的标题不谈,这个项目我是真觉得有东西,尤其是大数据量那块。
它专门做了大张量的场景,就是数据量大到一张卡的显存根本装不下。这种时候它和 fastkmeans 在 FP32、128 维上对比,数据点从 25 万一路怼到 2.68 亿,对应的 K 取根号 N。它的做法是数据先待在 CPU 的 pinned memory 里,一块一块传到 GPU 上算,传输和计算用双缓冲叠起来,这边在算那边在传,把数据搬运的时间藏到计算后面。
更狠的是多卡。大 N 场景下你 device 参数不填,它自动把你所有能看见的 GPU 都用上。数据按块切给每张卡,每张卡跑自己那份的双缓冲流水线,PCIe 带宽随卡数线性涨。算完各卡的质心要汇总,它居然没用 NCCL,而是手动 gather-reduce-broadcast。
我看到这个细节笑了一下。它的理由很实在,要汇总的数据小得可怜,质心部分和大概 4MB,计数才 32KB,这点量你上 NCCL 那套重型通信库,反而是杀鸡用牛刀,手动拷贝走 NVLink 还更快,而且全程留在一个进程里,干净。
你看,这就是被真实场景磨过的代码才有的判断力。教科书会告诉你多卡通信用 NCCL,但真到了这个数据量级,作者敢说一句「这里不用」。知道什么时候该守规矩,知道什么时候规矩反而是累赘,这个分寸,是踩出来的,不是学出来的。
用起来也朴素,没什么花活。
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)
它还给了一套跟 faiss、sklearn 长得差不多的接口,FlashKMeans 类,fit_predict 那一套你闭着眼都会用。这点挺识相的,没逼你学一套新黑话,能无痛替换,迁移成本几乎是零。一个开源项目肯在 API 上向老大哥看齐,而不是非要自立门户搞一套新概念,我对它的好感度直接加分。
聊到这我想说说,为什么我觉得一个程序员应该关心这么个看着很冷门的东西。
K-Means 你以为离你很远,其实它藏在你天天碰的东西底下。向量检索建索引,那个 IVF 的倒排桶就是拿 K-Means 分出来的;你搭 RAG 的向量库要分桶加速,背后常常是它;图像和特征做量化、推荐系统做召回、数据去重,扒开来看,里面蹲着的很可能还是这个 1957 年的老家伙。
它太基础了,基础到你已经默认它就该是现在这样、就该这么慢。CPU 上 sklearn 的 KMeans 数据一大就慢到没法用,换 GPU 方案又动不动 OOM,你心里早就接受了这是「本来如此」。
Flash-KMeans 干的事,就是告诉你这个「本来如此」是可以被掀翻的。精确解、省显存、还快,这三个你以前以为只能挑两个的东西,它说我全都要。
而它掀翻的方式,恰恰是最让我上头的地方,它没发明任何新数学。Lloyd 那套迭代,质心、分配、更新,一个字没改,算出来的结果跟你大一写的那版一模一样。它只是把数据在显卡里怎么流动这件事,重新想了一遍。
这事的启发,比一个 200 倍的数字大多了。
我们这行有个错觉,总觉得性能的尽头是更好的算法、更新的模型、更大的卡。但 FlashAttention 和现在这个 Flash-KMeans 都在反复讲同一句话,很多老算法在现代硬件上,压根没被人认真地、贴着硬件特性重写过一遍。它们还跑在十几年前的实现惯性里,守着「距离矩阵当然要整个算出来」这种早就该被质疑的假设。这里头的红利,大得很。
是谁来自山川湖海,却囿于昼夜厨房与爱。我有时候觉得这些老算法也挺像的,它们诞生时面对的是几十年前的硬件,能力本可以连着山川湖海,结果被困在那个时代留下的实现习惯里,一困就是几十年,没人回头看一眼。
Flash-KMeans 就是回头看了一眼的人。它没造神,它只是蹲下来,把一个所有人都觉得理所当然的老东西,从 GPU 数据流的角度重新拆开看了看,然后说,这里,其实可以更好。
所以下次你再看到「快 200 倍」这种标题,我还是劝你先去摸 baseline。但摸完之后,如果你发现它快的方式是把一个被忽略了很久的老假设掀了,那这个项目,值得你认真看两眼。
我已经把它 star 了。等手头有张趁手的卡,我想拿自己的向量库去试试,到时候真跑出来的数,再回来跟你说。