跳到主要內容

GPU KMeans 跑 10 秒,FlashLib 版本 0.28 秒——差距來自記憶體,不是算法

GPU K-Means 的瓶頸不在計算。一張 H200,理論算力是 1,979 TFLOPS,但標準實作裡真正花時間的不是距離計算——是把中間結果寫到 GPU 記憶體再讀回來這件事。

Flash-KMeans 2026 年初發布,用一個叫 FlashAssign 的設計讓 K-Means 比 FAISS 快 200 倍、比 cuML 快 33 倍。現在,同個團隊把這套 IO-aware 設計邏輯擴展到六個 ML 算法,做成一個 GPU 機器學習庫:FlashLib

最高加速幅度是 TruncatedSVD 的 208 倍。


為什麼 GPU K-Means 這麼慢

你有一張 GPU,計算力夠強,但跑 K-Means 的時候還是很慢。問題在哪?

標準流程是這樣的:為每個資料點計算到所有 cluster 的距離 → 把這個 N×K 距離矩陣寫到 HBM(GPU 的主記憶體)→ 再把它讀回來找最近的 cluster。

對 N=1M、K=64K 的工作負載,寫入這個矩陣要 23 毫秒。但實際的距離計算呢?2.6 毫秒。記憶體傳輸慢了將近 9 倍。

FlashAssign 解決這件事的方法:完全不寫那個矩陣

Loading diagram...

FlashAssign 把 centroid 分成 tile 逐批載入 on-chip buffer,在 chip 上同時計算距離和更新 argmin,整個過程只在最後寫一次結果到 HBM。中間那個 N×K 矩陣從來不存在。

同樣的問題在 centroid 更新階段也有:多個點同時更新同一個 centroid 時的 atomic scatter 衝突。FlashLib 用 sort-inverse update 建立反向映射,把原本的 high-contention write 改成 segment-level 的有序 reduction,centroid 更新快了 6.3 倍。


FlashLib:把同一招套到六個算法

Flash-KMeans 的發現是:這類 IO-aware 設計不只能用在 K-Means,任何「計算本身很快、但記憶體讀寫是瓶頸」的 ML 算法都可以這樣做。

FlashLib 的 benchmark 全部在 NVIDIA H200、cuML 25.10 條件下測試:

算法加速倍數 vs cuML
KMeans26x
KNN19x
HDBSCAN40x
TruncatedSVD208x
PCA47x
exact t-SNE147x

來源:flashml-org.github.io,CUDA 13.0,H200 SM90 150GB HBM3e

這些數字代表的不只是速度,而是硬體利用率的提升。Flash-KMeans 達到 H200 61% 的 peak FLOPs,Flash-KNN 達到 85.2% 的 peak HBM 頻寬。相比之下,標準 cuML 實作通常離 GPU 的理論上限差很遠。

FlashLib 另外做了一件事:Flash Informative API。在任何工作負載跑起來之前,可以在 ~5 微秒內(純 CPU,不需要 GPU profiling)預測這個任務的執行時間、記憶體用量、overhead。這讓排程和容量規劃變得可行。


怎麼裝

Flash-KMeans 獨立安裝:

pip install flash-kmeans

FlashLib 的安裝和文件在 flashml-org.github.io,支援 OpenAI / Anthropic API 風格的調用介面。

Flash-KMeans 的 Python API:

from flash_kmeans import flash_kmeans

# data: (N, D) tensor on GPU, FP16
# centroids: (K, D) tensor on GPU, FP16
labels, new_centroids = flash_kmeans(data, centroids, max_iters=20)

PyTorch tensor 直接丟進去,不需要特殊格式。資料超過 VRAM 的場景(out-of-core),FlashLib 自動從 CPU pinned memory 分批轉移,你不需要手動處理。

第一次執行有約 2.5 秒的 Triton JIT 編譯,這已經比原本的 325 秒 exhaustive autotuning 快很多了。


想在本地跑 FlashLib 而不花 cloud GPU 費用,24GB VRAM 的 麗臺 RTX PRO 4000 Blackwell 是搭配 H200 benchmark 工作負載以外的合理起點——FP16 的向量計算和 Triton kernel 都跑得起來。


FlashLib 和 cuML 的差異

很多人看到 FlashLib 的數字第一個念頭是:NVIDIA 的 cuML 不就在做這件事嗎?

是,也不是。

cuML 是 NVIDIA RAPIDS 生態的一部分,官方出品,目標是讓你的 scikit-learn 程式碼不改一行就在 GPU 上跑。它的加速比較基準是 CPU——在 H100 上,一般算法 10-50x,HDBSCAN 可以達到 175x vs CPU。一行指令就能啟用:

%load_ext cuml.accel
# 之後你的 sklearn 程式碼自動走 GPU

FlashLib 的基準不是 CPU,是 cuML 本身。它假設你已經在 GPU 上跑了,但覺得還不夠快。這兩個在解不同的問題:

cuMLFlashLib
解決的問題CPU 太慢,搬到 GPUGPU 還有餘裕,往硬體上限推
加速基準vs CPUvs cuML
API 相容性sklearn drop-in,零改動需要換 API 呼叫
算法覆蓋廣,涵蓋大部分 sklearn 算法目前 6-7 個
零改動啟用%load_ext cuml.accel不支援
出品方NVIDIA 官方學術團隊(MIT/UCB)

說直接一點:cuML 是遷移工具,FlashLib 是調優工具。

如果你的現有程式碼跑在 CPU 上,第一步是 cuML,不是 FlashLib。

如果你已經在 GPU 上、pipeline 裡的 K-Means 或 t-SNE 還是瓶頸,那時候 FlashLib 才有意義。兩者也可以共存——其他算法繼續走 cuML,把最慢的那幾個換成 FlashLib。

值得注意的是,cuML 對應的算法加速,其實也內建了部分 IO 優化,只是沒有 FlashLib 徹底。FlashLib 的 208x 是在 cuML 已經是 GPU 實作的基礎上再乘,換算到 vs CPU 的場景,差距更大。


這說明了一件事

GPU 的問題,從來不是算力不夠,是你沒辦法餵夠快。

Flash Attention 在 transformer 領域證明了這件事,Flash-KMeans 把它帶到 clustering,FlashLib 把它系統化成一套設計方法論套用到 classical ML。

208 倍聽起來離譜,但算一下就合理了:如果你的算法 90% 時間都在等記憶體,解掉這個瓶頸給你 10 倍不奇怪,解到底線有機會超過 100 倍。

下一個問題是:除了 KMeans 和 t-SNE,還有哪些算法的瓶頸也是這個?