11/*
2- * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION.
2+ * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved .
33 * SPDX-License-Identifier: Apache-2.0
44 */
55
66#include < cstdint>
77#include < dlpack/dlpack.h>
8+ #include < type_traits>
89
910#include < raft/core/error.hpp>
1011#include < raft/core/mdspan_types.hpp>
@@ -80,36 +81,50 @@ static cuvs::neighbors::all_neighbors::all_neighbors_params convert_params(
8081 return out;
8182}
8283
83- static void ensure_indices_dtype_and_device_compatibility (DLManagedTensor* indices)
84+ static void ensure_indices_dtype (DLManagedTensor* indices)
8485{
8586 auto dtype = indices->dl_tensor .dtype ;
8687 RAFT_EXPECTS (dtype.code == kDLInt && dtype.bits == 64 , " indices must be int64 output tensor" );
87- RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (indices->dl_tensor ),
88- " indices tensor must be device-compatible" );
8988}
9089
91- static void ensure_optional_distance_dtype_and_device_compatibility (DLManagedTensor* distances)
90+ static void ensure_optional_distance_dtype (DLManagedTensor* distances)
9291{
9392 if (distances == nullptr ) { return ; }
9493 auto dtype = distances->dl_tensor .dtype ;
9594 RAFT_EXPECTS (dtype.code == kDLFloat && dtype.bits == 32 ,
9695 " distances must be float32 output tensor" );
97- RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (distances->dl_tensor ),
98- " distances tensor must be device-compatible" );
9996}
10097
101- static void ensure_optional_core_distance_dtype_and_device_compatibility (
102- DLManagedTensor* core_distances)
98+ static void ensure_optional_core_distance_dtype (DLManagedTensor* core_distances)
10399{
104100 if (core_distances == nullptr ) { return ; }
105101 auto dtype = core_distances->dl_tensor .dtype ;
106102 RAFT_EXPECTS (dtype.code == kDLFloat && dtype.bits == 32 ,
107103 " core_distances must be float32 output tensor" );
108- RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (core_distances->dl_tensor ),
109- " core_distances tensor must be device-compatible" );
110104}
111105
112- template <typename T>
106+ // Validate that the outputs (indices/distances/core_distances) all live in the same memory space
107+ static bool validate_output_memory_space (DLManagedTensor* indices,
108+ DLManagedTensor* distances,
109+ DLManagedTensor* core_distances)
110+ {
111+ const bool host = cuvs::core::is_dlpack_host_compatible (indices->dl_tensor );
112+ RAFT_EXPECTS (host || cuvs::core::is_dlpack_device_compatible (indices->dl_tensor ),
113+ " indices tensor must be host- or device-compatible" );
114+ auto same_space = [&](DLManagedTensor* t, const char * name) {
115+ if (t == nullptr ) { return ; }
116+ RAFT_EXPECTS (cuvs::core::is_dlpack_host_compatible (t->dl_tensor ) == host,
117+ " %s tensor must be in the same memory space (host or device) as indices" ,
118+ name);
119+ };
120+ same_space (distances, " distances" );
121+ same_space (core_distances, " core_distances" );
122+ return host;
123+ }
124+
125+ // Build with a host-resident dataset. HostOutput selects whether the outputs live on host or
126+ // device.
127+ template <typename T, bool HostOutput>
113128void _build_host (cuvsResources_t res,
114129 cuvsAllNeighborsIndexParams_t params,
115130 DLManagedTensor* dataset_tensor,
@@ -124,24 +139,23 @@ void _build_host(cuvsResources_t res,
124139 RAFT_EXPECTS (cuvs::core::is_dlpack_host_compatible (dlt),
125140 " Host build expects host-compatible dataset tensor" );
126141
127- ensure_indices_dtype_and_device_compatibility (indices_tensor);
128- ensure_optional_distance_dtype_and_device_compatibility (distances_tensor);
129- ensure_optional_core_distance_dtype_and_device_compatibility (core_distances_tensor);
130-
131- // Check dependencies between parameters
132- if (core_distances_tensor != nullptr && distances_tensor == nullptr ) {
133- RAFT_FAIL (" distances tensor must be provided when core_distances tensor is provided" );
134- }
135-
136142 int64_t n_rows = dlt.shape [0 ];
137143 int64_t n_cols = dlt.shape [1 ];
138144
139145 auto cpp_params = convert_params (params, n_rows, n_cols);
140146
141- using dataset_mdspan_t = raft::host_matrix_view<const T, int64_t , raft::row_major>;
142- using indices_mdspan_t = raft::device_matrix_view<int64_t , int64_t , raft::row_major>;
143- using distances_mdspan_t = raft::device_matrix_view<float , int64_t , raft::row_major>;
144- using core_mdspan_t = raft::device_vector_view<float , int64_t >;
147+ using dataset_mdspan_t = raft::host_matrix_view<const T, int64_t , raft::row_major>;
148+ using indices_mdspan_t =
149+ std::conditional_t <HostOutput,
150+ raft::host_matrix_view<int64_t , int64_t , raft::row_major>,
151+ raft::device_matrix_view<int64_t , int64_t , raft::row_major>>;
152+ using distances_mdspan_t =
153+ std::conditional_t <HostOutput,
154+ raft::host_matrix_view<float , int64_t , raft::row_major>,
155+ raft::device_matrix_view<float , int64_t , raft::row_major>>;
156+ using core_mdspan_t = std::conditional_t <HostOutput,
157+ raft::host_vector_view<float , int64_t >,
158+ raft::device_vector_view<float , int64_t >>;
145159
146160 auto dataset = cuvs::core::from_dlpack<dataset_mdspan_t >(dataset_tensor);
147161 auto indices = cuvs::core::from_dlpack<indices_mdspan_t >(indices_tensor);
@@ -160,6 +174,7 @@ void _build_host(cuvsResources_t res,
160174 cpp_res, cpp_params, dataset, indices, distances, core_distances, alpha);
161175}
162176
177+ // Build with a device-resident dataset. Outputs are always device-resident.
163178template <typename T>
164179void _build_device (cuvsResources_t device_res,
165180 cuvsAllNeighborsIndexParams_t params,
@@ -175,15 +190,6 @@ void _build_device(cuvsResources_t device_res,
175190 RAFT_EXPECTS (cuvs::core::is_dlpack_device_compatible (dlt),
176191 " Device build expects device-compatible dataset tensor" );
177192
178- ensure_indices_dtype_and_device_compatibility (indices_tensor);
179- ensure_optional_distance_dtype_and_device_compatibility (distances_tensor);
180- ensure_optional_core_distance_dtype_and_device_compatibility (core_distances_tensor);
181-
182- // Check dependencies between parameters
183- if (core_distances_tensor != nullptr && distances_tensor == nullptr ) {
184- RAFT_FAIL (" distances tensor must be provided when core_distances tensor is provided" );
185- }
186-
187193 int64_t n_rows = dlt.shape [0 ];
188194 int64_t n_cols = dlt.shape [1 ];
189195
@@ -250,17 +256,33 @@ extern "C" cuvsError_t cuvsAllNeighborsBuild(cuvsResources_t res,
250256 return cuvs::core::translate_exceptions ([=] {
251257 auto dataset = dataset_tensor->dl_tensor ;
252258
259+ ensure_indices_dtype (indices_tensor);
260+ ensure_optional_distance_dtype (distances_tensor);
261+ ensure_optional_core_distance_dtype (core_distances_tensor);
262+ if (core_distances_tensor != nullptr && distances_tensor == nullptr ) {
263+ RAFT_FAIL (" distances tensor must be provided when core_distances tensor is provided" );
264+ }
265+
266+ const bool host_output =
267+ validate_output_memory_space (indices_tensor, distances_tensor, core_distances_tensor);
268+
253269 if (dataset.dtype .code == kDLFloat && dataset.dtype .bits == 32 ) {
254270 // Check if dataset is host-compatible or device-compatible
255271 if (cuvs::core::is_dlpack_host_compatible (dataset)) {
256- _build_host<float >(res,
257- params,
258- dataset_tensor,
259- indices_tensor,
260- distances_tensor,
261- core_distances_tensor,
262- alpha);
272+ // Host dataset supports both host- and device-resident outputs.
273+ if (host_output) {
274+ _build_host<float , true >(
275+ res, params, dataset_tensor, indices_tensor, distances_tensor, core_distances_tensor,
276+ alpha);
277+ } else {
278+ _build_host<float , false >(
279+ res, params, dataset_tensor, indices_tensor, distances_tensor, core_distances_tensor,
280+ alpha);
281+ }
263282 } else if (cuvs::core::is_dlpack_device_compatible (dataset)) {
283+ RAFT_EXPECTS (!host_output,
284+ " A device-resident dataset requires device-resident outputs; put the dataset "
285+ " on host to produce host-resident outputs." );
264286 _build_device<float >(res,
265287 params,
266288 dataset_tensor,
0 commit comments