From 3291efacf99f349f84a7ac76a4db2aef4ae785c9 Mon Sep 17 00:00:00 2001 From: Dmitry Vesnin Date: Fri, 7 Aug 2026 20:41:27 +0300 Subject: [PATCH 1/2] fix knn_merge_parts for inner product --- .../cuvs/neighbors/knn_merge_parts.hpp | 23 ++++ cpp/src/neighbors/detail/knn_merge_parts.cuh | 71 ++++++----- cpp/src/neighbors/knn_merge_parts.cu | 44 ++++++- cpp/src/neighbors/mg/snmg.cuh | 11 +- cpp/tests/CMakeLists.txt | 2 +- cpp/tests/neighbors/knn_merge_parts.cu | 111 ++++++++++++++++++ 6 files changed, 226 insertions(+), 36 deletions(-) create mode 100644 cpp/tests/neighbors/knn_merge_parts.cu diff --git a/cpp/include/cuvs/neighbors/knn_merge_parts.hpp b/cpp/include/cuvs/neighbors/knn_merge_parts.hpp index 2236e03733..de60d0415d 100644 --- a/cpp/include/cuvs/neighbors/knn_merge_parts.hpp +++ b/cpp/include/cuvs/neighbors/knn_merge_parts.hpp @@ -23,7 +23,16 @@ namespace neighbors { * @param outK merged knn distance matrix * @param outV merged knn index matrix * @param translations mapping of index offsets for each partition + * @param select_min whether to retain the smallest or largest values. Overloads + * without this parameter retain the smallest values. */ +void knn_merge_parts(raft::resources const& res, + raft::device_matrix_view inK, + raft::device_matrix_view inV, + raft::device_matrix_view outK, + raft::device_matrix_view outV, + raft::device_vector_view translations, + bool select_min); void knn_merge_parts(raft::resources const& res, raft::device_matrix_view inK, raft::device_matrix_view inV, @@ -36,11 +45,25 @@ void knn_merge_parts(raft::resources const& res, raft::device_matrix_view outK, raft::device_matrix_view outV, raft::device_vector_view translations); +void knn_merge_parts(raft::resources const& res, + raft::device_matrix_view inK, + raft::device_matrix_view inV, + raft::device_matrix_view outK, + raft::device_matrix_view outV, + raft::device_vector_view translations, + bool select_min); void knn_merge_parts(raft::resources const& res, raft::device_matrix_view inK, raft::device_matrix_view inV, raft::device_matrix_view outK, raft::device_matrix_view outV, raft::device_vector_view translations); +void knn_merge_parts(raft::resources const& res, + raft::device_matrix_view inK, + raft::device_matrix_view inV, + raft::device_matrix_view outK, + raft::device_matrix_view outV, + raft::device_vector_view translations, + bool select_min); } // namespace neighbors } // namespace CUVS_EXPORT cuvs diff --git a/cpp/src/neighbors/detail/knn_merge_parts.cuh b/cpp/src/neighbors/detail/knn_merge_parts.cuh index f8e1950f98..839f714374 100644 --- a/cpp/src/neighbors/detail/knn_merge_parts.cuh +++ b/cpp/src/neighbors/detail/knn_merge_parts.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023-2025, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION. * SPDX-License-Identifier: Apache-2.0 */ @@ -16,11 +16,7 @@ namespace cuvs::neighbors::detail { -template +template RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK, const value_idx* inV, value_t* outK, @@ -43,7 +39,7 @@ RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK, cuvs::neighbors::detail::faiss_select::BlockSelect< value_t, value_idx, - false, + !SelectMin, cuvs::neighbors::detail::faiss_select::Comparator, warp_q, thread_q, @@ -95,7 +91,7 @@ RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK, } } -template +template inline void knn_merge_parts_impl(const value_t* inK, const value_idx* inV, value_t* outK, @@ -111,9 +107,9 @@ inline void knn_merge_parts_impl(const value_t* inK, constexpr int n_threads = (warp_q < 1024) ? 128 : 64; auto block = dim3(n_threads); - auto kInit = std::numeric_limits::max(); + auto kInit = SelectMin ? raft::upper_bound() : raft::lower_bound(); auto vInit = -1; - knn_merge_parts_kernel + knn_merge_parts_kernel <<>>( inK, inV, outK, outV, n_samples, n_parts, kInit, vInit, k, translations); RAFT_CUDA_TRY(cudaPeekAtLastError()); @@ -133,39 +129,60 @@ inline void knn_merge_parts_impl(const value_t* inK, * @param stream CUDA stream to use * @param translations mapping of index offsets for each partition */ -template -inline void knn_merge_parts(const value_t* inK, - const value_idx* inV, - value_t* outK, - value_idx* outV, - size_t n_samples, - int n_parts, - int k, - cudaStream_t stream, - value_idx* translations) +template +inline void knn_merge_parts_dispatch(const value_t* inK, + const value_idx* inV, + value_t* outK, + value_idx* outV, + size_t n_samples, + int n_parts, + int k, + cudaStream_t stream, + value_idx* translations) { if (k == 1) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 32) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 64) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 128) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 256) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 512) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 1024) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else THROW("Unimplemented for k=%d, knn_merge_parts works for k<=1024", k); } + +template +inline void knn_merge_parts(const value_t* inK, + const value_idx* inV, + value_t* outK, + value_idx* outV, + size_t n_samples, + int n_parts, + int k, + cudaStream_t stream, + value_idx* translations, + bool select_min = true) +{ + if (select_min) { + knn_merge_parts_dispatch( + inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); + } else { + knn_merge_parts_dispatch( + inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); + } +} } // namespace cuvs::neighbors::detail diff --git a/cpp/src/neighbors/knn_merge_parts.cu b/cpp/src/neighbors/knn_merge_parts.cu index da7110a475..91694b098f 100644 --- a/cpp/src/neighbors/knn_merge_parts.cu +++ b/cpp/src/neighbors/knn_merge_parts.cu @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION. * SPDX-License-Identifier: Apache-2.0 */ @@ -15,7 +15,8 @@ void _knn_merge_parts(raft::resources const& res, raft::device_matrix_view inV, raft::device_matrix_view outK, raft::device_matrix_view outV, - raft::device_vector_view translations) + raft::device_vector_view translations, + bool select_min) { auto parts = translations.extent(0); auto rows = outK.extent(0); @@ -29,7 +30,8 @@ void _knn_merge_parts(raft::resources const& res, parts, k, raft::resource::get_cuda_stream(res), - translations.data_handle()); + translations.data_handle(), + select_min); } } // namespace @@ -40,7 +42,17 @@ void knn_merge_parts(raft::resources const& res, raft::device_matrix_view outV, raft::device_vector_view translations) { - _knn_merge_parts(res, inK, inV, outK, outV, translations); + _knn_merge_parts(res, inK, inV, outK, outV, translations, true); +} +void knn_merge_parts(raft::resources const& res, + raft::device_matrix_view inK, + raft::device_matrix_view inV, + raft::device_matrix_view outK, + raft::device_matrix_view outV, + raft::device_vector_view translations, + bool select_min) +{ + _knn_merge_parts(res, inK, inV, outK, outV, translations, select_min); } void knn_merge_parts(raft::resources const& res, raft::device_matrix_view inK, @@ -49,7 +61,17 @@ void knn_merge_parts(raft::resources const& res, raft::device_matrix_view outV, raft::device_vector_view translations) { - _knn_merge_parts(res, inK, inV, outK, outV, translations); + _knn_merge_parts(res, inK, inV, outK, outV, translations, true); +} +void knn_merge_parts(raft::resources const& res, + raft::device_matrix_view inK, + raft::device_matrix_view inV, + raft::device_matrix_view outK, + raft::device_matrix_view outV, + raft::device_vector_view translations, + bool select_min) +{ + _knn_merge_parts(res, inK, inV, outK, outV, translations, select_min); } void knn_merge_parts(raft::resources const& res, raft::device_matrix_view inK, @@ -58,6 +80,16 @@ void knn_merge_parts(raft::resources const& res, raft::device_matrix_view outV, raft::device_vector_view translations) { - _knn_merge_parts(res, inK, inV, outK, outV, translations); + _knn_merge_parts(res, inK, inV, outK, outV, translations, true); +} +void knn_merge_parts(raft::resources const& res, + raft::device_matrix_view inK, + raft::device_matrix_view inV, + raft::device_matrix_view outK, + raft::device_matrix_view outV, + raft::device_vector_view translations, + bool select_min) +{ + _knn_merge_parts(res, inK, inV, outK, outV, translations, select_min); } } // namespace cuvs::neighbors diff --git a/cpp/src/neighbors/mg/snmg.cuh b/cpp/src/neighbors/mg/snmg.cuh index 057a3e8272..7abf0b20f2 100644 --- a/cpp/src/neighbors/mg/snmg.cuh +++ b/cpp/src/neighbors/mg/snmg.cuh @@ -17,6 +17,7 @@ #include #include "../../core/omp_wrapper.hpp" +#include #include #include #include @@ -258,6 +259,8 @@ void sharded_search_with_direct_merge( int64_t n_neighbors, int64_t n_batches) { + const bool select_min = + cuvs::distance::is_min_close(index.ann_interfaces_.front().index_.value().metric()); const auto& root_handle = raft::resource::set_current_device_to_root_rank(clique); auto in_neighbors = raft::make_device_matrix( root_handle, index.num_ranks_ * n_rows_per_batch, n_neighbors); @@ -360,7 +363,8 @@ void sharded_search_with_direct_merge( in_neighbors.view(), out_distances.view(), out_neighbors.view(), - d_trans.view()); + d_trans.view(), + select_min); raft::copy( root_handle_, @@ -388,6 +392,8 @@ void sharded_search_with_tree_merge( int64_t n_neighbors, int64_t n_batches) { + const bool select_min = + cuvs::distance::is_min_close(index.ann_interfaces_.front().index_.value().metric()); for (int64_t batch_idx = 0; batch_idx < n_batches; batch_idx++) { int64_t offset = batch_idx * n_rows_per_batch; int64_t query_offset = offset * n_cols; @@ -487,7 +493,8 @@ void sharded_search_with_tree_merge( tmp_neighbors.view(), distances_merge_res.view(), neighbors_merge_res.view(), - d_trans.view()); + d_trans.view(), + select_min); raft::copy(dev_res, raft::make_device_vector_view(tmp_neighbors.data_handle(), part_size), raft::make_device_vector_view( diff --git a/cpp/tests/CMakeLists.txt b/cpp/tests/CMakeLists.txt index 437396b736..9b776b1c3f 100644 --- a/cpp/tests/CMakeLists.txt +++ b/cpp/tests/CMakeLists.txt @@ -109,7 +109,7 @@ endfunction() ConfigureTest( NAME NEIGHBORS_TEST PATH neighbors/brute_force.cu neighbors/brute_force_prefiltered.cu neighbors/sparse_brute_force.cu - neighbors/refine.cu neighbors/distance_nn.cu + neighbors/refine.cu neighbors/distance_nn.cu neighbors/knn_merge_parts.cu GPUS 1 PERCENT 100 ) diff --git a/cpp/tests/neighbors/knn_merge_parts.cu b/cpp/tests/neighbors/knn_merge_parts.cu new file mode 100644 index 0000000000..9f1d5d6346 --- /dev/null +++ b/cpp/tests/neighbors/knn_merge_parts.cu @@ -0,0 +1,111 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "knn_utils.cuh" + +#include +#include +#include +#include +#include + +#include + +#include +#include + +namespace cuvs::neighbors { +namespace { + +void run_merge(bool select_min, + const std::vector& expected_distances, + const std::vector& expected_neighbors) +{ + constexpr int64_t n_queries = 2; + constexpr int64_t n_parts = 2; + constexpr int64_t k = 3; + + ASSERT_EQ(expected_distances.size(), n_queries * k); + ASSERT_EQ(expected_neighbors.size(), n_queries * k); + + raft::resources res; + auto stream = raft::resource::get_cuda_stream(res); + + // Input layout is [part][query][neighbor]. Indices are local to each part. + const std::vector input_distances{ + 10.0f, 8.0f, 6.0f, -1.0f, -3.0f, -5.0f, 9.0f, 7.0f, 5.0f, 4.0f, 2.0f, 0.0f}; + const std::vector input_neighbors{0, 1, 2, 0, 1, 2, 0, 1, 2, 0, 1, 2}; + const std::vector translations{0, 100}; + + auto input_distances_device = + raft::make_device_matrix(res, n_parts * n_queries, k); + auto input_neighbors_device = + raft::make_device_matrix(res, n_parts * n_queries, k); + auto output_distances_device = raft::make_device_matrix(res, n_queries, k); + auto output_neighbors_device = raft::make_device_matrix(res, n_queries, k); + auto expected_distances_device = raft::make_device_matrix(res, n_queries, k); + auto expected_neighbors_device = raft::make_device_matrix(res, n_queries, k); + auto translations_device = raft::make_device_vector(res, n_parts); + + raft::update_device( + input_distances_device.data_handle(), input_distances.data(), input_distances.size(), stream); + raft::update_device( + input_neighbors_device.data_handle(), input_neighbors.data(), input_neighbors.size(), stream); + raft::update_device(expected_distances_device.data_handle(), + expected_distances.data(), + expected_distances.size(), + stream); + raft::update_device(expected_neighbors_device.data_handle(), + expected_neighbors.data(), + expected_neighbors.size(), + stream); + raft::update_device( + translations_device.data_handle(), translations.data(), translations.size(), stream); + + if (select_min) { + // Exercise the legacy overload, which must continue to select minimum distances. + knn_merge_parts(res, + input_distances_device.view(), + input_neighbors_device.view(), + output_distances_device.view(), + output_neighbors_device.view(), + translations_device.view()); + } else { + knn_merge_parts(res, + input_distances_device.view(), + input_neighbors_device.view(), + output_distances_device.view(), + output_neighbors_device.view(), + translations_device.view(), + false); + } + + ASSERT_TRUE(devArrMatchKnnPair(expected_neighbors_device.data_handle(), + output_neighbors_device.data_handle(), + expected_distances_device.data_handle(), + output_distances_device.data_handle(), + n_queries, + k, + 0.0f, + stream, + true)); +} + +TEST(KnnMergeParts, SelectsSmallestByDefault) +{ + run_merge(true, + {5.0f, 6.0f, 7.0f, -5.0f, -3.0f, -1.0f}, + {102, 2, 101, 2, 1, 0}); +} + +TEST(KnnMergeParts, SelectsLargestWhenRequested) +{ + run_merge(false, + {10.0f, 9.0f, 8.0f, 4.0f, 2.0f, 0.0f}, + {0, 100, 1, 100, 101, 102}); +} + +} // namespace +} // namespace cuvs::neighbors From 550693f2962c320bf8eddc8f56977567c856a4b3 Mon Sep 17 00:00:00 2001 From: Dmitry Vesnin Date: Fri, 7 Aug 2026 21:09:16 +0300 Subject: [PATCH 2/2] arg order + format --- cpp/src/neighbors/detail/knn_merge_parts.cuh | 29 +++++++++++++------- cpp/tests/neighbors/knn_merge_parts.cu | 14 ++++------ 2 files changed, 24 insertions(+), 19 deletions(-) diff --git a/cpp/src/neighbors/detail/knn_merge_parts.cuh b/cpp/src/neighbors/detail/knn_merge_parts.cuh index 839f714374..daf42296dd 100644 --- a/cpp/src/neighbors/detail/knn_merge_parts.cuh +++ b/cpp/src/neighbors/detail/knn_merge_parts.cuh @@ -16,7 +16,12 @@ namespace cuvs::neighbors::detail { -template +template RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK, const value_idx* inV, value_t* outK, @@ -91,7 +96,11 @@ RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK, } } -template +template inline void knn_merge_parts_impl(const value_t* inK, const value_idx* inV, value_t* outK, @@ -109,7 +118,7 @@ inline void knn_merge_parts_impl(const value_t* inK, auto kInit = SelectMin ? raft::upper_bound() : raft::lower_bound(); auto vInit = -1; - knn_merge_parts_kernel + knn_merge_parts_kernel <<>>( inK, inV, outK, outV, n_samples, n_parts, kInit, vInit, k, translations); RAFT_CUDA_TRY(cudaPeekAtLastError()); @@ -141,25 +150,25 @@ inline void knn_merge_parts_dispatch(const value_t* inK, value_idx* translations) { if (k == 1) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 32) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 64) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 128) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 256) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 512) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else if (k <= 1024) - knn_merge_parts_impl( + knn_merge_parts_impl( inK, inV, outK, outV, n_samples, n_parts, k, stream, translations); else THROW("Unimplemented for k=%d, knn_merge_parts works for k<=1024", k); diff --git a/cpp/tests/neighbors/knn_merge_parts.cu b/cpp/tests/neighbors/knn_merge_parts.cu index 9f1d5d6346..2607057adb 100644 --- a/cpp/tests/neighbors/knn_merge_parts.cu +++ b/cpp/tests/neighbors/knn_merge_parts.cu @@ -43,11 +43,11 @@ void run_merge(bool select_min, raft::make_device_matrix(res, n_parts * n_queries, k); auto input_neighbors_device = raft::make_device_matrix(res, n_parts * n_queries, k); - auto output_distances_device = raft::make_device_matrix(res, n_queries, k); - auto output_neighbors_device = raft::make_device_matrix(res, n_queries, k); + auto output_distances_device = raft::make_device_matrix(res, n_queries, k); + auto output_neighbors_device = raft::make_device_matrix(res, n_queries, k); auto expected_distances_device = raft::make_device_matrix(res, n_queries, k); auto expected_neighbors_device = raft::make_device_matrix(res, n_queries, k); - auto translations_device = raft::make_device_vector(res, n_parts); + auto translations_device = raft::make_device_vector(res, n_parts); raft::update_device( input_distances_device.data_handle(), input_distances.data(), input_distances.size(), stream); @@ -95,16 +95,12 @@ void run_merge(bool select_min, TEST(KnnMergeParts, SelectsSmallestByDefault) { - run_merge(true, - {5.0f, 6.0f, 7.0f, -5.0f, -3.0f, -1.0f}, - {102, 2, 101, 2, 1, 0}); + run_merge(true, {5.0f, 6.0f, 7.0f, -5.0f, -3.0f, -1.0f}, {102, 2, 101, 2, 1, 0}); } TEST(KnnMergeParts, SelectsLargestWhenRequested) { - run_merge(false, - {10.0f, 9.0f, 8.0f, 4.0f, 2.0f, 0.0f}, - {0, 100, 1, 100, 101, 102}); + run_merge(false, {10.0f, 9.0f, 8.0f, 4.0f, 2.0f, 0.0f}, {0, 100, 1, 100, 101, 102}); } } // namespace