diff --git a/cpp/src/neighbors/mg/snmg.cuh b/cpp/src/neighbors/mg/snmg.cuh index 057a3e8272..6aac75d7b5 100644 --- a/cpp/src/neighbors/mg/snmg.cuh +++ b/cpp/src/neighbors/mg/snmg.cuh @@ -13,10 +13,12 @@ #include #include #include +#include #include #include #include "../../core/omp_wrapper.hpp" +#include #include #include #include @@ -258,6 +260,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); @@ -355,6 +359,13 @@ void sharded_search_with_direct_merge( d_trans.view(), raft::make_host_vector_view(h_trans.data(), index.num_ranks_)); + if (!select_min) { + raft::linalg::map(root_handle_, + in_distances.view(), + raft::mul_const_op(-1), + raft::make_const_mdspan(in_distances.view())); + } + knn_merge_parts(root_handle_, in_distances.view(), in_neighbors.view(), @@ -362,6 +373,13 @@ void sharded_search_with_direct_merge( out_neighbors.view(), d_trans.view()); + if (!select_min) { + raft::linalg::map(root_handle_, + out_distances.view(), + raft::mul_const_op(-1), + raft::make_const_mdspan(out_distances.view())); + } + raft::copy( root_handle_, raft::make_host_vector_view(neighbors.data_handle() + output_offset, part_size), @@ -388,6 +406,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; @@ -417,6 +437,13 @@ void sharded_search_with_tree_merge( cuvs::neighbors::search( dev_res, ann_if, search_params, query_partition, neighbors_view, distances_view); + if (!select_min) { + raft::linalg::map(dev_res, + distances_view, + raft::mul_const_op(-1), + raft::make_const_mdspan(distances_view)); + } + searchIdxT translation_offset = 0; for (int r = 0; r < rank; r++) { translation_offset += index.ann_interfaces_[r].size(); @@ -499,6 +526,12 @@ void sharded_search_with_tree_merge( // If done, copy the final result if (remaining <= 1) { + if (!select_min) { + raft::linalg::map(dev_res, + distances_view, + raft::mul_const_op(-1), + raft::make_const_mdspan(distances_view)); + } raft::copy( dev_res, raft::make_host_vector_view(neighbors.data_handle() + output_offset, part_size), diff --git a/cpp/tests/CMakeLists.txt b/cpp/tests/CMakeLists.txt index 8421c3219e..b4a657c90c 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..65db1a08ce --- /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 +#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) { + raft::linalg::map(res, + input_distances_device.view(), + raft::mul_const_op(-1), + raft::make_const_mdspan(input_distances_device.view())); + } + + knn_merge_parts(res, + input_distances_device.view(), + input_neighbors_device.view(), + output_distances_device.view(), + output_neighbors_device.view(), + translations_device.view()); + + if (!select_min) { + raft::linalg::map(res, + output_distances_device.view(), + raft::mul_const_op(-1), + raft::make_const_mdspan(output_distances_device.view())); + } + + 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, SelectsLargestAfterNegation) +{ + 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