svg-project/flash-kmeans

Fast and memory-efficient exact kmeans

解決的問題

Flash-KMeans 提供了 K-Means 聚類演算法的高效率、記憶體節省的實作。它解決了在 GPU 上處理大型資料集(大 N)或高維資料(大 D)時常見的記憶體不足(OOM)錯誤與運算速度緩慢的問題,避免產生大型距離矩陣。

工作原理

本專案使用 Triton GPU 內核實作一種 IO 友好的批次處理 K-Means。根據資料維度採用兩種主要執行路徑:

  • 小 D 路徑:針對維度 $\le 512$ 的情況優化,使用針對特定 GPU 架構(H200、H100、A100、GB10)的手動調校啟發式方法。
  • Split-D 路徑:用於維度 $> 512$ 或共享記憶體受限的情況,透過分塊維度迴圈來維持 K-串流特性。

對於單一 GPU 無法容納的資料集,實作了雙緩衝串流設計,將資料從 CPU 分塊傳輸至 GPU。同時透過將資料分區至多個 GPU,並使用輕量級手動 gather-reduce-broadcast 機制進行中心點更新,支援多 GPU 擴展,避免對 NCCL 的依賴。

適用對象

專為處理大規模聚類任務的研究人員與開發者設計,特別適用於實作 Sparse VideoGen2 等系統,或任何需要跨多 GPU 擴展的快速、精確 K-Means 實作的使用者。

主要亮點

  • 基於 Triton 的加速:相比標準 PyTorch 及其他 Triton 實作,性能顯著提升。
  • 記憶體效率:透過避免產生完整距離矩陣,防止 OOM 錯誤。
  • 自動分派:根據輸入形狀與資料類型自動切換 Small-D 與 Split-D 內核。
  • 多 GPU 擴展:實現 PCIe 帶寬的線性擴展,並支援 H2D 傳輸與中心點歸約的重疊。
  • 廣泛的硬體支援:包含針對現代 NVIDIA GPU 的調校設定,並為未知架構提供保守回退方案。

相關

  • 專案
  • Dispatch
  • 專案
  • 專案
  • Dispatch