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
Original file line number Diff line number Diff line change
Expand Up @@ -21,6 +21,7 @@
#include <cuvs/distance/distance.hpp>
#include <raft/core/device_mdspan.hpp>
#include <raft/core/logger.hpp>
#include <raft/util/kernel_launch.hpp>
#include <rtcx/algorithm_launcher.hpp>

#include <cstddef>
Expand Down Expand Up @@ -102,55 +103,50 @@ void select_and_run(const dataset_descriptor_host<DataT, IndexT, DistanceT>& dat
// function The descriptor's state is managed by a shared_ptr internally, so no need to explicitly
// keep it alive

// Cast size_t/int64_t parameters to match kernel signature exactly
// The dispatch mechanism uses void* pointers, so parameter sizes must match exactly
// graph.extent(1) returns int64_t but kernel expects uint32_t
// traversed_hash_bitlen is int64_t but kernel expects uint32_t
// ps.itopk_size, ps.min_iterations, ps.max_iterations are size_t (8 bytes) but kernel expects
// uint32_t (4 bytes) ps.num_random_samplings is uint32_t but kernel expects unsigned - cast for
// consistency
// These arguments are wider than the kernel parameters they feed; the casts record the
// intentional narrowing. graph.extent() and traversed_hash_bitlen are int64_t, and
// ps.itopk_size / ps.min_iterations / ps.max_iterations are size_t.
const uint32_t graph_degree_u32 = static_cast<uint32_t>(graph.extent(1));
const uint32_t traversed_hash_bitlen_u32 = static_cast<uint32_t>(traversed_hash_bitlen);
const uint32_t itopk_size_u32 = static_cast<uint32_t>(ps.itopk_size);
const uint32_t min_iterations_u32 = static_cast<uint32_t>(ps.min_iterations);
const uint32_t max_iterations_u32 = static_cast<uint32_t>(ps.max_iterations);
const unsigned num_random_samplings_u = static_cast<unsigned>(ps.num_random_samplings);

auto kernel_launcher = [&]() -> void {
launcher->dispatch<
multi_cta_search::search_multi_cta_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>>(
stream,
grid_dims,
block_dims,
smem_size,
topk_indices_ptr,
topk_distances_ptr,
dev_desc,
queries_ptr,
graph.data_handle(),
max_elements,
graph_degree_u32,
source_indices_ptr,
num_random_samplings_u,
ps.rand_xor_mask,
dev_seed_ptr,
num_seeds,
visited_hash_bitlen,
traversed_hashmap_ptr,
traversed_hash_bitlen_u32,
itopk_size_u32,
min_iterations_u32,
max_iterations_u32,
num_executed_iterations,
static_cast<IndexT>(graph.extent(0)),
query_id_offset,
filter_payload);
};
cuvs::neighbors::detail::safely_launch_kernel_with_smem_size<
multi_cta_search::search_multi_cta_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>>(
smem_size, kernel_launcher, launcher->get_kernel());

RAFT_CUDA_TRY(cudaPeekAtLastError());
using kernel_t =
multi_cta_search::search_multi_cta_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>;
auto kernel = raft::kernel_ref<kernel_t>{launcher->get_kernel()};

cuvs::neighbors::detail::safely_launch_kernel_with_smem_size<kernel_t>(
smem_size,
[&] {
raft::launch_kernel({stream, smem_size},
grid_dims,
block_dims,
kernel,
topk_indices_ptr,
topk_distances_ptr,
dev_desc,
queries_ptr,
graph.data_handle(),
max_elements,
graph_degree_u32,
source_indices_ptr,
ps.num_random_samplings,
ps.rand_xor_mask,
dev_seed_ptr,
num_seeds,
visited_hash_bitlen,
traversed_hashmap_ptr,
traversed_hash_bitlen_u32,
itopk_size_u32,
min_iterations_u32,
max_iterations_u32,
num_executed_iterations,
static_cast<IndexT>(graph.extent(0)),
query_id_offset,
filter_payload);
},
kernel.handle);
}

// Multi-partition launcher. Drives `search_multi_cta_mp` with a 3D grid
Expand Down Expand Up @@ -216,41 +212,41 @@ void select_and_run_mp(const dataset_descriptor_host<DataT, IndexT, DistanceT>&
dim3 block_dims(block_size, 1, 1);
dim3 grid_dims(num_cta_per_query, num_queries, num_partitions);

const uint32_t max_graph_degree_u32 = static_cast<uint32_t>(max_graph_degree);
// traversed_hash_bitlen is int64_t and the ps iteration counts are size_t; the casts record the
// intentional narrowing to the kernel's uint32_t parameters.
const uint32_t traversed_hash_bitlen_u32 = static_cast<uint32_t>(traversed_hash_bitlen);
const uint32_t itopk_size_u32 = static_cast<uint32_t>(ps.itopk_size);
const uint32_t min_iterations_u32 = static_cast<uint32_t>(ps.min_iterations);
const uint32_t max_iterations_u32 = static_cast<uint32_t>(ps.max_iterations);
const unsigned num_random_samplings_u = static_cast<unsigned>(ps.num_random_samplings);

auto kernel_launcher = [&]() -> void {
launcher->dispatch<
multi_cta_search::search_multi_cta_mp_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>>(
stream,
grid_dims,
block_dims,
smem_size,
partition_descs,
intermediate_indices_ptr,
intermediate_distances_ptr,
queries_ptr,
max_elements,
max_graph_degree_u32,
num_random_samplings_u,
ps.rand_xor_mask,
visited_hash_bitlen,
traversed_hashmap_ptr,
traversed_hash_bitlen_u32,
itopk_size_u32,
min_iterations_u32,
max_iterations_u32,
query_id_offset);
};
cuvs::neighbors::detail::safely_launch_kernel_with_smem_size<
multi_cta_search::search_multi_cta_mp_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>>(
smem_size, kernel_launcher, launcher->get_kernel());

RAFT_CUDA_TRY(cudaPeekAtLastError());
using kernel_t =
multi_cta_search::search_multi_cta_mp_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>;
auto kernel = raft::kernel_ref<kernel_t>{launcher->get_kernel()};

cuvs::neighbors::detail::safely_launch_kernel_with_smem_size<kernel_t>(
smem_size,
[&] {
raft::launch_kernel({stream, smem_size},
grid_dims,
block_dims,
kernel,
partition_descs,
intermediate_indices_ptr,
intermediate_distances_ptr,
queries_ptr,
max_elements,
max_graph_degree,
ps.num_random_samplings,
ps.rand_xor_mask,
visited_hash_bitlen,
traversed_hashmap_ptr,
traversed_hash_bitlen_u32,
itopk_size_u32,
min_iterations_u32,
max_iterations_u32,
query_id_offset);
},
kernel.handle);
}

} // namespace cuvs::neighbors::cagra::detail::multi_cta_search
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
#include <cuvs/distance/distance.hpp>
#include <raft/core/device_mdspan.hpp>
#include <raft/core/logger.hpp>
#include <raft/util/kernel_launch.hpp>
#include <rtcx/algorithm_launcher.hpp>

