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