diff --git a/c/include/cuvs/core/dataset.h b/c/include/cuvs/core/dataset.h index 78d3547495..c474089543 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 PQ storage with f16 codebooks (CAGRA-Q search dataset). */ + CUVS_DATASET_LAYOUT_PQ_F16 = 2 } cuvsDatasetLayout_t; /** diff --git a/c/include/cuvs/neighbors/cagra.h b/c/include/cuvs/neighbors/cagra.h index 57d063ef44..81fbc1707a 100644 --- a/c/include/cuvs/neighbors/cagra.h +++ b/c/include/cuvs/neighbors/cagra.h @@ -282,6 +282,24 @@ CUVS_EXPORT cuvsError_t cuvsCagraCompressionParamsCreate(cuvsCagraCompressionPar */ CUVS_EXPORT cuvsError_t cuvsCagraCompressionParamsDestroy(cuvsCagraCompressionParams_t params); +/** + * @brief Train an owning device PQ (f16 codebook) dataset from a device-padded source. + * + * Used for CAGRA-Q: build a dense CAGRA index, train PQ 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 PQ compression params; NULL selects defaults + * @param[out] pq_dataset newly allocated owning PQ dataset handle + * @return cuvsError_t + */ +CUVS_EXPORT cuvsError_t cuvsDatasetMakePq(cuvsResources_t res, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* pq_dataset); + /** * @brief Allocate ACE params, and populate with default values * @@ -607,21 +625,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 PQ). + * + * 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 PQ_F16 dataset (from `cuvsDatasetMakePq`): if \p index is already PQ-typed, its + * dataset view is replaced in place; otherwise the graph is copied into a new PQ-typed index + * (CAGRA-Q). Search requires metric `L2Expanded`. The PQ 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 PQ_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 82dba5424d..2ce5af523c 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_pq_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_pq_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_pq_f16: { + // Intentionally not dispatched here: most C API helpers (serialize/extend/merge/...) do not + // support PQ. Call sites that need PQ (search, attach) handle device_pq_f16 explicitly. + RAFT_FAIL( + "%s: PQ (CAGRA-Q) index layout is not supported by this operation", null_handle_err); + } } } @@ -529,37 +543,44 @@ static void make_host_standard_dataset_view(raft::resources*, } template -static void update_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, - "cuvsCagraUpdateDataset: dataset must be device padded"); +static void make_device_pq_dataset(raft::resources* res_ptr, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* output_pq_dataset) +{ + RAFT_EXPECTS(source_dataset != nullptr, "cuvsDatasetMakePq: null source dataset"); + RAFT_EXPECTS(source_dataset->addr != 0, "cuvsDatasetMakePq: null source dataset storage"); + RAFT_EXPECTS(output_pq_dataset != nullptr, "cuvsDatasetMakePq: null output dataset"); + RAFT_EXPECTS(source_dataset->mem_type == CUVS_DATASET_MEM_TYPE_DEVICE && + source_dataset->layout == CUVS_DATASET_LAYOUT_PADDED, + "cuvsDatasetMakePq: 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::update_dataset(*res_ptr, std::move(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 pq = + cuvs::preprocessing::quantize::pq::make_device_pq_dataset(*res_ptr, ps, padded_view.view()); + using pq_owner_t = cuvs::neighbors::device_vpq_dataset; + auto* owned = new pq_owner_t{std::move(pq)}; + auto* out = new cuvsDataset{}; + out->addr = reinterpret_cast(owned); + out->destroy_addr = &destroy_typed_addr; + // PQ 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_PQ_F16; + out->is_owning = true; + *output_pq_dataset = out; }); } @@ -752,42 +773,53 @@ void _search(cuvsResources_t res, { auto res_ptr = reinterpret_cast(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_pq_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 @@ -1457,6 +1489,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 cuvsDatasetMakePq(cuvsResources_t res, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* pq_dataset) +{ + return cuvs::core::translate_exceptions([=] { + RAFT_EXPECTS(source_dataset != nullptr, "cuvsDatasetMakePq: null source dataset"); + RAFT_EXPECTS(pq_dataset != nullptr, "cuvsDatasetMakePq: null output dataset"); + auto* res_ptr = reinterpret_cast(res); + auto make_typed = [&]() { + make_device_pq_dataset(res_ptr, source_dataset, params, pq_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("cuvsDatasetMakePq: 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([=] { @@ -1504,90 +1601,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_PQ_F16, + "cuvsCagraUpdateDataset: dataset must be device-padded or " + "device PQ_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, std::move(*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"); + attach_and_rebind(idx); + 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_pq_f16: + RAFT_FAIL( + "cuvsCagraUpdateDataset: cannot attach a padded dataset to a PQ index; " + "pass a device PQ_F16 dataset from cuvsDatasetMakePq"); + } + }); + } else if (dataset->layout == CUVS_DATASET_LAYOUT_PQ_F16) { + RAFT_EXPECTS(dataset->is_owning, + "cuvsCagraUpdateDataset: PQ dataset handle must be owning " + "(from cuvsDatasetMakePq)"); + + 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, std::move(*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_pq_f16: { + auto* idx = + reinterpret_cast*>( + box->index_ptr); + RAFT_EXPECTS(idx != nullptr, "cuvsCagraUpdateDataset: null index handle"); + attach_and_rebind(idx); + 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_update_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) { - update_dataset(res_ptr, device_padded_dataset, index); + update.template operator()(); } else if (index->dtype.code == kDLFloat && index->dtype.bits == 16) { - update_dataset(res_ptr, device_padded_dataset, index); + update.template operator()(); } else if (index->dtype.code == kDLInt && index->dtype.bits == 8) { - update_dataset(res_ptr, device_padded_dataset, index); + update.template operator()(); } else if (index->dtype.code == kDLUInt && index->dtype.bits == 8) { - update_dataset(res_ptr, device_padded_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; } - return dispatch_update_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. @@ -1769,9 +1920,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_pq_f16, + "cuvsCagraSearch: index must be device-padded or device-PQ. Call " + "cuvsCagraUpdateDataset with a device-padded or owning PQ_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 ba0e3ee310..95ba85fc90 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 @@ -2014,3 +2015,109 @@ TEST(CagraC, SearchMultiPartitionMultiKernelRejected) } cuvsResourcesDestroy(res); } + +TEST(CagraC, BuildAttachPqSearch) +{ + // CAGRA-Q smoke test: dense build → MakePq → UpdateDataset(PQ) → 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 pq = nullptr; + ASSERT_EQ(cuvsDatasetMakePq(res, padded, compression, &pq), CUVS_SUCCESS); + { + cuvsDatasetLayout_t layout; + ASSERT_EQ(cuvsDatasetGetLayout(pq, &layout), CUVS_SUCCESS); + EXPECT_EQ(layout, CUVS_DATASET_LAYOUT_PQ_F16); + bool owning = false; + ASSERT_EQ(cuvsDatasetGetIsOwning(pq, &owning), CUVS_SUCCESS); + EXPECT_TRUE(owning); + } + + ASSERT_EQ(cuvsCagraUpdateDataset(res, pq, 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(pq); + cuvsCagraIndexDestroy(index); + cuvsCagraIndexParamsDestroy(build_params); + cuvsDatasetDestroy(padded); + cuvsResourcesDestroy(res); +} diff --git a/cpp/CMakeLists.txt b/cpp/CMakeLists.txt index 4d51f17836..da053f3c97 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" @@ -1382,6 +1389,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 src/neighbors/detail/cagra/cagra_merge_scaffold_shared.cu diff --git a/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h b/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h index ed067fac39..1f47c01432 100644 --- a/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h +++ b/cpp/bench/ann/src/cuvs/cuvs_cagra_wrapper.h @@ -407,11 +407,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_pq_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_pq_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 e5728d54df..24770aa604 100644 --- a/cpp/include/cuvs/neighbors/cagra.hpp +++ b/cpp/include/cuvs/neighbors/cagra.hpp @@ -4550,45 +4550,31 @@ struct fd_transfer { } // namespace detail /** - * @brief Convert a standard-device index into a padded-device index and attach padded dataset. + * @brief Update or attach a device dataset to a CAGRA index. * - * 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`. + * These overloads are the single C++ API entry point for changing an index 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 + * 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. */ -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"); - - using GraphIndexType = - typename index>:: - graph_index_type; - auto graph_host = raft::make_host_matrix(standard_idx.graph().extent(0), - standard_idx.graph().extent(1)); - if (standard_idx.graph().size() > 0) { - 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); - } - device_padded_index out( - res, standard_idx.metric(), padded_dataset, raft::make_const_mdspan(graph_host.view())); - if (standard_idx.source_indices().has_value()) { - out.update_source_indices(res, standard_idx.source_indices().value()); - } - return out; -} +#define CUVS_CAGRA_DECLARE_UPDATE_DATASET_OVERLOADS(T) \ + 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 auto update_dataset( raft::resources const& res, @@ -4666,7 +4652,6 @@ auto update_dataset( index>&& cagra_index, device_padded_dataset_view dataset) -> index>; - auto update_dataset( raft::resources const& res, index>&& cagra_index, @@ -4687,7 +4672,6 @@ auto update_dataset( index>&& cagra_index, device_standard_dataset_view dataset) -> index>; - auto update_dataset( raft::resources const& res, index>&& cagra_index, @@ -4707,7 +4691,63 @@ auto update_dataset( index>&& cagra_index, device_padded_dataset_view dataset) -> index>; - +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset( + raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; auto update_dataset( raft::resources const& res, index>&& cagra_index, @@ -4727,6 +4767,22 @@ auto update_dataset( index>&& cagra_index, device_vpq_dataset_view dataset) -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; +auto update_dataset(raft::resources const& res, + index>&& cagra_index, + device_vpq_dataset_view dataset) + -> index>; } // namespace cagra } // namespace neighbors diff --git a/cpp/include/cuvs/preprocessing/quantize/pq.hpp b/cpp/include/cuvs/preprocessing/quantize/pq.hpp index a633ac6672..77749174bf 100644 --- a/cpp/include/cuvs/preprocessing/quantize/pq.hpp +++ b/cpp/include/cuvs/preprocessing/quantize/pq.hpp @@ -267,7 +267,7 @@ namespace detail { } // namespace detail /** - * @brief Train VPQ storage (codebooks + encoded rows) from a row-major mdspan/mdarray/dataset. + * @brief Train PQ storage (codebooks + encoded rows) from a row-major mdspan/mdarray/dataset. * * Accepts either a row-major mdspan with `value_type`, `extent`, `stride`, and `data_handle` (same * pattern as `cuvs::neighbors::make_device_padded_dataset`), or any cuVS dense dataset / dataset @@ -279,8 +279,8 @@ namespace detail { * dense dataset is never staged on the device in full; they must be tightly packed. Empty sources * are rejected. The element type must be `float`, `half`, `int8_t` or `uint8_t`. * - * 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 + * Typical **CAGRA** usage: build the graph on dense vectors, then attach PQ for search (metric + * must remain `L2Expanded` for this path). Train PQ from the same CAGRA-padded device layout you * used for graph build, keep the `device_vpq_dataset` alive, and call * `cagra::update_dataset` with a non-owning view. * @@ -288,18 +288,18 @@ namespace detail { * #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); - * auto vpq_idx = - * cuvs::neighbors::cagra::update_dataset(res, std::move(idx), vpq.as_dataset_view()); + * cuvs::neighbors::vpq_params pq_params{}; + * auto pq = cuvs::preprocessing::quantize::pq::make_device_pq_dataset(res, pq_params, padded); + * auto pq_idx = + * cuvs::neighbors::cagra::update_dataset(res, std::move(idx), pq.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_pq_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. @@ -311,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_pq_dataset( res, params, raft::mdspan{ @@ -322,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_pq_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_pq_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_build_inst.cu.in b/cpp/src/neighbors/cagra_build_inst.cu.in index e81f8cae95..83d95858d1 100644 --- a/cpp/src/neighbors/cagra_build_inst.cu.in +++ b/cpp/src/neighbors/cagra_build_inst.cu.in @@ -90,6 +90,10 @@ CUVS_INST_CAGRA_UPDATE_DATASET(data_t, inst_device_padded_view_t, inst_device_padded_view_t); CUVS_INST_CAGRA_UPDATE_DATASET(data_t, index_t, inst_device_padded_view_t, inst_vpq_view_t); +CUVS_INST_CAGRA_UPDATE_DATASET(data_t, index_t, inst_vpq_view_t, inst_vpq_view_t); +CUVS_INST_CAGRA_UPDATE_DATASET(data_t, index_t, inst_host_standard_view_t, inst_vpq_view_t); +CUVS_INST_CAGRA_UPDATE_DATASET(data_t, index_t, inst_host_padded_view_t, inst_vpq_view_t); +CUVS_INST_CAGRA_UPDATE_DATASET(data_t, index_t, inst_device_standard_view_t, inst_vpq_view_t); #undef CUVS_INST_CAGRA_UPDATE_DATASET 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..ffda989774 --- /dev/null +++ b/cpp/src/neighbors/cagra_update_dataset_inst.cu.in @@ -0,0 +1,53 @@ +/* + * 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) \ + void update_dataset(raft::resources const& res, \ + device_standard_index& idx, \ + cuvs::neighbors::device_standard_dataset_view const& dataset) \ + { \ + idx = device_standard_index(res, std::move(idx), dataset); \ + } \ + void update_dataset(raft::resources const& res, \ + device_padded_index& idx, \ + cuvs::neighbors::device_padded_dataset_view const& dataset) \ + { \ + idx = device_padded_index(res, std::move(idx), dataset); \ + } \ + void update_dataset(raft::resources const& res, \ + vpq_f16_index& idx, \ + cuvs::neighbors::device_vpq_dataset_view const& dataset) \ + { \ + idx = vpq_f16_index(res, std::move(idx), 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/update_dataset.cuh b/cpp/src/neighbors/detail/cagra/update_dataset.cuh new file mode 100644 index 0000000000..f3fc3164fd --- /dev/null +++ b/cpp/src/neighbors/detail/cagra/update_dataset.cuh @@ -0,0 +1,47 @@ +/* + * 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 +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()); + } + update_dataset(res, out, padded_dataset); + return out; +} + +} // namespace cuvs::neighbors::cagra::detail diff --git a/cpp/src/neighbors/iface/iface.hpp b/cpp/src/neighbors/iface/iface.hpp index 472d9e9467..acf883d0c9 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 diff --git a/cpp/src/neighbors/tiered_index.cu b/cpp/src/neighbors/tiered_index.cu index fd7848454f..a404e6aeff 100644 --- a/cpp/src/neighbors/tiered_index.cu +++ b/cpp/src/neighbors/tiered_index.cu @@ -3,6 +3,7 @@ * SPDX-License-Identifier: Apache-2.0 */ +#include "detail/cagra/update_dataset.cuh" #include "detail/tiered_index.cuh" #include @@ -98,7 +99,7 @@ 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( + auto ann_padded_idx = cuvs::neighbors::cagra::detail::convert_standard_to_padded_index( res, *idx.state->ann_index, ann_padded_view); next_state->ann_index = std::make_shared>( diff --git a/cpp/src/preprocessing/quantize/pq.cu b/cpp/src/preprocessing/quantize/pq.cu index 20b8f21d36..5ffd8c44dc 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_pq_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_pq_dataset: unsupported dataset element type %d", + static_cast(dtype)); } } 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 5dccce21da..509cb271c5 100644 --- a/cpp/tests/neighbors/ann_cagra/test_float_uint32_t.cu +++ b/cpp/tests/neighbors/ann_cagra/test_float_uint32_t.cu @@ -153,4 +153,58 @@ TEST(AnnCagraMultiPartition, MixedGraphDegreeRejected) cagra::search_params{}); } +// CAGRA-Q smoke test: build the graph on dense rows, train PQ storage from the same padded rows, +// update the index to PQ storage, then search. +TEST(AnnCagraPq, BuildUpdatePqSearch) +{ + 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 pq_params{.pq_bits = 8, .pq_dim = 8}; + auto pq = + cuvs::preprocessing::quantize::pq::make_device_pq_dataset(handle, pq_params, padded.view); + raft::resource::sync_stream(handle); + + EXPECT_EQ(pq.n_rows(), n_rows); + EXPECT_EQ(pq.dim(), dim); + + auto pq_index = cagra::update_dataset(handle, std::move(dense_index), pq.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{}, + pq_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..825b43e74a 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_pq_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_pq_dataset(handle, params, padded); raft::resource::sync_stream(handle); EXPECT_EQ(vpq.n_rows(), n_rows); diff --git a/go/cagra/cagra.go b/go/cagra/cagra.go index bb7103b111..6a41a0f5e6 100644 --- a/go/cagra/cagra.go +++ b/go/cagra/cagra.go @@ -21,11 +21,22 @@ type PaddedDataset struct { dataset C.cuvsDataset_t } +// Owning PQ dataset handle for CAGRA-Q search. +type PqDataset 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 PQ). +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 PQ 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 } +// MakePqDataset trains an owning device PQ 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 MakePqDataset(Resources cuvs.Resource, source PaddedDatasetHandle, params *CompressionParams) (*PqDataset, 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 pqDataset C.cuvsDataset_t + err := cuvs.CheckCuvs(cuvs.CuvsError(C.cuvsDatasetMakePq( + C.cuvsResources_t(Resources.Resource), + source.datasetHandle(), + cParams, + &pqDataset, + ))) + if err != nil { + return nil, err + } + return &PqDataset{dataset: pqDataset}, nil +} + +func (dataset *PqDataset) datasetHandle() C.cuvsDataset_t { + if dataset == nil { + return nil + } + return dataset.dataset +} + +// Close destroys an owning PQ dataset handle. +func (dataset *PqDataset) 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..eed0a6dd34 100644 --- a/go/cagra/cagra_test.go +++ b/go/cagra/cagra_test.go @@ -128,6 +128,134 @@ func TestCagra(t *testing.T) { } } +func TestCagraPqBuildUpdateSearch(t *testing.T) { + // CAGRA-Q smoke: dense build → MakePqDataset → 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) + } + + pq, err := MakePqDataset(resource, padded, compression) + if err != nil { + t.Fatalf("error creating PQ dataset: %v", err) + } + defer pq.Close() + + if err := UpdateDataset(resource, pq, index); err != nil { + t.Fatalf("error updating index with PQ 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 PQ 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..6214075ac3 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 PQ 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 PQ 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..fd75a35281 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 PQ training via + * {@link CagraIndex#makePqDataset(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 96e31431f8..05fd26e25a 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 PQ dataset handle for CAGRA-Q. Keep this alive for as long as any index using it + * remains in use. + */ + final class PqDataset extends OwningDataset { + public PqDataset() {} + } + /** * 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 PQ dataset (CAGRA-Q). Keep {@code pqDataset} + * alive while this index uses it. Metric must remain L2Expanded. + */ + void updateDataset(PqDataset pqDataset) throws Throwable; + + /** + * Train an owning device PQ dataset (CAGRA-Q) from a device-padded source. + * + * @param paddedDataset device-padded source dataset, owned or viewed + * @param compressionParams PQ training parameters; may be {@code null} for defaults */ - void updateDataset(PaddedDataset dataset) throws Throwable; + PqDataset makePqDataset( + PaddedDatasetHandle paddedDataset, CagraCompressionParams compressionParams) throws Throwable; /** Returns the CAGRA graph * @@ -376,7 +383,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 9506b130cd..7e942c8164 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); @@ -498,23 +498,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.PqDataset pqDataset) throws Throwable { checkNotDestroyed(); - Objects.requireNonNull(dataset); - if (!dataset.isPresent()) { - throw new IllegalArgumentException("dataset is uninitialized"); + Objects.requireNonNull(pqDataset); + if (!pqDataset.isPresent()) { + throw new IllegalArgumentException("pqDataset is uninitialized"); } - updateDataset(dataset.nativeHandleAddress()); + updateDataset(pqDataset.nativeHandleAddress()); } private void updateDataset(long datasetHandleAddress) { @@ -529,6 +529,55 @@ private void updateDataset(long datasetHandleAddress) { } } + @Override + public CagraIndex.PqDataset makePqDataset( + 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 pqDatasetPtr = localArena.allocate(cuvsDataset_t); + var returnValue = + cuvsDatasetMakePq( + cuvsRes, + MemorySegment.ofAddress(paddedDataset.nativeHandleAddress()), + paramsSeg, + pqDatasetPtr); + checkCuVSError(returnValue, "cuvsDatasetMakePq"); + MemorySegment pqDataset = pqDatasetPtr.get(cuvsDataset_t, 0); + + var out = new CagraIndex.PqDataset(); + out.setDelegate(new DatasetCloseDelegate(pqDataset), pqDataset.address()); + return out; + } finally { + if (compressionHandle != null) { + compressionHandle.close(); + } + } + } + } + @Override public void serialize(OutputStream outputStream) throws Throwable { Path path = @@ -674,16 +723,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") @@ -996,7 +1039,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; @@ -1013,7 +1056,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 500ef451f3..21d990f5ec 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 @@ -164,6 +164,77 @@ public void testIndexingAndSearchingFlow() throws Throwable { } } + /** + * CAGRA-Q smoke: dense build → makePqDataset → updateDataset → search. + * + * PQ 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 testPqBuildUpdateSearch() 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 pq = index.makePqDataset(padded, compressionParams); + var queryVectors = CuVSMatrix.ofArray(queries)) { + assertTrue(padded.isPresent()); + assertTrue(pq.isPresent()); + index.updateDataset(pq); + + 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..ac4a100efd 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_PQ_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..92f4a95328 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_PQ_F16: + return "pq_f16" return "standard" @property diff --git a/python/cuvs/cuvs/neighbors/cagra/__init__.py b/python/cuvs/cuvs/neighbors/cagra/__init__.py index 60811a23eb..55c5f2b761 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_pq_dataset, save, search, update_dataset, @@ -21,6 +23,7 @@ __all__ = [ "AceParams", + "CompressionParams", "Dataset", "ExtendParams", "Index", @@ -30,6 +33,7 @@ "extend", "from_graph", "load", + "make_pq_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..5c0c3999e4 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 cuvsDatasetMakePq( + cuvsResources_t res, + cuvsDataset_t source_dataset, + cuvsCagraCompressionParams_t params, + cuvsDataset_t* pq_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..4e0a5c44d9 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 PQ compression (CAGRA-Q). + + Train a PQ dataset with :func:`make_pq_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 PQ dataset. - Accepts a ``Dataset`` or array. The index becomes search-ready in padded layout. + Accepts a ``Dataset`` (padded or ``pq_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", "pq_f16"): + raise TypeError("dataset must have padded or pq_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_pq_dataset(padded_dataset, compression_params=None, resources=None): + """ + Train an owning device PQ dataset (CAGRA-Q) from a device-padded dataset. + + Parameters + ---------- + padded_dataset : Dataset or array + Device-padded source used to train PQ. Arrays are converted via + :func:`cuvs.common.dataset.make_device_padded_dataset`. + compression_params : CompressionParams, optional + PQ training parameters. Defaults are used when omitted. + {resources_docstring} + + Returns + ------- + Dataset + Owning PQ 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 pq = Dataset() + cdef cuvsResources_t res = resources.get_c_obj() + cdef cuvsDataset_t source_handle = _cagra_dataset_handle(dataset_obj) + with cuda_interruptible(): + check_cuvs(cuvsDatasetMakePq( + res, + source_handle, + params_ptr, + &pq.dataset + )) + return pq + + 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..95c996d79a 100644 --- a/python/cuvs/cuvs/tests/test_cagra.py +++ b/python/cuvs/cuvs/tests/test_cagra.py @@ -226,6 +226,35 @@ def test_cagra_build_from_dataset_handle( assert distances.shape == (n_queries, k) +def test_cagra_pq_build_update_search(): + """CAGRA-Q smoke: dense build → make_pq_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) + pq = cagra.make_pq_dataset(dataset_device, compression_params=compression) + assert pq.layout == "pq_f16" + assert pq.is_owning is True + + index = cagra.update_dataset(index, pq) + + 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..461c8d2074 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_PQ_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 cuvsDatasetMakePq( + res: cuvsResources_t, + source_dataset: cuvsDataset_t, + params: cuvsCagraCompressionParams_t, + pq_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..73c57f5144 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 PQ (f16 codebook) dataset for CAGRA-Q search. + DevicePqF16, } 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_PQ_F16, + ) => Self::DevicePqF16, + (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 PQ dataset (f16 codebooks) for CAGRA-Q search. +/// +/// Prefer [`crate::neighbors::cagra::make_pq_dataset`] which accepts +/// [`crate::neighbors::cagra::CompressionParams`]. Keep this owner alive while +/// any index uses it. +#[derive(Debug)] +pub struct PqDataset { + handle: ffi::cuvsDataset_t, +} + +impl PqDataset { + /// Train PQ 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!( + "PQ training requires a device-padded dataset, got {:?}", + kind + ))); + } + unsafe { + let handle = init_handle(|out| { + ffi::cuvsDatasetMakePq(res.handle(), source.raw_dataset_handle(), params, out) + })?; + Ok(Self { handle }) + } + } +} + +impl Drop for PqDataset { + fn drop(&mut self) { + if let Err(e) = check_cuvs(unsafe { ffi::cuvsDatasetDestroy(self.handle) }) { + report_drop_failure("pq dataset", &e); + } + } +} + +impl private::Sealed for PqDataset { + fn raw_dataset_handle(&self) -> ffi::cuvsDataset_t { + self.handle + } +} + +impl CuvsDataset for PqDataset {} + /// 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..9392a99420 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 PQ 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::DevicePqF16 { return Err(CagraError::Validation(format!( - "CAGRA dataset update requires a device-padded view, got {:?}", + "CAGRA dataset update requires a device-padded or device PQ_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 PQ 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::DevicePqF16 { return Err(CagraError::Validation(format!( - "CAGRA dataset update requires a device-padded view, got {:?}", + "CAGRA dataset update requires a device-padded or device PQ_F16 view, got {:?}", kind ))); } @@ -483,6 +483,35 @@ mod tests { test_cagra(build_params); } + /// CAGRA-Q smoke: dense build → make_pq_dataset → update_dataset → search. + #[test] + fn test_cagra_pq_build_update_search() { + use crate::neighbors::cagra::{CompressionParams, make_pq_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 pq = make_pq_dataset(&res, &padded, Some(&compression)).expect("make_pq_dataset"); + assert_eq!(pq.dataset_kind().unwrap(), DatasetKind::DevicePqF16); + + let index = index.update_dataset(&res, &pq).expect("update_dataset with PQ"); + 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..6aa1fc47c2 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, PqDataset, +}; 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 PQ 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_pq_dataset( + res: &Resources, + source: &impl CuvsDataset, + params: Option<&CompressionParams>, +) -> Result { + let params_ptr = params.map(CompressionParams::as_ptr).unwrap_or(std::ptr::null_mut()); + PqDataset::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 // ---------------------------------------------------------------------------