diff --git a/c/include/cuvs/core/dataset.h b/c/include/cuvs/core/dataset.h index 78d3547495..41445e2677 100644 --- a/c/include/cuvs/core/dataset.h +++ b/c/include/cuvs/core/dataset.h @@ -20,7 +20,9 @@ extern "C" { */ typedef enum { CUVS_DATASET_LAYOUT_STANDARD = 0, - CUVS_DATASET_LAYOUT_PADDED = 1 + CUVS_DATASET_LAYOUT_PADDED = 1, + /** Device VPQ storage with f16 codebooks (CAGRA-Q search dataset). */ + CUVS_DATASET_LAYOUT_VPQ_F16 = 2 } cuvsDatasetLayout_t; /** diff --git a/c/include/cuvs/neighbors/cagra.h b/c/include/cuvs/neighbors/cagra.h index 350711d069..d97c6c9d65 100644 --- a/c/include/cuvs/neighbors/cagra.h +++ b/c/include/cuvs/neighbors/cagra.h @@ -255,6 +255,24 @@ CUVS_EXPORT cuvsError_t cuvsCagraCompressionParamsCreate(cuvsCagraCompressionPar */ CUVS_EXPORT cuvsError_t cuvsCagraCompressionParamsDestroy(cuvsCagraCompressionParams_t params); +/** + * @brief Train an owning device VPQ (f16 codebook) dataset from a device-padded source. + * + * Used for CAGRA-Q: build a dense CAGRA index, train VPQ with this factory, then attach via + * `cuvsCagraUpdateDataset`. Caller owns the returned dataset and must keep it alive while any + * index uses it. Metric for subsequent search must remain `L2Expanded`. + * + * @param[in] res cuvs resources + * @param[in] source_dataset device-padded dataset (owning or view) + * @param[in] params VPQ compression params; NULL selects defaults + * @param[out] vpq_dataset newly allocated owning VPQ dataset handle + * @return cuvsError_t + */ +CUVS_EXPORT cuvsError_t cuvsDatasetMakeVpq(cuvsResources_t res, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* vpq_dataset); + /** * @brief Allocate ACE params, and populate with default values * @@ -580,21 +598,25 @@ CUVS_EXPORT cuvsError_t cuvsCagraIndexGetDataset(cuvsCagraIndex_t index, DLManag CUVS_EXPORT cuvsError_t cuvsCagraIndexGetGraph(cuvsCagraIndex_t index, DLManagedTensor* graph); /** - * @brief Update a CAGRA index with a device-padded dataset. + * @brief Update a CAGRA index with a device dataset (padded or VPQ). + * + * This is the centralized dataset update/attach operation for C callers. + * + * - Device-padded dataset: if \p index is already device-padded, its dataset view is replaced in + * place (same index object); otherwise the index is converted via attach and rebound. + * - Device VPQ_F16 dataset (from `cuvsDatasetMakeVpq`): if \p index is already VPQ-typed, its + * dataset view is replaced in place; otherwise the graph is copied into a new VPQ-typed index + * (CAGRA-Q). Search requires metric `L2Expanded`. The VPQ handle must be owning. * - * This is the centralized dataset update operation for C callers. If \p index - * is already device-padded, its dataset view is replaced in place. Otherwise, - * the index is converted and its opaque handle is rebound to a search-ready - * device-padded index. Caller retains ownership of - * \p device_padded_dataset and must keep it alive while \p index uses it. + * Caller retains ownership of \p dataset and must keep it alive while \p index uses it. * - * @param[in] res cuvsResources_t opaque C handle - * @param[in] device_padded_dataset owning or non-owning device-padded dataset handle - * @param[inout] index CAGRA index handle + * @param[in] res cuvsResources_t opaque C handle + * @param[in] dataset device-padded or owning device VPQ_F16 dataset handle + * @param[inout] index CAGRA index handle * @return cuvsError_t */ CUVS_EXPORT cuvsError_t cuvsCagraUpdateDataset(cuvsResources_t res, - cuvsDataset_t device_padded_dataset, + cuvsDataset_t dataset, cuvsCagraIndex_t index); /** diff --git a/c/src/neighbors/cagra.cpp b/c/src/neighbors/cagra.cpp index 99e456e23c..758b2c6cb6 100644 --- a/c/src/neighbors/cagra.cpp +++ b/c/src/neighbors/cagra.cpp @@ -29,7 +29,7 @@ #include #include #include -#include +#include #include "../core/exceptions.hpp" #include "../core/interop.hpp" @@ -52,7 +52,13 @@ struct cuvs_cagra_c_api_index_lifetime_holder { /** Owns how to delete co-located index storage; `cuvsCagraIndex::addr` points here. */ struct sg_cagra_c_api_index_box { void* index_ptr; - enum class dataset_layout : uint8_t { device_padded, device_standard, host_padded, host_standard } layout; + enum class dataset_layout : uint8_t { + device_padded, + device_standard, + host_padded, + host_standard, + device_vpq_f16 + } layout; cuvs::neighbors::c_api::detail::owner_record owner_rec; }; @@ -65,6 +71,8 @@ constexpr auto sg_cagra_index_layout_from_view() return sg_cagra_c_api_index_box::dataset_layout::device_padded; } else if constexpr (cuvs::neighbors::is_host_standard_dataset_view_v) { return sg_cagra_c_api_index_box::dataset_layout::host_standard; + } else if constexpr (cuvs::neighbors::is_device_vpq_f16_dataset_view_v) { + return sg_cagra_c_api_index_box::dataset_layout::device_vpq_f16; } else { return sg_cagra_c_api_index_box::dataset_layout::host_padded; } @@ -110,6 +118,12 @@ static void with_index_by_layout(sg_cagra_c_api_index_box* box, } break; } + case sg_cagra_c_api_index_box::dataset_layout::device_vpq_f16: { + // Intentionally not dispatched here: most C API helpers (serialize/extend/merge/...) do not + // support VPQ. Call sites that need VPQ (search, attach) handle device_vpq_f16 explicitly. + RAFT_FAIL( + "%s: VPQ (CAGRA-Q) index layout is not supported by this operation", null_handle_err); + } } } @@ -524,86 +538,47 @@ static void make_host_standard_dataset_view(raft::resources*, } template -static void attach_dataset(raft::resources* res_ptr, - cuvsDataset_t device_padded_dataset, - cuvsCagraIndex_t index) -{ - RAFT_EXPECTS(device_padded_dataset != nullptr, "cuvsCagraUpdateDataset: null padded dataset"); - RAFT_EXPECTS(index != nullptr, "cuvsCagraUpdateDataset: null index handle"); - RAFT_EXPECTS(index->addr != 0, "cuvsCagraUpdateDataset: null index storage"); - RAFT_EXPECTS(device_padded_dataset->addr != 0, - "cuvsCagraUpdateDataset: null padded dataset storage"); - - auto* box = reinterpret_cast(index->addr); - RAFT_EXPECTS(device_padded_dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE && - device_padded_dataset->layout == CUVS_DATASET_LAYOUT_PADDED, - "cuvsCagraAttachDataset: dataset must be device padded"); +static void make_device_vpq_dataset(raft::resources* res_ptr, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* output_vpq_dataset) +{ + RAFT_EXPECTS(source_dataset != nullptr, "cuvsDatasetMakeVpq: null source dataset"); + RAFT_EXPECTS(source_dataset->addr != 0, "cuvsDatasetMakeVpq: null source dataset storage"); + RAFT_EXPECTS(output_vpq_dataset != nullptr, "cuvsDatasetMakeVpq: null output dataset"); + RAFT_EXPECTS(source_dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE && + source_dataset->layout == CUVS_DATASET_LAYOUT_PADDED, + "cuvsDatasetMakeVpq: source must be a device-padded dataset"); + + cuvs::neighbors::vpq_params ps{}; + if (params != nullptr) { + ps.pq_bits = params->pq_bits; + ps.pq_dim = params->pq_dim; + ps.vq_n_centers = params->vq_n_centers; + ps.kmeans_n_iters = params->kmeans_n_iters; + ps.vq_kmeans_trainset_fraction = params->vq_kmeans_trainset_fraction; + ps.pq_kmeans_trainset_fraction = params->pq_kmeans_trainset_fraction; + } using owner_t = cuvs::neighbors::device_padded_dataset; using view_t = cuvs::neighbors::device_padded_dataset_view; - with_dataset_view(device_padded_dataset, [&](auto const& padded_view) { - with_index_by_layout( - box, - "cuvsCagraUpdateDataset: null index handle", - "cuvsCagraUpdateDataset: host index layout is allowed for this operation", - [&](auto& idx) { - auto padded_idx = cuvs::neighbors::cagra::attach_dataset(*res_ptr, idx, padded_view); - auto* holder = - new cuvs_cagra_c_api_index_lifetime_holder{std::move(padded_idx)}; - destroy_sg_cagra_c_api_box(index->addr); - index->addr = 0; - bind_index_lifetime_holder_to_C_index(index, index->dtype, holder); - }); + with_dataset_view(source_dataset, [&](auto const& padded_view) { + auto vpq = + cuvs::preprocessing::quantize::pq::make_device_vpq_dataset(*res_ptr, ps, padded_view.view()); + using vpq_owner_t = cuvs::neighbors::device_vpq_dataset; + auto* owned = new vpq_owner_t{std::move(vpq)}; + auto* out = new cuvsDataset{}; + out->addr = reinterpret_cast(owned); + out->destroy_addr = &destroy_typed_addr; + // VPQ codebooks use f16 math type; source element type lives on the index dtype. + out->dtype = DLDataType{.code = kDLFloat, .bits = 16, .lanes = 1}; + out->mem_type = CUVS_DATASET_MEM_TYPE_DEVICE; + out->layout = CUVS_DATASET_LAYOUT_VPQ_F16; + out->is_owning = true; + *output_vpq_dataset = out; }); } -template -static void update_device_dataset_same_layout(raft::resources* res_ptr, - cuvsDataset_t device_dataset, - cuvsCagraIndex_t index) -{ - RAFT_EXPECTS(device_dataset != nullptr, "cuvsCagraUpdateDataset: null dataset"); - RAFT_EXPECTS(index != nullptr, "cuvsCagraUpdateDataset: null index handle"); - RAFT_EXPECTS(index->addr != 0, "cuvsCagraUpdateDataset: null index storage"); - RAFT_EXPECTS(device_dataset->addr != 0, "cuvsCagraUpdateDataset: null dataset storage"); - - auto* box = reinterpret_cast(index->addr); - if (box->layout == sg_cagra_c_api_index_box::dataset_layout::device_padded) { - RAFT_EXPECTS(device_dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE && - device_dataset->layout == CUVS_DATASET_LAYOUT_PADDED, - "cuvsCagraUpdateDeviceDatasetSameLayout: device-padded index " - "requires a " - "device-padded dataset"); - using owner_t = cuvs::neighbors::device_padded_dataset; - using view_t = cuvs::neighbors::device_padded_dataset_view; - with_dataset_view(device_dataset, [&](auto const& dataset_view) { - auto* idx = - reinterpret_cast*>(box->index_ptr); - RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); - idx->update_device_dataset_same_layout(*res_ptr, dataset_view); - }); - } else if (box->layout == sg_cagra_c_api_index_box::dataset_layout::device_standard) { - RAFT_EXPECTS(device_dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE && - device_dataset->layout == CUVS_DATASET_LAYOUT_STANDARD, - "cuvsCagraUpdateDeviceDatasetSameLayout: device-standard " - "index requires a " - "device-standard dataset"); - using owner_t = cuvs::neighbors::device_standard_dataset; - using view_t = cuvs::neighbors::device_standard_dataset_view; - with_dataset_view(device_dataset, [&](auto const& dataset_view) { - auto* idx = - reinterpret_cast*>(box->index_ptr); - RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); - idx->update_device_dataset_same_layout(*res_ptr, dataset_view); - }); - } else { - RAFT_FAIL( - "cuvsCagraUpdateDataset: C++ " - "update_device_dataset_same_layout " - "requires a device index and dataset"); - } -} - static void _set_graph_build_params( std::variant(res); auto* box = reinterpret_cast(index.addr); + + auto run_search = [&](auto& idx) { + auto search_params = cuvs::neighbors::cagra::search_params(); + convert_c_search_params(params, &search_params); + + using queries_mdspan_type = raft::device_matrix_view; + using neighbors_mdspan_type = raft::device_matrix_view; + using distances_mdspan_type = raft::device_matrix_view; + auto queries_mds = cuvs::core::from_dlpack(queries_tensor); + auto neighbors_mds = cuvs::core::from_dlpack(neighbors_tensor); + auto distances_mds = cuvs::core::from_dlpack(distances_tensor); + if (filter.type == NO_FILTER) { + cuvs::neighbors::cagra::search( + *res_ptr, search_params, idx, queries_mds, neighbors_mds, distances_mds); + } else if (filter.type == BITSET) { + using filter_mdspan_type = raft::device_vector_view; + auto removed_indices_tensor = reinterpret_cast(filter.addr); + auto removed_indices = cuvs::core::from_dlpack(removed_indices_tensor); + cuvs::core::bitset_view removed_indices_bitset( + removed_indices, idx.dataset().n_rows()); + auto bitset_filter_obj = cuvs::neighbors::filtering::bitset_filter(removed_indices_bitset); + cuvs::neighbors::cagra::search(*res_ptr, + search_params, + idx, + queries_mds, + neighbors_mds, + distances_mds, + bitset_filter_obj); + } else { + RAFT_FAIL("Unsupported filter type: BITMAP"); + } + }; + + if (box->layout == sg_cagra_c_api_index_box::dataset_layout::device_vpq_f16) { + auto* idx = + reinterpret_cast*>(box->index_ptr); + RAFT_EXPECTS(idx != nullptr, "cuvsCagraSearch: null index handle"); + run_search(*idx); + return; + } + with_index_by_layout( box, "cuvsCagraSearch: null index handle", "cuvsCagraSearch: host index must be converted to device first via " "cuvsCagraUpdateDataset with a device padded dataset view", - [&](auto& idx) { - auto search_params = cuvs::neighbors::cagra::search_params(); - convert_c_search_params(params, &search_params); - - using queries_mdspan_type = raft::device_matrix_view; - using neighbors_mdspan_type = raft::device_matrix_view; - using distances_mdspan_type = raft::device_matrix_view; - auto queries_mds = cuvs::core::from_dlpack(queries_tensor); - auto neighbors_mds = cuvs::core::from_dlpack(neighbors_tensor); - auto distances_mds = cuvs::core::from_dlpack(distances_tensor); - if (filter.type == NO_FILTER) { - cuvs::neighbors::cagra::search( - *res_ptr, search_params, idx, queries_mds, neighbors_mds, distances_mds); - } else if (filter.type == BITSET) { - using filter_mdspan_type = raft::device_vector_view; - auto removed_indices_tensor = reinterpret_cast(filter.addr); - auto removed_indices = cuvs::core::from_dlpack(removed_indices_tensor); - cuvs::core::bitset_view removed_indices_bitset( - removed_indices, idx.dataset().n_rows()); - auto bitset_filter_obj = cuvs::neighbors::filtering::bitset_filter(removed_indices_bitset); - cuvs::neighbors::cagra::search(*res_ptr, - search_params, - idx, - queries_mds, - neighbors_mds, - distances_mds, - bitset_filter_obj); - } else { - RAFT_FAIL("Unsupported filter type: BITMAP"); - } - }); + run_search); } template @@ -1495,6 +1481,71 @@ extern "C" cuvsError_t cuvsDatasetMakePaddedView(cuvsResources_t res, }); } +extern "C" cuvsError_t cuvsDatasetMakeStandardView(cuvsResources_t res, + DLManagedTensor* dataset_tensor, + cuvsDataset_t* standard_dataset) +{ + return cuvs::core::translate_exceptions([=] { + RAFT_EXPECTS(dataset_tensor != nullptr, "cuvsDatasetMakeStandardView: null input tensor"); + RAFT_EXPECTS(standard_dataset != nullptr, "cuvsDatasetMakeStandardView: null output view"); + *standard_dataset = nullptr; + auto dataset = dataset_tensor->dl_tensor; + auto* res_ptr = reinterpret_cast(res); + auto make_typed = [&]() { + if (cuvs::core::is_dlpack_device_compatible(dataset)) { + make_device_standard_dataset_view(res_ptr, dataset_tensor, standard_dataset); + } else if (cuvs::core::is_dlpack_host_compatible(dataset)) { + make_host_standard_dataset_view(res_ptr, dataset_tensor, standard_dataset); + } else { + RAFT_FAIL("cuvsDatasetMakeStandardView: unsupported tensor memory type"); + } + }; + + if (dataset.dtype.code == kDLFloat && dataset.dtype.bits == 32) { + make_typed.template operator()(); + } else if (dataset.dtype.code == kDLFloat && dataset.dtype.bits == 16) { + make_typed.template operator()(); + } else if (dataset.dtype.code == kDLInt && dataset.dtype.bits == 8) { + make_typed.template operator()(); + } else if (dataset.dtype.code == kDLUInt && dataset.dtype.bits == 8) { + make_typed.template operator()(); + } else { + RAFT_FAIL("Unsupported dataset DLtensor dtype: %d and bits: %d", + dataset.dtype.code, + dataset.dtype.bits); + } + }); +} + +extern "C" cuvsError_t cuvsDatasetMakeVpq(cuvsResources_t res, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* vpq_dataset) +{ + return cuvs::core::translate_exceptions([=] { + RAFT_EXPECTS(source_dataset != nullptr, "cuvsDatasetMakeVpq: null source dataset"); + RAFT_EXPECTS(vpq_dataset != nullptr, "cuvsDatasetMakeVpq: null output dataset"); + auto* res_ptr = reinterpret_cast(res); + auto make_typed = [&]() { + make_device_vpq_dataset(res_ptr, source_dataset, params, vpq_dataset); + }; + + if (source_dataset->dtype.code == kDLFloat && source_dataset->dtype.bits == 32) { + make_typed.template operator()(); + } else if (source_dataset->dtype.code == kDLFloat && source_dataset->dtype.bits == 16) { + make_typed.template operator()(); + } else if (source_dataset->dtype.code == kDLInt && source_dataset->dtype.bits == 8) { + make_typed.template operator()(); + } else if (source_dataset->dtype.code == kDLUInt && source_dataset->dtype.bits == 8) { + make_typed.template operator()(); + } else { + RAFT_FAIL("cuvsDatasetMakeVpq: unsupported source dtype: %d and bits: %d", + source_dataset->dtype.code, + source_dataset->dtype.bits); + } + }); +} + extern "C" cuvsError_t cuvsDatasetDestroy(cuvsDataset_t dataset) { return cuvs::core::translate_exceptions([=] { @@ -1542,122 +1593,144 @@ extern "C" cuvsError_t cuvsDatasetGetDtype(cuvsDataset_t dataset, DLDataType* dt }); } -extern "C" cuvsError_t cuvsDatasetMakeStandardView(cuvsResources_t res, - DLManagedTensor* dataset_tensor, - cuvsDataset_t* standard_dataset) -{ +extern "C" cuvsError_t cuvsCagraUpdateDataset(cuvsResources_t res, + cuvsDataset_t dataset, + cuvsCagraIndex_t index) { return cuvs::core::translate_exceptions([=] { - RAFT_EXPECTS(dataset_tensor != nullptr, "cuvsDatasetMakeStandardView: null input tensor"); - RAFT_EXPECTS(standard_dataset != nullptr, "cuvsDatasetMakeStandardView: null output view"); - *standard_dataset = nullptr; - auto dataset = dataset_tensor->dl_tensor; - auto* res_ptr = reinterpret_cast(res); - auto make_typed = [&]() { - if (cuvs::core::is_dlpack_device_compatible(dataset)) { - make_device_standard_dataset_view(res_ptr, dataset_tensor, standard_dataset); - } else if (cuvs::core::is_dlpack_host_compatible(dataset)) { - make_host_standard_dataset_view(res_ptr, dataset_tensor, standard_dataset); - } else { - RAFT_FAIL("cuvsDatasetMakeStandardView: unsupported tensor memory type"); + RAFT_EXPECTS(index != nullptr, "cuvsCagraUpdateDataset: null index handle"); + RAFT_EXPECTS(index->addr != 0, + "cuvsCagraUpdateDataset: null index storage"); + RAFT_EXPECTS(dataset != nullptr, + "cuvsCagraUpdateDataset: null dataset view"); + RAFT_EXPECTS(dataset->addr != 0, + "cuvsCagraUpdateDataset: null dataset view storage"); + RAFT_EXPECTS(dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE, + "cuvsCagraUpdateDataset: dataset must be device-resident"); + RAFT_EXPECTS(dataset->layout == CUVS_DATASET_LAYOUT_PADDED || + dataset->layout == CUVS_DATASET_LAYOUT_VPQ_F16, + "cuvsCagraUpdateDataset: dataset must be device-padded or " + "device VPQ_F16"); + + auto *res_ptr = reinterpret_cast(res); + auto *box = reinterpret_cast(index->addr); + + using layout_t = sg_cagra_c_api_index_box::dataset_layout; + + auto update = [&]() { + if (dataset->layout == CUVS_DATASET_LAYOUT_PADDED) { + RAFT_EXPECTS(index->dtype.code == dataset->dtype.code && + index->dtype.bits == dataset->dtype.bits, + "cuvsCagraUpdateDataset: dtype mismatch between index and dataset"); + + using owner_t = cuvs::neighbors::device_padded_dataset; + using view_t = cuvs::neighbors::device_padded_dataset_view; + with_dataset_view(dataset, [&](auto const& dataset_view) { + auto attach_and_rebind = [&](auto* idx) { + RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); + auto updated_idx = cuvs::neighbors::cagra::update_dataset(*res_ptr, *idx, dataset_view); + auto* holder = + new cuvs_cagra_c_api_index_lifetime_holder{std::move(updated_idx)}; + destroy_sg_cagra_c_api_box(index->addr); + index->addr = 0; + bind_index_lifetime_holder_to_C_index(index, index->dtype, holder); + }; + + switch (box->layout) { + case layout_t::device_padded: { + auto* idx = reinterpret_cast< + cuvs::neighbors::cagra::device_padded_index*>(box->index_ptr); + RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); + cuvs::neighbors::cagra::update_dataset(*res_ptr, *idx, dataset_view); + break; + } + case layout_t::device_standard: + attach_and_rebind(reinterpret_cast< + cuvs::neighbors::cagra::device_standard_index*>( + box->index_ptr)); + break; + case layout_t::host_standard: + attach_and_rebind( + reinterpret_cast*>( + box->index_ptr)); + break; + case layout_t::host_padded: + attach_and_rebind( + reinterpret_cast*>( + box->index_ptr)); + break; + case layout_t::device_vpq_f16: + RAFT_FAIL( + "cuvsCagraUpdateDataset: cannot attach a padded dataset to a VPQ index; " + "pass a device VPQ_F16 dataset from cuvsDatasetMakeVpq"); + } + }); + } else if (dataset->layout == CUVS_DATASET_LAYOUT_VPQ_F16) { + RAFT_EXPECTS(dataset->is_owning, + "cuvsCagraUpdateDataset: VPQ dataset handle must be owning " + "(from cuvsDatasetMakeVpq)"); + + using owner_t = cuvs::neighbors::device_vpq_dataset; + using view_t = cuvs::neighbors::device_vpq_dataset_view; + with_dataset_view(dataset, [&](auto const& dataset_view) { + auto attach_and_rebind = [&](auto* idx) { + RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); + auto updated_idx = cuvs::neighbors::cagra::update_dataset(*res_ptr, *idx, dataset_view); + auto* holder = + new cuvs_cagra_c_api_index_lifetime_holder{std::move(updated_idx)}; + destroy_sg_cagra_c_api_box(index->addr); + index->addr = 0; + bind_index_lifetime_holder_to_C_index(index, index->dtype, holder); + }; + + switch (box->layout) { + case layout_t::device_vpq_f16: { + auto* idx = + reinterpret_cast*>( + box->index_ptr); + RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); + cuvs::neighbors::cagra::update_dataset(*res_ptr, *idx, dataset_view); + break; + } + case layout_t::device_padded: + attach_and_rebind(reinterpret_cast< + cuvs::neighbors::cagra::device_padded_index*>( + box->index_ptr)); + break; + case layout_t::device_standard: + attach_and_rebind(reinterpret_cast< + cuvs::neighbors::cagra::device_standard_index*>( + box->index_ptr)); + break; + case layout_t::host_standard: + attach_and_rebind( + reinterpret_cast*>( + box->index_ptr)); + break; + case layout_t::host_padded: + attach_and_rebind( + reinterpret_cast*>( + box->index_ptr)); + break; + } + }); } }; - if (dataset.dtype.code == kDLFloat && dataset.dtype.bits == 32) { - make_typed.template operator()(); - } else if (dataset.dtype.code == kDLFloat && dataset.dtype.bits == 16) { - make_typed.template operator()(); - } else if (dataset.dtype.code == kDLInt && dataset.dtype.bits == 8) { - make_typed.template operator()(); - } else if (dataset.dtype.code == kDLUInt && dataset.dtype.bits == 8) { - make_typed.template operator()(); - } else { - RAFT_FAIL("Unsupported dataset DLtensor dtype: %d and bits: %d", - dataset.dtype.code, - dataset.dtype.bits); - } - }); -} - -static cuvsError_t dispatch_attach_dataset(cuvsResources_t res, - cuvsDataset_t device_padded_dataset, - cuvsCagraIndex_t index) -{ - return cuvs::core::translate_exceptions([=] { - auto* res_ptr = reinterpret_cast(res); - RAFT_EXPECTS(index != nullptr, "cuvsCagraUpdateDataset: null index handle"); - RAFT_EXPECTS(device_padded_dataset != nullptr, "cuvsCagraUpdateDataset: null dataset view"); - RAFT_EXPECTS(device_padded_dataset->layout == CUVS_DATASET_LAYOUT_PADDED, - "cuvsCagraUpdateDataset: dataset handle layout must be PADDED"); - RAFT_EXPECTS(index->dtype.code == device_padded_dataset->dtype.code && - index->dtype.bits == device_padded_dataset->dtype.bits, - "cuvsCagraUpdateDataset: dtype mismatch between index and dataset"); - if (index->dtype.code == kDLFloat && index->dtype.bits == 32) { - attach_dataset(res_ptr, device_padded_dataset, index); - } else if (index->dtype.code == kDLFloat && index->dtype.bits == 16) { - attach_dataset(res_ptr, device_padded_dataset, index); - } else if (index->dtype.code == kDLInt && index->dtype.bits == 8) { - attach_dataset(res_ptr, device_padded_dataset, index); - } else if (index->dtype.code == kDLUInt && index->dtype.bits == 8) { - attach_dataset(res_ptr, device_padded_dataset, index); - } else { - RAFT_FAIL("Unsupported index dtype: %d and bits: %d", index->dtype.code, index->dtype.bits); - } - }); -} - -static cuvsError_t dispatch_update_device_dataset_same_layout(cuvsResources_t res, - cuvsDataset_t device_dataset, - cuvsCagraIndex_t index) -{ - return cuvs::core::translate_exceptions([=] { - auto* res_ptr = reinterpret_cast(res); - RAFT_EXPECTS(index != nullptr, "cuvsCagraUpdateDataset: null index handle"); - RAFT_EXPECTS(device_dataset != nullptr, - "cuvsCagraUpdateDataset: null dataset view"); - RAFT_EXPECTS(index->dtype.code == device_dataset->dtype.code && - index->dtype.bits == device_dataset->dtype.bits, - "cuvsCagraUpdateDataset: dtype mismatch " - "between index and dataset"); if (index->dtype.code == kDLFloat && index->dtype.bits == 32) { - update_device_dataset_same_layout(res_ptr, device_dataset, index); + update.template operator()(); } else if (index->dtype.code == kDLFloat && index->dtype.bits == 16) { - update_device_dataset_same_layout(res_ptr, device_dataset, index); + update.template operator()(); } else if (index->dtype.code == kDLInt && index->dtype.bits == 8) { - update_device_dataset_same_layout(res_ptr, device_dataset, index); + update.template operator()(); } else if (index->dtype.code == kDLUInt && index->dtype.bits == 8) { - update_device_dataset_same_layout(res_ptr, device_dataset, index); + update.template operator()(); } else { - RAFT_FAIL("Unsupported index dtype: %d and bits: %d", index->dtype.code, index->dtype.bits); + RAFT_FAIL("Unsupported index dtype: %d and bits: %d", index->dtype.code, + index->dtype.bits); } }); } -extern "C" cuvsError_t cuvsCagraUpdateDataset(cuvsResources_t res, - cuvsDataset_t device_padded_dataset, - cuvsCagraIndex_t index) -{ - auto status = cuvs::core::translate_exceptions([=] { - RAFT_EXPECTS(index != nullptr, "cuvsCagraUpdateDataset: null index handle"); - RAFT_EXPECTS(index->addr != 0, "cuvsCagraUpdateDataset: null index storage"); - RAFT_EXPECTS(device_padded_dataset != nullptr, "cuvsCagraUpdateDataset: null dataset view"); - RAFT_EXPECTS(device_padded_dataset->addr != 0, - "cuvsCagraUpdateDataset: null dataset view storage"); - RAFT_EXPECTS(device_padded_dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE && - device_padded_dataset->layout == CUVS_DATASET_LAYOUT_PADDED, - "cuvsCagraUpdateDataset: dataset view must be device padded"); - RAFT_EXPECTS(index->dtype.code == device_padded_dataset->dtype.code && - index->dtype.bits == device_padded_dataset->dtype.bits, - "cuvsCagraUpdateDataset: dtype mismatch between index and dataset"); - }); - if (status != CUVS_SUCCESS) { return status; } - - auto* box = reinterpret_cast(index->addr); - if (box->layout == sg_cagra_c_api_index_box::dataset_layout::device_padded) { - return dispatch_update_device_dataset_same_layout(res, device_padded_dataset, index); - } - return dispatch_attach_dataset(res, device_padded_dataset, index); -} - /** * Build from an already-constructed C++ dataset view. `DatasetViewT` selects the * `cuvs::neighbors::cagra::build` overload, and therefore the resulting index type. @@ -1839,9 +1912,10 @@ extern "C" cuvsError_t cuvsCagraSearch(cuvsResources_t res, auto index = *index_c_ptr; auto* box = reinterpret_cast(index.addr); RAFT_EXPECTS(box != nullptr, "cuvsCagraSearch: null index handle"); - RAFT_EXPECTS(box->layout == sg_cagra_c_api_index_box::dataset_layout::device_padded, - "cuvsCagraSearch: index must be device-padded. For standard indices, call " - "cuvsCagraUpdateDataset first."); + RAFT_EXPECTS(box->layout == sg_cagra_c_api_index_box::dataset_layout::device_padded || + box->layout == sg_cagra_c_api_index_box::dataset_layout::device_vpq_f16, + "cuvsCagraSearch: index must be device-padded or device-VPQ. Call " + "cuvsCagraUpdateDataset with a device-padded or owning VPQ_F16 dataset."); RAFT_EXPECTS(queries.dtype.code == index.dtype.code, "type mismatch between index and queries"); if (queries.dtype.code == kDLFloat && queries.dtype.bits == 32) { diff --git a/c/tests/neighbors/ann_cagra_c.cu b/c/tests/neighbors/ann_cagra_c.cu index 483e34dcb7..17046ad7e7 100644 --- a/c/tests/neighbors/ann_cagra_c.cu +++ b/c/tests/neighbors/ann_cagra_c.cu @@ -19,6 +19,7 @@ #include #include #include +#include #include #include @@ -2008,3 +2009,109 @@ TEST(CagraC, SearchMultiPartitionMultiKernelRejected) } cuvsResourcesDestroy(res); } + +TEST(CagraC, BuildAttachVpqSearch) +{ + // CAGRA-Q smoke test: dense build → MakeVpq → UpdateDataset(VPQ) → Search. + constexpr int64_t n_rows = 256; + constexpr int64_t dim = 32; + constexpr int64_t n_queries = 4; + constexpr int64_t k = 1; + + cuvsResources_t res; + ASSERT_EQ(cuvsResourcesCreate(&res), CUVS_SUCCESS); + cudaStream_t stream; + ASSERT_EQ(cuvsStreamGet(res, &stream), CUVS_SUCCESS); + + rmm::device_uvector dataset_d(n_rows * dim, stream); + { + std::vector host(n_rows * dim); + for (int64_t i = 0; i < n_rows * dim; ++i) { + host[i] = static_cast((i % 17) + 1); + } + raft::copy(dataset_d.data(), host.data(), host.size(), stream); + } + + // dim=32 float already matches CAGRA padded row width; MakePadded refuses a + // no-op device copy — wrap with MakePaddedView instead. + DLManagedTensor dataset_tensor{}; + dataset_tensor.dl_tensor.data = dataset_d.data(); + dataset_tensor.dl_tensor.device.device_type = kDLCUDA; + dataset_tensor.dl_tensor.ndim = 2; + dataset_tensor.dl_tensor.dtype = {kDLFloat, 32, 1}; + int64_t dataset_shape[2] = {n_rows, dim}; + dataset_tensor.dl_tensor.shape = dataset_shape; + dataset_tensor.dl_tensor.strides = nullptr; + + cuvsDataset_t padded; + ASSERT_EQ(cuvsDatasetMakePaddedView(res, &dataset_tensor, &padded), CUVS_SUCCESS); + + cuvsCagraIndexParams_t build_params; + ASSERT_EQ(cuvsCagraIndexParamsCreate(&build_params), CUVS_SUCCESS); + cuvsCagraIndex_t index; + ASSERT_EQ(cuvsCagraIndexCreate(&index), CUVS_SUCCESS); + ASSERT_EQ(cuvsCagraBuild(res, build_params, padded, index), CUVS_SUCCESS); + + cuvsCagraCompressionParams_t compression; + ASSERT_EQ(cuvsCagraCompressionParamsCreate(&compression), CUVS_SUCCESS); + compression->pq_bits = 8; + compression->pq_dim = 8; + + cuvsDataset_t vpq = nullptr; + ASSERT_EQ(cuvsDatasetMakeVpq(res, padded, compression, &vpq), CUVS_SUCCESS); + { + cuvsDatasetLayout_t layout; + ASSERT_EQ(cuvsDatasetGetLayout(vpq, &layout), CUVS_SUCCESS); + EXPECT_EQ(layout, CUVS_DATASET_LAYOUT_VPQ_F16); + bool owning = false; + ASSERT_EQ(cuvsDatasetGetIsOwning(vpq, &owning), CUVS_SUCCESS); + EXPECT_TRUE(owning); + } + + ASSERT_EQ(cuvsCagraUpdateDataset(res, vpq, index), CUVS_SUCCESS); + + rmm::device_uvector queries_d(n_queries * dim, stream); + raft::copy(queries_d.data(), dataset_d.data(), n_queries * dim, stream); + DLManagedTensor queries_tensor{}; + queries_tensor.dl_tensor.data = queries_d.data(); + queries_tensor.dl_tensor.device.device_type = kDLCUDA; + queries_tensor.dl_tensor.ndim = 2; + queries_tensor.dl_tensor.dtype = {kDLFloat, 32, 1}; + int64_t queries_shape[2] = {n_queries, dim}; + queries_tensor.dl_tensor.shape = queries_shape; + + rmm::device_uvector neighbors_d(n_queries * k, stream); + DLManagedTensor neighbors_tensor{}; + neighbors_tensor.dl_tensor.data = neighbors_d.data(); + neighbors_tensor.dl_tensor.device.device_type = kDLCUDA; + neighbors_tensor.dl_tensor.ndim = 2; + neighbors_tensor.dl_tensor.dtype = {kDLUInt, 32, 1}; + int64_t neighbors_shape[2] = {n_queries, k}; + neighbors_tensor.dl_tensor.shape = neighbors_shape; + + rmm::device_uvector distances_d(n_queries * k, stream); + DLManagedTensor distances_tensor{}; + distances_tensor.dl_tensor.data = distances_d.data(); + distances_tensor.dl_tensor.device.device_type = kDLCUDA; + distances_tensor.dl_tensor.ndim = 2; + distances_tensor.dl_tensor.dtype = {kDLFloat, 32, 1}; + int64_t distances_shape[2] = {n_queries, k}; + distances_tensor.dl_tensor.shape = distances_shape; + + cuvsFilter filter; + filter.type = NO_FILTER; + filter.addr = (uintptr_t)NULL; + cuvsCagraSearchParams_t search_params; + ASSERT_EQ(cuvsCagraSearchParamsCreate(&search_params), CUVS_SUCCESS); + ASSERT_EQ(cuvsCagraSearch( + res, search_params, index, &queries_tensor, &neighbors_tensor, &distances_tensor, filter), + CUVS_SUCCESS); + + cuvsCagraSearchParamsDestroy(search_params); + cuvsCagraCompressionParamsDestroy(compression); + cuvsDatasetDestroy(vpq); + cuvsCagraIndexDestroy(index); + cuvsCagraIndexParamsDestroy(build_params); + cuvsDatasetDestroy(padded); + cuvsResourcesDestroy(res); +} diff --git a/cpp/CMakeLists.txt b/cpp/CMakeLists.txt index 4cff54e011..e808982983 100644 --- a/cpp/CMakeLists.txt +++ b/cpp/CMakeLists.txt @@ -1186,6 +1186,13 @@ if(NOT BUILD_CPU_ONLY) OUTPUT_FILE_FORMAT "${CMAKE_CURRENT_BINARY_DIR}/src/neighbors/cagra_extend_inst_data_@data_abbrev@_index_@index_abbrev@.cu" ) + generate_inst_matrix( + cagra_update_dataset_inst_files + MATRIX_JSON_FILE "${CMAKE_CURRENT_SOURCE_DIR}/src/neighbors/cagra_update_dataset_matrix.json" + INPUT_FILE "${CMAKE_CURRENT_SOURCE_DIR}/src/neighbors/cagra_update_dataset_inst.cu.in" + OUTPUT_FILE_FORMAT + "${CMAKE_CURRENT_BINARY_DIR}/src/neighbors/cagra_update_dataset_inst_data_@data_abbrev@_index_@index_abbrev@.cu" + ) generate_inst_matrix( cagra_serialize_inst_files MATRIX_JSON_FILE "${CMAKE_CURRENT_SOURCE_DIR}/src/neighbors/cagra_serialize_matrix.json" @@ -1387,6 +1394,7 @@ if(NOT BUILD_CPU_ONLY) src/neighbors/cagra.cpp ${cagra_build_inst_files} ${cagra_extend_inst_files} + ${cagra_update_dataset_inst_files} src/neighbors/cagra_optimize.cu src/neighbors/detail/cagra/graph_shared.cu ${cagra_serialize_inst_files} diff --git a/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h b/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h index 23eeabb00f..800ae4e6a1 100644 --- a/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h +++ b/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h @@ -400,11 +400,12 @@ void cuvs_cagra::compress_dataset(const T* dataset, size_t nrow) "cagra: compression_* (CAGRA-Q) requires the graph in memory; it cannot be combined " "with a disk-resident (ACE) graph."); auto rows = static_cast(nrow); - // make_vpq_dataset() reads the rows wherever they are: host-resident ones are subsampled and - // encoded in bounded batches instead of being staged on the device. + // make_device_vpq_dataset() reads the rows wherever they are: host-resident ones are subsampled + // and encoded in bounded batches instead of being staged on the device. auto src = raft::make_device_matrix_view(dataset, rows, dim_); vpq_dataset_ = std::make_shared>( - cuvs::preprocessing::quantize::pq::make_vpq_dataset(handle_, *index_params_.compression, src)); + cuvs::preprocessing::quantize::pq::make_device_vpq_dataset( + handle_, *index_params_.compression, src)); vpq_index_ = std::make_shared>( handle_, parse_metric_type(metric_), vpq_dataset_->as_dataset_view(), index_->graph()); diff --git a/cpp/include/cuvs/neighbors/cagra.hpp b/cpp/include/cuvs/neighbors/cagra.hpp index 43ae7a6235..8eabf19e02 100644 --- a/cpp/include/cuvs/neighbors/cagra.hpp +++ b/cpp/include/cuvs/neighbors/cagra.hpp @@ -203,14 +203,13 @@ struct index_params : cuvs::neighbors::index_params { * as the index is used. A device-backed index is ready to search immediately; a host-backed index * retains the dataset for operations such as serialization but is not searchable. * - `false` means `build` only builds the graph and the caller is expected to attach a dataset - * separately via `cuvs::neighbors::cagra::index::update_device_dataset_same_layout` before - * searching. + * separately via `cuvs::neighbors::cagra::update_dataset` before searching. * * Unlike the legacy behavior, no copy of the dataset is made: the index always stores a view. * Setting `attach_dataset_on_build = false` is useful when the caller needs to apply specific * memory placement or transformation (e.g. moving to managed memory) before attaching. * - * Host indexes are not directly searchable. Call `attach_dataset` with a user-provided + * Host indexes are not directly searchable. Call `update_dataset` with a user-provided * device-padded dataset view to obtain a search-ready `device_padded_index`. Disk-based ACE * builds manage file-backed dataset state separately and ignore this flag. * @@ -222,7 +221,7 @@ struct index_params : cuvs::neighbors::index_params { * auto index = cagra::build(res, index_params, dataset->as_dataset_view()); * // ASSERT(index.size() == 0); // no dataset yet * // Attach with a view (storage owned by `dataset`). - * index.update_device_dataset_same_layout(res, dataset->as_dataset_view()); + * cagra::update_dataset(res, index, dataset->as_dataset_view()); * cagra::search(res, search_params, index, queries, neighbors, distances); * @endcode */ @@ -941,7 +940,7 @@ using cagra_index_t = index - requires cuvs::neighbors::is_host_dataset_view_v -auto convert_host_to_device_index(raft::resources const& res, index const& src) - -> index> -{ - using DeviceViewT = cuvs::neighbors::device_counterpart_t; - using GraphIndexType = typename index::graph_index_type; - index out(res, src.metric()); - if (src.graph().size() > 0) { - // The graph lives in device memory owned by `src`. `update_graph(device_view)` would only - // store a view (no ownership transfer), leaving `out` with a dangling pointer once `src` - // is destroyed. Copy device→host→device so that `out` owns its graph memory. - auto graph_host = - raft::make_host_matrix(src.graph().extent(0), src.graph().extent(1)); - raft::copy(graph_host.data_handle(), - src.graph().data_handle(), - src.graph().size(), - raft::resource::get_cuda_stream(res)); - raft::resource::sync_stream(res); - out.update_graph(res, raft::make_const_mdspan(graph_host.view())); // host overload: copies H→D - } - return out; -} - } // namespace detail /** - * @brief Convert a standard-device index into a padded-device index and attach padded dataset. - * - * CAGRA search requires padded device layout. This helper copies graph/source-indices from - * `standard_idx` into a new `device_padded_index` and attaches `padded_dataset`. - * - * @param[in] res RAFT resources - * @param[in] standard_idx index returned by `build` with a standard device dataset view - * @param[in] padded_dataset device padded dataset view (caller owns underlying memory) - * @return device padded index with graph and dataset ready for search - */ -template -auto convert_standard_to_padded_index( - raft::resources const& res, - index> const& standard_idx, - cuvs::neighbors::device_padded_dataset_view const& padded_dataset) - -> device_padded_index -{ - RAFT_EXPECTS(padded_dataset.n_rows() == standard_idx.size(), - "Padded dataset row count must match the index size"); - - device_padded_index out(res, standard_idx.metric()); - if (standard_idx.graph().extent(0) > 0) { - using GraphIndexType = - typename index>:: - graph_index_type; - auto graph_host = raft::make_host_matrix( - standard_idx.graph().extent(0), standard_idx.graph().extent(1)); - raft::copy(graph_host.data_handle(), - standard_idx.graph().data_handle(), - standard_idx.graph().size(), - raft::resource::get_cuda_stream(res)); - raft::resource::sync_stream(res); - out.update_graph(res, raft::make_const_mdspan(graph_host.view())); - } - if (standard_idx.source_indices().has_value()) { - out.update_source_indices(res, standard_idx.source_indices().value()); - } - out.update_device_dataset_same_layout(res, padded_dataset); - return out; -} - -/** - * @brief Attach a device-padded dataset and return a search-ready padded-device index. - * - * The dataset is provided by the caller and must already be device-padded. - * - * For host/standard index layouts, this function converts to and returns a new - * `device_padded_index`. - * - * If `idx` is already a `device_padded_index`, call `idx.update_device_dataset_same_layout(res, - * device_padded_dataset)` directly to avoid an unnecessary copy path. - * - * @param[in] res RAFT resources - * @param[in] idx CAGRA index in any host/device + standard/padded layout - * @param[in] device_padded_dataset caller-owned device-padded dataset view - * @return search-ready padded-device CAGRA index - */ -template - requires cuvs::neighbors::ann_dataset_view -auto attach_dataset( - raft::resources const& res, - index const& idx, - cuvs::neighbors::device_padded_dataset_view const& device_padded_dataset) - -> device_padded_index -{ - RAFT_EXPECTS(device_padded_dataset.n_rows() == idx.size(), - "Padded dataset row count must match the index size"); - - if constexpr (cuvs::neighbors::is_host_standard_dataset_view_v) { - auto dev_std = detail::convert_host_to_device_index(res, idx); - return convert_standard_to_padded_index(res, dev_std, device_padded_dataset); - } else if constexpr (cuvs::neighbors::is_host_padded_dataset_view_v) { - auto dev_pad = detail::convert_host_to_device_index(res, idx); - dev_pad.update_device_dataset_same_layout(res, device_padded_dataset); - return dev_pad; - } else if constexpr (cuvs::neighbors::is_device_standard_dataset_view_v) { - return convert_standard_to_padded_index(res, idx, device_padded_dataset); - } else if constexpr (cuvs::neighbors::is_device_padded_dataset_view_v) { - RAFT_LOG_WARN( - "cagra::attach_dataset called with an already device-padded index. " - "To avoid an unnecessary index copy, call " - "index.update_device_dataset_same_layout(res, device_padded_dataset) " - "directly on the original index."); - RAFT_FAIL( - "cagra::attach_dataset: device_padded_index input is not supported in this overload. " - "Call index.update_device_dataset_same_layout(res, device_padded_dataset) directly."); - } else { - static_assert(!sizeof(IndexViewT), "Unsupported CAGRA index dataset view type"); - } -} + * @brief Update or attach a device dataset to a CAGRA index. + * + * These overloads are the single C++ API entry point for changing an index dataset. + * + * When the dataset layout changes, the input index is immutable and the overload returns a new + * index whose C++ type reflects the new layout. When the layout is unchanged, the overload accepts + * a mutable index and updates it in place. Dataset storage remains owned by the caller. + */ +#define CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS(T) \ + auto update_dataset(raft::resources const& res, \ + host_standard_index const& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + -> device_padded_index; \ + auto update_dataset(raft::resources const& res, \ + host_padded_index const& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + -> device_padded_index; \ + auto update_dataset(raft::resources const& res, \ + device_standard_index const& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + -> device_padded_index; \ + auto update_dataset(raft::resources const& res, \ + host_standard_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index; \ + auto update_dataset(raft::resources const& res, \ + host_padded_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index; \ + auto update_dataset(raft::resources const& res, \ + device_standard_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index; \ + auto update_dataset(raft::resources const& res, \ + device_padded_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index; \ + void update_dataset(raft::resources const& res, \ + device_standard_index& idx, \ + cuvs::neighbors::device_standard_dataset_view const& dataset); \ + void update_dataset(raft::resources const& res, \ + device_padded_index& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset); \ + void update_dataset(raft::resources const& res, \ + vpq_f16_index& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) + +CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS(float); +CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS(half); +CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS(int8_t); +CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS(uint8_t); + +#undef CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS } // namespace cagra } // namespace neighbors diff --git a/cpp/include/cuvs/preprocessing/quantize/pq.hpp b/cpp/include/cuvs/preprocessing/quantize/pq.hpp index 112341f2ad..a216cb50bf 100644 --- a/cpp/include/cuvs/preprocessing/quantize/pq.hpp +++ b/cpp/include/cuvs/preprocessing/quantize/pq.hpp @@ -281,24 +281,25 @@ namespace detail { * * Typical **CAGRA** usage: build the graph on dense vectors, then attach VPQ for search (metric * must remain `L2Expanded` for this path). Train VPQ from the same CAGRA-padded device layout you - * used for graph build, keep the `device_vpq_dataset` alive, and call - * `index::update_device_dataset_same_layout` with a non-owning view. + * used for graph build, keep the `device_vpq_dataset` alive, and attach it with + * `cagra::update_dataset` (returns a `vpq_f16_index`). * * @code{.cpp} * #include * #include * - * // `idx` is a `cagra::index` with graph built on dense rows. + * // `idx` is a dense CAGRA index with graph built on padded rows. * // `padded` is a `device_padded_dataset_view` view of those same rows. * cuvs::neighbors::vpq_params vpq_params{}; - * auto vpq = cuvs::preprocessing::quantize::pq::make_vpq_dataset(res, vpq_params, padded); - * idx.update_device_dataset_same_layout(res, vpq.as_dataset_view()); + * auto vpq = cuvs::preprocessing::quantize::pq::make_device_vpq_dataset( + * res, vpq_params, padded); + * auto vpq_idx = cagra::update_dataset(res, idx, vpq.as_dataset_view()); * @endcode */ template -[[nodiscard]] auto make_vpq_dataset(raft::resources const& res, - cuvs::neighbors::vpq_params const& params, - SrcT const& src) +[[nodiscard]] auto make_device_vpq_dataset(raft::resources const& res, + cuvs::neighbors::vpq_params const& params, + SrcT const& src) -> cuvs::neighbors::device_vpq_dataset { // A cuVS dataset keeps its logical width in `dim()` while `view()` spans the full row pitch. @@ -310,7 +311,7 @@ template auto const rows = src.view(); using value_type = typename decltype(rows)::value_type; using extents_type = raft::matrix_extent; - return make_vpq_dataset( + return make_device_vpq_dataset( res, params, raft::mdspan{ @@ -321,11 +322,11 @@ template using value_type = typename SrcT::value_type; static_assert(std::is_same_v || std::is_same_v || std::is_same_v || std::is_same_v, - "make_vpq_dataset: element type must be float, half, int8_t or uint8_t"); + "make_device_vpq_dataset: element type must be float, half, int8_t or uint8_t"); const int64_t n_rows = src.extent(0); const int64_t dim = src.extent(1); const int64_t stride = src.stride(0) > 0 ? src.stride(0) : dim; - RAFT_EXPECTS(n_rows > 0, "make_vpq_dataset: dataset is empty"); + RAFT_EXPECTS(n_rows > 0, "make_device_vpq_dataset: dataset is empty"); return detail::vpq_train_from_rows( res, params, src.data_handle(), raft::get_cuda_data_type(), n_rows, dim, stride); } diff --git a/cpp/src/neighbors/cagra.cuh b/cpp/src/neighbors/cagra.cuh index 138b899fc5..5a8a0b1887 100644 --- a/cpp/src/neighbors/cagra.cuh +++ b/cpp/src/neighbors/cagra.cuh @@ -286,7 +286,7 @@ void optimize( * stored in the returned index as a non-owning view — no copy is made. The caller must keep the * underlying storage alive for the lifetime of the index. * - * Host-backed indexes cannot be searched; call `attach_dataset` with a device-padded dataset to get + * Host-backed indexes cannot be searched; call `update_dataset` with a device-padded dataset to get * a search-ready device index. */ template @@ -300,7 +300,7 @@ auto build(raft::resources const& res, const index_params& params, DatasetViewT using IdxT = uint32_t; // Dense paths build the graph and optionally attach the input dataset view. Host indexes remain - // non-searchable until attach_dataset(...) supplies a device-padded dataset. + // non-searchable until update_dataset(...) supplies a device-padded dataset. if constexpr (cuvs::neighbors::is_device_vpq_dataset_view_v) { RAFT_FAIL("cagra::build: VPQ-compressed dataset cannot be used for dense graph construction."); } else if constexpr (cuvs::neighbors::is_dense_row_major_device_dataset_view_v) { diff --git a/cpp/src/neighbors/cagra_update_dataset_inst.cu.in b/cpp/src/neighbors/cagra_update_dataset_inst.cu.in new file mode 100644 index 0000000000..3e0fd19319 --- /dev/null +++ b/cpp/src/neighbors/cagra_update_dataset_inst.cu.in @@ -0,0 +1,102 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include +#include +#include + +namespace { + +using data_t = @data_type@; +using index_t = @index_type@; +using inst_device_padded_view_t = cuvs::neighbors::device_padded_dataset_view; +using inst_device_standard_view_t = cuvs::neighbors::device_standard_dataset_view; +using inst_device_vpq_view_t = cuvs::neighbors::device_vpq_dataset_view; + +} // namespace + +namespace cuvs::neighbors::cagra { + +extern template void index::compute_dataset_norms_( + raft::resources const&); +extern template void index::compute_dataset_norms_( + raft::resources const&); +extern template void index::compute_dataset_norms_( + raft::resources const&); + +#define CUVS_CAGRA_DEFINE_UPDATE_DATASET_OVERLOADS(T) \ + auto update_dataset(raft::resources const& res, \ + host_standard_index const& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + -> device_padded_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + auto update_dataset(raft::resources const& res, \ + host_padded_index const& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + -> device_padded_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + auto update_dataset(raft::resources const& res, \ + device_standard_index const& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + -> device_padded_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + auto update_dataset(raft::resources const& res, \ + host_standard_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + auto update_dataset(raft::resources const& res, \ + host_padded_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + auto update_dataset(raft::resources const& res, \ + device_standard_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + auto update_dataset(raft::resources const& res, \ + device_padded_index const& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + -> vpq_f16_index \ + { \ + return detail::attach_dataset(res, idx, dataset); \ + } \ + void update_dataset(raft::resources const& res, \ + device_standard_index& idx, \ + cuvs::neighbors::device_standard_dataset_view const& dataset) \ + { \ + idx.update_device_dataset_same_layout(res, dataset); \ + } \ + void update_dataset(raft::resources const& res, \ + device_padded_index& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + { \ + idx.update_device_dataset_same_layout(res, dataset); \ + } \ + void update_dataset(raft::resources const& res, \ + vpq_f16_index& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + { \ + idx.update_device_dataset_same_layout(res, dataset); \ + } + +CUVS_CAGRA_DEFINE_UPDATE_DATASET_OVERLOADS(data_t) + +#undef CUVS_CAGRA_DEFINE_UPDATE_DATASET_OVERLOADS + +} // namespace cuvs::neighbors::cagra diff --git a/cpp/src/neighbors/cagra_update_dataset_matrix.json b/cpp/src/neighbors/cagra_update_dataset_matrix.json new file mode 100644 index 0000000000..a7995005c4 --- /dev/null +++ b/cpp/src/neighbors/cagra_update_dataset_matrix.json @@ -0,0 +1,26 @@ +{ + "_data": [ + { + "data_type": "float", + "data_abbrev": "f" + }, + { + "data_type": "half", + "data_abbrev": "h" + }, + { + "data_type": "int8_t", + "data_abbrev": "i8" + }, + { + "data_type": "uint8_t", + "data_abbrev": "u8" + } + ], + "_index": [ + { + "index_type": "uint32_t", + "index_abbrev": "u32" + } + ] +} diff --git a/cpp/src/neighbors/detail/cagra/cagra_search.cuh b/cpp/src/neighbors/detail/cagra/cagra_search.cuh index 165e478337..2dd4438047 100644 --- a/cpp/src/neighbors/detail/cagra/cagra_search.cuh +++ b/cpp/src/neighbors/detail/cagra/cagra_search.cuh @@ -233,7 +233,7 @@ void search_main(raft::resources const& res, if constexpr (cuvs::neighbors::is_empty_dataset_view_v) { RAFT_FAIL( "Attempted to search without a dataset. Please call " - "index.update_device_dataset_same_layout(...) first."); + "cagra::update_dataset(...) first."); } else if constexpr (cuvs::neighbors::is_device_vpq_f32_dataset_view_v) { RAFT_FAIL("FP32 VPQ dataset support is coming soon"); } else if constexpr (cuvs::neighbors::is_device_vpq_f16_dataset_view_v) { @@ -260,13 +260,13 @@ void search_main(raft::resources const& res, RAFT_FAIL( "CAGRA search requires a padded device dataset. Build from a standard dataset view, then " "call " - "cagra::attach_dataset(res, index, padded_view) before search."); + "cagra::update_dataset(res, index, padded_view) before search."); } else if constexpr (cuvs::neighbors::is_device_padded_dataset_view_v) { run_strided_like(index.dataset()); } else if constexpr (cuvs::neighbors::is_host_dataset_view_v) { static_assert(sizeof(DatasetViewT) == 0, "search requires a device-resident dataset. " - "Call cagra::attach_dataset(res, index, padded_view) " + "Call cagra::update_dataset(res, index, padded_view) " "to convert/attach into a search-ready device padded index before searching."); } else { static_assert(sizeof(DatasetViewT) == 0, "search: unsupported dataset view type"); diff --git a/cpp/src/neighbors/detail/cagra/update_dataset.cuh b/cpp/src/neighbors/detail/cagra/update_dataset.cuh new file mode 100644 index 0000000000..0c79b98044 --- /dev/null +++ b/cpp/src/neighbors/detail/cagra/update_dataset.cuh @@ -0,0 +1,156 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include + +#include +#include +#include + +namespace cuvs::neighbors::cagra::detail { + +template + requires cuvs::neighbors::is_host_dataset_view_v +CUVS_HIDDEN auto convert_host_to_device_index(raft::resources const& res, + index const& src) + -> index> +{ + using device_view_type = cuvs::neighbors::device_counterpart_t; + using graph_index_type = typename index::graph_index_type; + + index out(res, src.metric()); + if (src.graph().size() > 0) { + auto graph_host = raft::make_host_matrix(src.graph().extent(0), + src.graph().extent(1)); + raft::copy(graph_host.data_handle(), + src.graph().data_handle(), + src.graph().size(), + raft::resource::get_cuda_stream(res)); + raft::resource::sync_stream(res); + out.update_graph(res, raft::make_const_mdspan(graph_host.view())); + } + return out; +} + +template +CUVS_HIDDEN auto convert_standard_to_padded_index( + raft::resources const& res, + index> const& standard_idx, + cuvs::neighbors::device_padded_dataset_view const& padded_dataset) + -> device_padded_index +{ + RAFT_EXPECTS(padded_dataset.n_rows() == standard_idx.size(), + "Padded dataset row count must match the index size"); + + device_padded_index out(res, standard_idx.metric()); + if (standard_idx.graph().extent(0) > 0) { + using graph_index_type = + typename index>:: + graph_index_type; + auto graph_host = raft::make_host_matrix( + standard_idx.graph().extent(0), standard_idx.graph().extent(1)); + raft::copy(graph_host.data_handle(), + standard_idx.graph().data_handle(), + standard_idx.graph().size(), + raft::resource::get_cuda_stream(res)); + raft::resource::sync_stream(res); + out.update_graph(res, raft::make_const_mdspan(graph_host.view())); + } + if (standard_idx.source_indices().has_value()) { + out.update_source_indices(res, standard_idx.source_indices().value()); + } + out.update_device_dataset_same_layout(res, padded_dataset); + return out; +} + +template + requires cuvs::neighbors::ann_dataset_view +CUVS_HIDDEN auto convert_dense_to_vpq_f16_index( + raft::resources const& res, + index const& src, + cuvs::neighbors::device_vpq_dataset_view const& vpq_dataset) + -> vpq_f16_index +{ + RAFT_EXPECTS(vpq_dataset.n_rows() == src.size(), + "VPQ dataset row count must match the index size"); + + vpq_f16_index out(res, src.metric()); + if (src.graph().extent(0) > 0) { + using graph_index_type = typename index::graph_index_type; + auto graph_host = raft::make_host_matrix(src.graph().extent(0), + src.graph().extent(1)); + raft::copy(graph_host.data_handle(), + src.graph().data_handle(), + src.graph().size(), + raft::resource::get_cuda_stream(res)); + raft::resource::sync_stream(res); + out.update_graph(res, raft::make_const_mdspan(graph_host.view())); + } + if (src.source_indices().has_value()) { + out.update_source_indices(res, src.source_indices().value()); + } + out.update_device_dataset_same_layout(res, vpq_dataset); + return out; +} + +template + requires cuvs::neighbors::ann_dataset_view +CUVS_HIDDEN auto attach_dataset( + raft::resources const& res, + index const& idx, + cuvs::neighbors::device_padded_dataset_view const& device_padded_dataset) + -> device_padded_index +{ + RAFT_EXPECTS(device_padded_dataset.n_rows() == idx.size(), + "Padded dataset row count must match the index size"); + + if constexpr (cuvs::neighbors::is_host_standard_dataset_view_v) { + auto dev_std = convert_host_to_device_index(res, idx); + return convert_standard_to_padded_index(res, dev_std, device_padded_dataset); + } else if constexpr (cuvs::neighbors::is_host_padded_dataset_view_v) { + auto dev_pad = convert_host_to_device_index(res, idx); + dev_pad.update_device_dataset_same_layout(res, device_padded_dataset); + return dev_pad; + } else if constexpr (cuvs::neighbors::is_device_standard_dataset_view_v) { + return convert_standard_to_padded_index(res, idx, device_padded_dataset); + } else if constexpr (cuvs::neighbors::is_device_padded_dataset_view_v) { + RAFT_LOG_WARN( + "cagra::attach_dataset called with an already device-padded index. " + "To avoid an unnecessary index copy, call " + "index.update_device_dataset_same_layout(res, device_padded_dataset) " + "directly on the original index."); + RAFT_FAIL( + "cagra::attach_dataset: device_padded_index input is not supported in this overload. " + "Call index.update_device_dataset_same_layout(res, device_padded_dataset) directly."); + } else { + static_assert(!sizeof(IndexViewT), "Unsupported CAGRA index dataset view type"); + } +} + +template + requires cuvs::neighbors::ann_dataset_view +CUVS_HIDDEN auto attach_dataset( + raft::resources const& res, + index const& idx, + cuvs::neighbors::device_vpq_dataset_view const& vpq_dataset) + -> vpq_f16_index +{ + if constexpr (cuvs::neighbors::is_device_vpq_f16_dataset_view_v) { + RAFT_LOG_WARN( + "cagra::attach_dataset called with an already vpq_f16 index. " + "To avoid an unnecessary index copy, call " + "index.update_device_dataset_same_layout(res, vpq_dataset) " + "directly on the original index."); + RAFT_FAIL( + "cagra::attach_dataset: vpq_f16_index input is not supported in this overload. " + "Call index.update_device_dataset_same_layout(res, vpq_dataset) directly."); + } else { + return convert_dense_to_vpq_f16_index(res, idx, vpq_dataset); + } +} + +} // namespace cuvs::neighbors::cagra::detail diff --git a/cpp/src/neighbors/iface/iface.hpp b/cpp/src/neighbors/iface/iface.hpp index 5d5ee76406..aa43cf7c4e 100644 --- a/cpp/src/neighbors/iface/iface.hpp +++ b/cpp/src/neighbors/iface/iface.hpp @@ -10,6 +10,7 @@ #include #include #include +#include #include #include #include @@ -109,7 +110,7 @@ void build(const raft::resources& handle, auto host_idx = cuvs::neighbors::cagra::build(handle, cagra_params, host_padded); auto padded_r = cuvs::neighbors::make_device_padded_dataset(handle, index_dataset); auto device_idx = - cuvs::neighbors::cagra::attach_dataset(handle, host_idx, padded_r->as_dataset_view()); + cuvs::neighbors::cagra::update_dataset(handle, host_idx, padded_r->as_dataset_view()); interface.cagra_owned_padded_dataset_ = std::move(padded_r); interface.cagra_owned_standard_dataset_.reset(); interface.index_.emplace(std::move(device_idx)); diff --git a/cpp/src/neighbors/mg/mg_cagra_inst.cu.in b/cpp/src/neighbors/mg/mg_cagra_inst.cu.in index 78d6b109ef..61959af4bc 100644 --- a/cpp/src/neighbors/mg/mg_cagra_inst.cu.in +++ b/cpp/src/neighbors/mg/mg_cagra_inst.cu.in @@ -92,7 +92,7 @@ void distribute_padded_dataset( res, idx, padded_dataset, [&](const raft::resources& dev_res, int rank, auto dataset) { \ const auto& in_if = idx.ann_interfaces_[rank]; \ auto& out_if = out.ann_interfaces_[rank]; \ - auto padded_idx = cuvs::neighbors::cagra::attach_dataset( \ + auto padded_idx = cuvs::neighbors::cagra::update_dataset( \ dev_res, in_if.index_.value(), dataset->as_dataset_view()); \ out_if.cagra_owned_padded_dataset_ = std::move(dataset); \ out_if.cagra_owned_standard_dataset_.reset(); \ diff --git a/cpp/src/neighbors/tiered_index.cu b/cpp/src/neighbors/tiered_index.cu index f5942001d7..764df0cc69 100644 --- a/cpp/src/neighbors/tiered_index.cu +++ b/cpp/src/neighbors/tiered_index.cu @@ -98,8 +98,8 @@ auto convert_standard_to_padded_index( padded_mds.data_handle(), ann_rows, static_cast(padded_mds.extent(1))); auto ann_padded_view = cuvs::neighbors::device_padded_dataset_view(ann_mds, padded_dataset.dim()); - auto ann_padded_idx = cuvs::neighbors::cagra::convert_standard_to_padded_index( - res, *idx.state->ann_index, ann_padded_view); + auto ann_padded_idx = + cuvs::neighbors::cagra::update_dataset(res, *idx.state->ann_index, ann_padded_view); next_state->ann_index = std::make_shared>( std::move(ann_padded_idx)); diff --git a/cpp/src/preprocessing/quantize/pq.cu b/cpp/src/preprocessing/quantize/pq.cu index 20b8f21d36..aebe6aa3ba 100644 --- a/cpp/src/preprocessing/quantize/pq.cu +++ b/cpp/src/preprocessing/quantize/pq.cu @@ -5,6 +5,8 @@ #include "./detail/pq.cuh" +#include +#include #include #include @@ -92,7 +94,7 @@ auto train_from_rows(raft::resources const& res, if (device_ptr == nullptr) { // A host mdspan makes training subsample the rows and encoding stream them in bounded batches, // so the dense dataset is never staged on the device. - RAFT_EXPECTS(stride == dim, "make_vpq_dataset: host input must be tightly packed"); + RAFT_EXPECTS(stride == dim, "make_device_vpq_dataset: host input must be tightly packed"); auto row_view = raft::make_host_matrix_view(src_ptr, n_rows, dim); return detail::vpq_build_half(res, params, row_view); } @@ -132,7 +134,8 @@ auto vpq_train_from_rows(raft::resources const& res, return train_from_rows( res, params, static_cast(src_ptr), n_rows, dim, stride); default: - RAFT_FAIL("make_vpq_dataset: unsupported dataset element type %d", static_cast(dtype)); + RAFT_FAIL("make_device_vpq_dataset: unsupported dataset element type %d", + static_cast(dtype)); } } diff --git a/cpp/tests/neighbors/ann_cagra.cuh b/cpp/tests/neighbors/ann_cagra.cuh index 529b8fe038..89bcd8874c 100644 --- a/cpp/tests/neighbors/ann_cagra.cuh +++ b/cpp/tests/neighbors/ann_cagra.cuh @@ -72,11 +72,11 @@ void cagra_build_into_index( *ace_host_dataset, static_cast(ace_host_dataset->extent(1))); auto host_idx = cagra::build(res, params, host_view); // In-memory ACE returns graph-only; attach device padded storage for search. - index = cagra::attach_dataset(res, host_idx, padded); + index = cagra::update_dataset(res, host_idx, padded); return; } index = cagra::build(res, params, padded); - index.update_device_dataset_same_layout(res, padded); + cagra::update_dataset(res, index, padded); } struct test_cagra_sample_filter { diff --git a/cpp/tests/neighbors/ann_cagra/test_float_uint32_t.cu b/cpp/tests/neighbors/ann_cagra/test_float_uint32_t.cu index de06bef34a..b813d98db3 100644 --- a/cpp/tests/neighbors/ann_cagra/test_float_uint32_t.cu +++ b/cpp/tests/neighbors/ann_cagra/test_float_uint32_t.cu @@ -152,4 +152,58 @@ TEST(AnnCagraMultiPartition, MixedGraphDegreeRejected) cagra::search_params{}); } +// CAGRA-Q smoke test: build the graph on dense rows, train VPQ storage from the same padded rows, +// update the index to VPQ storage, then search. +TEST(AnnCagraVpq, BuildUpdateVpqSearch) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + + constexpr int n_rows = 256, dim = 32, n_queries = 4, k = 1; + + auto dataset = raft::make_device_matrix(handle, n_rows, dim); + raft::random::RngState r(1234ULL); + InitDataset( + handle, dataset.data_handle(), n_rows, dim, cuvs::distance::DistanceType::L2Expanded, r); + raft::resource::sync_stream(handle); + + cuvs::neighbors::test::padded_device_matrix_for_cagra padded( + handle, raft::make_const_mdspan(dataset.view())); + + cagra::index_params index_params; + index_params.metric = cuvs::distance::DistanceType::L2Expanded; + auto dense_index = cagra::build(handle, index_params, padded.view); + + cuvs::neighbors::vpq_params vpq_params{.pq_bits = 8, .pq_dim = 8}; + auto vpq = + cuvs::preprocessing::quantize::pq::make_device_vpq_dataset(handle, vpq_params, padded.view); + raft::resource::sync_stream(handle); + + EXPECT_EQ(vpq.n_rows(), n_rows); + EXPECT_EQ(vpq.dim(), dim); + + auto vpq_index = cagra::update_dataset(handle, dense_index, vpq.as_dataset_view()); + + auto queries = raft::make_device_matrix(handle, n_queries, dim); + raft::copy(queries.data_handle(), dataset.data_handle(), queries.size(), stream); + + auto neighbors = raft::make_device_matrix(handle, n_queries, k); + auto distances = raft::make_device_matrix(handle, n_queries, k); + cagra::search(handle, + cagra::search_params{}, + vpq_index, + raft::make_const_mdspan(queries.view()), + neighbors.view(), + distances.view()); + + auto neighbors_h = raft::make_host_matrix(n_queries, k); + raft::copy(neighbors_h.data_handle(), neighbors.data_handle(), neighbors.size(), stream); + raft::resource::sync_stream(handle); + + // Queries are exact dataset rows, so the top hit must be the row itself. + for (int i = 0; i < n_queries; i++) { + EXPECT_EQ(neighbors_h(i, 0), static_cast(i)); + } +} + } // namespace cuvs::neighbors::cagra diff --git a/cpp/tests/preprocessing/product_quantization.cu b/cpp/tests/preprocessing/product_quantization.cu index a392e7e1db..d9db154d26 100644 --- a/cpp/tests/preprocessing/product_quantization.cu +++ b/cpp/tests/preprocessing/product_quantization.cu @@ -324,7 +324,7 @@ TEST(ProductQuantizationTestF, MakeVpqDatasetFromHost) cuvs::neighbors::vpq_params params{ .pq_bits = 4, .pq_dim = 4, .vq_n_centers = 1, .kmeans_n_iters = 2}; - auto vpq = make_vpq_dataset(handle, params, raft::make_const_mdspan(dataset.view())); + auto vpq = make_device_vpq_dataset(handle, params, raft::make_const_mdspan(dataset.view())); raft::resource::sync_stream(handle); EXPECT_EQ(vpq.n_rows(), n_rows); @@ -357,7 +357,7 @@ TEST(ProductQuantizationTestF, MakeVpqDatasetFromPaddedView) cuvs::neighbors::vpq_params params{ .pq_bits = 4, .pq_dim = 4, .vq_n_centers = 1, .kmeans_n_iters = 2}; - auto vpq = make_vpq_dataset(handle, params, padded); + auto vpq = make_device_vpq_dataset(handle, params, padded); raft::resource::sync_stream(handle); EXPECT_EQ(vpq.n_rows(), n_rows); diff --git a/examples/cpp/src/cagra_hnsw_ace_example.cu b/examples/cpp/src/cagra_hnsw_ace_example.cu index 05d448fbc1..326fff955b 100644 --- a/examples/cpp/src/cagra_hnsw_ace_example.cu +++ b/examples/cpp/src/cagra_hnsw_ace_example.cu @@ -100,7 +100,7 @@ void cagra_build_search_ace(raft::device_resources const& dev_resources, // attach it before from_cagra builds the HNSW hierarchy in memory. padded_owner = cuvs::neighbors::make_device_padded_dataset(dev_resources, dataset_host_view); auto device_index = - cagra::attach_dataset(dev_resources, ace_host_index, padded_owner->as_dataset_view()); + cagra::update_dataset(dev_resources, ace_host_index, padded_owner->as_dataset_view()); hnsw_index = hnsw::from_cagra(dev_resources, hnsw_params, device_index, dataset_host_view); } diff --git a/go/cagra/cagra.go b/go/cagra/cagra.go index bb7103b111..fdd2f4dec7 100644 --- a/go/cagra/cagra.go +++ b/go/cagra/cagra.go @@ -21,11 +21,22 @@ type PaddedDataset struct { dataset C.cuvsDataset_t } +// Owning VPQ dataset handle for CAGRA-Q search. +type VpqDataset struct { + dataset C.cuvsDataset_t +} + // PaddedDatasetHandle is an owning padded dataset or non-owning padded dataset view. type PaddedDatasetHandle interface { datasetHandle() C.cuvsDataset_t } +// DatasetHandle is any CAGRA dataset handle accepted by UpdateDataset +// (device-padded or device VPQ). +type DatasetHandle interface { + datasetHandle() C.cuvsDataset_t +} + // Non-owning padded dataset view handle. type PaddedDatasetView struct { view C.cuvsDataset_t @@ -184,18 +195,18 @@ func (view *StandardDatasetView) Close() error { return nil } -// UpdateDataset updates any CAGRA index layout with a caller-provided padded -// dataset or view and leaves the same handle search-ready. -func UpdateDataset(Resources cuvs.Resource, paddedDataset PaddedDatasetHandle, index *CagraIndex) error { +// UpdateDataset updates any CAGRA index layout with a caller-provided device +// padded or VPQ dataset/view and leaves the same handle search-ready. +func UpdateDataset(Resources cuvs.Resource, dataset DatasetHandle, index *CagraIndex) error { if !index.trained { return errors.New("index needs to be built before attaching dataset") } - if paddedDataset == nil || paddedDataset.datasetHandle() == nil { - return errors.New("padded dataset is nil") + if dataset == nil || dataset.datasetHandle() == nil { + return errors.New("dataset is nil") } err := cuvs.CheckCuvs(cuvs.CuvsError(C.cuvsCagraUpdateDataset( C.cuvsResources_t(Resources.Resource), - paddedDataset.datasetHandle(), + dataset.datasetHandle(), index.index, ))) if err != nil { @@ -204,6 +215,49 @@ func UpdateDataset(Resources cuvs.Resource, paddedDataset PaddedDatasetHandle, i return nil } +// MakeVpqDataset trains an owning device VPQ dataset (CAGRA-Q) from a device-padded source. +// params may be nil to use library defaults. Keep the returned dataset alive while any index uses it. +func MakeVpqDataset(Resources cuvs.Resource, source PaddedDatasetHandle, params *CompressionParams) (*VpqDataset, error) { + if source == nil || source.datasetHandle() == nil { + return nil, errors.New("source padded dataset is nil") + } + var cParams C.cuvsCagraCompressionParams_t + if params != nil { + cParams = params.params + } + var vpqDataset C.cuvsDataset_t + err := cuvs.CheckCuvs(cuvs.CuvsError(C.cuvsDatasetMakeVpq( + C.cuvsResources_t(Resources.Resource), + source.datasetHandle(), + cParams, + &vpqDataset, + ))) + if err != nil { + return nil, err + } + return &VpqDataset{dataset: vpqDataset}, nil +} + +func (dataset *VpqDataset) datasetHandle() C.cuvsDataset_t { + if dataset == nil { + return nil + } + return dataset.dataset +} + +// Close destroys an owning VPQ dataset handle. +func (dataset *VpqDataset) Close() error { + if dataset == nil || dataset.dataset == nil { + return nil + } + err := cuvs.CheckCuvs(cuvs.CuvsError(C.cuvsDatasetDestroy(dataset.dataset))) + if err != nil { + return err + } + dataset.dataset = nil + return nil +} + // Creates a new empty Cagra Index func CreateIndex() (*CagraIndex, error) { var index C.cuvsCagraIndex_t diff --git a/go/cagra/cagra_test.go b/go/cagra/cagra_test.go index 2f087196a7..9661d001a9 100644 --- a/go/cagra/cagra_test.go +++ b/go/cagra/cagra_test.go @@ -128,6 +128,134 @@ func TestCagra(t *testing.T) { } } +func TestCagraVpqBuildUpdateSearch(t *testing.T) { + // CAGRA-Q smoke: dense build → MakeVpqDataset → UpdateDataset → Search. + const ( + nDataPoints = 256 + nFeatures = 32 + nQueries = 4 + k = 1 + ) + r := rand.New(rand.NewPCG(42, 0)) + + resource, err := cuvs.NewResource(nil) + if err != nil { + t.Fatalf("error creating resource: %v", err) + } + defer resource.Close() + + testDataset := make([][]float32, nDataPoints) + for i := range testDataset { + testDataset[i] = make([]float32, nFeatures) + for j := range testDataset[i] { + testDataset[i][j] = r.Float32() + } + } + + dataset, err := cuvs.NewTensor(testDataset) + if err != nil { + t.Fatalf("error creating dataset tensor: %v", err) + } + defer dataset.Close() + + if _, err := dataset.ToDevice(&resource); err != nil { + t.Fatalf("error moving dataset to device: %v", err) + } + + indexParams, err := CreateIndexParams() + if err != nil { + t.Fatalf("error creating index params: %v", err) + } + defer indexParams.Close() + + index, err := CreateIndex() + if err != nil { + t.Fatalf("error creating index: %v", err) + } + defer index.Close() + + if err := BuildIndex(resource, indexParams, &dataset, index); err != nil { + t.Fatalf("error building index: %v", err) + } + + // dim=32 float already matches CAGRA padded row width; wrap with a view. + padded, err := MakePaddedDatasetView(resource, &dataset) + if err != nil { + t.Fatalf("error creating padded dataset view: %v", err) + } + defer padded.Close() + + compression, err := CreateCompressionParams() + if err != nil { + t.Fatalf("error creating compression params: %v", err) + } + defer compression.Close() + if _, err := compression.SetPQBits(8); err != nil { + t.Fatalf("error setting pq_bits: %v", err) + } + if _, err := compression.SetPQDim(8); err != nil { + t.Fatalf("error setting pq_dim: %v", err) + } + + vpq, err := MakeVpqDataset(resource, padded, compression) + if err != nil { + t.Fatalf("error creating VPQ dataset: %v", err) + } + defer vpq.Close() + + if err := UpdateDataset(resource, vpq, index); err != nil { + t.Fatalf("error updating index with VPQ dataset: %v", err) + } + + queries, err := cuvs.NewTensor(testDataset[:nQueries]) + if err != nil { + t.Fatalf("error creating queries tensor: %v", err) + } + defer queries.Close() + if _, err := queries.ToDevice(&resource); err != nil { + t.Fatalf("error moving queries to device: %v", err) + } + + neighbors, err := cuvs.NewTensorOnDevice[uint32](&resource, []int64{int64(nQueries), int64(k)}) + if err != nil { + t.Fatalf("error creating neighbors tensor: %v", err) + } + defer neighbors.Close() + + distances, err := cuvs.NewTensorOnDevice[float32](&resource, []int64{int64(nQueries), int64(k)}) + if err != nil { + t.Fatalf("error creating distances tensor: %v", err) + } + defer distances.Close() + + searchParams, err := CreateSearchParams() + if err != nil { + t.Fatalf("error creating search params: %v", err) + } + defer searchParams.Close() + + if err := SearchIndex(resource, searchParams, index, &queries, &neighbors, &distances, nil); err != nil { + t.Fatalf("error searching VPQ index: %v", err) + } + + if _, err := neighbors.ToHost(&resource); err != nil { + t.Fatalf("error moving neighbors to host: %v", err) + } + if err := resource.Sync(); err != nil { + t.Fatalf("error syncing resource: %v", err) + } + + neighborsSlice, err := neighbors.Slice() + if err != nil { + t.Fatalf("error getting neighbors slice: %v", err) + } + for i := range neighborsSlice { + if neighborsSlice[i][0] != uint32(i) { + t.Errorf("wrong neighbor for query %d: expected %d, got %d", i, i, neighborsSlice[i][0]) + } + } +} + func TestCagraFiltering(t *testing.T) { const ( nDataPoints = 1024 diff --git a/go/cagra/index_params.go b/go/cagra/index_params.go index bf2268df8b..3861c22539 100644 --- a/go/cagra/index_params.go +++ b/go/cagra/index_params.go @@ -13,6 +13,11 @@ type IndexParams struct { params C.cuvsCagraIndexParams_t } +// CompressionParams holds VPQ training parameters for CAGRA-Q. +type CompressionParams struct { + params C.cuvsCagraCompressionParams_t +} + type BuildAlgo int const ( @@ -27,6 +32,71 @@ var cBuildAlgos = map[BuildAlgo]int{ AutoSelect: C.AUTO_SELECT, } +// CreateCompressionParams creates VPQ compression params with library defaults. +func CreateCompressionParams() (*CompressionParams, error) { + var params C.cuvsCagraCompressionParams_t + + err := cuvs.CheckCuvs(cuvs.CuvsError(C.cuvsCagraCompressionParamsCreate(¶ms))) + if err != nil { + return nil, err + } + + if params == nil { + return nil, errors.New("memory allocation failed") + } + + return &CompressionParams{params: params}, nil +} + +// SetPQBits sets the bit length of the vector element after PQ compression. +func (p *CompressionParams) SetPQBits(pq_bits uint32) (*CompressionParams, error) { + p.params.pq_bits = C.uint32_t(pq_bits) + return p, nil +} + +// SetPQDim sets the dimensionality after PQ compression (0 = heuristic). +func (p *CompressionParams) SetPQDim(pq_dim uint32) (*CompressionParams, error) { + p.params.pq_dim = C.uint32_t(pq_dim) + return p, nil +} + +// SetVQNCenters sets the VQ codebook size (0 = heuristic). +func (p *CompressionParams) SetVQNCenters(vq_n_centers uint32) (*CompressionParams, error) { + p.params.vq_n_centers = C.uint32_t(vq_n_centers) + return p, nil +} + +// SetKMeansNIters sets kmeans iterations for VQ and PQ phases. +func (p *CompressionParams) SetKMeansNIters(kmeans_n_iters uint32) (*CompressionParams, error) { + p.params.kmeans_n_iters = C.uint32_t(kmeans_n_iters) + return p, nil +} + +// SetVQKMeansTrainsetFraction sets the VQ kmeans trainset fraction (0 = heuristic). +func (p *CompressionParams) SetVQKMeansTrainsetFraction(vq_kmeans_trainset_fraction float64) (*CompressionParams, error) { + p.params.vq_kmeans_trainset_fraction = C.double(vq_kmeans_trainset_fraction) + return p, nil +} + +// SetPQKMeansTrainsetFraction sets the PQ kmeans trainset fraction (0 = heuristic). +func (p *CompressionParams) SetPQKMeansTrainsetFraction(pq_kmeans_trainset_fraction float64) (*CompressionParams, error) { + p.params.pq_kmeans_trainset_fraction = C.double(pq_kmeans_trainset_fraction) + return p, nil +} + +// Close destroys CompressionParams. +func (p *CompressionParams) Close() error { + if p == nil || p.params == nil { + return nil + } + err := cuvs.CheckCuvs(cuvs.CuvsError(C.cuvsCagraCompressionParamsDestroy(p.params))) + if err != nil { + return err + } + p.params = nil + return nil +} + // Creates a new IndexParams func CreateIndexParams() (*IndexParams, error) { var params C.cuvsCagraIndexParams_t diff --git a/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraCompressionParams.java b/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraCompressionParams.java index a10d6f6725..28c356f8b9 100644 --- a/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraCompressionParams.java +++ b/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraCompressionParams.java @@ -1,11 +1,12 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ package com.nvidia.cuvs; /** - * Supplemental compression parameters to build CAGRA Index. + * Supplemental compression parameters for CAGRA-Q VPQ training via + * {@link CagraIndex#makeVpqDataset(CagraIndex.PaddedDataset, CagraCompressionParams)}. * * @since 25.02 */ diff --git a/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java b/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java index 51403982dd..bd05e79050 100644 --- a/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java +++ b/java/cuvs-java/src/main/java/com/nvidia/cuvs/CagraIndex.java @@ -23,8 +23,11 @@ * @since 25.02 */ public interface CagraIndex extends AutoCloseable { - /** Caller-owned non-owning dataset view handle. */ - abstract class DatasetView implements AutoCloseable { + /** + * Base class for the native dataset handles CAGRA hands back, whether they own their storage + * or merely view storage the caller owns. + */ + abstract class DatasetHandle implements AutoCloseable { private AutoCloseable delegate; private long handleAddress; @@ -37,7 +40,7 @@ public final void setDelegate(AutoCloseable delegate, long handleAddress) { } /** - * Returns true when this view has a native handle. + * Returns true when this handle refers to native dataset storage. */ public final boolean isPresent() { return delegate != null && handleAddress != 0; @@ -60,76 +63,71 @@ public void close() throws Exception { } } - /** Caller-owned padded dataset view. */ - final class PaddedDatasetView extends DatasetView { - public PaddedDatasetView() {} - } - - /** Caller-owned standard dataset view. */ - final class StandardDatasetView extends DatasetView { - public StandardDatasetView() {} - } + /** Non-owning view of dataset storage the caller owns and must keep alive. */ + abstract class DatasetView extends DatasetHandle {} /** - * Caller-owned dataset handle populated by explicit deserialization or created by - * {@link #makePaddedDataset(CuVSMatrix)}. + * Dataset handle that owns its native storage and releases it on {@link #close()}. */ - abstract class DeserializeDataset implements AutoCloseable { - private AutoCloseable delegate; - private long handleAddress; - + abstract class OwningDataset extends DatasetHandle { /** * Internal wiring hook used by the Java wrapper implementation. */ public final void setDelegate(AutoCloseable delegate) { setDelegate(delegate, 0); } + } - /** - * Internal wiring hook used by the Java wrapper implementation. - */ - public final void setDelegate(AutoCloseable delegate, long handleAddress) { - this.delegate = delegate; - this.handleAddress = handleAddress; - } + /** + * Owning storage holding vectors in a dense row-major layout, either padded to CAGRA's + * required row stride or standard. This is the storage a serialized index carries, so + * deserialization populates one of these. + */ + abstract class DenseOwningDataset extends OwningDataset {} - /** - * Returns true when this handle owns native dataset storage. - */ - public final boolean isPresent() { - return delegate != null && handleAddress != 0; - } + /** + * A device-padded dataset, either owned ({@link PaddedDataset}) or viewed + * ({@link PaddedDatasetView}). Operations that only read padded storage accept either form. + */ + interface PaddedDatasetHandle { + /** Returns true when this handle refers to native dataset storage. */ + boolean isPresent(); - /** - * Internal accessor for native handle address. - */ - public final long nativeHandleAddress() { - return handleAddress; - } + /** Internal accessor for native handle address. */ + long nativeHandleAddress(); + } - @Override - public void close() throws Exception { - if (delegate != null) { - delegate.close(); - delegate = null; - } - handleAddress = 0; - } + /** Caller-owned padded dataset view. */ + final class PaddedDatasetView extends DatasetView implements PaddedDatasetHandle { + public PaddedDatasetView() {} + } + + /** Caller-owned standard dataset view. */ + final class StandardDatasetView extends DatasetView { + public StandardDatasetView() {} } /** * Owning padded dataset handle. Keep this alive for as long as any index using it remains in * use. */ - final class PaddedDataset extends DeserializeDataset { + final class PaddedDataset extends DenseOwningDataset implements PaddedDatasetHandle { public PaddedDataset() {} } /** Owning standard dataset handle populated by deserialization. */ - final class StandardDataset extends DeserializeDataset { + final class StandardDataset extends DenseOwningDataset { public StandardDataset() {} } + /** + * Owning VPQ dataset handle for CAGRA-Q. Keep this alive for as long as any index using it + * remains in use. + */ + final class VpqDataset extends OwningDataset { + public VpqDataset() {} + } + /** * Invokes the native destroy_cagra_index to de-allocate the CAGRA index */ @@ -164,17 +162,26 @@ public StandardDataset() {} StandardDatasetView makeStandardDatasetView(CuVSMatrix dataset) throws Throwable; /** - * Update this index with a caller-provided padded device dataset view and leave it - * search-ready in padded-device layout. The caller retains ownership of the underlying - * padded storage and must keep it alive while this index uses it. + * Update this index with a padded device dataset and leave it search-ready in padded-device + * layout. The caller retains ownership of the underlying padded storage and must keep it alive + * while this index uses it. */ - void updateDataset(PaddedDatasetView datasetView) throws Throwable; + void updateDataset(PaddedDatasetHandle dataset) throws Throwable; /** - * Update this index with a caller-owned padded device dataset. The dataset must remain alive - * while this index uses it. + * Update this index with a caller-owned device VPQ dataset (CAGRA-Q). Keep {@code vpqDataset} + * alive while this index uses it. Metric must remain L2Expanded. + */ + void updateDataset(VpqDataset vpqDataset) throws Throwable; + + /** + * Train an owning device VPQ dataset (CAGRA-Q) from a device-padded source. + * + * @param paddedDataset device-padded source dataset, owned or viewed + * @param compressionParams VPQ training parameters; may be {@code null} for defaults */ - void updateDataset(PaddedDataset dataset) throws Throwable; + VpqDataset makeVpqDataset( + PaddedDatasetHandle paddedDataset, CagraCompressionParams compressionParams) throws Throwable; /** Returns the CAGRA graph * @@ -359,7 +366,7 @@ interface Builder { * @param outDataset an empty {@link PaddedDataset} or {@link StandardDataset} * @return an instance of this Builder */ - Builder from(InputStream inputStream, DeserializeDataset outDataset); + Builder from(InputStream inputStream, DenseOwningDataset outDataset); /** * Sets a CAGRA graph instance to re-create an index from a diff --git a/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java b/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java index 691c99e93f..b7dbb78d5a 100644 --- a/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java +++ b/java/cuvs-java/src/main/java22/com/nvidia/cuvs/internal/CagraIndexImpl.java @@ -81,7 +81,7 @@ private CagraIndexImpl(InputStream inputStream, CuVSResources resources) throws } private CagraIndexImpl( - InputStream inputStream, CuVSResources resources, CagraIndex.DeserializeDataset outDataset) + InputStream inputStream, CuVSResources resources, CagraIndex.DenseOwningDataset outDataset) throws Throwable { this.resources = resources; this.cagraIndexReference = deserialize(inputStream, outDataset); @@ -485,23 +485,23 @@ public CagraIndex.StandardDatasetView makeStandardDatasetView(CuVSMatrix dataset } @Override - public void updateDataset(CagraIndex.PaddedDatasetView datasetView) throws Throwable { + public void updateDataset(CagraIndex.PaddedDatasetHandle dataset) throws Throwable { checkNotDestroyed(); - Objects.requireNonNull(datasetView); - if (!datasetView.isPresent()) { - throw new IllegalArgumentException("datasetView is uninitialized"); + Objects.requireNonNull(dataset); + if (!dataset.isPresent()) { + throw new IllegalArgumentException("dataset is uninitialized"); } - updateDataset(datasetView.nativeHandleAddress()); + updateDataset(dataset.nativeHandleAddress()); } @Override - public void updateDataset(CagraIndex.PaddedDataset dataset) throws Throwable { + public void updateDataset(CagraIndex.VpqDataset vpqDataset) throws Throwable { checkNotDestroyed(); - Objects.requireNonNull(dataset); - if (!dataset.isPresent()) { - throw new IllegalArgumentException("dataset is uninitialized"); + Objects.requireNonNull(vpqDataset); + if (!vpqDataset.isPresent()) { + throw new IllegalArgumentException("vpqDataset is uninitialized"); } - updateDataset(dataset.nativeHandleAddress()); + updateDataset(vpqDataset.nativeHandleAddress()); } private void updateDataset(long datasetHandleAddress) { @@ -516,6 +516,55 @@ private void updateDataset(long datasetHandleAddress) { } } + @Override + public CagraIndex.VpqDataset makeVpqDataset( + CagraIndex.PaddedDatasetHandle paddedDataset, CagraCompressionParams compressionParams) + throws Throwable { + checkNotDestroyed(); + Objects.requireNonNull(paddedDataset); + if (!paddedDataset.isPresent()) { + throw new IllegalArgumentException("paddedDataset is uninitialized"); + } + + try (var localArena = Arena.ofConfined(); + var resourcesAccessor = resources.access()) { + var cuvsRes = resourcesAccessor.handle(); + MemorySegment paramsSeg = MemorySegment.NULL; + CloseableHandle compressionHandle = null; + try { + if (compressionParams != null) { + compressionHandle = createCagraCompressionParams(); + paramsSeg = compressionHandle.handle(); + cuvsCagraCompressionParams.pq_bits(paramsSeg, compressionParams.getPqBits()); + cuvsCagraCompressionParams.pq_dim(paramsSeg, compressionParams.getPqDim()); + cuvsCagraCompressionParams.vq_n_centers(paramsSeg, compressionParams.getVqNCenters()); + cuvsCagraCompressionParams.kmeans_n_iters(paramsSeg, compressionParams.getKmeansNIters()); + cuvsCagraCompressionParams.vq_kmeans_trainset_fraction( + paramsSeg, compressionParams.getVqKmeansTrainsetFraction()); + cuvsCagraCompressionParams.pq_kmeans_trainset_fraction( + paramsSeg, compressionParams.getPqKmeansTrainsetFraction()); + } + MemorySegment vpqDatasetPtr = localArena.allocate(cuvsDataset_t); + var returnValue = + cuvsDatasetMakeVpq( + cuvsRes, + MemorySegment.ofAddress(paddedDataset.nativeHandleAddress()), + paramsSeg, + vpqDatasetPtr); + checkCuVSError(returnValue, "cuvsDatasetMakeVpq"); + MemorySegment vpqDataset = vpqDatasetPtr.get(cuvsDataset_t, 0); + + var out = new CagraIndex.VpqDataset(); + out.setDelegate(new DatasetCloseDelegate(vpqDataset), vpqDataset.address()); + return out; + } finally { + if (compressionHandle != null) { + compressionHandle.close(); + } + } + } + } + @Override public void serialize(OutputStream outputStream) throws Throwable { Path path = @@ -661,16 +710,10 @@ public void serializeToHNSW(OutputStream outputStream, Path tempFile, int buffer * @return an instance of {@link IndexReference} */ private IndexReference deserialize( - InputStream inputStream, CagraIndex.DeserializeDataset outDataset) throws Throwable { + InputStream inputStream, CagraIndex.DenseOwningDataset outDataset) throws Throwable { if (outDataset != null && outDataset.isPresent()) { throw new IllegalArgumentException("outDataset must be empty before deserialization"); } - if (outDataset != null - && !(outDataset instanceof CagraIndex.PaddedDataset) - && !(outDataset instanceof CagraIndex.StandardDataset)) { - throw new IllegalArgumentException( - "outDataset must be CagraIndex.PaddedDataset or CagraIndex.StandardDataset"); - } Path tmpIndexFile = Files.createTempFile(resources.tempDirectory(), UUID.randomUUID().toString(), ".cag") @@ -983,7 +1026,7 @@ public static class Builder implements CagraIndex.Builder { private CuVSMatrix dataset; private InputStream inputStream; - private CagraIndex.DeserializeDataset outDataset; + private CagraIndex.DenseOwningDataset outDataset; private CagraIndexParams cagraIndexParams; private final CuVSResources cuvsResources; private CuVSMatrix graph; @@ -1000,7 +1043,7 @@ public Builder from(InputStream inputStream) { } @Override - public Builder from(InputStream inputStream, CagraIndex.DeserializeDataset outDataset) { + public Builder from(InputStream inputStream, CagraIndex.DenseOwningDataset outDataset) { this.inputStream = inputStream; this.outDataset = Objects.requireNonNull(outDataset); return this; diff --git a/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java b/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java index e2287c0a22..fd539b9925 100644 --- a/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java +++ b/java/cuvs-java/src/test/java/com/nvidia/cuvs/CagraBuildAndSearchIT.java @@ -163,6 +163,77 @@ public void testIndexingAndSearchingFlow() throws Throwable { } } + /** + * CAGRA-Q smoke: dense build → makeVpqDataset → updateDataset → search. + * + * VPQ search requires a CAGRA-aligned dim, so dim=32 is used here and wrapped with + * {@link CagraIndex#makePaddedDatasetView}; an unaligned dim would leave the attached index + * reporting the padded row stride as its dimensionality. + */ + @Test + public void testVpqBuildUpdateSearch() throws Throwable { + final int nRows = 256; + final int nCols = 32; + final int nQueries = 4; + final int topK = 1; + + float[][] dataset = generateData(random, nRows, nCols); + float[][] queries = Arrays.copyOf(dataset, nQueries); + + CagraIndexParams indexParams = + new CagraIndexParams.Builder() + .withCagraGraphBuildAlgo(CagraGraphBuildAlgo.NN_DESCENT) + .withGraphDegree(32) + .withIntermediateGraphDegree(64) + .withMetric(CuvsDistanceType.L2Expanded) + .build(); + + CagraCompressionParams compressionParams = + new CagraCompressionParams.Builder().withPqBits(8).withPqDim(8).build(); + + CagraSearchParams searchParams = + new CagraSearchParams.Builder().withAlgo(CagraSearchParams.SearchAlgo.SINGLE_CTA).build(); + + try (CuVSResources resources = CheckedCuVSResources.create(); + var hostVectors = CuVSMatrix.ofArray(dataset); + var deviceVectors = hostVectors.toDevice(resources); + var index = + CagraIndex.newBuilder(resources) + .withDataset(hostVectors) + .withIndexParams(indexParams) + .build(); + var padded = index.makePaddedDatasetView(deviceVectors); + var vpq = index.makeVpqDataset(padded, compressionParams); + var queryVectors = CuVSMatrix.ofArray(queries)) { + assertTrue(padded.isPresent()); + assertTrue(vpq.isPresent()); + index.updateDataset(vpq); + + CagraQuery query = + new CagraQuery.Builder(resources) + .withTopK(topK) + .withSearchParams(searchParams) + .withQueryVectors(queryVectors) + .withMapping(SearchResults.IDENTITY_MAPPING) + .build(); + + SearchResults results = index.search(query); + List> rows = results.getResults(); + assertEquals(nQueries, rows.size()); + for (int i = 0; i < nQueries; i++) { + Integer topNeighbor = + rows.get(i).entrySet().stream() + .min(Map.Entry.comparingByValue()) + .map(Map.Entry::getKey) + .orElseThrow(); + assertEquals( + "query " + i + " should find itself as top-1 neighbor", + Integer.valueOf(i), + topNeighbor); + } + } + } + @Test public void testDeserializeReturnsCallerOwnedStandardDataset() throws Throwable { float[][] dataset = createSampleData(); diff --git a/python/cuvs/cuvs/common/dataset.pxd b/python/cuvs/cuvs/common/dataset.pxd index ac2d76ec18..96d0afa94f 100644 --- a/python/cuvs/cuvs/common/dataset.pxd +++ b/python/cuvs/cuvs/common/dataset.pxd @@ -14,6 +14,7 @@ cdef extern from "cuvs/core/dataset.h" nogil: ctypedef enum cuvsDatasetLayout_t: CUVS_DATASET_LAYOUT_STANDARD CUVS_DATASET_LAYOUT_PADDED + CUVS_DATASET_LAYOUT_VPQ_F16 ctypedef enum cuvsDatasetMemType_t: CUVS_DATASET_MEM_TYPE_HOST diff --git a/python/cuvs/cuvs/common/dataset.pyx b/python/cuvs/cuvs/common/dataset.pyx index 0c83633d13..813fe0b138 100644 --- a/python/cuvs/cuvs/common/dataset.pyx +++ b/python/cuvs/cuvs/common/dataset.pyx @@ -48,6 +48,8 @@ cdef class Dataset: check_cuvs(cuvsDatasetGetLayout(self.dataset, &layout)) if layout == CUVS_DATASET_LAYOUT_PADDED: return "padded" + if layout == CUVS_DATASET_LAYOUT_VPQ_F16: + return "vpq_f16" return "standard" @property diff --git a/python/cuvs/cuvs/neighbors/cagra/__init__.py b/python/cuvs/cuvs/neighbors/cagra/__init__.py index 60811a23eb..7bb3d8f502 100644 --- a/python/cuvs/cuvs/neighbors/cagra/__init__.py +++ b/python/cuvs/cuvs/neighbors/cagra/__init__.py @@ -6,6 +6,7 @@ from .cagra import ( AceParams, + CompressionParams, ExtendParams, Index, IndexParams, @@ -14,6 +15,7 @@ extend, from_graph, load, + make_vpq_dataset, save, search, update_dataset, @@ -21,6 +23,7 @@ __all__ = [ "AceParams", + "CompressionParams", "Dataset", "ExtendParams", "Index", @@ -30,6 +33,7 @@ "extend", "from_graph", "load", + "make_vpq_dataset", "save", "search", "update_dataset", diff --git a/python/cuvs/cuvs/neighbors/cagra/cagra.pxd b/python/cuvs/cuvs/neighbors/cagra/cagra.pxd index 9e4dbdb6f3..0b2806d3ad 100644 --- a/python/cuvs/cuvs/neighbors/cagra/cagra.pxd +++ b/python/cuvs/cuvs/neighbors/cagra/cagra.pxd @@ -144,9 +144,31 @@ cdef extern from "cuvs/neighbors/cagra.h" nogil: cuvsFilter filter) cuvsError_t cuvsCagraUpdateDataset( cuvsResources_t res, - cuvsDataset_t device_padded_dataset, + cuvsDataset_t dataset, cuvsCagraIndex_t index) + ctypedef struct cuvsCagraCompressionParams: + uint32_t pq_bits + uint32_t pq_dim + uint32_t vq_n_centers + uint32_t kmeans_n_iters + double vq_kmeans_trainset_fraction + double pq_kmeans_trainset_fraction + + ctypedef cuvsCagraCompressionParams* cuvsCagraCompressionParams_t + + cuvsError_t cuvsCagraCompressionParamsCreate( + cuvsCagraCompressionParams_t* params) + + cuvsError_t cuvsCagraCompressionParamsDestroy( + cuvsCagraCompressionParams_t params) + + cuvsError_t cuvsDatasetMakeVpq( + cuvsResources_t res, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* vpq_dataset) + cuvsError_t cuvsCagraSerializeGraph(cuvsResources_t res, const char * filename, cuvsCagraIndex_t index) diff --git a/python/cuvs/cuvs/neighbors/cagra/cagra.pyx b/python/cuvs/cuvs/neighbors/cagra/cagra.pyx index dd481df259..1fbf81fa14 100644 --- a/python/cuvs/cuvs/neighbors/cagra/cagra.pyx +++ b/python/cuvs/cuvs/neighbors/cagra/cagra.pyx @@ -56,6 +56,83 @@ from cuvs.neighbors import ivf_pq from cuvs.neighbors.filters import no_filter +cdef class CompressionParams: + """ + Parameters for VPQ compression (CAGRA-Q). + + Train a VPQ dataset with :func:`make_vpq_dataset`, then attach it with + :func:`update_dataset`. Metric must remain ``sqeuclidean`` / L2Expanded. + + Parameters + ---------- + pq_bits: int + The bit length of the vector element after compression by PQ. + Possible values: [4, 5, 6, 7, 8]. The smaller the 'pq_bits', the + smaller the index size and the better the search performance, but + the lower the recall. + pq_dim: int + The dimensionality of the vector after compression by PQ. When zero, + an optimal value is selected using a heuristic. + vq_n_centers: int + Vector Quantization (VQ) codebook size - number of "coarse cluster + centers". When zero, an optimal value is selected using a heuristic. + kmeans_n_iters: int + The number of iterations searching for kmeans centers (both VQ & PQ + phases). + vq_kmeans_trainset_fraction: float + The fraction of data to use during iterative kmeans building (VQ + phase). When zero, an optimal value is selected using a heuristic. + pq_kmeans_trainset_fraction: float + The fraction of data to use during iterative kmeans building (PQ + phase). When zero, an optimal value is selected using a heuristic. + """ + cdef cuvsCagraCompressionParams * params + + def __cinit__(self): + check_cuvs(cuvsCagraCompressionParamsCreate(&self.params)) + + def __dealloc__(self): + check_cuvs(cuvsCagraCompressionParamsDestroy(self.params)) + + def __init__(self, *, + pq_bits=8, + pq_dim=0, + vq_n_centers=0, + kmeans_n_iters=25, + vq_kmeans_trainset_fraction=0.0, + pq_kmeans_trainset_fraction=0.0): + self.params.pq_bits = pq_bits + self.params.pq_dim = pq_dim + self.params.vq_n_centers = vq_n_centers + self.params.kmeans_n_iters = kmeans_n_iters + self.params.vq_kmeans_trainset_fraction = vq_kmeans_trainset_fraction + self.params.pq_kmeans_trainset_fraction = pq_kmeans_trainset_fraction + + @property + def pq_bits(self): + return self.params.pq_bits + + @property + def pq_dim(self): + return self.params.pq_dim + + @property + def vq_n_centers(self): + return self.params.vq_n_centers + + @property + def kmeans_n_iters(self): + return self.params.kmeans_n_iters + + @property + def vq_kmeans_trainset_fraction(self): + return self.params.vq_kmeans_trainset_fraction + + @property + def pq_kmeans_trainset_fraction(self): + return self.params.pq_kmeans_trainset_fraction + + cdef class AceParams: """ Parameters for ACE (Augmented Core Extraction) graph building algorithm. @@ -579,27 +656,28 @@ def build(IndexParams index_params, dataset, resources=None): @auto_sync_resources -def update_dataset(Index index, padded_dataset, resources=None): +def update_dataset(Index index, dataset, resources=None): """ - Update any CAGRA index layout with a padded dataset. + Update/attach a CAGRA index with a device-padded or device VPQ dataset. - Accepts a ``Dataset`` or array. The index becomes search-ready in padded layout. + Accepts a ``Dataset`` (padded or ``vpq_f16``) or array (promoted to padded). + The index becomes search-ready in the matching layout. """ if not index.trained: raise ValueError("Index needs to be built before attaching dataset.") cdef Dataset dataset_obj source_array = None - if isinstance(padded_dataset, Dataset): - dataset_obj = padded_dataset + if isinstance(dataset, Dataset): + dataset_obj = dataset else: - source_array = padded_dataset - dataset_obj = make_device_padded_dataset(padded_dataset, resources=resources) + source_array = dataset + dataset_obj = make_device_padded_dataset(dataset, resources=resources) - cdef cuvsDataset_t dataset_handle = _cagra_dataset_handle(dataset_obj) - if dataset_obj.layout != "padded": - raise TypeError("padded_dataset must have padded layout") + if dataset_obj.layout not in ("padded", "vpq_f16"): + raise TypeError("dataset must have padded or vpq_f16 layout") + cdef cuvsDataset_t dataset_handle = _cagra_dataset_handle(dataset_obj) cdef cuvsResources_t res = resources.get_c_obj() with cuda_interruptible(): check_cuvs(cuvsCagraUpdateDataset( @@ -611,6 +689,55 @@ def update_dataset(Index index, padded_dataset, resources=None): return index +@auto_sync_resources +def make_vpq_dataset(padded_dataset, compression_params=None, resources=None): + """ + Train an owning device VPQ dataset (CAGRA-Q) from a device-padded dataset. + + Parameters + ---------- + padded_dataset : Dataset or array + Device-padded source used to train VPQ. Arrays are converted via + :func:`cuvs.common.dataset.make_device_padded_dataset`. + compression_params : CompressionParams, optional + VPQ training parameters. Defaults are used when omitted. + {resources_docstring} + + Returns + ------- + Dataset + Owning VPQ dataset handle. Keep it alive while any index uses it. + """ + cdef Dataset dataset_obj + if isinstance(padded_dataset, Dataset): + dataset_obj = padded_dataset + else: + dataset_obj = make_device_padded_dataset(padded_dataset, resources=resources) + + if dataset_obj.layout != "padded" or dataset_obj.memory_type != "device": + raise TypeError("padded_dataset must be a device-padded Dataset") + + cdef CompressionParams params_obj = None + cdef cuvsCagraCompressionParams_t params_ptr = NULL + if compression_params is not None: + if not isinstance(compression_params, CompressionParams): + raise TypeError("compression_params must be a CompressionParams") + params_obj = compression_params + params_ptr = params_obj.params + + cdef Dataset vpq = Dataset() + cdef cuvsResources_t res = resources.get_c_obj() + cdef cuvsDataset_t source_handle = _cagra_dataset_handle(dataset_obj) + with cuda_interruptible(): + check_cuvs(cuvsDatasetMakeVpq( + res, + source_handle, + params_ptr, + &vpq.dataset + )) + return vpq + + def build_index(IndexParams index_params, dataset, resources=None): warnings.warn("cagra.build_index is deprecated, use cagra.build instead", FutureWarning) diff --git a/python/cuvs/cuvs/tests/test_cagra.py b/python/cuvs/cuvs/tests/test_cagra.py index 25893a16f7..9e15ad76d3 100644 --- a/python/cuvs/cuvs/tests/test_cagra.py +++ b/python/cuvs/cuvs/tests/test_cagra.py @@ -226,6 +226,37 @@ def test_cagra_build_from_dataset_handle( assert distances.shape == (n_queries, k) +def test_cagra_vpq_build_update_search(): + """CAGRA-Q smoke: dense build → make_vpq_dataset → update_dataset → search.""" + n_rows, n_cols, n_queries, k = 256, 32, 4, 1 + dataset = generate_data((n_rows, n_cols), np.float32) + dataset_device = device_ndarray(dataset) + + index = cagra.build( + cagra.IndexParams(metric="sqeuclidean"), + dataset_device, + ) + compression = cagra.CompressionParams(pq_bits=8, pq_dim=8) + vpq = cagra.make_vpq_dataset( + dataset_device, compression_params=compression + ) + assert vpq.layout == "vpq_f16" + assert vpq.is_owning is True + + index = cagra.update_dataset(index, vpq) + + queries_device = device_ndarray(dataset[:n_queries]) + distances, neighbors = cagra.search( + cagra.SearchParams(), + index, + queries_device, + k, + ) + neighbors_h = neighbors.copy_to_host() + for i in range(n_queries): + assert neighbors_h[i, 0] == i + + @pytest.mark.parametrize("sparsity", [0.2, 0.5, 0.7, 1.0]) def test_filtered_cagra(sparsity): run_filtered_search_test(cagra, sparsity) diff --git a/rust/cuvs-sys/src/bindings.rs b/rust/cuvs-sys/src/bindings.rs index e723abaaea..b1689760ed 100644 --- a/rust/cuvs-sys/src/bindings.rs +++ b/rust/cuvs-sys/src/bindings.rs @@ -264,6 +264,7 @@ unsafe extern "C" { pub enum cuvsDatasetLayout_t { CUVS_DATASET_LAYOUT_STANDARD = 0, CUVS_DATASET_LAYOUT_PADDED = 1, + CUVS_DATASET_LAYOUT_VPQ_F16 = 2, } #[repr(u32)] #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] @@ -1356,10 +1357,19 @@ unsafe extern "C" { #[must_use] pub fn cuvsCagraUpdateDataset( res: cuvsResources_t, - device_padded_dataset: cuvsDataset_t, + dataset: cuvsDataset_t, index: cuvsCagraIndex_t, ) -> cuvsError_t; } +unsafe extern "C" { + #[must_use] + pub fn cuvsDatasetMakeVpq( + res: cuvsResources_t, + source_dataset: cuvsDataset_t, + params: cuvsCagraCompressionParams_t, + vpq_dataset: *mut cuvsDataset_t, + ) -> cuvsError_t; +} unsafe extern "C" { #[must_use] pub fn cuvsCagraBuild( diff --git a/rust/cuvs/src/dataset.rs b/rust/cuvs/src/dataset.rs index 35f9c4faf9..5a1f7c60e0 100644 --- a/rust/cuvs/src/dataset.rs +++ b/rust/cuvs/src/dataset.rs @@ -25,6 +25,8 @@ pub enum DatasetKind { HostPadded, /// Host-resident rows with a standard, unpadded width. HostStandard, + /// Device-resident VPQ (f16 codebook) dataset for CAGRA-Q search. + DeviceVpqF16, } impl DatasetKind { @@ -48,6 +50,16 @@ impl DatasetKind { ffi::cuvsDatasetMemType_t::CUVS_DATASET_MEM_TYPE_HOST, ffi::cuvsDatasetLayout_t::CUVS_DATASET_LAYOUT_STANDARD, ) => Self::HostStandard, + ( + ffi::cuvsDatasetMemType_t::CUVS_DATASET_MEM_TYPE_DEVICE, + ffi::cuvsDatasetLayout_t::CUVS_DATASET_LAYOUT_VPQ_F16, + ) => Self::DeviceVpqF16, + (mem, layout) => { + return Err(CagraError::Validation(format!( + "unsupported dataset mem_type/layout pair: {:?}/{:?}", + mem, layout + ))); + } }) } } @@ -211,6 +223,57 @@ impl private::Sealed for PaddedDataset { impl CuvsDataset for PaddedDataset {} +/// Owning device VPQ dataset (f16 codebooks) for CAGRA-Q search. +/// +/// Prefer [`crate::neighbors::cagra::make_vpq_dataset`] which accepts +/// [`crate::neighbors::cagra::CompressionParams`]. Keep this owner alive while +/// any index uses it. +#[derive(Debug)] +pub struct VpqDataset { + handle: ffi::cuvsDataset_t, +} + +impl VpqDataset { + /// Train VPQ storage from a device-padded dataset. + /// + /// `params` may be null to use library defaults. + pub(crate) fn train_raw( + res: &Resources, + source: &impl CuvsDataset, + params: ffi::cuvsCagraCompressionParams_t, + ) -> Result { + let kind = source.dataset_kind()?; + if kind != DatasetKind::DevicePadded { + return Err(CagraError::Validation(format!( + "VPQ training requires a device-padded dataset, got {:?}", + kind + ))); + } + unsafe { + let handle = init_handle(|out| { + ffi::cuvsDatasetMakeVpq(res.handle(), source.raw_dataset_handle(), params, out) + })?; + Ok(Self { handle }) + } + } +} + +impl Drop for VpqDataset { + fn drop(&mut self) { + if let Err(e) = check_cuvs(unsafe { ffi::cuvsDatasetDestroy(self.handle) }) { + report_drop_failure("vpq dataset", &e); + } + } +} + +impl private::Sealed for VpqDataset { + fn raw_dataset_handle(&self) -> ffi::cuvsDataset_t { + self.handle + } +} + +impl CuvsDataset for VpqDataset {} + /// Owning dataset storage returned by CAGRA deserialization. /// /// The allocation preserves the serialized host/device residency and diff --git a/rust/cuvs/src/neighbors/cagra/index.rs b/rust/cuvs/src/neighbors/cagra/index.rs index a93c90ca4d..fd09ae11f5 100644 --- a/rust/cuvs/src/neighbors/cagra/index.rs +++ b/rust/cuvs/src/neighbors/cagra/index.rs @@ -101,15 +101,15 @@ impl<'d> Index<'d> { Ok(handle) } - /// Attach a device-padded dataset and return a search-ready index borrowing it. + /// Attach a device-padded or device VPQ dataset and return a search-ready index borrowing it. pub fn update_dataset<'a, D>(self, res: &Resources, dataset: &'a D) -> Result> where D: CuvsDataset + ?Sized, { let kind = dataset.dataset_kind()?; - if kind != DatasetKind::DevicePadded { + if kind != DatasetKind::DevicePadded && kind != DatasetKind::DeviceVpqF16 { return Err(CagraError::Validation(format!( - "CAGRA dataset update requires a device-padded view, got {:?}", + "CAGRA dataset update requires a device-padded or device VPQ_F16 view, got {:?}", kind ))); } @@ -275,15 +275,15 @@ impl DeserializedIndex { serialize_to_hnswlib_impl(&self.handle, res, filename.as_ref()) } - /// Replace the deserialized storage with a caller-owned device-padded view. + /// Replace the deserialized storage with a caller-owned device-padded or VPQ view. pub fn update_dataset<'a, T>(self, res: &Resources, dataset: &'a T) -> Result> where T: CuvsDataset + ?Sized, { let kind = dataset.dataset_kind()?; - if kind != DatasetKind::DevicePadded { + if kind != DatasetKind::DevicePadded && kind != DatasetKind::DeviceVpqF16 { return Err(CagraError::Validation(format!( - "CAGRA dataset update requires a device-padded view, got {:?}", + "CAGRA dataset update requires a device-padded or device VPQ_F16 view, got {:?}", kind ))); } @@ -483,6 +483,35 @@ mod tests { test_cagra(build_params); } + /// CAGRA-Q smoke: dense build → make_vpq_dataset → update_dataset → search. + #[test] + fn test_cagra_vpq_build_update_search() { + use crate::neighbors::cagra::{CompressionParams, make_vpq_dataset}; + + const N_ROWS: usize = 256; + const N_COLS: usize = 32; + const N_QUERIES: usize = 4; + const K: usize = 1; + + let res = Resources::new().unwrap(); + let dataset = + ndarray::Array::::random((N_ROWS, N_COLS), Uniform::new(0., 1.0).unwrap()); + let dataset_device = DeviceTensor::from_host(&res, &dataset).unwrap(); + let index = Index::build(&res, &IndexParams::builder().build().unwrap(), &dataset_device) + .expect("failed to build dense cagra index"); + + // dim=32 float already matches CAGRA padded row width → padded view. + let padded = DatasetView::new(&res, &dataset_device).unwrap(); + assert_eq!(padded.dataset_kind().unwrap(), DatasetKind::DevicePadded); + + let compression = CompressionParams::new().unwrap().set_pq_bits(8).set_pq_dim(8); + let vpq = make_vpq_dataset(&res, &padded, Some(&compression)).expect("make_vpq_dataset"); + assert_eq!(vpq.dataset_kind().unwrap(), DatasetKind::DeviceVpqF16); + + let index = index.update_dataset(&res, &vpq).expect("update_dataset with VPQ"); + search_and_verify_self_neighbors(&res, &index, &dataset, N_QUERIES, K); + } + #[test] fn explicit_views_classify_and_build_all_supported_kinds() { let res = Resources::new().unwrap(); diff --git a/rust/cuvs/src/neighbors/cagra/mod.rs b/rust/cuvs/src/neighbors/cagra/mod.rs index c8ab865e02..d3f7c97429 100644 --- a/rust/cuvs/src/neighbors/cagra/mod.rs +++ b/rust/cuvs/src/neighbors/cagra/mod.rs @@ -20,13 +20,29 @@ mod index; mod params; -pub use crate::dataset::{CuvsDataset, Dataset, DatasetKind, DatasetView, PaddedDataset}; +pub use crate::dataset::{ + CuvsDataset, Dataset, DatasetKind, DatasetView, PaddedDataset, VpqDataset, +}; pub use crate::neighbors::filters::{Bitset, Filter}; pub use index::{DeserializedIndex, Index}; -pub use params::{IndexParams, SearchParams}; +pub use params::{CompressionParams, IndexParams, SearchParams}; use crate::dlpack::DLPackError; use crate::error::LibraryError; +use crate::resources::Resources; + +/// Train an owning device VPQ dataset (CAGRA-Q) from a device-padded source. +/// +/// `params` may be `None` to use library defaults. Keep the returned dataset +/// alive while any index uses it, then attach with [`Index::update_dataset`]. +pub fn make_vpq_dataset( + res: &Resources, + source: &impl CuvsDataset, + params: Option<&CompressionParams>, +) -> Result { + let params_ptr = params.map(CompressionParams::as_ptr).unwrap_or(std::ptr::null_mut()); + VpqDataset::train_raw(res, source, params_ptr) +} /// Algorithm for building the internal k-NN graph. #[derive(Debug, Copy, Clone, Hash, PartialEq, Eq)] diff --git a/rust/cuvs/src/neighbors/cagra/params.rs b/rust/cuvs/src/neighbors/cagra/params.rs index 4fcc3d18af..da919764f3 100644 --- a/rust/cuvs/src/neighbors/cagra/params.rs +++ b/rust/cuvs/src/neighbors/cagra/params.rs @@ -211,6 +211,88 @@ impl Drop for IndexParams { } } +// --------------------------------------------------------------------------- +// CompressionParams (CAGRA-Q / VPQ training) +// --------------------------------------------------------------------------- + +/// Parameters for VPQ compression used by CAGRA-Q. +pub struct CompressionParams { + handle: ffi::cuvsCagraCompressionParams_t, +} + +impl CompressionParams { + /// Allocate compression params with library defaults. + pub fn new() -> Result { + let mut handle: ffi::cuvsCagraCompressionParams_t = ptr::null_mut(); + check_cuvs(unsafe { ffi::cuvsCagraCompressionParamsCreate(&mut handle) })?; + Ok(Self { handle }) + } + + pub(crate) fn as_ptr(&self) -> ffi::cuvsCagraCompressionParams_t { + self.handle + } + + /// Bit length of each PQ code element. Valid values: 4..=8. + pub fn set_pq_bits(self, pq_bits: u32) -> Self { + unsafe { + (*self.handle).pq_bits = pq_bits; + } + self + } + + /// Dimensionality after PQ compression (`0` = heuristic). + pub fn set_pq_dim(self, pq_dim: u32) -> Self { + unsafe { + (*self.handle).pq_dim = pq_dim; + } + self + } + + /// VQ codebook size (`0` = heuristic). + pub fn set_vq_n_centers(self, vq_n_centers: u32) -> Self { + unsafe { + (*self.handle).vq_n_centers = vq_n_centers; + } + self + } + + /// KMeans iterations for VQ and PQ phases. + pub fn set_kmeans_n_iters(self, kmeans_n_iters: u32) -> Self { + unsafe { + (*self.handle).kmeans_n_iters = kmeans_n_iters; + } + self + } + + /// Fraction of data used for VQ kmeans (`0` = heuristic). + pub fn set_vq_kmeans_trainset_fraction(self, fraction: f64) -> Self { + unsafe { + (*self.handle).vq_kmeans_trainset_fraction = fraction; + } + self + } + + /// Fraction of data used for PQ kmeans (`0` = heuristic). + pub fn set_pq_kmeans_trainset_fraction(self, fraction: f64) -> Self { + unsafe { + (*self.handle).pq_kmeans_trainset_fraction = fraction; + } + self + } +} + +impl fmt::Debug for CompressionParams { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_tuple("CompressionParams").field(unsafe { &*self.handle }).finish() + } +} + +impl Drop for CompressionParams { + fn drop(&mut self) { + let _ = unsafe { ffi::cuvsCagraCompressionParamsDestroy(self.handle) }; + } +} + // --------------------------------------------------------------------------- // SearchParams // ---------------------------------------------------------------------------