跳到主要內容
GPU KMeans 跑 10 秒,FlashLib 版本 0.28 秒——差距來自記憶體,不是算法

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

Flash-KMeans 已經做到讓 KMeans 比 FAISS 快 200 倍、比 cuML 快 33 倍。現在同個團隊推出 FlashLib,把同樣的 IO-aware 設計套到六個算法:KMeans 26x、t-SNE 147x、TruncatedSVD 208x。瓶頸從來不在算法,在記憶體。

目錄+

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

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

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


Flash Attention 是怎麼走到這裡的

要理解 FlashLib 為什麼能做到這個數字,要先理解它的設計哲學從哪來。

2022 年,Stanford 的 Tri Dao 提出 Flash Attention,解決了 transformer self-attention 的效能瓶頸。標準 attention 的問題和 K-Means 一樣:計算過程會產生一個 N×N 的 attention matrix,要寫進 HBM 再讀出來。Flash Attention 的做法是把矩陣分成小塊,用 tiling 的方式讓計算全程在 on-chip SRAM 裡進行,完全繞過這個大型中間矩陣。

這個設計讓 attention 的記憶體複雜度從 O(N²) 降到 O(N),速度快了 2-4 倍,記憶體用量縮小了 5-20 倍。更重要的是,它沒有改變任何算法——輸入輸出完全一樣,只是換了一個更聰明的執行方式。

FlashLib 的核心主張是:這不是 attention 專屬的技巧,而是一個可以套用到任何「中間結果是瓶頸」算法的工程方法論。

K-Means、KNN、HDBSCAN、TruncatedSVD、PCA、t-SNE——這些算法在 GPU 上跑慢,不是因為計算量大,而是因為都有類似的問題:中間產生的矩陣或向量需要寫進 HBM,然後再讀回來。把這個環節消掉,就能大幅提升速度。


為什麼 GPU K-Means 這麼慢

具體來說,K-Means 的標準 GPU 實作是這樣的:為每個資料點計算到所有 cluster centroid 的距離,把這個 N×K 距離矩陣寫到 HBM,再讀回來找最近的 cluster。

聽起來合理,但測量一下就知道問題在哪。

在 H200 上,N=1M、K=64K 的工作負載:

  • 距離計算本身:2.6 毫秒
  • 把 N×K 矩陣寫入 HBM:23 毫秒

記憶體傳輸比實際計算慢了將近 9 倍。這不是邊際問題,是整個 pipeline 的主要瓶頸。

FlashAssign 的解法:完全不寫那個矩陣。

Loading diagram...

FlashAssign 把 centroid 集合分成小塊(tile),逐批載入 on-chip SRAM buffer。在 SRAM 上同時計算當前 tile 的距離,並維護一個 running minimum——不需要等整個矩陣算完,也不需要把中間結果存到 HBM。整個過程只在最後把最終的 cluster assignment 寫一次到 HBM。

在 centroid 更新階段,FlashLib 解決了另一個問題:多個資料點同時更新同一個 centroid 時,GPU 需要用 atomic scatter 做累加,這會造成大量的 write contention。

Sort-Inverse Update 的做法是先對資料點按照 cluster 排序,建立一個反向映射,把原本的 high-contention atomic write 轉換成 segment-level 的有序 reduction。相同 cluster 的點連續存放,GPU 可以用高頻寬的向量化操作處理,而不是一堆互相競爭的 atomic。這讓 centroid 更新快了 6.3 倍。


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

Flash-KMeans 的設計被驗證之後,同個團隊把 IO-aware kernel 設計系統化,擴展到六個算法。

FlashLib 的 benchmark 全部在單張 NVIDIA H200(SM90, 150 GB HBM3e)、CUDA 13.0、PyTorch 2.11、Triton 3.6、cuML 25.10 條件下測試:

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

來源:flashml-org.github.io

KMeans 和 KNN

這兩個算法都有「計算距離矩陣 → 找最近鄰」的結構,是 FlashAssign 最直接的應用。KNN 的 Flash-KNN 在 H200 上達到 85.2% 的 peak HBM 頻寬——接近硬體理論上限。Flash-KMeans 則達到 61% 的 peak FLOPs。這兩個數字很重要:代表 FlashLib 不只是比競品快,而是真正在逼近這張 GPU 的物理極限。

