Skip to content
Merged
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
1 change: 1 addition & 0 deletions cpp/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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}
Expand Down
178 changes: 11 additions & 167 deletions cpp/src/neighbors/detail/cagra/graph_core.cuh
Original file line number Diff line number Diff line change
@@ -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 <raft/core/copy.cuh>
Expand All @@ -23,7 +24,6 @@

#include <cuvs/neighbors/cagra.hpp>

#include <raft/util/bitonic_sort.cuh>
#include <raft/util/cuda_rt_essentials.hpp>
#include <raft/util/integer_utils.hpp>

Expand All @@ -39,6 +39,7 @@
#include <iostream>
#include <memory>
#include <random>
#include <type_traits>

namespace cg = cooperative_groups;

Expand All @@ -53,127 +54,6 @@ inline double cur_time(void)
return ((double)tv.tv_sec + (double)tv.tv_usec * 1e-6);
}

template <typename T>
__device__ inline void swap(T& val1, T& val2)
{
T val0 = val1;
val1 = val2;
val2 = val0;
}

template <typename K, typename V>
__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<K>(key1, key2);
swap<V>(val1, val2);
return true;
}
return false;
}

template <class DATA_T, class IdxT, int numElementsPerThread>
__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<uint64_t>(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<float>{}(
dataset[d + static_cast<uint64_t>(dataset_dim) * dstNode]);
dist -= cuvs::spatial::knn::detail::utils::mapping<float>{}(
dataset[d + static_cast<uint64_t>(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<float>{}(
dataset[d + static_cast<uint64_t>(dataset_dim) * srcNode]) -
cuvs::spatial::knn::detail::utils::mapping<float>{}(
dataset[d + static_cast<uint64_t>(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<float>{}(
dataset[d + static_cast<uint64_t>(dataset_dim) * srcNode]) -
cuvs::spatial::knn::detail::utils::mapping<float>{}(
dataset[d + static_cast<uint64_t>(dataset_dim) * dstNode]);
dist += raft::abs(diff);
}
} else if (metric == cuvs::distance::DistanceType::BitwiseHamming) {
if constexpr (std::is_integral_v<DATA_T>) {
for (int d = lane_id; d < dataset_dim; d += raft::WarpSize) {
dist += __popc(
static_cast<uint32_t>(dataset[d + static_cast<uint64_t>(dataset_dim) * srcNode] ^
dataset[d + static_cast<uint64_t>(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<float>();
my_vals[k / raft::WarpSize] = utils::get_max_value<IdxT>();
}
}

// Sort by RAFT bitonic sort
raft::util::bitonic<numElementsPerThread>(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<uint64_t>(graph_degree) * srcNode)] = my_vals[i];
}
}
}

template <typename IdxT, typename OutputMatrixView>
__global__ void kern_make_rev_graph_k(
OutputMatrixView output_graph, // [graph_size, degree]
Expand Down Expand Up @@ -983,6 +863,7 @@ void sort_knn_graph(
raft::mdspan<const DataT, raft::matrix_extent<int64_t>, raft::row_major, d_accessor> dataset,
raft::mdspan<IdxT, raft::matrix_extent<int64_t>, raft::row_major, g_accessor> knn_graph)
{
static_assert(std::is_same_v<IdxT, uint32_t>, "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(
Expand Down Expand Up @@ -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<DataT, IdxT, numElementsPerThread>;
} else if (input_graph_degree <= 64) {
constexpr int numElementsPerThread = 2;
kernel_sort = kern_sort<DataT, IdxT, numElementsPerThread>;
} else if (input_graph_degree <= 128) {
constexpr int numElementsPerThread = 4;
kernel_sort = kern_sort<DataT, IdxT, numElementsPerThread>;
} else if (input_graph_degree <= 256) {
constexpr int numElementsPerThread = 8;
kernel_sort = kern_sort<DataT, IdxT, numElementsPerThread>;
} else if (input_graph_degree <= 512) {
constexpr int numElementsPerThread = 16;
kernel_sort = kern_sort<DataT, IdxT, numElementsPerThread>;
} else if (input_graph_degree <= 1024) {
constexpr int numElementsPerThread = 32;
kernel_sort = kern_sort<DataT, IdxT, numElementsPerThread>;
} 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<<<grid_size, block_size, 0, raft::resource::get_cuda_stream(res)>>>(
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<uint32_t>(dataset_size),
static_cast<uint32_t>(dataset_dim),
d_input_graph.data_handle(),
static_cast<uint32_t>(input_graph_degree));
raft::resource::sync_stream(res);
RAFT_LOG_DEBUG(".");
raft::copy(res, knn_graph, raft::make_const_mdspan(d_input_graph.view()));
Expand Down
Loading
Loading