@@ -122,6 +122,7 @@ static void merge_indices_for_layout(
122122 cuvs::neighbors::cagra::index_params const & params_cpp,
123123 std::vector<cuvs::neighbors::cagra::index<T, uint32_t , DatasetViewT>*>& index_ptrs,
124124 cuvsFilter filter,
125+ cuvs::neighbors::cagra::merge_params const & merge_params,
125126 cuvsDataset_t merged_dataset,
126127 cuvsCagraIndex_t output_index)
127128{
@@ -151,7 +152,8 @@ static void merge_indices_for_layout(
151152 auto owner = std::make_unique<owner_t >(std::move (matrix), dim);
152153 auto view = owner->as_dataset_view ();
153154 auto merged_idx =
154- cuvs::neighbors::cagra::merge (*res_ptr, params_cpp, index_ptrs, view, row_filter);
155+ cuvs::neighbors::cagra::merge (
156+ *res_ptr, params_cpp, index_ptrs, view, merge_params, row_filter);
155157 auto * holder =
156158 new cuvs_cagra_c_api_index_lifetime_holder<T, DatasetViewT>{std::move (merged_idx)};
157159 bind_index_lifetime_holder_to_C_index<T, DatasetViewT>(
@@ -164,7 +166,10 @@ static void merge_indices_for_layout(
164166 merged_dataset->layout = output_layout;
165167 merged_dataset->is_owning = true ;
166168 return ;
167- } catch (std::bad_alloc const &) {
169+ } catch (std::bad_alloc const & failure) {
170+ if (merge_params.algo == cuvs::neighbors::cagra::merge_algo::FASTENER ) {
171+ RAFT_FAIL (" FASTENER cagra::merge could not allocate device memory: %s" , failure.what ());
172+ }
168173 // Filtered merge gathers rows with device-only primitives, matching the restriction on the
169174 // legacy host fallback.
170175 RAFT_EXPECTS (filter.type == NO_FILTER ,
@@ -1132,6 +1137,7 @@ void _merge(cuvsResources_t res,
11321137 cuvsCagraIndex_t* indices,
11331138 size_t num_indices,
11341139 cuvsFilter filter,
1140+ const cuvs::neighbors::cagra::merge_params& merge_params,
11351141 cuvsDataset_t merged_dataset,
11361142 cuvsCagraIndex_t output_index)
11371143{
@@ -1179,13 +1185,13 @@ void _merge(cuvsResources_t res,
11791185 convert_opaque_indices_to_concrete_types<T, cuvs::neighbors::device_padded_dataset_view<T, int64_t >>(
11801186 indices, num_indices);
11811187 merge_indices_for_layout<T, cuvs::neighbors::device_padded_dataset_view<T, int64_t >>(
1182- res_ptr, params_cpp, index_ptrs, filter, merged_dataset, output_index);
1188+ res_ptr, params_cpp, index_ptrs, filter, merge_params, merged_dataset, output_index);
11831189 } else {
11841190 auto index_ptrs =
11851191 convert_opaque_indices_to_concrete_types<T, cuvs::neighbors::device_standard_dataset_view<T, int64_t >>(
11861192 indices, num_indices);
11871193 merge_indices_for_layout<T, cuvs::neighbors::device_standard_dataset_view<T, int64_t >>(
1188- res_ptr, params_cpp, index_ptrs, filter, merged_dataset, output_index);
1194+ res_ptr, params_cpp, index_ptrs, filter, merge_params, merged_dataset, output_index);
11891195 }
11901196}
11911197
@@ -1919,15 +1925,41 @@ extern "C" cuvsError_t cuvsCagraMerge(cuvsResources_t res,
19191925 cuvsFilter filter,
19201926 cuvsDataset_t merged_dataset,
19211927 cuvsCagraIndex_t output_index)
1928+ {
1929+ return cuvsCagraMergeWithParams (
1930+ res, params, nullptr , indices, num_indices, filter, merged_dataset, output_index);
1931+ }
1932+
1933+ extern " C" cuvsError_t cuvsCagraMergeWithParams (cuvsResources_t res,
1934+ cuvsCagraIndexParams_t params,
1935+ cuvsCagraMergeParams_t merge_params,
1936+ cuvsCagraIndex_t* indices,
1937+ size_t num_indices,
1938+ cuvsFilter filter,
1939+ cuvsDataset_t merged_dataset,
1940+ cuvsCagraIndex_t output_index)
19221941{
19231942 return cuvs::core::translate_exceptions ([=] {
1924- // Basic checks on inputs
19251943 RAFT_EXPECTS (indices != nullptr && num_indices > 0 , " indices array cannot be null or empty" );
19261944 RAFT_EXPECTS (params != nullptr , " params cannot be null" );
19271945 RAFT_EXPECTS (indices[0 ] != nullptr && indices[0 ]->addr != 0 ,
19281946 " All input indices must be built (non-empty)" );
19291947
1930- // Use first index dtype as reference
1948+ auto merge_params_cpp = cuvs::neighbors::cagra::merge_params{};
1949+ if (merge_params != nullptr ) {
1950+ RAFT_EXPECTS (merge_params->algo >= CUVS_CAGRA_MERGE_AUTO &&
1951+ merge_params->algo <= CUVS_CAGRA_MERGE_REBUILD ,
1952+ " Unsupported CAGRA merge algorithm" );
1953+ merge_params_cpp = {
1954+ .algo = static_cast <cuvs::neighbors::cagra::merge_algo>(merge_params->algo ),
1955+ .levels = merge_params->levels ,
1956+ .root_fanout = merge_params->root_fanout ,
1957+ .lower_fanout = merge_params->lower_fanout ,
1958+ .leader_fraction = merge_params->leader_fraction ,
1959+ .max_leaders = merge_params->max_leaders ,
1960+ .leaf_size = merge_params->leaf_size ,
1961+ .leaf_degree = merge_params->leaf_degree };
1962+ }
19311963 auto dtype = (*indices[0 ]).dtype ;
19321964 for (size_t i = 1 ; i < num_indices; ++i) {
19331965 RAFT_EXPECTS (indices[i] != nullptr && indices[i]->addr != 0 ,
@@ -1943,13 +1975,17 @@ extern "C" cuvsError_t cuvsCagraMerge(cuvsResources_t res,
19431975 output_index->addr = 0 ;
19441976 // Dispatch based on data type
19451977 if (dtype.code == kDLFloat && dtype.bits == 32 ) {
1946- _merge<float >(res, *params, indices, num_indices, filter, merged_dataset, output_index);
1978+ _merge<float >(
1979+ res, *params, indices, num_indices, filter, merge_params_cpp, merged_dataset, output_index);
19471980 } else if (dtype.code == kDLFloat && dtype.bits == 16 ) {
1948- _merge<half>(res, *params, indices, num_indices, filter, merged_dataset, output_index);
1981+ _merge<half>(
1982+ res, *params, indices, num_indices, filter, merge_params_cpp, merged_dataset, output_index);
19491983 } else if (dtype.code == kDLInt && dtype.bits == 8 ) {
1950- _merge<int8_t >(res, *params, indices, num_indices, filter, merged_dataset, output_index);
1984+ _merge<int8_t >(
1985+ res, *params, indices, num_indices, filter, merge_params_cpp, merged_dataset, output_index);
19511986 } else if (dtype.code == kDLUInt && dtype.bits == 8 ) {
1952- _merge<uint8_t >(res, *params, indices, num_indices, filter, merged_dataset, output_index);
1987+ _merge<uint8_t >(
1988+ res, *params, indices, num_indices, filter, merge_params_cpp, merged_dataset, output_index);
19531989 } else {
19541990 RAFT_FAIL (" Unsupported index data type: code=%d, bits=%d" , dtype.code , dtype.bits );
19551991 }
@@ -1995,6 +2031,26 @@ extern "C" cuvsError_t cuvsCagraIndexParamsDestroy(cuvsCagraIndexParams_t params
19952031 });
19962032}
19972033
2034+ extern " C" cuvsError_t cuvsCagraMergeParamsCreate (cuvsCagraMergeParams_t* params)
2035+ {
2036+ return cuvs::core::translate_exceptions ([=] {
2037+ auto defaults = cuvs::neighbors::cagra::merge_params{};
2038+ *params = new cuvsCagraMergeParams{.algo = CUVS_CAGRA_MERGE_AUTO ,
2039+ .levels = defaults.levels ,
2040+ .root_fanout = defaults.root_fanout ,
2041+ .lower_fanout = defaults.lower_fanout ,
2042+ .leader_fraction = defaults.leader_fraction ,
2043+ .max_leaders = defaults.max_leaders ,
2044+ .leaf_size = defaults.leaf_size ,
2045+ .leaf_degree = defaults.leaf_degree };
2046+ });
2047+ }
2048+
2049+ extern " C" cuvsError_t cuvsCagraMergeParamsDestroy (cuvsCagraMergeParams_t params)
2050+ {
2051+ return cuvs::core::translate_exceptions ([=] { delete params; });
2052+ }
2053+
19982054extern " C" cuvsError_t cuvsCagraCompressionParamsCreate (cuvsCagraCompressionParams_t* params)
19992055{
20002056 return cuvs::core::translate_exceptions ([=] {
0 commit comments