#include <cstddef>
Expand Down Expand Up @@ -58,15 +59,14 @@ void random_pickup_jit(const dataset_descriptor_host<DataT, IndexT, DistanceT>&
// Get the device descriptor pointer
const auto* dev_desc = dataset_desc.dev_ptr(cuda_stream);

// Cast size_t parameters to match kernel signature exactly
// The dispatch mechanism uses void* pointers, so parameter sizes must match exactly
// `ldr` is wider than the kernel parameter; the cast records the intentional narrowing.
const uint32_t ldr_u32 = static_cast<uint32_t>(ldr);

launcher->dispatch<random_pickup_kernel_func_t<DataT, IndexT, DistanceT>>(
cuda_stream,
raft::launch_kernel(
{cuda_stream, dataset_desc.smem_ws_size_in_bytes},
grid_size,
dim3(block_size, 1, 1),
dataset_desc.smem_ws_size_in_bytes,
raft::kernel_ref<random_pickup_kernel_func_t<DataT, IndexT, DistanceT>>{launcher->get_kernel()},
dev_desc,
queries_ptr,
num_pickup,
Expand All @@ -80,8 +80,6 @@ void random_pickup_jit(const dataset_descriptor_host<DataT, IndexT, DistanceT>&
visited_hashmap_ptr,
hash_bitlen,
graph_size);

RAFT_CUDA_TRY(cudaPeekAtLastError());
}

// JIT version of compute_distance_to_child_nodes
Expand Down Expand Up @@ -121,12 +119,13 @@ void compute_distance_to_child_nodes_jit(
// Get the device descriptor pointer
const auto* dev_desc = dataset_desc.dev_ptr(cuda_stream);

launcher->dispatch<
compute_distance_to_child_nodes_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>>(
cuda_stream,
raft::launch_kernel(
{cuda_stream, dataset_desc.smem_ws_size_in_bytes},
grid_size,
dim3(block_size, 1, 1),
dataset_desc.smem_ws_size_in_bytes,
raft::kernel_ref<
compute_distance_to_child_nodes_kernel_func_t<DataT, IndexT, DistanceT, SourceIndexT>>{
launcher->get_kernel()},
parent_node_list,
parent_candidates_ptr,
parent_distance_ptr,
Expand All @@ -143,8 +142,6 @@ void compute_distance_to_child_nodes_jit(
result_distances_ptr,
ldd,
filter_payload);

RAFT_CUDA_TRY(cudaPeekAtLastError());
}

// JIT version of apply_filter
Expand All @@ -166,24 +163,20 @@ void apply_filter_jit(const SourceIndexT* source_indices_ptr,
const std::uint32_t block_size = 256;
const std::uint32_t grid_size = raft::ceildiv(num_queries * result_buffer_size, block_size);

// Alias avoids nested `dispatch< alias_template<...>>` which NVCC can misparse as
// comparison/shift.
using apply_filter_kernel_func_t = apply_filter_kernel_func_t<INDEX_T, DISTANCE_T, SourceIndexT>;
// `template` required: in template code, `->dispatch<...>` is otherwise parsed as `dispatch <` …
launcher->template dispatch<apply_filter_kernel_func_t>(cuda_stream,
dim3(grid_size, 1, 1),
dim3(block_size, 1, 1),
0,
source_indices_ptr,
result_indices_ptr,
result_distances_ptr,
lds,
result_buffer_size,
num_queries,
effective_query_id_offset,
filter_payload);

RAFT_CUDA_TRY(cudaPeekAtLastError());
raft::launch_kernel(
cuda_stream,
dim3(grid_size, 1, 1),
dim3(block_size, 1, 1),
raft::kernel_ref<apply_filter_kernel_func_t<INDEX_T, DISTANCE_T, SourceIndexT>>{
launcher->get_kernel()},
source_indices_ptr,
result_indices_ptr,
result_distances_ptr,
lds,
result_buffer_size,
num_queries,
effective_query_id_offset,
filter_payload);
}

} // namespace cuvs::neighbors::cagra::detail::multi_kernel_search
Loading
Loading