svg-project/flash-kmeans

Fast and memory-efficient exact kmeans

何を解決するか

Flash-KMeans は、K-Means クラスタリングアルゴリズムの高性能かつメモリ効率の良い実装を提供します。GPU上で大規模なデータセット(N が大きい)や高次元データ(D が大きい)を扱う際に発生する一般的な問題、すなわちメモリ不足(OOM)エラーと遅い計算速度を回避し、大きな距離行列の生成を避けています。

動作方法

このプロジェクトは、Triton GPU カーネルを使用して、IOに配慮したバッチ処理 K-Means を実装しています。データの次元に応じて、2つの主要な実行パスを採用しています:

  • Small-D パス:次元数が $\le 512$ の場合に最適化され、特定のGPUアーキテクチャ(H200、H100、A100、GB10)向けに手動でチューニングされたヒューリスティクスを使用します。
  • Split-D パス:次元数が $> 512$ または共有メモリが限界に達した場合に使用され、次元ループをタイル化して K-ストリーミング性を維持します。

単一のGPUに収まらないデータセットに対しては、CPUからGPUへデータをチャンク単位で転送する二重バッファリングストリーミング設計を実装しています。また、データをGPU間で分割し、重心更新に軽量な手動の gather-reduce-broadcast メカニズムを使用することで、NCCL依存を回避しながらマルチGPUスケーリングをサポートしています。

対象ユーザー

大規模なクラスタリングタスクに取り組む研究者や開発者向けに設計されています。特に Sparse VideoGen2 のようなシステムの実装に取り組む人、または複数GPUにスケーリング可能な高速かつ正確な K-Means 実装が必要な人にとって最適です。

特徴

  • Tritonベースの加速:標準の PyTorch や他の Triton 実装と比較して顕著なパフォーマンス向上。
  • メモリ効率:距離行列の完全な生成を回避することで、OOMを防止。
  • 自動ディスパッチ:入力の形状と dtype に基づいて、Small-D と Split-D カーネルの自動切り替え。
  • マルチGPUスケーリング:PCIe帯域幅の線形スケーリングと、H2D転送と重心削減のオーバーラップ。
  • 広範なハードウェア対応:最新の NVIDIA GPU 用にチューニングされた設定を含み、未知のアーキテクチャには保守的なフォールバックを備えています。

関連

  • プロジェクト
  • Dispatch
  • プロジェクト
  • プロジェクト
  • Dispatch