Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 23 additions & 0 deletions cpp/include/cuvs/neighbors/knn_merge_parts.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<const float, int64_t> inK,
raft::device_matrix_view<const int64_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<int64_t, int64_t> outV,
raft::device_vector_view<int64_t, int64_t> translations,
bool select_min);
void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<const float, int64_t> inK,
raft::device_matrix_view<const int64_t, int64_t> inV,
Expand All @@ -36,11 +45,25 @@ void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<uint32_t, int64_t> outV,
raft::device_vector_view<uint32_t, int64_t> translations);
void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<const float, int64_t> inK,
raft::device_matrix_view<const uint32_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<uint32_t, int64_t> outV,
raft::device_vector_view<uint32_t, int64_t> translations,
bool select_min);
void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<const float, int64_t> inK,
raft::device_matrix_view<const int32_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<int32_t, int64_t> outV,
raft::device_vector_view<int32_t, int64_t> translations);
void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<const float, int64_t> inK,
raft::device_matrix_view<const int32_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<int32_t, int64_t> outV,
raft::device_vector_view<int32_t, int64_t> translations,
bool select_min);
} // namespace neighbors
} // namespace CUVS_EXPORT cuvs
72 changes: 49 additions & 23 deletions cpp/src/neighbors/detail/knn_merge_parts.cuh
Original file line number Diff line number Diff line change
@@ -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
*/

Expand All @@ -20,7 +20,8 @@ template <typename value_idx = std::int64_t,
typename value_t = float,
int warp_q,
int thread_q,
int tpb>
int tpb,
bool SelectMin = true>
RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK,
const value_idx* inV,
value_t* outK,
Expand All @@ -43,7 +44,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<value_t>,
warp_q,
thread_q,
Expand Down Expand Up @@ -95,7 +96,11 @@ RAFT_KERNEL knn_merge_parts_kernel(const value_t* inK,
}
}

template <typename value_idx = std::int64_t, typename value_t = float, int warp_q, int thread_q>
template <typename value_idx = std::int64_t,
typename value_t = float,
int warp_q,
int thread_q,
bool SelectMin = true>
inline void knn_merge_parts_impl(const value_t* inK,
const value_idx* inV,
value_t* outK,
Expand All @@ -111,9 +116,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<value_t>::max();
auto kInit = SelectMin ? raft::upper_bound<value_t>() : raft::lower_bound<value_t>();
auto vInit = -1;
knn_merge_parts_kernel<value_idx, value_t, warp_q, thread_q, n_threads>
knn_merge_parts_kernel<value_idx, value_t, warp_q, thread_q, n_threads, SelectMin>
<<<grid, block, 0, stream>>>(
inK, inV, outK, outV, n_samples, n_parts, kInit, vInit, k, translations);
RAFT_CUDA_TRY(cudaPeekAtLastError());
Expand All @@ -133,39 +138,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 <typename value_idx = std::int64_t, typename value_t = float>
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 <bool SelectMin, typename value_idx = std::int64_t, typename value_t = float>
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<value_idx, value_t, 1, 1>(
knn_merge_parts_impl<value_idx, value_t, 1, 1, SelectMin>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
else if (k <= 32)
knn_merge_parts_impl<value_idx, value_t, 32, 2>(
knn_merge_parts_impl<value_idx, value_t, 32, 2, SelectMin>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
else if (k <= 64)
knn_merge_parts_impl<value_idx, value_t, 64, 3>(
knn_merge_parts_impl<value_idx, value_t, 64, 3, SelectMin>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
else if (k <= 128)
knn_merge_parts_impl<value_idx, value_t, 128, 3>(
knn_merge_parts_impl<value_idx, value_t, 128, 3, SelectMin>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
else if (k <= 256)
knn_merge_parts_impl<value_idx, value_t, 256, 4>(
knn_merge_parts_impl<value_idx, value_t, 256, 4, SelectMin>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
else if (k <= 512)
knn_merge_parts_impl<value_idx, value_t, 512, 8>(
knn_merge_parts_impl<value_idx, value_t, 512, 8, SelectMin>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
else if (k <= 1024)
knn_merge_parts_impl<value_idx, value_t, 1024, 8>(
knn_merge_parts_impl<value_idx, value_t, 1024, 8, SelectMin>(
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 <typename value_idx = std::int64_t, typename value_t = float>
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<true>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
} else {
knn_merge_parts_dispatch<false>(
inK, inV, outK, outV, n_samples, n_parts, k, stream, translations);
}
}
} // namespace cuvs::neighbors::detail
44 changes: 38 additions & 6 deletions cpp/src/neighbors/knn_merge_parts.cu
Original file line number Diff line number Diff line change
@@ -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
*/

Expand All @@ -15,7 +15,8 @@ void _knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<const IdxT, int64_t> inV,
raft::device_matrix_view<T, int64_t> outK,
raft::device_matrix_view<IdxT, int64_t> outV,
raft::device_vector_view<IdxT, int64_t> translations)
raft::device_vector_view<IdxT, int64_t> translations,
bool select_min)
{
auto parts = translations.extent(0);
auto rows = outK.extent(0);
Expand All @@ -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

Expand All @@ -40,7 +42,17 @@ void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<int64_t, int64_t> outV,
raft::device_vector_view<int64_t, int64_t> 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<const float, int64_t> inK,
raft::device_matrix_view<const int64_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<int64_t, int64_t> outV,
raft::device_vector_view<int64_t, int64_t> 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<const float, int64_t> inK,
Expand All @@ -49,7 +61,17 @@ void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<uint32_t, int64_t> outV,
raft::device_vector_view<uint32_t, int64_t> 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<const float, int64_t> inK,
raft::device_matrix_view<const uint32_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<uint32_t, int64_t> outV,
raft::device_vector_view<uint32_t, int64_t> 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<const float, int64_t> inK,
Expand All @@ -58,6 +80,16 @@ void knn_merge_parts(raft::resources const& res,
raft::device_matrix_view<int32_t, int64_t> outV,
raft::device_vector_view<int32_t, int64_t> 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<const float, int64_t> inK,
raft::device_matrix_view<const int32_t, int64_t> inV,
raft::device_matrix_view<float, int64_t> outK,
raft::device_matrix_view<int32_t, int64_t> outV,
raft::device_vector_view<int32_t, int64_t> translations,
bool select_min)
{
_knn_merge_parts(res, inK, inV, outK, outV, translations, select_min);
}
} // namespace cuvs::neighbors
11 changes: 9 additions & 2 deletions cpp/src/neighbors/mg/snmg.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
#include <raft/util/cuda_dev_essentials.cuh>

#include "../../core/omp_wrapper.hpp"
#include <cuvs/distance/distance.hpp>
#include <cuvs/neighbors/cagra.hpp>
#include <cuvs/neighbors/common.hpp>
#include <cuvs/neighbors/ivf_flat.hpp>
Expand Down Expand Up @@ -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<searchIdxT, int64_t, row_major>(
root_handle, index.num_ranks_ * n_rows_per_batch, n_neighbors);
Expand Down Expand Up @@ -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_,
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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<const searchIdxT>(
Expand Down
2 changes: 1 addition & 1 deletion cpp/tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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
)
Expand Down
Loading
Loading