CUDA kernels for batched Euclidean K-Means, focused on the fp16 assignment hot
path on NVIDIA GPUs. This project is a libtorch/PyTorch CUDA implementation that
mirrors the public flash-kmeans Euclidean API while replacing the Triton
assignment/update path with hand-written CUDA kernels.
The current tuning target is RTX 4090 / Ada (sm_89), fp16, D=128, large
N and K.
This is an optimization-oriented CUDA port. The Euclidean fp16 path is the main supported fast path. Cosine, dot-product, and large-N CPU streaming are still provided by the original Triton project, not this CUDA port.
Recent RTX 4090 fp16 results, median-style local runs:
Shape (B, N, K, D) | Assign: CUDA | Assign: Triton | Assign Speedup |
|---|---|---|---|
(1, 32K, 256, 128) | 92 TFLOPS | 38 TFLOPS | 2.40x |
(1, 131K, 2048, 128) | 171 TFLOPS | 120 TFLOPS | 1.42x |
(1, 262K, 4096, 128) | 176 TFLOPS | 126 TFLOPS | 1.39x |
(1, 524K, 8192, 128) | 185 TFLOPS | 126 TFLOPS | 1.46x |
End-to-end K-Means includes assignment, sorting, centroid update, and finalize, so speedups are lower and vary by shape:
Shape (B, N, K, D) | CUDA ms/iter | Triton ms/iter | Speedup |
|---|---|---|---|
(1, 32K, 256, 128) | 0.129 | 0.508 | 3.94x |
(1, 131K, 2048, 128) | 0.489 | 1.056 | 2.16x |
(1, 262K, 4096, 128) | 1.667 | 2.786 | 1.67x |
(1, 524K, 8192, 128) | 6.179 | 9.157 | 1.48x |
Quality checks compare final objective/inertia, centroid drift, and label
disagreement against the Triton reference. The latest sampled exact inertia
deltas were within about +/-0.03% on the large fp16 shapes.
The implementation has three layers.
flash_kmeans_cuda/: Python package and CUDA/C++ sources.benchmarks/: assign, end-to-end, and quality comparison scripts.tests/: CUDA correctness and shape coverage.scripts/windows/: local Windows build, benchmark, and profiling helpers.docs/: C++ shared-library and maintainer notes.cmake/: CMake package config template.third_party/flash-kmeans: upstream Triton reference submodule.flash_kmeans_cuda/kmeans.py provides:
from flash_kmeans_cuda import batch_kmeans_Euclid
The public shape convention is (B, N, D) for points and (B, K, D) for
centroids. The loop mirrors the original Triton implementation:
The loop preallocates buffers across iterations, skips centroid-shift work when
tol <= 0, and skips x_sq setup for the D=128 fp16 raw-distance assignment
path where x_sq is row-constant and cannot affect argmin.
flash_kmeans_cuda/csrc/assign/assign_sm80.cu is the main tensor-core
assignment kernel for fp16/bf16 on Ampere+ GPUs.
Key points:
mma.sync.m16n8k16 tensor-core tiles.ldmatrix.x4 loads for both operands.cp.async staging of point and centroid tiles into shared memory.K>=256.N_TILES=2 for med/big/huge and
routes mega K>=8192 to N_TILES=4.assign_safe.cu is the non-tensor-core fallback for fp32 or unsupported shapes.
update_sorted.cu accumulates centroid sums/counts from sorted cluster IDs. For
fp16 D=128 and K>=256, the indexed update path consumes the sorted
permutation directly and avoids materializing a full sorted copy of x.
update_finalize.cu computes new_centroid = sum / count, preserving old
centroids for empty clusters.
There are two build surfaces:
setup.py builds flash_kmeans_cuda._C with nanobind.CMakeLists.txt builds flash_kmeans_cuda.dll /
libflash_kmeans_cuda.so with a libtorch-based public API.The C++ API is declared in:
#include <flash_kmeans_cuda/flash_kmeans_cuda.h>
It accepts and returns at::Tensor objects.
Requirements used for the current Windows development environment:
2.11.* CUDA wheeluvFrom a fresh checkout:
git clone --recursive https://github.com/OpsiClear/flash-kmeans-cuda.git
cd flash-kmeans-cuda
uv sync --locked --python 3.12
For an existing checkout, initialize the reference implementation submodule:
git submodule update --init --recursive
Build the Python extension on Windows:
cmd /c scripts\windows\run_exp_t.bat
Run tests:
uv run python -m pytest tests/ -q
This project exposes library interfaces rather than a graphical UI.
Python usage:
import torch
from flash_kmeans_cuda import batch_kmeans_Euclid
x = torch.randn(1, 32768, 128, device="cuda", dtype=torch.float16)
labels, centroids, n_iters = batch_kmeans_Euclid(
x,
n_clusters=256,
max_iters=10,
tol=0.0,
)
The C++ build is documented in docs/cpp_shared_library.md.
Short Windows build:
$env:TORCH_CUDA_ARCH_LIST = "8.9"
$TorchPrefix = uv run python -c "import torch; print(torch.utils.cmake_prefix_path)"
cmake -S . -B build-shared -G "Visual Studio 17 2022" -A x64 -DCMAKE_PREFIX_PATH="$TorchPrefix" -DCMAKE_CUDA_ARCHITECTURES=89
cmake --build build-shared --config Release --target flash_kmeans_cuda
cmake --install build-shared --config Release --prefix build-shared\install
External CMake project:
find_package(Torch REQUIRED)
find_package(flash_kmeans_cuda CONFIG REQUIRED)
add_executable(my_app main.cpp)
target_link_libraries(my_app PRIVATE flash_kmeans_cuda::flash_kmeans_cuda)
Example C++ call:
#include <flash_kmeans_cuda/flash_kmeans_cuda.h>
auto ids = fkc::euclid_assign(x, centroids, x_sq, c_sq);
For Linux releases, prefer an explicit ABI/versioned artifact such as
linux-x86_64-torch2.11-cu13-sm80_86_89_90. The shared library links against
libtorch, so consumers must match the Torch/CUDA runtime family and C++ ABI.
Assign-only benchmark:
uv run python benchmarks/bench_assign_vs_triton.py --shape mega --check-accuracy
End-to-end benchmark:
uv run python benchmarks/bench_vs_triton.py --batch-size 1 --num-points 524288 --num-clusters 8192 --dim 128 --max-iters 1 --dtype fp16
Quality comparison over several iterations:
uv run python benchmarks\quality_compare.py --shapes med big huge mega --iters 10 --dtype fp16 --sample-points 4096
Available benchmark shapes in the local scripts:
| Name | Shape (B, N, K, D) |
|---|---|
med | (1, 32768, 256, 128) |
big | (1, 131072, 2048, 128) |
huge | (1, 262144, 4096, 128) |
mega | (1, 524288, 8192, 128) |
CI runs on pushes to main / optimize/**, on pull requests, and on manual
dispatch. It checks Python packaging, scans source distributions for generated
artifacts, and compiles both the Linux CUDA Python wheel and Linux C++ shared
library package.
GitHub Actions builds release artifacts when a v* tag is pushed:
git tag v0.1.0
git push origin v0.1.0
The release workflow builds:
sm80/86/89/90It then creates or updates the matching GitHub Release using gh and the
repository GITHUB_TOKEN.
Manual rebuild/publish:
gh workflow run release.yml -f tag=v0.1.0 -f publish=true
Environment flags used by the launcher:
FKC_ASSIGN_FORCE_SAFE=1: force safe non-MMA assignment.FKC_NTILES=1|2|4: override persistent N-tile routing.FKC_ASSIGN_DEEP_TILE=1: force deep tile fallback.FKC_NARROW=1, FKC_WIDE3=1, FKC_W4=1: experimental tile variants.These flags are mainly for benchmarking and correctness A/B checks.
This repository uses the upstream Triton project as the API and correctness
reference. The vendored reference lives under third_party/flash-kmeans.
Original project:
If this CUDA port is useful in your work, cite the original Flash-KMeans work:
@article{yang2026flash,
title={Flash-KMeans: Fast and Memory-Efficient Exact K-Means},
author={Yang, Shuo and Xi, Haocheng and Zhao, Yilong and Li, Muyang and Fan, Xiaoze and Zhang, Jintao and Cai, Han and Lin, Yujun and Li, Xiuyu and Keutzer, Kurt and others},
journal={arXiv preprint arXiv:2603.09229},
year={2026}
}
@article{yang2025sparse,
title={Sparse VideoGen2: Accelerate Video Generation with Sparse Attention via Semantic-Aware Permutation},
author={Yang, Shuo and Xi, Haocheng and Zhao, Yilong and Li, Muyang and Zhang, Jintao and Cai, Han and Lin, Yujun and Li, Xiuyu and Xu, Chenfeng and Peng, Kelly and others},
journal={arXiv preprint arXiv:2505.18875},
year={2025}
}
The upstream Flash-KMeans project is MIT licensed. Check the repository license files before redistributing binaries that bundle or depend on PyTorch, CUDA, or third-party components.
44 commits
Cuda
47.2%
Python
34.1%
C++
13.1%
CMake
3.0%
Batchfile
2.2%
CUDA kernels for batched Euclidean K-Means, focused on the fp16 assignment hot
path on NVIDIA GPUs. This project is a libtorch/PyTorch CUDA implementation that
mirrors the public flash-kmeans Euclidean API while replacing the Triton
assignment/update path with hand-written CUDA kernels.
The current tuning target is RTX 4090 / Ada (sm_89), fp16, D=128, large
N and K.
This is an optimization-oriented CUDA port. The Euclidean fp16 path is the main supported fast path. Cosine, dot-product, and large-N CPU streaming are still provided by the original Triton project, not this CUDA port.
Recent RTX 4090 fp16 results, median-style local runs:
Shape (B, N, K, D) | Assign: CUDA | Assign: Triton | Assign Speedup |
|---|---|---|---|
(1, 32K, 256, 128) | 92 TFLOPS | 38 TFLOPS | 2.40x |
(1, 131K, 2048, 128) | 171 TFLOPS | 120 TFLOPS | 1.42x |
(1, 262K, 4096, 128) | 176 TFLOPS | 126 TFLOPS | 1.39x |
(1, 524K, 8192, 128) | 185 TFLOPS | 126 TFLOPS | 1.46x |
End-to-end K-Means includes assignment, sorting, centroid update, and finalize, so speedups are lower and vary by shape:
Shape (B, N, K, D) | CUDA ms/iter | Triton ms/iter | Speedup |
|---|---|---|---|
(1, 32K, 256, 128) | 0.129 | 0.508 | 3.94x |
(1, 131K, 2048, 128) | 0.489 | 1.056 | 2.16x |
(1, 262K, 4096, 128) | 1.667 | 2.786 | 1.67x |
(1, 524K, 8192, 128) | 6.179 | 9.157 | 1.48x |
Quality checks compare final objective/inertia, centroid drift, and label
disagreement against the Triton reference. The latest sampled exact inertia
deltas were within about +/-0.03% on the large fp16 shapes.
The implementation has three layers.
flash_kmeans_cuda/: Python package and CUDA/C++ sources.benchmarks/: assign, end-to-end, and quality comparison scripts.tests/: CUDA correctness and shape coverage.scripts/windows/: local Windows build, benchmark, and profiling helpers.docs/: C++ shared-library and maintainer notes.cmake/: CMake package config template.third_party/flash-kmeans: upstream Triton reference submodule.flash_kmeans_cuda/kmeans.py provides:
from flash_kmeans_cuda import batch_kmeans_Euclid
The public shape convention is (B, N, D) for points and (B, K, D) for
centroids. The loop mirrors the original Triton implementation:
The loop preallocates buffers across iterations, skips centroid-shift work when
tol <= 0, and skips x_sq setup for the D=128 fp16 raw-distance assignment
path where x_sq is row-constant and cannot affect argmin.
flash_kmeans_cuda/csrc/assign/assign_sm80.cu is the main tensor-core
assignment kernel for fp16/bf16 on Ampere+ GPUs.
Key points:
mma.sync.m16n8k16 tensor-core tiles.ldmatrix.x4 loads for both operands.cp.async staging of point and centroid tiles into shared memory.K>=256.N_TILES=2 for med/big/huge and
routes mega K>=8192 to N_TILES=4.assign_safe.cu is the non-tensor-core fallback for fp32 or unsupported shapes.
update_sorted.cu accumulates centroid sums/counts from sorted cluster IDs. For
fp16 D=128 and K>=256, the indexed update path consumes the sorted
permutation directly and avoids materializing a full sorted copy of x.
update_finalize.cu computes new_centroid = sum / count, preserving old
centroids for empty clusters.
There are two build surfaces:
setup.py builds flash_kmeans_cuda._C with nanobind.CMakeLists.txt builds flash_kmeans_cuda.dll /
libflash_kmeans_cuda.so with a libtorch-based public API.The C++ API is declared in:
#include <flash_kmeans_cuda/flash_kmeans_cuda.h>
It accepts and returns at::Tensor objects.
Requirements used for the current Windows development environment:
2.11.* CUDA wheeluvFrom a fresh checkout:
git clone --recursive https://github.com/OpsiClear/flash-kmeans-cuda.git
cd flash-kmeans-cuda
uv sync --locked --python 3.12
For an existing checkout, initialize the reference implementation submodule:
git submodule update --init --recursive
Build the Python extension on Windows:
cmd /c scripts\windows\run_exp_t.bat
Run tests:
uv run python -m pytest tests/ -q
This project exposes library interfaces rather than a graphical UI.
Python usage:
import torch
from flash_kmeans_cuda import batch_kmeans_Euclid
x = torch.randn(1, 32768, 128, device="cuda", dtype=torch.float16)
labels, centroids, n_iters = batch_kmeans_Euclid(
x,
n_clusters=256,
max_iters=10,
tol=0.0,
)
The C++ build is documented in docs/cpp_shared_library.md.
Short Windows build:
$env:TORCH_CUDA_ARCH_LIST = "8.9"
$TorchPrefix = uv run python -c "import torch; print(torch.utils.cmake_prefix_path)"
cmake -S . -B build-shared -G "Visual Studio 17 2022" -A x64 -DCMAKE_PREFIX_PATH="$TorchPrefix" -DCMAKE_CUDA_ARCHITECTURES=89
cmake --build build-shared --config Release --target flash_kmeans_cuda
cmake --install build-shared --config Release --prefix build-shared\install
External CMake project:
find_package(Torch REQUIRED)
find_package(flash_kmeans_cuda CONFIG REQUIRED)
add_executable(my_app main.cpp)
target_link_libraries(my_app PRIVATE flash_kmeans_cuda::flash_kmeans_cuda)
Example C++ call:
#include <flash_kmeans_cuda/flash_kmeans_cuda.h>
auto ids = fkc::euclid_assign(x, centroids, x_sq, c_sq);
For Linux releases, prefer an explicit ABI/versioned artifact such as
linux-x86_64-torch2.11-cu13-sm80_86_89_90. The shared library links against
libtorch, so consumers must match the Torch/CUDA runtime family and C++ ABI.
Assign-only benchmark:
uv run python benchmarks/bench_assign_vs_triton.py --shape mega --check-accuracy
End-to-end benchmark:
uv run python benchmarks/bench_vs_triton.py --batch-size 1 --num-points 524288 --num-clusters 8192 --dim 128 --max-iters 1 --dtype fp16
Quality comparison over several iterations:
uv run python benchmarks\quality_compare.py --shapes med big huge mega --iters 10 --dtype fp16 --sample-points 4096
Available benchmark shapes in the local scripts:
| Name | Shape (B, N, K, D) |
|---|---|
med | (1, 32768, 256, 128) |
big | (1, 131072, 2048, 128) |
huge | (1, 262144, 4096, 128) |
mega | (1, 524288, 8192, 128) |
CI runs on pushes to main / optimize/**, on pull requests, and on manual
dispatch. It checks Python packaging, scans source distributions for generated
artifacts, and compiles both the Linux CUDA Python wheel and Linux C++ shared
library package.
GitHub Actions builds release artifacts when a v* tag is pushed:
git tag v0.1.0
git push origin v0.1.0
The release workflow builds:
sm80/86/89/90It then creates or updates the matching GitHub Release using gh and the
repository GITHUB_TOKEN.
Manual rebuild/publish:
gh workflow run release.yml -f tag=v0.1.0 -f publish=true
Environment flags used by the launcher:
FKC_ASSIGN_FORCE_SAFE=1: force safe non-MMA assignment.FKC_NTILES=1|2|4: override persistent N-tile routing.FKC_ASSIGN_DEEP_TILE=1: force deep tile fallback.FKC_NARROW=1, FKC_WIDE3=1, FKC_W4=1: experimental tile variants.These flags are mainly for benchmarking and correctness A/B checks.
This repository uses the upstream Triton project as the API and correctness
reference. The vendored reference lives under third_party/flash-kmeans.
Original project:
If this CUDA port is useful in your work, cite the original Flash-KMeans work:
@article{yang2026flash,
title={Flash-KMeans: Fast and Memory-Efficient Exact K-Means},
author={Yang, Shuo and Xi, Haocheng and Zhao, Yilong and Li, Muyang and Fan, Xiaoze and Zhang, Jintao and Cai, Han and Lin, Yujun and Li, Xiuyu and Keutzer, Kurt and others},
journal={arXiv preprint arXiv:2603.09229},
year={2026}
}
@article{yang2025sparse,
title={Sparse VideoGen2: Accelerate Video Generation with Sparse Attention via Semantic-Aware Permutation},
author={Yang, Shuo and Xi, Haocheng and Zhao, Yilong and Li, Muyang and Zhang, Jintao and Cai, Han and Lin, Yujun and Li, Xiuyu and Xu, Chenfeng and Peng, Kelly and others},
journal={arXiv preprint arXiv:2505.18875},
year={2025}
}
The upstream Flash-KMeans project is MIT licensed. Check the repository license files before redistributing binaries that bundle or depend on PyTorch, CUDA, or third-party components.
44 commits
Cuda
47.2%
Python
34.1%
C++
13.1%
CMake
3.0%
Batchfile
2.2%