From 92e9755c097a4287709320d11182ec3f4a9705b8 Mon Sep 17 00:00:00 2001 From: Ben Landrum Date: Tue, 4 Aug 2026 21:15:01 +0000 Subject: [PATCH 1/2] Factor CAGRA graph kernels into shared translation unit --- cpp/CMakeLists.txt | 1 + cpp/src/neighbors/detail/cagra/graph_core.cuh | 178 +--------------- .../neighbors/detail/cagra/graph_shared.cu | 190 ++++++++++++++++++ .../neighbors/detail/cagra/graph_shared.cuh | 46 +++++ 4 files changed, 248 insertions(+), 167 deletions(-) create mode 100644 cpp/src/neighbors/detail/cagra/graph_shared.cu create mode 100644 cpp/src/neighbors/detail/cagra/graph_shared.cuh diff --git a/cpp/CMakeLists.txt b/cpp/CMakeLists.txt index 1a125fa71e..088c5f689c 100644 --- a/cpp/CMakeLists.txt +++ b/cpp/CMakeLists.txt @@ -1386,6 +1386,7 @@ if(NOT BUILD_CPU_ONLY) ${cagra_build_inst_files} ${cagra_extend_inst_files} src/neighbors/cagra_optimize.cu + src/neighbors/detail/cagra/graph_shared.cu ${cagra_serialize_inst_files} ${cagra_merge_inst_files} ${iface_cagra_inst_files} diff --git a/cpp/src/neighbors/detail/cagra/graph_core.cuh b/cpp/src/neighbors/detail/cagra/graph_core.cuh index 52b4542798..5c5c2d5920 100644 --- a/cpp/src/neighbors/detail/cagra/graph_core.cuh +++ b/cpp/src/neighbors/detail/cagra/graph_core.cuh @@ -1,10 +1,11 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ #pragma once #include "cagra_helpers.hpp" +#include "graph_shared.cuh" #include "utils.hpp" #include @@ -23,7 +24,6 @@ #include -#include #include #include @@ -39,6 +39,7 @@ #include #include #include +#include namespace cg = cooperative_groups; @@ -53,127 +54,6 @@ inline double cur_time(void) return ((double)tv.tv_sec + (double)tv.tv_usec * 1e-6); } -template -__device__ inline void swap(T& val1, T& val2) -{ - T val0 = val1; - val1 = val2; - val2 = val0; -} - -template -__device__ inline bool swap_if_needed(K& key1, K& key2, V& val1, V& val2, bool ascending) -{ - if (key1 == key2) { return false; } - if ((key1 > key2) == ascending) { - swap(key1, key2); - swap(val1, val2); - return true; - } - return false; -} - -template -__global__ void kern_sort(const DATA_T* const dataset, // [dataset_chunk_size, dataset_dim] - const IdxT dataset_size, - const uint32_t dataset_dim, - IdxT* const knn_graph, // [graph_chunk_size, graph_degree] - const uint32_t graph_size, - const uint32_t graph_degree, - const cuvs::distance::DistanceType metric) -{ - const IdxT srcNode = (blockDim.x * blockIdx.x + threadIdx.x) / raft::WarpSize; - if (srcNode >= graph_size) { return; } - - const uint32_t lane_id = threadIdx.x % raft::WarpSize; - - float my_keys[numElementsPerThread]; - IdxT my_vals[numElementsPerThread]; - - // Compute distance from a src node to its neighbors - for (int k = 0; k < graph_degree; k++) { - const IdxT dstNode = knn_graph[k + static_cast(graph_degree) * srcNode]; - float dist = 0; - float norm2_dst = 0; - if (metric == cuvs::distance::DistanceType::InnerProduct || - metric == cuvs::distance::DistanceType::CosineExpanded) { - for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { - auto elem_b = cuvs::spatial::knn::detail::utils::mapping{}( - dataset[d + static_cast(dataset_dim) * dstNode]); - dist -= cuvs::spatial::knn::detail::utils::mapping{}( - dataset[d + static_cast(dataset_dim) * srcNode]) * - elem_b; - - if (metric == cuvs::distance::DistanceType::CosineExpanded) { - norm2_dst += elem_b * elem_b; - } - } - } else if (metric == cuvs::distance::DistanceType::L2Expanded) { - // L2Expanded - for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { - float diff = cuvs::spatial::knn::detail::utils::mapping{}( - dataset[d + static_cast(dataset_dim) * srcNode]) - - cuvs::spatial::knn::detail::utils::mapping{}( - dataset[d + static_cast(dataset_dim) * dstNode]); - dist += diff * diff; - } - } else if (metric == cuvs::distance::DistanceType::L1) { - for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { - float diff = cuvs::spatial::knn::detail::utils::mapping{}( - dataset[d + static_cast(dataset_dim) * srcNode]) - - cuvs::spatial::knn::detail::utils::mapping{}( - dataset[d + static_cast(dataset_dim) * dstNode]); - dist += raft::abs(diff); - } - } else if (metric == cuvs::distance::DistanceType::BitwiseHamming) { - if constexpr (std::is_integral_v) { - for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { - dist += __popc( - static_cast(dataset[d + static_cast(dataset_dim) * srcNode] ^ - dataset[d + static_cast(dataset_dim) * dstNode]) & - 0xffu); - } - } - } - dist += __shfl_xor_sync(0xffffffff, dist, 1); - dist += __shfl_xor_sync(0xffffffff, dist, 2); - dist += __shfl_xor_sync(0xffffffff, dist, 4); - dist += __shfl_xor_sync(0xffffffff, dist, 8); - dist += __shfl_xor_sync(0xffffffff, dist, 16); - - if (metric == cuvs::distance::DistanceType::CosineExpanded) { - norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 1); - norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 2); - norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 4); - norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 8); - norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 16); - if (lane_id == (k % raft::WarpSize)) { dist /= sqrt(norm2_dst); } - } - - if (lane_id == (k % raft::WarpSize)) { - my_keys[k / raft::WarpSize] = dist; - my_vals[k / raft::WarpSize] = dstNode; - } - } - for (int k = graph_degree; k < raft::WarpSize * numElementsPerThread; k++) { - if (lane_id == k % raft::WarpSize) { - my_keys[k / raft::WarpSize] = utils::get_max_value(); - my_vals[k / raft::WarpSize] = utils::get_max_value(); - } - } - - // Sort by RAFT bitonic sort - raft::util::bitonic(true).sort(my_keys, my_vals); - - // Update knn_graph - for (int i = 0; i < numElementsPerThread; i++) { - const int k = i * raft::WarpSize + lane_id; - if (k < graph_degree) { - knn_graph[k + (static_cast(graph_degree) * srcNode)] = my_vals[i]; - } - } -} - template __global__ void kern_make_rev_graph_k( OutputMatrixView output_graph, // [graph_size, degree] @@ -983,6 +863,7 @@ void sort_knn_graph( raft::mdspan, raft::row_major, d_accessor> dataset, raft::mdspan, raft::row_major, g_accessor> knn_graph) { + static_assert(std::is_same_v, "CAGRA graph indices must be uint32_t"); RAFT_EXPECTS(dataset.extent(0) == knn_graph.extent(0), "dataset size is expected to have the same number of graph index size"); RAFT_EXPECTS( @@ -1018,51 +899,14 @@ void sort_knn_graph( raft::copy(res, d_input_graph.view(), knn_graph); - void (*kernel_sort)(const DataT* const, - const IdxT, - const uint32_t, - IdxT* const, - const uint32_t, - const uint32_t, - const cuvs::distance::DistanceType); - if (input_graph_degree <= 32) { - constexpr int numElementsPerThread = 1; - kernel_sort = kern_sort; - } else if (input_graph_degree <= 64) { - constexpr int numElementsPerThread = 2; - kernel_sort = kern_sort; - } else if (input_graph_degree <= 128) { - constexpr int numElementsPerThread = 4; - kernel_sort = kern_sort; - } else if (input_graph_degree <= 256) { - constexpr int numElementsPerThread = 8; - kernel_sort = kern_sort; - } else if (input_graph_degree <= 512) { - constexpr int numElementsPerThread = 16; - kernel_sort = kern_sort; - } else if (input_graph_degree <= 1024) { - constexpr int numElementsPerThread = 32; - kernel_sort = kern_sort; - } else { - RAFT_FAIL( - "The degree of input knn graph is too large (%lu). " - "It must be equal to or smaller than %d.", - input_graph_degree, - 1024); - } - const auto block_size = 256; - const auto num_warps_per_block = block_size / raft::WarpSize; - const auto grid_size = (graph_size + num_warps_per_block - 1) / num_warps_per_block; - RAFT_LOG_DEBUG("."); - kernel_sort<<>>( - d_dataset.data_handle(), - dataset_size, - dataset_dim, - d_input_graph.data_handle(), - graph_size, - input_graph_degree, - metric); + launch_sort_knn_graph(res, + metric, + d_dataset.data_handle(), + static_cast(dataset_size), + static_cast(dataset_dim), + d_input_graph.data_handle(), + static_cast(input_graph_degree)); raft::resource::sync_stream(res); RAFT_LOG_DEBUG("."); raft::copy(res, knn_graph, raft::make_const_mdspan(d_input_graph.view())); diff --git a/cpp/src/neighbors/detail/cagra/graph_shared.cu b/cpp/src/neighbors/detail/cagra/graph_shared.cu new file mode 100644 index 0000000000..3ed3d9c4b9 --- /dev/null +++ b/cpp/src/neighbors/detail/cagra/graph_shared.cu @@ -0,0 +1,190 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "graph_core.cuh" +#include "graph_shared.cuh" +#include "utils.hpp" + +// TODO: This shouldn't be invoking anything from spatial/knn +#include "../ann_utils.cuh" + +#include +#include +#include + +#include + +namespace cuvs::neighbors::cagra::detail::graph { +namespace { + +template +__global__ void kern_sort(const DATA_T* const dataset, // [dataset_chunk_size, dataset_dim] + const uint32_t dataset_dim, + uint32_t* const knn_graph, // [graph_chunk_size, graph_degree] + const uint32_t graph_size, + const uint32_t graph_degree, + const cuvs::distance::DistanceType metric) +{ + const uint32_t srcNode = (blockDim.x * blockIdx.x + threadIdx.x) / raft::WarpSize; + if (srcNode >= graph_size) { return; } + + const uint32_t lane_id = threadIdx.x % raft::WarpSize; + + float my_keys[numElementsPerThread]; + uint32_t my_vals[numElementsPerThread]; + + // Compute distance from a src node to its neighbors + for (int k = 0; k < graph_degree; k++) { + const uint32_t dstNode = knn_graph[k + static_cast(graph_degree) * srcNode]; + float dist = 0; + float norm2_dst = 0; + if (metric == cuvs::distance::DistanceType::InnerProduct || + metric == cuvs::distance::DistanceType::CosineExpanded) { + for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { + auto elem_b = cuvs::spatial::knn::detail::utils::mapping{}( + dataset[d + static_cast(dataset_dim) * dstNode]); + dist -= cuvs::spatial::knn::detail::utils::mapping{}( + dataset[d + static_cast(dataset_dim) * srcNode]) * + elem_b; + + if (metric == cuvs::distance::DistanceType::CosineExpanded) { + norm2_dst += elem_b * elem_b; + } + } + } else if (metric == cuvs::distance::DistanceType::L2Expanded) { + for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { + float diff = cuvs::spatial::knn::detail::utils::mapping{}( + dataset[d + static_cast(dataset_dim) * srcNode]) - + cuvs::spatial::knn::detail::utils::mapping{}( + dataset[d + static_cast(dataset_dim) * dstNode]); + dist += diff * diff; + } + } else if (metric == cuvs::distance::DistanceType::L1) { + for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { + float diff = cuvs::spatial::knn::detail::utils::mapping{}( + dataset[d + static_cast(dataset_dim) * srcNode]) - + cuvs::spatial::knn::detail::utils::mapping{}( + dataset[d + static_cast(dataset_dim) * dstNode]); + dist += raft::abs(diff); + } + } else if (metric == cuvs::distance::DistanceType::BitwiseHamming) { + if constexpr (std::is_integral_v) { + for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) { + dist += __popc( + static_cast(dataset[d + static_cast(dataset_dim) * srcNode] ^ + dataset[d + static_cast(dataset_dim) * dstNode]) & + 0xffu); + } + } + } + dist += __shfl_xor_sync(0xffffffff, dist, 1); + dist += __shfl_xor_sync(0xffffffff, dist, 2); + dist += __shfl_xor_sync(0xffffffff, dist, 4); + dist += __shfl_xor_sync(0xffffffff, dist, 8); + dist += __shfl_xor_sync(0xffffffff, dist, 16); + + if (metric == cuvs::distance::DistanceType::CosineExpanded) { + norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 1); + norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 2); + norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 4); + norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 8); + norm2_dst += __shfl_xor_sync(0xffffffff, norm2_dst, 16); + if (lane_id == (k % raft::WarpSize)) { dist /= sqrt(norm2_dst); } + } + + if (lane_id == (k % raft::WarpSize)) { + my_keys[k / raft::WarpSize] = dist; + my_vals[k / raft::WarpSize] = dstNode; + } + } + for (int k = graph_degree; k < raft::WarpSize * numElementsPerThread; k++) { + if (lane_id == k % raft::WarpSize) { + my_keys[k / raft::WarpSize] = utils::get_max_value(); + my_vals[k / raft::WarpSize] = utils::get_max_value(); + } + } + + raft::util::bitonic(true).sort(my_keys, my_vals); + + for (int i = 0; i < numElementsPerThread; i++) { + const int k = i * raft::WarpSize + lane_id; + if (k < graph_degree) { + knn_graph[k + (static_cast(graph_degree) * srcNode)] = my_vals[i]; + } + } +} + +constexpr int kMaxSortElementsPerThread = 32; + +template +using sort_kernel_type = + void (*)(DataT const*, uint32_t, uint32_t*, uint32_t, uint32_t, cuvs::distance::DistanceType); + +template +auto select_sort_kernel(uint32_t degree) -> sort_kernel_type +{ + if (degree <= raft::WarpSize * 1) { return kern_sort; } + if (degree <= raft::WarpSize * 2) { return kern_sort; } + if (degree <= raft::WarpSize * 4) { return kern_sort; } + if (degree <= raft::WarpSize * 8) { return kern_sort; } + if (degree <= raft::WarpSize * 16) { return kern_sort; } + if (degree <= kMaxSortDegree) { return kern_sort; } + RAFT_FAIL( + "The degree of input knn graph is too large (%u). It must be equal to or smaller than %lu.", + degree, + kMaxSortDegree); +} + +template +void launch_sort_knn_graph_impl(raft::resources const& res, + cuvs::distance::DistanceType metric, + DataT const* dataset, + uint32_t dataset_size, + uint32_t dataset_dim, + uint32_t* knn_graph, + uint32_t graph_degree) +{ + auto kernel = select_sort_kernel(graph_degree); + + constexpr uint32_t block_size = 256; + auto const warps = block_size / raft::WarpSize; + auto const blocks = (dataset_size + warps - 1) / warps; + kernel<<>>( + dataset, dataset_dim, knn_graph, dataset_size, graph_degree, metric); + RAFT_CUDA_TRY(cudaGetLastError()); +} + +} // namespace + +#define CUVS_DEFINE_CAGRA_GRAPH_SORT(DataT) \ + void launch_sort_knn_graph(raft::resources const& res, \ + cuvs::distance::DistanceType metric, \ + DataT const* dataset, \ + uint32_t dataset_size, \ + uint32_t dataset_dim, \ + uint32_t* knn_graph, \ + uint32_t graph_degree) \ + { \ + launch_sort_knn_graph_impl( \ + res, metric, dataset, dataset_size, dataset_dim, knn_graph, graph_degree); \ + } + +CUVS_DEFINE_CAGRA_GRAPH_SORT(float) +CUVS_DEFINE_CAGRA_GRAPH_SORT(half) +CUVS_DEFINE_CAGRA_GRAPH_SORT(int8_t) +CUVS_DEFINE_CAGRA_GRAPH_SORT(uint8_t) + +#undef CUVS_DEFINE_CAGRA_GRAPH_SORT + +void optimize_device_graph( + raft::resources const& res, + raft::device_matrix_view knn_graph, + raft::device_matrix_view output_graph, + bool guarantee_connectivity) +{ + optimize(res, knn_graph, output_graph, guarantee_connectivity); +} + +} // namespace cuvs::neighbors::cagra::detail::graph diff --git a/cpp/src/neighbors/detail/cagra/graph_shared.cuh b/cpp/src/neighbors/detail/cagra/graph_shared.cuh new file mode 100644 index 0000000000..0f2bc0f7dc --- /dev/null +++ b/cpp/src/neighbors/detail/cagra/graph_shared.cuh @@ -0,0 +1,46 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include +#include + +#include +#include + +#include + +#include + +namespace cuvs::neighbors::cagra::detail::graph { + +// kern_sort capacity is WarpSize * numElementsPerThread; largest specialization uses 32. +inline constexpr uint64_t kMaxSortDegree = 32 * 32; + +#define CUVS_DECL_CAGRA_GRAPH_SORT(DataT) \ + CUVS_EXPORT void launch_sort_knn_graph(raft::resources const& res, \ + cuvs::distance::DistanceType, \ + DataT const* dataset, \ + uint32_t dataset_size, \ + uint32_t dataset_dim, \ + uint32_t* knn_graph, \ + uint32_t graph_degree) + +CUVS_DECL_CAGRA_GRAPH_SORT(float); +CUVS_DECL_CAGRA_GRAPH_SORT(half); +CUVS_DECL_CAGRA_GRAPH_SORT(int8_t); +CUVS_DECL_CAGRA_GRAPH_SORT(uint8_t); + +#undef CUVS_DECL_CAGRA_GRAPH_SORT + +/** Run the existing CAGRA optimizer through one compiled instantiation instead of rematerializing + * its reverse-graph, prune, merge, and MST kernels in every Fastener dtype TU. */ +CUVS_EXPORT void optimize_device_graph( + raft::resources const& res, + raft::device_matrix_view knn_graph, + raft::device_matrix_view output_graph, + bool guarantee_connectivity); + +} // namespace cuvs::neighbors::cagra::detail::graph From 1bd1636a2541f86c52a58c86b44ee4969bb61b66 Mon Sep 17 00:00:00 2001 From: Ben Landrum Date: Tue, 4 Aug 2026 22:29:42 +0000 Subject: [PATCH 2/2] removed needless exports --- .../neighbors/detail/cagra/graph_shared.cuh | 19 +++++++++---------- 1 file changed, 9 insertions(+), 10 deletions(-) diff --git a/cpp/src/neighbors/detail/cagra/graph_shared.cuh b/cpp/src/neighbors/detail/cagra/graph_shared.cuh index 0f2bc0f7dc..08ccc3108f 100644 --- a/cpp/src/neighbors/detail/cagra/graph_shared.cuh +++ b/cpp/src/neighbors/detail/cagra/graph_shared.cuh @@ -4,7 +4,6 @@ */ #pragma once -#include #include #include @@ -19,14 +18,14 @@ namespace cuvs::neighbors::cagra::detail::graph { // kern_sort capacity is WarpSize * numElementsPerThread; largest specialization uses 32. inline constexpr uint64_t kMaxSortDegree = 32 * 32; -#define CUVS_DECL_CAGRA_GRAPH_SORT(DataT) \ - CUVS_EXPORT void launch_sort_knn_graph(raft::resources const& res, \ - cuvs::distance::DistanceType, \ - DataT const* dataset, \ - uint32_t dataset_size, \ - uint32_t dataset_dim, \ - uint32_t* knn_graph, \ - uint32_t graph_degree) +#define CUVS_DECL_CAGRA_GRAPH_SORT(DataT) \ + void launch_sort_knn_graph(raft::resources const& res, \ + cuvs::distance::DistanceType, \ + DataT const* dataset, \ + uint32_t dataset_size, \ + uint32_t dataset_dim, \ + uint32_t* knn_graph, \ + uint32_t graph_degree) CUVS_DECL_CAGRA_GRAPH_SORT(float); CUVS_DECL_CAGRA_GRAPH_SORT(half); @@ -37,7 +36,7 @@ CUVS_DECL_CAGRA_GRAPH_SORT(uint8_t); /** Run the existing CAGRA optimizer through one compiled instantiation instead of rematerializing * its reverse-graph, prune, merge, and MST kernels in every Fastener dtype TU. */ -CUVS_EXPORT void optimize_device_graph( +void optimize_device_graph( raft::resources const& res, raft::device_matrix_view knn_graph, raft::device_matrix_view output_graph,