HDBSCAN

密度聚類算法,需要計算 core distance 和 mutual reachability distance。傳統 GPU 實作在這個步驟會產生大型中間矩陣,FlashLib 同樣用 tiling 策略把這部分留在 on-chip。40x vs cuML 的加速讓 HDBSCAN 在生產環境的大規模 log 分析和異常偵測場景變得更實際。

TruncatedSVD 和 PCA

這兩個算法的瓶頸在矩陣分解過程中的大型中間矩陣。TruncatedSVD 的 208x 是六個算法裡最大的加速倍數,原因是傳統實作在這個步驟對 HBM 的讀寫特別頻繁,FlashLib 的 IO-aware 設計在這裡空間最大。

PCA 的 47x 也很顯著。在大規模 feature 壓縮、embedding 降維的 pipeline 裡,PCA 往往不是你第一個想到要優化的地方,但它確實可以是瓶頸。

exact t-SNE

t-SNE 特別有意思。標準 GPU 實作通常用近似版(Barnes-Hut)來降低計算量,因為 exact t-SNE 需要計算所有點對之間的距離,複雜度是 O(N²)。FlashLib 的 exact t-SNE 在 147x 的加速下,讓「對百萬個 embedding 跑精確 t-SNE」這件事從幾乎不可行變成幾分鐘可以完成的任務。


Flash Informative API:在跑之前就知道要花多久

FlashLib 還做了一件其他 GPU ML 庫沒做的事:讓你在任何工作負載真正執行之前,就能預測它要花多久、用多少記憶體。

Flash Informative API 在約 5 微秒內(純 CPU 執行,不需要 GPU profiling)回傳:

  • 預估執行時間
  • 預估記憶體用量
  • 預估 overhead

這在生產系統裡意義很大。如果你的任務排程器需要決定哪個工作負載先跑、要配幾張 GPU,但每次都要先跑一遍才知道要花多久,那個 overhead 很可觀。Flash Informative API 讓排程決策可以在毫秒內完成。


怎麼裝

Flash-KMeans 可以獨立安裝:

pip install flash-kmeans

Flash-KMeans 的 Python API,輸入是 PyTorch tensor,直接丟進去:

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)

資料超過 VRAM 的場景(out-of-core),FlashLib 自動從 CPU pinned memory 分批轉移到 GPU,你不需要手動處理 chunking。實測 N=10⁹、K=32768 的極端工作負載,Flash-KMeans 完成一次 iteration 需要 41.4 秒,同場景的 fastkmeans baseline 是 261.8 秒。

完整 FlashLib 的安裝和文件在 flashml-org.github.io。

第一次執行有約 2.5 秒的 Triton JIT 編譯 overhead——這已經比原本需要 325 秒的 exhaustive autotuning 快很多了。後續呼叫會使用 cache,不會重新編譯。


適用場景與限制

FlashLib 在哪些情況最有用?

最理想的場景:

  • 大規模向量資料庫建索引(KMeans、KNN)
  • 影片生成 pipeline 的 token 聚類(Flash-KMeans 是 Sparse VideoGen2 的官方實作)
  • 百萬以上 embedding 的降維視覺化(t-SNE、PCA)
  • 生產環境需要頻繁重跑 ML 前處理的系統

限制要知道:

FlashLib 的 kernel 目前只針對以下 GPU 有完整調優:H200、H100、A100、GB10。如果你的 GPU 不在這個清單裡,FlashLib 會退回一個保守的 fallback 設定——功能不會出錯,但效能會比最優情況差。

另外,benchmark 數字是在 FP16 精度下測試的。FP32 場景的加速幅度會有所不同。

如果你的資料量不大、或 ML 步驟不是你 pipeline 的瓶頸,換 FlashLib 帶來的收益不會很明顯。先 profile 你的 pipeline,確認瓶頸真的在這些算法,再考慮是否值得遷移。

想在本地跑 FlashLib 不花 cloud GPU 費用,麗臺 RTX PRO 4000 Blackwell 的 24GB VRAM 支援 FP16 的 Triton kernel 執行,適合在非 H200 環境下測試和開發。


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 算法目前 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,還有哪些算法的瓶頸也是這個?