diff --git a/cpp/include/cuvs/cluster/kmeans.hpp b/cpp/include/cuvs/cluster/kmeans.hpp index a805aff0a4..6ac1b881f6 100644 --- a/cpp/include/cuvs/cluster/kmeans.hpp +++ b/cpp/include/cuvs/cluster/kmeans.hpp @@ -216,6 +216,19 @@ struct balanced_params : base_params { * average cluster size; in that mode, `balance_upper_tolerance` does not control donor selection. */ balanced_donor_selection donor_selection = balanced_donor_selection::SizeSorted; + + /** + * If true, treats uint8_t input data as bit-packed binary data where each byte contains 8 bits. + * Bits are expanded on-the-fly, least-significant bit first, to {-1, +1} floats + * during training and prediction. Other input types are rejected when this flag is set. + * When enabled: + * - Input data dimension represents packed dimension (actual_dim / 8) + * - Output centroids dimension is expanded (packed_dim * 8) + * - The metric operates on the expanded floating-point vectors (for example L2Expanded), + * not on the packed bytes; BitwiseHamming is not a balanced k-means training metric. + * - CosineExpanded is not supported. + */ + bool is_packed_binary = false; }; /** @@ -717,6 +730,12 @@ void fit(const raft::resources& handle, /** * @brief Find balanced clusters with k-means algorithm. * + * @note When `params.is_packed_binary` is true, `X.extent(1)` counts packed bytes, + * and centroids must have `8 * X.extent(1)` floating-point coordinates. Bits are + * expanded least-significant bit first to {-1, +1}; the selected metric operates + * on those expanded vectors. CosineExpanded is not supported in packed binary mode. + * With the flag disabled, uint8_t values are numeric. + * * @code{.cpp} * #include * #include @@ -1215,6 +1234,12 @@ void predict(const raft::resources& handle, /** * @brief Predict the closest cluster each sample in X belongs to. * + * @note When `params.is_packed_binary` is true, `X.extent(1)` counts packed bytes, + * and centroids must have `8 * X.extent(1)` floating-point coordinates. Bits are + * expanded least-significant bit first to {-1, +1}; the selected metric operates + * on those expanded vectors. CosineExpanded is not supported in packed binary mode. + * With the flag disabled, uint8_t values are numeric. + * * @code{.cpp} * #include * #include diff --git a/cpp/include/cuvs/detail/jit_lto/common_fragments.hpp b/cpp/include/cuvs/detail/jit_lto/common_fragments.hpp index 56180e3434..4f7ec74b15 100644 --- a/cpp/include/cuvs/detail/jit_lto/common_fragments.hpp +++ b/cpp/include/cuvs/detail/jit_lto/common_fragments.hpp @@ -12,6 +12,7 @@ struct tag_h {}; struct tag_d {}; struct tag_i8 {}; struct tag_u8 {}; +struct tag_u32 {}; struct tag_filter_none {}; struct tag_filter_bitset {}; struct tag_filter_bloom_filter {}; diff --git a/cpp/include/cuvs/detail/jit_lto/ivf_flat/interleaved_scan_fragments.hpp b/cpp/include/cuvs/detail/jit_lto/ivf_flat/interleaved_scan_fragments.hpp index 6f65837741..4c3230c920 100644 --- a/cpp/include/cuvs/detail/jit_lto/ivf_flat/interleaved_scan_fragments.hpp +++ b/cpp/include/cuvs/detail/jit_lto/ivf_flat/interleaved_scan_fragments.hpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -16,6 +16,7 @@ struct tag_acc_u32 {}; // Tag types for distance metrics struct tag_metric_euclidean {}; struct tag_metric_inner_product {}; +struct tag_metric_bitwise_hamming {}; struct tag_metric_custom_udf {}; // Tag types for post-processing diff --git a/cpp/include/cuvs/detail/jit_lto/pairwise_matrix/pairwise_matrix_fragments.hpp b/cpp/include/cuvs/detail/jit_lto/pairwise_matrix/pairwise_matrix_fragments.hpp index d9b7f4ec8d..acd431e61b 100644 --- a/cpp/include/cuvs/detail/jit_lto/pairwise_matrix/pairwise_matrix_fragments.hpp +++ b/cpp/include/cuvs/detail/jit_lto/pairwise_matrix/pairwise_matrix_fragments.hpp @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -13,6 +13,7 @@ struct tag_layout_col {}; struct tag_fin_op_identity {}; struct tag_fin_op_rbf {}; +struct tag_distance_bitwise_hamming {}; struct tag_distance_canberra {}; struct tag_distance_correlation {}; struct tag_distance_cosine {}; diff --git a/cpp/include/cuvs/neighbors/ivf_flat.hpp b/cpp/include/cuvs/neighbors/ivf_flat.hpp index 61fb10b1dc..29859bb4f0 100644 --- a/cpp/include/cuvs/neighbors/ivf_flat.hpp +++ b/cpp/include/cuvs/neighbors/ivf_flat.hpp @@ -38,11 +38,11 @@ struct index_params : cuvs::neighbors::index_params { * from scratch after invoking (`ivf_flat::extend`) a few times with new data, the distribution of * which is no longer representative of the original training set. * - * The alternative behavior (adaptive_centers = true) is to update the cluster centers for new - * data when it is added. In this case, `index.centers()` are always exactly the centroids of the - * data in the corresponding clusters. The drawback of this behavior is that the centroids depend - * on the order of adding new data (through the classification of the added data); that is, - * `index.centers()` "drift" together with the changing distribution of the newly added data. + * The alternative behavior (adaptive_centers = true) is to update the cluster centers when new + * data is added. For BitwiseHamming, centers are packed bitwise majorities of the data in each + * cluster, with ties resolved to zero. For other metrics, centers are floating-point means of + * the data in each cluster. Cluster assignments and centers depend on the order of adding new + * data, so the centers drift with the changing distribution of the newly added data. */ bool adaptive_centers = false; /** @@ -208,10 +208,38 @@ struct index : cuvs::neighbors::index { raft::device_vector_view list_sizes() noexcept; raft::device_vector_view list_sizes() const noexcept; - /** k-means cluster centers corresponding to the lists [n_lists, dim] */ + /** Floating-point k-means centers [n_lists, dim]; empty for binary indexes. */ raft::device_matrix_view centers() noexcept; raft::device_matrix_view centers() const noexcept; + /** + * @brief Packed binary cluster centers, with `dim()` bytes per center. + * @return A mutable device view of shape [n_lists, dim], or an empty view for nonbinary indexes. + */ + raft::device_matrix_view binary_centers() noexcept; + + /** + * @brief Packed binary cluster centers, with `dim()` bytes per center. + * @return A read-only device view of shape [n_lists, dim], or an empty view for nonbinary + * indexes. + */ + raft::device_matrix_view binary_centers() const noexcept; + + /** + * @brief Exact per-bit one-counts for adaptive binary centers. + * Together with list_sizes(), these retain majority statistics across extensions. + * @return A mutable device view of shape [n_lists, dim * 8], or an empty view when unused. + */ + raft::device_matrix_view binary_center_counts() noexcept; + + /** + * @brief Exact per-bit one-counts for adaptive binary centers. + * Together with list_sizes(), these retain majority statistics across extensions. + * @return A read-only device view of shape [n_lists, dim * 8], or an empty view when unused. + */ + raft::device_matrix_view binary_center_counts() + const noexcept; + /** * (Optional) Precomputed norms of the `centers` w.r.t. the chosen distance metric [n_lists]. * @@ -237,7 +265,10 @@ struct index : cuvs::neighbors::index { /** Total length of the index. */ IdxT size() const noexcept; - /** Dimensionality of the data. */ + /** Dimensionality of the data. + * @note For binary index, this returns the dimensionality of the byte dataset, which is the + * number of bits / 8. + */ uint32_t dim() const noexcept; /** Number of clusters/inverted lists. */ @@ -263,6 +294,9 @@ struct index : cuvs::neighbors::index { void check_consistency(); + /** Whether the index uses byte-packed vectors and BitwiseHamming distance. */ + bool binary_index() const noexcept; + private: /** * TODO: in theory, we can lift this to the template parameter and keep it at hardware maximum @@ -275,7 +309,10 @@ struct index : cuvs::neighbors::index { std::vector>> lists_; raft::device_vector list_sizes_; raft::device_matrix centers_; + raft::device_matrix binary_centers_; + raft::device_matrix binary_center_counts_; std::optional> center_norms_; + bool binary_index_; // Computed members raft::device_vector data_ptrs_; @@ -310,6 +347,7 @@ struct index : cuvs::neighbors::index { * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -339,6 +377,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -369,6 +408,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -398,6 +438,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -428,6 +469,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -457,6 +499,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -487,6 +530,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -516,6 +560,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Usage example: * @code{.cpp} @@ -546,6 +591,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -582,6 +628,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -619,6 +666,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -655,6 +703,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -692,6 +741,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -728,6 +778,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -765,6 +816,7 @@ void build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy @@ -801,6 +853,7 @@ auto build(raft::resources const& handle, * - L2Unexpanded * - InnerProduct * - CosineExpanded + * - BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) * * Note, if index_params.add_data_on_build is set to true, the user can set a * stream pool in the input raft::resource with at least one stream to enable kernel and copy diff --git a/cpp/src/cluster/detail/kmeans_balanced.cuh b/cpp/src/cluster/detail/kmeans_balanced.cuh index bb96c9115b..321acc9895 100644 --- a/cpp/src/cluster/detail/kmeans_balanced.cuh +++ b/cpp/src/cluster/detail/kmeans_balanced.cuh @@ -10,6 +10,7 @@ #include "../../core/nvtx.hpp" #include "../../distance/distance.cuh" +#include "../../distance/fused_distance_nn.cuh" #include #include @@ -38,9 +39,11 @@ #include #include +#include #include #include +#include "../../neighbors/detail/ann_utils.cuh" #include #include #include @@ -51,6 +54,47 @@ namespace cuvs::cluster::kmeans::detail { +/** Validate the input type and calculate the width used by the floating-point centers. */ +template +IdxT centers_dim(IdxT dim, bool is_packed_binary) +{ + RAFT_EXPECTS(dim > 0, "The number of features must be strictly positive"); + if (!is_packed_binary) { return dim; } + if constexpr (std::is_same_v) { + RAFT_EXPECTS(dim <= std::numeric_limits::max() / 8, + "The chosen index type cannot represent the expanded binary dimension"); + return dim * 8; + } else { + RAFT_FAIL("Packed binary mode is only supported for uint8_t data type"); + } +} + +inline void validate_packed_binary_metric(const cuvs::cluster::kmeans::balanced_params& params) +{ + RAFT_EXPECTS( + !params.is_packed_binary || params.metric != cuvs::distance::DistanceType::CosineExpanded, + "CosineExpanded is not supported for packed binary input"); +} + +/** + * @brief Create a transform iterator for on-the-fly bit expansion + * + * This helper function creates a thrust transform iterator that expands packed + * uint8_t data into float values on-the-fly (bit 1 → +1.0f, bit 0 → -1.0f), + * + * @tparam IdxT index type + * + * @param packed_data Pointer to row-major packed uint8_t data + * @return A transform iterator that yields float values for each bit + */ +template +auto make_bitwise_expanded_iterator(const uint8_t* packed_data) +{ + auto counting_iter = thrust::make_counting_iterator(0); + auto decoder = cuvs::spatial::knn::detail::utils::bitwise_decode_op(packed_data); + return thrust::make_transform_iterator(counting_iter, decoder); +} + /** * @brief Predict labels for the dataset; floating-point types only. * @@ -150,6 +194,86 @@ inline std::enable_if_t> predict_core( } } +/** + * @brief Predict labels for the dataset; uint8_t only (specialization for BitwiseHamming). + */ +template +inline void predict_bitwise_hamming(const raft::resources& handle, + const cuvs::cluster::kmeans::balanced_params& params, + const uint8_t* centers, + IdxT n_clusters, + IdxT dim, + const uint8_t* dataset, + const uint8_t* dataset_norm, + IdxT n_rows, + LabelT* labels, + rmm::device_async_resource_ref mr) +{ + RAFT_EXPECTS(params.metric == cuvs::distance::DistanceType::BitwiseHamming, + "uint8_t data only supports BitwiseHamming distance"); + + RAFT_EXPECTS(n_clusters > 0 && dim > 0, "Centers must have nonzero dimensions"); + if (n_rows == 0) { return; } + + auto workspace = raft::make_device_mdarray( + handle, mr, raft::make_extents((sizeof(int)) * n_rows)); + + auto minClusterAndDistance = raft::make_device_mdarray, IdxT>( + handle, mr, raft::make_extents(n_rows)); + raft::KeyValuePair initial_value(0, std::numeric_limits::max()); + raft::matrix::fill(handle, minClusterAndDistance.view(), initial_value); + + cuvs::distance::fusedDistanceNNMinReduce, IdxT>( + handle, + minClusterAndDistance.data_handle(), + dataset, + centers, + nullptr, + nullptr, + n_rows, + n_clusters, + dim, + (void*)workspace.data_handle(), + false, + false, + true, + params.metric, + 0.0f); + + raft::linalg::map(handle, + raft::make_const_mdspan(minClusterAndDistance.view()), + raft::make_device_vector_view(labels, n_rows), + raft::compose_op, raft::key_op>()); +} + +template +inline void predict_bitwise_hamming(const raft::resources& handle, + raft::device_matrix_view dataset, + raft::device_matrix_view centers, + raft::device_vector_view labels) +{ + RAFT_EXPECTS(dataset.extent(1) == centers.extent(1), + "Number of features in dataset and centroids are different"); + RAFT_EXPECTS(dataset.extent(0) == labels.extent(0), + "Number of rows in dataset and labels are different"); + RAFT_EXPECTS(static_cast(centers.extent(0)) <= + static_cast(std::numeric_limits::max()), + "The chosen label type cannot represent all cluster labels"); + cuvs::cluster::kmeans::balanced_params params; + params.metric = cuvs::distance::DistanceType::BitwiseHamming; + + predict_bitwise_hamming(handle, + params, + centers.data_handle(), + centers.extent(0), + centers.extent(1), + dataset.data_handle(), + nullptr, + dataset.extent(0), + labels.data_handle(), + raft::resource::get_workspace_resource_ref(handle)); +} + /** * @brief Suggest a minibatch size for kmeans prediction. * @@ -195,25 +319,30 @@ auto calc_minibatch_size(const raft::resources& handle, } // If we need to convert to MathT, space required for the converted batch. - if (!needs_conversion) { mem_per_row += sizeof(MathT) * dim; } + if (needs_conversion) { mem_per_row += sizeof(MathT) * size_t(dim); } - // Heuristic: calculate the minibatch size in order to use at most 80% or 512MB workspace memory. - // We go below 1GB here as the allocation is mostly done in a single chunk which - // is problematic if e.g. a pool allocator manages its own chunks <= 1GB. + // Include row norms in the workspace estimate. Cap large contiguous allocations at + // 80% of free workspace or 512 MiB, preserving upstream's pool-friendly allocation policy. + mem_per_row += sizeof(MathT); const auto free_ws_size = raft::resource::get_workspace_free_bytes(handle); const auto available_ws_size = - std::min((free_ws_size * size_t{8}) / size_t{10}, size_t{1} << 29); - - IdxT minibatch_size = std::max(IdxT{1}, static_cast(available_ws_size / mem_per_row)); - - minibatch_size = raft::round_down_safe(minibatch_size, IdxT{64}); - minibatch_size = std::min(minibatch_size, n_rows); - return std::make_tuple(minibatch_size, mem_per_row); + std::min((free_ws_size / size_t{10}) * size_t{8}, size_t{1} << 29); + size_t minibatch_size = std::max(1, available_ws_size / mem_per_row); + // Always make progress, including when fewer than 64 rows fit the budget. + if (minibatch_size >= 64) { minibatch_size = (minibatch_size / 64) * 64; } + IdxT batch_size = static_cast(std::min(minibatch_size, size_t(n_rows))); + return std::make_tuple(batch_size, mem_per_row); } /** * @brief Given the data and labels, calculate cluster centers and sizes in one sweep. * + * This function supports two modes: + * 1. Regular mode: Works with any data type T with optional type conversion via mapping_op + * 2. Packed binary mode: When T=uint8_t and is_packed_binary=true, treats data as bit-packed + * and expands bits on-the-fly (bit 1 → +1, bit 0 → -1) into float centers. + * In this mode, dim represents the packed dimension (dim_expanded / 8). + * * @note all pointers must be accessible on the device. * * @tparam T element type @@ -224,10 +353,10 @@ auto calc_minibatch_size(const raft::resources& handle, * @tparam MappingOpT type of the mapping operation * * @param[in] handle The raft handle. - * @param[inout] centers Pointer to the output [n_clusters, dim] + * @param[inout] centers Pointer to the output [n_clusters, dim] or [n_clusters, dim*8] if packed * @param[inout] cluster_sizes Number of rows in each cluster [n_clusters] * @param[in] n_clusters Number of clusters/centers - * @param[in] dim Dimensionality of the data + * @param[in] dim Dimensionality of the data (or packed dim if is_packed_binary=true) * @param[in] dataset Pointer to the data [n_rows, dim] * @param[in] n_rows Number of samples in the `dataset` * @param[in] labels Output predictions [n_rows] @@ -236,6 +365,8 @@ auto calc_minibatch_size(const raft::resources& handle, * the weighted average principle. * @param[in] mapping_op Mapping operation from T to MathT * @param[inout] mr (optional) Memory resource to use for temporary allocations on the device + * @param[in] is_packed_binary If true and T=uint8_t, treats data as bit-packed and expands + * on-the-fly */ template (centers, n_clusters, dim); + // For packed binary, dim is packed dimension, centers are in expanded dimension (dim * 8) + IdxT centers_dim = detail::centers_dim(dim, is_packed_binary); + + auto centersView = raft::make_device_matrix_view(centers, n_clusters, centers_dim); auto clusterSizesView = raft::make_device_vector_view(cluster_sizes, n_clusters); if (!reset_counters) { @@ -276,8 +411,26 @@ void calc_centers_and_sizes(const raft::resources& handle, temp_sizes = temp_cluster_sizes.data(); } + // Handle packed binary data with on-the-fly bit expansion + if (is_packed_binary) { + if constexpr (std::is_same_v) { + auto decoded_dataset_iter = make_bitwise_expanded_iterator(dataset); + raft::linalg::reduce_rows_by_key(decoded_dataset_iter, + centers_dim, + labels, + nullptr, + n_rows, + centers_dim, + n_clusters, + centers, + stream.get(), + reset_counters); + } else { + RAFT_FAIL("Packed binary mode is only supported for uint8_t data type"); + } + } // Apply mapping only when the data and math types are different. - if constexpr (std::is_same_v) { + else if constexpr (std::is_same_v) { raft::linalg::reduce_rows_by_key(dataset, dim, labels, @@ -533,18 +686,27 @@ void predict(const raft::resources& handle, auto stream = raft::resource::get_cuda_stream(handle); raft::common::nvtx::range fun_scope( "predict(%zu, %u)", static_cast(n_rows), n_clusters); - auto mem_res = mr.value_or(raft::resource::get_workspace_resource_ref(handle)); + auto mem_res = mr.value_or(raft::resource::get_workspace_resource_ref(handle)); + IdxT transformed_dim = centers_dim(dim, params.is_packed_binary); + if (n_rows == 0) { return; } auto [max_minibatch_size, _mem_per_row] = calc_minibatch_size( - handle, n_clusters, n_rows, dim, params.metric, std::is_same_v); + handle, n_clusters, n_rows, transformed_dim, params.metric, !std::is_same_v); rmm::device_uvector cur_dataset( - std::is_same_v ? 0 : max_minibatch_size * dim, stream, mem_res); + std::is_same_v ? 0 : max_minibatch_size * transformed_dim, stream, mem_res); constexpr bool native_half = std::is_same_v && std::is_same_v; - bool need_compute_norm = + bool need_norm = dataset_norm == nullptr && (params.metric == cuvs::distance::DistanceType::L2Expanded || params.metric == cuvs::distance::DistanceType::L2SqrtExpanded || params.metric == cuvs::distance::DistanceType::CosineExpanded); + bool need_compute_norm = need_norm && !params.is_packed_binary; rmm::device_uvector cur_dataset_norm( - need_compute_norm || native_half ? max_minibatch_size : 0, stream, mem_res); + need_norm || native_half ? max_minibatch_size : 0, stream, mem_res); + if (need_norm && params.is_packed_binary) { + raft::matrix::fill( + handle, + raft::make_device_matrix_view(cur_dataset_norm.data(), max_minibatch_size, 1), + static_cast(transformed_dim)); + } const auto native_centers_size = native_half ? static_cast(n_clusters) * static_cast(dim) : 0; std::optional native_half_scratch; @@ -552,8 +714,9 @@ void predict(const raft::resources& handle, native_half_scratch.emplace( native_centers_size, n_clusters, max_minibatch_size, stream, mem_res); } - const MathT* dataset_norm_ptr = nullptr; - auto cur_dataset_ptr = cur_dataset.data(); + const MathT* dataset_norm_ptr = + need_norm && params.is_packed_binary ? cur_dataset_norm.data() : nullptr; + auto cur_dataset_ptr = cur_dataset.data(); for (IdxT offset = 0; offset < n_rows; offset += max_minibatch_size) { IdxT minibatch_size = std::min(max_minibatch_size, n_rows - offset); @@ -575,6 +738,14 @@ void predict(const raft::resources& handle, } if constexpr (std::is_same_v) { cur_dataset_ptr = const_cast(dataset + offset * dim); + } else if (params.is_packed_binary) { + if constexpr (std::is_same_v) { + raft::linalg::map_offset(handle, + raft::make_device_matrix_view( + cur_dataset_ptr, minibatch_size, transformed_dim), + cuvs::spatial::knn::detail::utils::bitwise_decode_op( + dataset + offset * dim)); + } } else { raft::linalg::map( handle, @@ -612,7 +783,7 @@ void predict(const raft::resources& handle, params, centers, n_clusters, - dim, + transformed_dim, cur_dataset_ptr, dataset_norm_ptr, minibatch_size, @@ -622,7 +793,7 @@ void predict(const raft::resources& handle, } template average_size * balance_upper_tolerance * balance_upper_tolerance > 1 * @param[in] centroid_offset offset from the donor cluster centroid towards a donor point - * @param[in] mapping_op Mapping operation from T to MathT + * @param[in] mapping_op Mapping operation from dataset values to MathT * @param[inout] device_memory memory resource to use for temporary allocations * * @return whether any of the centers has been updated (and thus, `labels` need to be recalculated). */ -template (dim, params.is_packed_binary); for (uint32_t iter = 0; iter < n_iters; iter++) { // Balancing step - move the centers around to equalize cluster sizes // (but not on the first iteration) - if (iter > 0 && adjust_centers(handle, - cluster_centers, - n_clusters, - dim, - dataset, - n_rows, - cluster_labels, - cluster_sizes, - balance_lower_tolerance, - balance_upper_tolerance, - static_cast(params.centroid_offset), - params.donor_selection, - mapping_op, - device_memory)) { + bool did_adjust = false; + if (iter > 0) { + auto adjust = [&](auto data, auto data_mapping) { + return adjust_centers(handle, + cluster_centers, + n_clusters, + transformed_dim, + data, + n_rows, + cluster_labels, + cluster_sizes, + balance_lower_tolerance, + balance_upper_tolerance, + static_cast(params.centroid_offset), + params.donor_selection, + data_mapping, + device_memory); + }; + if (params.is_packed_binary) { + if constexpr (std::is_same_v) { + did_adjust = + adjust(make_bitwise_expanded_iterator(dataset), raft::identity_op{}); + } + } else { + did_adjust = adjust(dataset, mapping_op); + } + } + if (did_adjust) { if (balancing_counter++ >= balancing_pullback) { balancing_counter -= balancing_pullback; n_iters++; @@ -993,9 +1179,9 @@ void balancing_em_iters(const raft::resources& handle, case cuvs::distance::DistanceType::CosineExpanded: case cuvs::distance::DistanceType::CorrelationExpanded: { auto clusters_in_view = raft::make_device_matrix_view( - cluster_centers, n_clusters, dim); + cluster_centers, n_clusters, transformed_dim); auto clusters_out_view = raft::make_device_matrix_view( - cluster_centers, n_clusters, dim); + cluster_centers, n_clusters, transformed_dim); raft::linalg::row_normalize( handle, clusters_in_view, clusters_out_view); break; @@ -1024,6 +1210,7 @@ void balancing_em_iters(const raft::resources& handle, n_rows, cluster_labels, true, + params.is_packed_binary, mapping_op, device_memory); } @@ -1066,6 +1253,7 @@ void build_clusters(const raft::resources& handle, n_rows, cluster_labels, true, + params.is_packed_binary, mapping_op, device_memory); @@ -1190,7 +1378,8 @@ auto build_fine_clusters(const raft::resources& handle, rmm::device_async_resource_ref managed_memory, rmm::device_async_resource_ref device_memory) -> IdxT { - auto stream = raft::resource::get_cuda_stream(handle); + auto stream = raft::resource::get_cuda_stream(handle); + IdxT transformed_dim = centers_dim(dim, params.is_packed_binary); rmm::device_uvector mc_trainset_ids_buf(mesocluster_size_max, stream, managed_memory); // for small cluster counts the maximum mesocluster size is proportional to the number of rows, so // we use large workspace @@ -1205,7 +1394,7 @@ auto build_fine_clusters(const raft::resources& handle, rmm::device_uvector mc_trainset_labels(mesocluster_size_max, stream, device_memory); rmm::device_uvector mc_trainset_ccenters( - fine_clusters_nums_max * dim, stream, device_memory); + fine_clusters_nums_max * transformed_dim, stream, device_memory); // number of vectors in each cluster rmm::device_uvector mc_trainset_csizes_tmp( fine_clusters_nums_max, stream, device_memory); @@ -1237,11 +1426,18 @@ auto build_fine_clusters(const raft::resources& handle, if (params.metric == cuvs::distance::DistanceType::L2Expanded || params.metric == cuvs::distance::DistanceType::L2SqrtExpanded || params.metric == cuvs::distance::DistanceType::CosineExpanded) { - thrust::gather(raft::resource::get_thrust_policy(handle), - mc_trainset_ids, - mc_trainset_ids + k, - dataset_norm_mptr, - mc_trainset_norm); + if (params.is_packed_binary) { + // Expanded bits have a constant squared row norm. + raft::matrix::fill(handle, + raft::make_device_matrix_view(mc_trainset_norm, k, 1), + static_cast(transformed_dim)); + } else { + thrust::gather(raft::resource::get_thrust_policy(handle), + mc_trainset_ids, + mc_trainset_ids + k, + dataset_norm_mptr, + mc_trainset_norm); + } } build_clusters(handle, @@ -1256,12 +1452,12 @@ auto build_fine_clusters(const raft::resources& handle, mapping_op, device_memory, mc_trainset_norm); - - raft::copy(handle, - raft::make_device_vector_view(cluster_centers + (dim * fine_clusters_csum[i]), - fine_clusters_nums[i] * dim), - raft::make_device_vector_view(mc_trainset_ccenters.data(), - fine_clusters_nums[i] * dim)); + raft::copy( + handle, + raft::make_device_vector_view(cluster_centers + (transformed_dim * fine_clusters_csum[i]), + fine_clusters_nums[i] * transformed_dim), + raft::make_device_vector_view(mc_trainset_ccenters.data(), + fine_clusters_nums[i] * transformed_dim)); raft::resource::sync_stream(handle, stream); n_clusters_done += fine_clusters_nums[i]; } @@ -1299,8 +1495,9 @@ void build_hierarchical(const raft::resources& handle, MappingOpT mapping_op, MathT* inertia = nullptr) { - auto stream = raft::resource::get_cuda_stream(handle); - using LabelT = uint32_t; + auto stream = raft::resource::get_cuda_stream(handle); + using LabelT = uint32_t; + IdxT transformed_dim = centers_dim(dim, params.is_packed_binary); raft::common::nvtx::range fun_scope( "build_hierarchical(%zu, %u)", static_cast(n_rows), n_clusters); @@ -1312,21 +1509,22 @@ void build_hierarchical(const raft::resources& handle, rmm::mr::managed_memory_resource managed_memory; rmm::device_async_resource_ref device_memory = raft::resource::get_workspace_resource_ref(handle); auto [max_minibatch_size, mem_per_row] = calc_minibatch_size( - handle, n_clusters, n_rows, dim, params.metric, std::is_same_v); + handle, n_clusters, n_rows, transformed_dim, params.metric, !std::is_same_v); // Precompute the L2 norm of the dataset if relevant and not yet computed. rmm::device_uvector dataset_norm_buf(0, stream, device_memory); const MathT* dataset_norm = nullptr; if ((params.metric == cuvs::distance::DistanceType::L2Expanded || params.metric == cuvs::distance::DistanceType::L2SqrtExpanded || - params.metric == cuvs::distance::DistanceType::CosineExpanded)) { + params.metric == cuvs::distance::DistanceType::CosineExpanded) && + !params.is_packed_binary) { dataset_norm_buf.resize(n_rows, stream); for (IdxT offset = 0; offset < n_rows; offset += max_minibatch_size) { IdxT minibatch_size = std::min(max_minibatch_size, n_rows - offset); if (params.metric == cuvs::distance::DistanceType::CosineExpanded) compute_norm(handle, dataset_norm_buf.data() + offset, - dataset + dim * offset, + dataset + offset * dim, dim, minibatch_size, mapping_op, @@ -1335,14 +1533,23 @@ void build_hierarchical(const raft::resources& handle, else compute_norm(handle, dataset_norm_buf.data() + offset, - dataset + dim * offset, + dataset + offset * dim, dim, minibatch_size, mapping_op, raft::identity_op{}, device_memory); } - dataset_norm = (const MathT*)dataset_norm_buf.data(); + dataset_norm = dataset_norm_buf.data(); + } else if (params.is_packed_binary && + (params.metric == cuvs::distance::DistanceType::L2Expanded || + params.metric == cuvs::distance::DistanceType::L2SqrtExpanded)) { + dataset_norm_buf.resize(n_rows, stream); + raft::matrix::fill( + handle, + raft::make_device_matrix_view(dataset_norm_buf.data(), n_rows, 1), + static_cast(transformed_dim)); + dataset_norm = dataset_norm_buf.data(); } /* Temporary workaround to cub::DeviceHistogram not supporting any type that isn't natively @@ -1354,7 +1561,8 @@ void build_hierarchical(const raft::resources& handle, rmm::device_uvector mesocluster_labels_buf(n_rows, stream, managed_memory); rmm::device_uvector mesocluster_sizes_buf(n_mesoclusters, stream, managed_memory); { - rmm::device_uvector mesocluster_centers_buf(n_mesoclusters * dim, stream, device_memory); + rmm::device_uvector mesocluster_centers_buf( + n_mesoclusters * transformed_dim, stream, device_memory); build_clusters(handle, params, dim, diff --git a/cpp/src/cluster/kmeans_balanced.cuh b/cpp/src/cluster/kmeans_balanced.cuh index 68c4925252..9f2833f50b 100644 --- a/cpp/src/cluster/kmeans_balanced.cuh +++ b/cpp/src/cluster/kmeans_balanced.cuh @@ -73,10 +73,12 @@ void fit(const raft::resources& handle, MappingOpT mapping_op = raft::identity_op(), std::optional> inertia = std::nullopt) { - RAFT_EXPECTS(X.extent(1) == centroids.extent(1), + cuvs::cluster::kmeans::detail::validate_packed_binary_metric(params); + auto centers_dim = + cuvs::cluster::kmeans::detail::centers_dim(X.extent(1), params.is_packed_binary); + RAFT_EXPECTS(centers_dim == centroids.extent(1), "Number of features in dataset and centroids are different"); - RAFT_EXPECTS(static_cast(X.extent(0)) * static_cast(X.extent(1)) <= - static_cast(std::numeric_limits::max()), + RAFT_EXPECTS(X.extent(0) <= std::numeric_limits::max() / centers_dim, "The chosen index type cannot represent all indices for the given dataset"); RAFT_EXPECTS(centroids.extent(0) > IndexT{0} && centroids.extent(0) <= X.extent(0), "The number of centroids must be strictly positive and cannot exceed the number of " @@ -135,12 +137,16 @@ void predict(const raft::resources& handle, raft::device_vector_view labels, MappingOpT mapping_op = raft::identity_op()) { + cuvs::cluster::kmeans::detail::validate_packed_binary_metric(params); RAFT_EXPECTS(X.extent(0) == labels.extent(0), "Number of rows in dataset and labels are different"); - RAFT_EXPECTS(X.extent(1) == centroids.extent(1), + auto centers_dim = + cuvs::cluster::kmeans::detail::centers_dim(X.extent(1), params.is_packed_binary); + RAFT_EXPECTS(centers_dim == centroids.extent(1), "Number of features in dataset and centroids are different"); - RAFT_EXPECTS(static_cast(X.extent(0)) * static_cast(X.extent(1)) <= - static_cast(std::numeric_limits::max()), + RAFT_EXPECTS(centroids.extent(0) > 0, "The number of centroids must be strictly positive"); + RAFT_EXPECTS(X.extent(0) <= std::numeric_limits::max() / centers_dim && + centroids.extent(0) <= std::numeric_limits::max() / centers_dim, "The chosen index type cannot represent all indices for the given dataset"); RAFT_EXPECTS(static_cast(centroids.extent(0)) <= static_cast(std::numeric_limits::max()), @@ -266,13 +272,16 @@ EXTERN_TEMPLATE_BUILD_CLUSTERS( * @param[in] X Dataset for which to calculate cluster centers. The data must be in * row-major format. [dim = n_samples x n_features] * @param[in] labels The input labels [dim = n_samples] - * @param[out] centroids The output centroids [dim = n_clusters x n_features] + * @param[out] centroids The output centroids + * [dim = n_clusters x + * (is_packed_binary ? 8 * n_features : n_features)] * @param[out] cluster_sizes Size of each cluster [dim = n_clusters] * @param[in] reset_counters Whether to clear the output arrays before calculating. * When set to `false`, this function may be used to update existing * centers and sizes using the weighted average principle. * @param[in] mapping_op (optional) Functor to convert from the input datatype to the * arithmetic datatype. If DataT == MathT, this must be the identity. + * @param[in] is_packed_binary Treat uint8_t X as packed bits. Requires DataT == uint8_t. */ template centroids, raft::device_vector_view cluster_sizes, bool reset_counters = true, - MappingOpT mapping_op = raft::identity_op()) + MappingOpT mapping_op = raft::identity_op(), + bool is_packed_binary = false) { RAFT_EXPECTS(X.extent(0) == labels.extent(0), "Number of rows in dataset and labels are different"); - RAFT_EXPECTS(X.extent(1) == centroids.extent(1), + auto centers_dim = + cuvs::cluster::kmeans::detail::centers_dim(X.extent(1), is_packed_binary); + RAFT_EXPECTS(centers_dim == centroids.extent(1), "Number of features in dataset and centroids are different"); + RAFT_EXPECTS(X.extent(0) <= std::numeric_limits::max() / centers_dim && + centroids.extent(0) <= std::numeric_limits::max() / centers_dim, + "The chosen index type cannot represent all indices for the given dataset"); RAFT_EXPECTS(centroids.extent(0) == cluster_sizes.extent(0), - "Number of rows in centroids and clusyer_sizes are different"); + "Number of rows in centroids and cluster_sizes are different"); cuvs::cluster::kmeans::detail::calc_centers_and_sizes( handle, @@ -305,6 +320,7 @@ void calc_centers_and_sizes(const raft::resources& handle, X.extent(0), labels.data_handle(), reset_counters, + is_packed_binary, mapping_op, raft::resource::get_workspace_resource_ref(handle)); } diff --git a/cpp/src/cluster/kmeans_balanced_build_clusters_impl.cuh b/cpp/src/cluster/kmeans_balanced_build_clusters_impl.cuh index 2bce856c6c..e160cc710f 100644 --- a/cpp/src/cluster/kmeans_balanced_build_clusters_impl.cuh +++ b/cpp/src/cluster/kmeans_balanced_build_clusters_impl.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -50,10 +50,19 @@ void build_clusters(const raft::resources& handle, MappingOpT mapping_op, std::optional> X_norm) { + cuvs::cluster::kmeans::detail::validate_packed_binary_metric(params); + RAFT_EXPECTS(centroids.extent(0) > IndexT{0}, + "The number of centroids must be strictly positive"); RAFT_EXPECTS(X.extent(0) == labels.extent(0), "Number of rows in dataset and labels are different"); - RAFT_EXPECTS(X.extent(1) == centroids.extent(1), + auto centers_dim = + cuvs::cluster::kmeans::detail::centers_dim(X.extent(1), params.is_packed_binary); + RAFT_EXPECTS(centers_dim == centroids.extent(1), "Number of features in dataset and centroids are different"); + RAFT_EXPECTS(X.extent(0) <= std::numeric_limits::max() / centers_dim, + "The chosen index type cannot represent all indices for the given dataset"); + RAFT_EXPECTS(!X_norm.has_value() || X_norm->extent(0) == X.extent(0), + "Number of rows in dataset and norms are different"); RAFT_EXPECTS(centroids.extent(0) == cluster_sizes.extent(0), "Number of rows in centroids and clusyer_sizes are different"); diff --git a/cpp/src/distance/detail/distance_ops/all_ops.cuh b/cpp/src/distance/detail/distance_ops/all_ops.cuh index f0a3984eb6..104ce00f01 100644 --- a/cpp/src/distance/detail/distance_ops/all_ops.cuh +++ b/cpp/src/distance/detail/distance_ops/all_ops.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -9,6 +9,7 @@ #include "cutlass.cuh" // The distance operations: +#include "../distance_ops/bitwise_hamming.cuh" #include "../distance_ops/canberra.cuh" #include "../distance_ops/correlation.cuh" #include "../distance_ops/cosine.cuh" diff --git a/cpp/src/distance/detail/distance_ops/bitwise_hamming.cuh b/cpp/src/distance/detail/distance_ops/bitwise_hamming.cuh new file mode 100644 index 0000000000..5f1808143e --- /dev/null +++ b/cpp/src/distance/detail/distance_ops/bitwise_hamming.cuh @@ -0,0 +1,70 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include + +#include +#include +#include +#include + +namespace cuvs::distance::detail::ops { + +/** + * @brief the Bitwise Hamming distance matrix calculation + * It computes the following equation: + * + * c_ij = sum_k popcount(x_ik XOR y_kj) + * + * where x and y are binary data packed as uint8_t + */ +template +struct bitwise_hamming_distance_op { + using DataT = DataType; + using AccT = AccType; + using IdxT = IdxType; + + IdxT k; + + bitwise_hamming_distance_op(IdxT k_) : k(k_) + { + static_assert(std::is_same_v, "BitwiseHamming only supports uint8_t"); + static_assert(std::is_same_v, "BitwiseHamming requires a uint32_t accumulator"); + RAFT_EXPECTS(k >= 0 && static_cast(k) <= std::numeric_limits::max() / 8, + "BitwiseHamming dimension exceeds the uint32_t accumulator range"); + } + + static constexpr bool use_norms = false; + static constexpr bool expensive_inner_loop = false; + + template + static constexpr size_t shared_mem_size() + { + return Policy::SmemSize; + } + + __device__ __forceinline__ void core(AccT& acc, DataT& x, DataT& y) const + { + static_assert(std::is_same_v, "BitwiseHamming only supports uint8_t"); + // Ensure proper masking and casting to avoid undefined behavior + uint32_t xor_val = static_cast(static_cast(x ^ y)); + uint32_t masked_val = xor_val & 0xffu; + int popcount = __popc(masked_val); + acc += static_cast(popcount); + } + + template + __device__ __forceinline__ void epilog(AccT acc[Policy::AccRowsPerTh][Policy::AccColsPerTh], + AccT* regxn, + AccT* regyn, + IdxT gridStrideX, + IdxT gridStrideY) const + { + } +}; + +} // namespace cuvs::distance::detail::ops diff --git a/cpp/src/distance/detail/fused_distance_nn.cuh b/cpp/src/distance/detail/fused_distance_nn.cuh index 632780fc33..e100ace651 100644 --- a/cpp/src/distance/detail/fused_distance_nn.cuh +++ b/cpp/src/distance/detail/fused_distance_nn.cuh @@ -10,6 +10,7 @@ #include "fused_distance_nn/cutile/fused_1nn_tile.hpp" #endif #include "fused_distance_nn/cutlass_base.cuh" +#include "fused_distance_nn/fused_bitwise_hamming_nn.cuh" #include "fused_distance_nn/fused_cosine_nn.cuh" #include "fused_distance_nn/fused_l2_nn.cuh" #include "fused_distance_nn/helper_structs.cuh" @@ -180,28 +181,46 @@ void fusedDistanceNNImpl(raft::resources const& handle, dim3 blk(P::Nthreads); auto nblks = raft::ceildiv(m, P::Nthreads); - constexpr auto maxVal = std::numeric_limits::max(); - typedef raft::KeyValuePair KVPair; + using AccT = std::conditional_t, uint32_t, DataT>; + constexpr auto maxVal = std::numeric_limits::max(); RAFT_CUDA_TRY(cudaMemsetAsync(workspace, 0, sizeof(int) * m, stream.get())); if (initOutBuffer) { - initKernel + initKernel <<>>(min, m, maxVal, redOp); RAFT_CUDA_TRY(cudaGetLastError()); } + // An empty candidate set leaves the initialized (or supplied) result unchanged. + if (n == 0) { return; } + switch (metric) { case cuvs::distance::DistanceType::CosineExpanded: - fusedCosineNN( - min, x, y, xn, yn, m, n, k, workspace, redOp, pairRedOp, sqrt, stream.get()); + if constexpr (std::is_same_v || std::is_same_v) { + RAFT_FAIL("Cosine distance is not supported for uint8_t/int8_t data types"); + } else { + fusedCosineNN( + min, x, y, xn, yn, m, n, k, workspace, redOp, pairRedOp, sqrt, stream.get()); + } break; case cuvs::distance::DistanceType::L2SqrtExpanded: case cuvs::distance::DistanceType::L2Expanded: - // initOutBuffer is take care by fusedDistanceNNImpl() so we set it false to fusedL2NNImpl. - fusedL2NNImpl( - min, x, y, xn, yn, m, n, k, workspace, redOp, pairRedOp, sqrt, false, stream.get()); + if constexpr (std::is_same_v || std::is_same_v) { + RAFT_FAIL("L2 distance is not supported for uint8_t/int8_t data types"); + } else { + fusedL2NNImpl( + min, x, y, xn, yn, m, n, k, workspace, redOp, pairRedOp, sqrt, false, stream.get()); + } + break; + case cuvs::distance::DistanceType::BitwiseHamming: + if constexpr (std::is_same_v) { + fusedBitwiseHammingNN( + min, x, y, xn, yn, m, n, k, workspace, redOp, pairRedOp, sqrt, stream.get()); + } else { + RAFT_FAIL("BitwiseHamming distance only supports uint8_t data type"); + } break; - default: RAFT_FAIL("Only cosine and L2 metrics are supported by fusedDistanceNN"); + default: RAFT_FAIL("only cosine/l2/bitwise hamming metric is supported with fusedDistanceNN"); } } diff --git a/cpp/src/distance/detail/fused_distance_nn/fused_bitwise_hamming_nn.cuh b/cpp/src/distance/detail/fused_distance_nn/fused_bitwise_hamming_nn.cuh new file mode 100644 index 0000000000..e0daa55a8d --- /dev/null +++ b/cpp/src/distance/detail/fused_distance_nn/fused_bitwise_hamming_nn.cuh @@ -0,0 +1,81 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include "../distance_ops/bitwise_hamming.cuh" // ops::bitwise_hamming_distance_op +#include "../pairwise_distance_base.cuh" // PairwiseDistances +#include "helper_structs.cuh" +#include "simt_kernel.cuh" + +namespace cuvs { +namespace distance { +namespace detail { + +/** + * @brief Fused BitwiseHamming distance and 1-nearest-neighbor + * + * This implementation is only meaningful for uint8_t data type. + * The if constexpr in fusedDistanceNNImpl ensures it's only called for uint8_t. + */ +template +void fusedBitwiseHammingNN(OutT* min, + const DataT* x, + const DataT* y, + const DataT* xn, + const DataT* yn, + IdxT m, + IdxT n, + IdxT k, + int* workspace, + ReduceOpT redOp, + KVPReduceOpT pairRedOp, + bool sqrt, + cudaStream_t stream) +{ + typedef Policy P; + + dim3 blk(P::Nthreads); + constexpr auto maxVal = std::numeric_limits::max(); + using distance_op_type = ops::bitwise_hamming_distance_op; + distance_op_type distance_op{k}; + auto kernel = fusedDistanceNNkernel; + + constexpr size_t shmemSize = P::SmemSize; + + dim3 grid = launchConfigGenerator

(m, n, shmemSize, kernel); + + kernel<<>>(min, + x, + y, + nullptr, + nullptr, + m, + n, + k, + maxVal, + workspace, + redOp, + pairRedOp, + distance_op, + raft::identity_op{}); + + RAFT_CUDA_TRY(cudaGetLastError()); +} + +} // namespace detail +} // namespace distance +} // namespace cuvs diff --git a/cpp/src/distance/detail/fused_distance_nn/helper_structs.cuh b/cpp/src/distance/detail/fused_distance_nn/helper_structs.cuh index 762c720568..eeccdd5e45 100644 --- a/cpp/src/distance/detail/fused_distance_nn/helper_structs.cuh +++ b/cpp/src/distance/detail/fused_distance_nn/helper_structs.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2021-2024, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2021-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -27,8 +27,15 @@ namespace detail { template struct KVPMinReduceImpl { typedef raft::KeyValuePair KVP; - DI KVP operator()(LabelT rit, const KVP& a, const KVP& b) { return b.value < a.value ? b : a; } - DI KVP operator()(const KVP& a, const KVP& b) { return b.value < a.value ? b : a; } + // Use index as tiebreaker for consistent behavior when distances are equal + DI KVP operator()(LabelT rit, const KVP& a, const KVP& b) + { + return (b.value < a.value || (b.value == a.value && b.key < a.key)) ? b : a; + } + DI KVP operator()(const KVP& a, const KVP& b) + { + return (b.value < a.value || (b.value == a.value && b.key < a.key)) ? b : a; + } }; // KVPMinReduce @@ -38,14 +45,16 @@ struct MinAndDistanceReduceOpImpl { DI void operator()(LabelT rid, KVP* out, const KVP& other) const { - if (other.value < out->value) { + // Use index as tiebreaker for consistent behavior when distances are equal + if (other.value < out->value || (other.value == out->value && other.key < out->key)) { out->key = other.key; out->value = other.value; } } DI void operator()(LabelT rid, volatile KVP* out, const KVP& other) const { - if (other.value < out->value) { + // Use index as tiebreaker for consistent behavior when distances are equal + if (other.value < out->value || (other.value == out->value && other.key < out->key)) { out->key = other.key; out->value = other.value; } @@ -75,7 +84,7 @@ struct MinAndDistanceReduceOpImpl { DI void init(KVP* out, DataT maxVal) const { out->value = maxVal; - out->key = 0xfffffff0; + out->key = std::numeric_limits::max(); } DI void init_key(DataT& out, LabelT idx) const { return; } diff --git a/cpp/src/distance/detail/fused_distance_nn/simt_kernel.cuh b/cpp/src/distance/detail/fused_distance_nn/simt_kernel.cuh index 4211ff653d..eccdc3de67 100644 --- a/cpp/src/distance/detail/fused_distance_nn/simt_kernel.cuh +++ b/cpp/src/distance/detail/fused_distance_nn/simt_kernel.cuh @@ -1,14 +1,15 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2021-2024, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2021-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ #pragma once -#include "../distance_ops/l2_exp.cuh" // ops::l2_exp_distance_op -#include "../pairwise_distance_base.cuh" // PairwiseDistances -#include // raft::KeyValuePair -#include // Policy +#include "../distance_ops/bitwise_hamming.cuh" // ops::bitwise_hamming_distance_op +#include "../distance_ops/l2_exp.cuh" // ops::l2_exp_distance_op +#include "../pairwise_distance_base.cuh" // PairwiseDistances +#include // raft::KeyValuePair +#include // Policy #include // size_t #include // std::numeric_limits @@ -64,111 +65,121 @@ __launch_bounds__(P::Nthreads, 2) RAFT_KERNEL fusedDistanceNNkernel(OutT* min, IdxT m, IdxT n, IdxT k, - DataT maxVal, + typename OpT::AccT maxVal, int* mutex, ReduceOpT redOp, KVPReduceOpT pairRedOp, OpT distance_op, FinalLambda fin_op) { -// compile only if below non-ampere arch. -#if __CUDA_ARCH__ < 800 - extern __shared__ char smem[]; + // For hamming-like distances, we need this kernel on all architectures + // For other distances, only use for pre-ampere architectures + +#if __CUDA_ARCH__ >= 800 + static constexpr bool compile = + std::is_same_v>; +#else + static constexpr bool compile = true; +#endif - typedef raft::KeyValuePair KVPair; - KVPair val[P::AccRowsPerTh]; -#pragma unroll - for (int i = 0; i < P::AccRowsPerTh; ++i) { - val[i] = {0, maxVal}; - } + if constexpr (compile) { + extern __shared__ char smem[]; - // epilogue operation lambda for final value calculation - auto epilog_lambda = [n, pairRedOp, &val, maxVal] __device__( - DataT acc[P::AccRowsPerTh][P::AccColsPerTh], - DataT * regxn, - DataT * regyn, - IdxT gridStrideX, - IdxT gridStrideY) { - KVPReduceOpT pairRed_op(pairRedOp); - - // intra thread reduce - const auto acccolid = threadIdx.x % P::AccThCols; - const auto accrowid = threadIdx.x / P::AccThCols; + using AccT = typename OpT::AccT; + typedef raft::KeyValuePair KVPair; + KVPair val[P::AccRowsPerTh]; #pragma unroll for (int i = 0; i < P::AccRowsPerTh; ++i) { -#pragma unroll - for (int j = 0; j < P::AccColsPerTh; ++j) { - auto tmpkey = acccolid + j * P::AccThCols + gridStrideX; - KVPair tmp = {tmpkey, acc[i][j]}; - if (tmpkey < n) { - val[i] = pairRed_op(accrowid + i * P::AccThRows + gridStrideY, tmp, val[i]); - } - } + val[i] = {0, maxVal}; } - }; - auto rowEpilog_lambda = - [m, mutex, min, pairRedOp, redOp, &val, maxVal] __device__(IdxT gridStrideY) { + // epilogue operation lambda for final value calculation + auto epilog_lambda = [n, pairRedOp, &val, maxVal] __device__( + AccT acc[P::AccRowsPerTh][P::AccColsPerTh], + AccT * regxn, + AccT * regyn, + IdxT gridStrideX, + IdxT gridStrideY) { KVPReduceOpT pairRed_op(pairRedOp); - ReduceOpT red_op(redOp); + // intra thread reduce + const auto acccolid = threadIdx.x % P::AccThCols; const auto accrowid = threadIdx.x / P::AccThCols; - const auto lid = raft::laneId(); - - // reduce #pragma unroll for (int i = 0; i < P::AccRowsPerTh; ++i) { #pragma unroll - for (int j = P::AccThCols / 2; j > 0; j >>= 1) { - // Actually, the srcLane (lid +j) should be (lid +j) % P:AccThCols, - // but the shfl op applies the modulo internally. - auto tmpkey = raft::shfl(val[i].key, lid + j, P::AccThCols); - auto tmpvalue = raft::shfl(val[i].value, lid + j, P::AccThCols); - KVPair tmp = {tmpkey, tmpvalue}; - val[i] = pairRed_op(accrowid + i * P::AccThRows + gridStrideY, tmp, val[i]); + for (int j = 0; j < P::AccColsPerTh; ++j) { + auto tmpkey = acccolid + j * P::AccThCols + gridStrideX; + KVPair tmp = {tmpkey, acc[i][j]}; + if (tmpkey < n) { + val[i] = pairRed_op(accrowid + i * P::AccThRows + gridStrideY, tmp, val[i]); + } } } + }; - updateReducedVal(mutex, min, val, red_op, m, gridStrideY); + auto rowEpilog_lambda = + [m, mutex, min, pairRedOp, redOp, &val, maxVal] __device__(IdxT gridStrideY) { + KVPReduceOpT pairRed_op(pairRedOp); + ReduceOpT red_op(redOp); - // reset the val array. + const auto accrowid = threadIdx.x / P::AccThCols; + const auto lid = raft::laneId(); + + // reduce #pragma unroll - for (int i = 0; i < P::AccRowsPerTh; ++i) { - val[i] = {0, maxVal}; - } - }; + for (int i = 0; i < P::AccRowsPerTh; ++i) { +#pragma unroll + for (int j = P::AccThCols / 2; j > 0; j >>= 1) { + // Actually, the srcLane (lid +j) should be (lid +j) % P:AccThCols, + // but the shfl op applies the modulo internally. + auto tmpkey = raft::shfl(val[i].key, lid + j, P::AccThCols); + auto tmpvalue = raft::shfl(val[i].value, lid + j, P::AccThCols); + KVPair tmp = {tmpkey, tmpvalue}; + val[i] = pairRed_op(accrowid + i * P::AccThRows + gridStrideY, tmp, val[i]); + } + } - IdxT lda = k, ldb = k, ldd = n; - constexpr bool row_major = true; - constexpr bool write_out = false; - PairwiseDistances - obj(x, - y, - m, - n, - k, - lda, - ldb, - ldd, - xn, - yn, - nullptr, // Output pointer - smem, - distance_op, - epilog_lambda, - fin_op, - rowEpilog_lambda); - obj.run(); -#endif + updateReducedVal(mutex, min, val, red_op, m, gridStrideY); + + // reset the val array. +#pragma unroll + for (int i = 0; i < P::AccRowsPerTh; ++i) { + val[i] = {0, maxVal}; + } + }; + + IdxT lda = k, ldb = k, ldd = n; + constexpr bool row_major = true; + constexpr bool write_out = false; + PairwiseDistances + obj(x, + y, + m, + n, + k, + lda, + ldb, + ldd, + reinterpret_cast(xn), + reinterpret_cast(yn), + nullptr, // Output pointer + smem, + distance_op, + epilog_lambda, + fin_op, + rowEpilog_lambda); + obj.run(); + } } } // namespace detail diff --git a/cpp/src/distance/detail/pairwise_matrix/dispatch-ext.cuh b/cpp/src/distance/detail/pairwise_matrix/dispatch-ext.cuh index c93a2f3f2b..e6878c80b2 100644 --- a/cpp/src/distance/detail/pairwise_matrix/dispatch-ext.cuh +++ b/cpp/src/distance/detail/pairwise_matrix/dispatch-ext.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ #pragma once @@ -67,6 +67,11 @@ void pairwise_matrix_dispatch(OpT distance_op, instantiate_cuvs_distance_detail_pairwise_matrix_dispatch( \ OpT, half, float, float, FinOpT, IdxT); +#define instantiate_cuvs_distance_detail_pairwise_matrix_dispatch_by_algo_bitwise_hamming(OpT, \ + IdxT) \ + instantiate_cuvs_distance_detail_pairwise_matrix_dispatch( \ + OpT, uint8_t, uint32_t, uint32_t, raft::identity_op, IdxT); + /* * Hierarchy of instantiations: * @@ -112,5 +117,9 @@ instantiate_cuvs_distance_detail_pairwise_matrix_dispatch_by_algo( instantiate_cuvs_distance_detail_pairwise_matrix_dispatch_by_algo_default( cuvs::distance::detail::ops::l2_exp_distance_op, int64_t); +instantiate_cuvs_distance_detail_pairwise_matrix_dispatch_by_algo_bitwise_hamming( + cuvs::distance::detail::ops::bitwise_hamming_distance_op, int64_t); + +#undef instantiate_cuvs_distance_detail_pairwise_matrix_dispatch_by_algo_bitwise_hamming #undef instantiate_cuvs_distance_detail_pairwise_matrix_dispatch_by_algo #undef instantiate_cuvs_distance_detail_pairwise_matrix_dispatch diff --git a/cpp/src/distance/detail/pairwise_matrix/dispatch-inl.cuh b/cpp/src/distance/detail/pairwise_matrix/dispatch-inl.cuh index aff78d87f9..691f319b72 100644 --- a/cpp/src/distance/detail/pairwise_matrix/dispatch-inl.cuh +++ b/cpp/src/distance/detail/pairwise_matrix/dispatch-inl.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ #pragma once @@ -79,6 +79,9 @@ void pairwise_matrix_dispatch(OpT distance_op, cudaStream_t stream, bool is_row_major) { + RAFT_EXPECTS(m >= 0 && n >= 0 && k >= 0, "Distance dimensions must be non-negative"); + if (m == 0 || n == 0) { return; } + // Create kernel parameter struct. Flip x and y if column major. IdxT ldx = is_row_major ? k : m; IdxT ldy = is_row_major ? k : n; diff --git a/cpp/src/distance/detail/pairwise_matrix/dispatch_matrix.json b/cpp/src/distance/detail/pairwise_matrix/dispatch_matrix.json index bf6be6bed6..d7cdb07668 100644 --- a/cpp/src/distance/detail/pairwise_matrix/dispatch_matrix.json +++ b/cpp/src/distance/detail/pairwise_matrix/dispatch_matrix.json @@ -124,5 +124,26 @@ ], "index_type": "int64_t", "index_abbrev": "i64" + }, + { + "_data": [ + { + "data_type": "uint8_t", + "data_abbrev": "u8", + "acc_type": "uint32_t", + "acc_abbrev": "u32", + "out_type": "uint32_t", + "out_abbrev": "u32" + } + ], + "_op": [ + { + "op_type": "cuvs::distance::detail::ops::bitwise_hamming_distance_op", + "op_abbrev": "bitwise_hamming", + "arch_includes": "" + } + ], + "index_type": "int64_t", + "index_abbrev": "i64" } ] diff --git a/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_epilog_matrix.json b/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_epilog_matrix.json index 3c03a381f7..91b6e15b6f 100644 --- a/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_epilog_matrix.json +++ b/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_epilog_matrix.json @@ -258,5 +258,64 @@ "index_abbrev": "i64" } ] + }, + { + "_distance": [ + { + "distance_name": "bitwise_hamming", + "distance_abbrev": "bitwise_hamming", + "op_type": "cuvs::distance::detail::ops::bitwise_hamming_distance_op", + "header_file": "distance/detail/distance_ops/bitwise_hamming.cuh" + } + ], + "_data": [ + { + "data_type": "uint8_t", + "data_abbrev": "u8", + "type_abbrev": "u8", + "acc_type": "uint32_t", + "acc_abbrev": "u32", + "out_type": "uint32_t", + "out_abbrev": "u32" + } + ], + "_index": [ + { + "index_type": "int64_t", + "index_abbrev": "i64" + } + ], + "_policy": [ + { + "policy_type": "Policy", + "layout_abbrev": "row", + "veclen": "1" + }, + { + "policy_type": "ColPolicy", + "layout_abbrev": "col", + "veclen": "1" + }, + { + "policy_type": "Policy", + "layout_abbrev": "row", + "veclen": "2" + }, + { + "policy_type": "ColPolicy", + "layout_abbrev": "col", + "veclen": "2" + }, + { + "policy_type": "Policy", + "layout_abbrev": "row", + "veclen": "4" + }, + { + "policy_type": "ColPolicy", + "layout_abbrev": "col", + "veclen": "4" + } + ] } ] diff --git a/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_matrix.json b/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_matrix.json index 0c3b0c4164..91c9a89678 100644 --- a/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_matrix.json +++ b/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/compute_distance_matrix.json @@ -151,5 +151,32 @@ "index_abbrev": "i64" } ] + }, + { + "_distance": [ + { + "distance_name": "bitwise_hamming", + "distance_abbrev": "bitwise_hamming", + "op_type": "cuvs::distance::detail::ops::bitwise_hamming_distance_op", + "header_file": "distance/detail/distance_ops/bitwise_hamming.cuh" + } + ], + "_data": [ + { + "data_type": "uint8_t", + "data_abbrev": "u8", + "type_abbrev": "u8", + "acc_type": "uint32_t", + "acc_abbrev": "u32", + "out_type": "uint32_t", + "out_abbrev": "u32" + } + ], + "_index": [ + { + "index_type": "int64_t", + "index_abbrev": "i64" + } + ] } ] diff --git a/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/pairwise_matrix_jit.cuh b/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/pairwise_matrix_jit.cuh index 17526e75b7..d8d665b112 100644 --- a/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/pairwise_matrix_jit.cuh +++ b/cpp/src/distance/detail/pairwise_matrix/jit_lto_kernels/pairwise_matrix_jit.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -35,6 +35,10 @@ constexpr auto get_pairwise_scalar_type_tag() return cuvs::neighbors::detail::tag_f{}; } else if constexpr (std::is_same_v) { return cuvs::neighbors::detail::tag_d{}; + } else if constexpr (std::is_same_v) { + return cuvs::neighbors::detail::tag_u8{}; + } else if constexpr (std::is_same_v) { + return cuvs::neighbors::detail::tag_u32{}; } else if constexpr (std::is_same_v || std::is_same_v) { return cuvs::neighbors::detail::tag_h{}; } else { @@ -84,6 +88,11 @@ struct pairwise_distance_op_tag { "Pairwise matrix JIT LTO does not have a distance-op tag for this type"); }; +template +struct pairwise_distance_op_tag> { + using type = tag_distance_bitwise_hamming; +}; + template struct pairwise_distance_op_tag> { using type = tag_distance_canberra; diff --git a/cpp/src/distance/fused_distance_nn-inl.cuh b/cpp/src/distance/fused_distance_nn-inl.cuh index bbf3bd1e68..15e9a6bb3c 100644 --- a/cpp/src/distance/fused_distance_nn-inl.cuh +++ b/cpp/src/distance/fused_distance_nn-inl.cuh @@ -69,7 +69,8 @@ namespace distance { * @param[in] initOutBuffer whether to initialize the output buffer before the * main kernel launch * @param[in] isRowMajor whether the input/output is row or column major. - * @param[in] metric Distance metric to be used (supports L2, cosine) + * @param[in] metric Distance metric to be used (supports L2, cosine, and packed uint8_t + * Hamming) * @param[in] metric_arg power argument for distances like Minkowski (not supported for now) * @param[in] handle RAFT resources containing the caller-provided CUDA stream */ @@ -92,110 +93,123 @@ void fusedDistanceNN(raft::resources const& handle, cuvs::distance::DistanceType metric, float metric_arg) { - ASSERT(isRowMajor, "fusedDistanceNN only supports row major inputs"); + RAFT_EXPECTS(isRowMajor, "fusedDistanceNN only supports row major inputs"); + RAFT_EXPECTS(m >= 0 && n >= 0 && k >= 0, "Distance dimensions must be non-negative"); + if constexpr (std::is_same_v) { + RAFT_EXPECTS(metric == cuvs::distance::DistanceType::BitwiseHamming, + "uint8_t fused distance only supports BitwiseHamming"); + RAFT_EXPECTS(static_cast(k) <= std::numeric_limits::max() / 8, + "BitwiseHamming distance exceeds the uint32_t accumulator range"); + } else { + RAFT_EXPECTS(metric == cuvs::distance::DistanceType::CosineExpanded || + metric == cuvs::distance::DistanceType::L2Expanded || + metric == cuvs::distance::DistanceType::L2SqrtExpanded, + "Floating-point fused distance only supports cosine and L2 metrics"); + } + if (m == 0) { return; } + // When k is smaller than 32, the Policy4x4 results in redundant calculations // as it uses tiles that have k=32. Therefore, use a "skinny" policy instead // that uses tiles with a smaller value of k. bool is_skinny = k < 32; + // Packed bytes use at most four elements per vectorized load. + constexpr int veclen16 = std::is_same_v ? 4 : 16 / sizeof(DataT); + constexpr int veclen8 = std::is_same_v ? 4 : 8 / sizeof(DataT); size_t bytes = sizeof(DataT) * k; auto px = reinterpret_cast(x); auto py = reinterpret_cast(y); if (16 % sizeof(DataT) == 0 && bytes % 16 == 0 && px % 16 == 0 && py % 16 == 0) { if (is_skinny) { - detail::fusedDistanceNNImpl< - DataT, - OutT, - IdxT, - typename raft::linalg::Policy4x4Skinny::Policy, - ReduceOpT>(handle, - min, - x, - y, - xn, - yn, - m, - n, - k, - (int*)workspace, - redOp, - pairRedOp, - sqrt, - initOutBuffer, - isRowMajor, - metric, - metric_arg); + detail::fusedDistanceNNImpl::Policy, + ReduceOpT>(handle, + min, + x, + y, + xn, + yn, + m, + n, + k, + (int*)workspace, + redOp, + pairRedOp, + sqrt, + initOutBuffer, + isRowMajor, + metric, + metric_arg); } else { - detail::fusedDistanceNNImpl< - DataT, - OutT, - IdxT, - typename raft::linalg::Policy4x4::Policy, - ReduceOpT>(handle, - min, - x, - y, - xn, - yn, - m, - n, - k, - (int*)workspace, - redOp, - pairRedOp, - sqrt, - initOutBuffer, - isRowMajor, - metric, - metric_arg); + detail::fusedDistanceNNImpl::Policy, + ReduceOpT>(handle, + min, + x, + y, + xn, + yn, + m, + n, + k, + (int*)workspace, + redOp, + pairRedOp, + sqrt, + initOutBuffer, + isRowMajor, + metric, + metric_arg); } } else if (8 % sizeof(DataT) == 0 && bytes % 8 == 0 && px % 8 == 0 && py % 8 == 0) { if (is_skinny) { - detail::fusedDistanceNNImpl< - DataT, - OutT, - IdxT, - typename raft::linalg::Policy4x4Skinny::Policy, - ReduceOpT>(handle, - min, - x, - y, - xn, - yn, - m, - n, - k, - (int*)workspace, - redOp, - pairRedOp, - sqrt, - initOutBuffer, - isRowMajor, - metric, - metric_arg); + detail::fusedDistanceNNImpl::Policy, + ReduceOpT>(handle, + min, + x, + y, + xn, + yn, + m, + n, + k, + (int*)workspace, + redOp, + pairRedOp, + sqrt, + initOutBuffer, + isRowMajor, + metric, + metric_arg); } else { - detail::fusedDistanceNNImpl< - DataT, - OutT, - IdxT, - typename raft::linalg::Policy4x4::Policy, - ReduceOpT>(handle, - min, - x, - y, - xn, - yn, - m, - n, - k, - (int*)workspace, - redOp, - pairRedOp, - sqrt, - initOutBuffer, - isRowMajor, - metric, - metric_arg); + detail::fusedDistanceNNImpl::Policy, + ReduceOpT>(handle, + min, + x, + y, + xn, + yn, + m, + n, + k, + (int*)workspace, + redOp, + pairRedOp, + sqrt, + initOutBuffer, + isRowMajor, + metric, + metric_arg); } } else { if (is_skinny) { @@ -273,7 +287,8 @@ void fusedDistanceNN(raft::resources const& handle, * @param[in] initOutBuffer whether to initialize the output buffer before the * main kernel launch * @param[in] isRowMajor whether the input/output is row or column major. - * @param[in] metric Distance metric to be used (supports L2, cosine) + * @param[in] metric Distance metric to be used (supports L2, cosine, and packed uint8_t + * Hamming) * @param[in] metric_arg power argument for distances like Minkowski (not supported for now) * @param[in] handle RAFT resources containing the caller-provided CUDA stream */ @@ -294,30 +309,56 @@ void fusedDistanceNNMinReduce(raft::resources const& handle, cuvs::distance::DistanceType metric, float metric_arg) { - static_assert( - std::is_same_v> || std::is_same_v, - "fusedDistanceNNMinReduce supports KVP or scalar distance output"); - detail::Top1nnTuning tuning{}; - const auto workspace_bytes = - top_1_nn_workspace_size(m, n, k, tuning, detail::Top1nnBackend::Cutlass); - top_1_nn(handle, - min, - x, - y, - xn, - yn, - m, - n, - k, - tuning, - workspace, - workspace_bytes, - sqrt, - initOutBuffer, - isRowMajor, - metric, - metric_arg, - detail::Top1nnBackend::Cutlass); + if constexpr (std::is_same_v) { + using AccT = uint32_t; + static_assert( + std::is_same_v> || std::is_same_v, + "BitwiseHamming supports uint32_t KVP or scalar distance output"); + MinAndDistanceReduceOp red_op; + KVPMinReduce pair_red_op; + fusedDistanceNN(handle, + min, + x, + y, + xn, + yn, + m, + n, + k, + workspace, + red_op, + pair_red_op, + sqrt, + initOutBuffer, + isRowMajor, + metric, + metric_arg); + } else { + static_assert( + std::is_same_v> || std::is_same_v, + "fusedDistanceNNMinReduce supports KVP or scalar distance output"); + detail::Top1nnTuning tuning{}; + const auto workspace_bytes = + top_1_nn_workspace_size(m, n, k, tuning, detail::Top1nnBackend::Cutlass); + top_1_nn(handle, + min, + x, + y, + xn, + yn, + m, + n, + k, + tuning, + workspace, + workspace_bytes, + sqrt, + initOutBuffer, + isRowMajor, + metric, + metric_arg, + detail::Top1nnBackend::Cutlass); + } } namespace detail { diff --git a/cpp/src/neighbors/detail/ann_utils.cuh b/cpp/src/neighbors/detail/ann_utils.cuh index deddc47dcb..06ba30a1a3 100644 --- a/cpp/src/neighbors/detail/ann_utils.cuh +++ b/cpp/src/neighbors/detail/ann_utils.cuh @@ -216,6 +216,18 @@ HDI constexpr auto mapping::operator()(const float& x) const -> int8_t return static_cast(std::clamp(x * 128.0f, -128.0f, 127.0f)); } +template +struct bitwise_decode_op { + explicit bitwise_decode_op(const uint8_t* binary_vecs) : binary_vecs(binary_vecs) {} + const uint8_t* binary_vecs; + HDI constexpr auto operator()(const IdxT& i) const -> OutT + { + // Rows contain complete bytes, so flattened bit offsets directly address the packed input. + // Avoid multiplying the packed dimension (or row offset) in a potentially narrow index type. + return ((binary_vecs[i >> 3] >> (i & 7)) & 1) ? OutT{1} : OutT{-1}; + } +}; + /** * @brief Sets the first num bytes of the block of memory pointed by ptr to the specified value. * diff --git a/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_impl.cuh b/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_impl.cuh index 252daaf2aa..b9ea405f2e 100644 --- a/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_impl.cuh +++ b/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_impl.cuh @@ -1,5 +1,5 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ @@ -48,4 +48,20 @@ __device__ void compute_dist_inner_product_impl(AccT& acc, AccT x, AccT y) } } +// Bitwise Hamming over byte-packed binary vectors (uint8_t data, uint32_t acc). +// `x` and `y` are already promoted to uint32_t by the load fragment: for Veclen >= 4 they pack +// 4 source bytes per accumulator word, for Veclen == 2 the upper 16 bits are zero from the +// uint16->uint32 zero-extension, and for Veclen == 1 only the low 8 bits carry a byte. +template +__device__ void compute_dist_bitwise_hamming_impl(AccT& acc, AccT x, AccT y) +{ + static_assert(std::is_same_v && std::is_same_v, + "compute_dist_bitwise_hamming_impl is only valid for uint8_t/uint32_t"); + if constexpr (Veclen > 1) { + acc += __popc(x ^ y); + } else { + acc += __popc(static_cast(x ^ y) & 0xffu); + } +} + } // namespace cuvs::neighbors::ivf_flat::detail diff --git a/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_matrix.json b/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_matrix.json index 5670ee0d84..5db309adb7 100644 --- a/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_matrix.json +++ b/cpp/src/neighbors/ivf_flat/detail/jit_lto_kernels/metric_matrix.json @@ -1,33 +1,47 @@ -{ - "metric_name": ["euclidean", "inner_product"], - "_data_type": [ - { - "data_type": "float", - "data_abbrev": "f", - "acc_type": "float", - "acc_abbrev": "f", - "veclen": ["1", "4"] - }, - { - "data_type": "__half", - "data_abbrev": "h", - "acc_type": "__half", - "acc_abbrev": "h", - "veclen": ["1", "8"] - }, - { - "data_type": "uint8_t", - "data_abbrev": "u8", - "acc_type": "uint32_t", - "acc_abbrev": "u32", - "veclen": ["1", "16"] - }, - { - "data_type": "int8_t", - "data_abbrev": "i8", - "acc_type": "int32_t", - "acc_abbrev": "i32", - "veclen": ["1", "16"] - } - ] -} +[ + { + "metric_name": ["euclidean", "inner_product"], + "_data_type": [ + { + "data_type": "float", + "data_abbrev": "f", + "acc_type": "float", + "acc_abbrev": "f", + "veclen": ["1", "4"] + }, + { + "data_type": "__half", + "data_abbrev": "h", + "acc_type": "__half", + "acc_abbrev": "h", + "veclen": ["1", "8"] + }, + { + "data_type": "uint8_t", + "data_abbrev": "u8", + "acc_type": "uint32_t", + "acc_abbrev": "u32", + "veclen": ["1", "16"] + }, + { + "data_type": "int8_t", + "data_abbrev": "i8", + "acc_type": "int32_t", + "acc_abbrev": "i32", + "veclen": ["1", "16"] + } + ] + }, + { + "metric_name": ["bitwise_hamming"], + "_data_type": [ + { + "data_type": "uint8_t", + "data_abbrev": "u8", + "acc_type": "uint32_t", + "acc_abbrev": "u32", + "veclen": ["1", "16"] + } + ] + } +] diff --git a/cpp/src/neighbors/ivf_flat/ivf_flat_build.cuh b/cpp/src/neighbors/ivf_flat/ivf_flat_build.cuh index 294ae91f5b..a8b52122ee 100644 --- a/cpp/src/neighbors/ivf_flat/ivf_flat_build.cuh +++ b/cpp/src/neighbors/ivf_flat/ivf_flat_build.cuh @@ -12,6 +12,9 @@ #include #include #include +#include +#include +#include #include "../../cluster/kmeans_balanced.cuh" #include "../detail/ann_utils.cuh" @@ -55,7 +58,13 @@ auto clone(const raft::resources& res, const index& source) -> index +RAFT_KERNEL accumulate_binary_center_counts( + const uint8_t* data, const LabelT* labels, uint32_t* counts, uint64_t n_elements, uint32_t dim) +{ + const uint64_t offset = uint64_t(blockIdx.x) * blockDim.x + threadIdx.x; + if (offset >= n_elements) { return; } + const auto row = offset / dim; + const auto col = offset % dim; + auto* output = counts + (uint64_t(labels[row]) * dim + col) * 8; + const uint8_t value = data[offset]; + for (int bit = 0; bit < 8; ++bit) { + if ((value >> bit) & 1) { atomicAdd(output + bit, 1u); } + } +} + /** See cuvs::neighbors::ivf_flat::extend docs */ template void extend(raft::resources const& handle, @@ -182,8 +207,6 @@ void extend(raft::resources const& handle, raft::make_extents(n_rows)); cuvs::cluster::kmeans::balanced_params kmeans_params; kmeans_params.metric = index->metric(); - auto orig_centroids_view = - raft::make_device_matrix_view(index->centers().data_handle(), n_lists, dim); // Calculate the batch size for the input data if it's not accessible directly from the device constexpr size_t kReasonableMaxBatchSize = 65536; size_t max_batch_size = std::min(n_rows, kReasonableMaxBatchSize); @@ -198,6 +221,7 @@ void extend(raft::resources const& handle, copy_stream = raft::resource::get_stream_from_stream_pool(handle); } } + // Predict the cluster labels for the new data, in batches if necessary auto vec_batches = utils::make_batch_load_iterator(handle, @@ -215,8 +239,26 @@ void extend(raft::resources const& handle, raft::make_device_matrix_view(batch.data(), batch.size(), index->dim()); auto batch_labels_view = raft::make_device_vector_view( new_labels.data_handle() + batch.offset(), batch.size()); - cuvs::cluster::kmeans::predict( - handle, kmeans_params, batch_data_view, orig_centroids_view, batch_labels_view); + if (index->binary_index()) { + if constexpr (std::is_same_v) { + auto centroids_view = raft::make_device_matrix_view( + index->binary_centers().data_handle(), n_lists, dim); + cuvs::cluster::kmeans::detail::predict_bitwise_hamming( + handle, batch_data_view, centroids_view, batch_labels_view); + } else { + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t data type, got %s", + typeid(T).name()); + } + } else { + auto orig_centroids_view = raft::make_device_matrix_view( + index->centers().data_handle(), n_lists, dim); + cuvs::cluster::kmeans_balanced::predict(handle, + kmeans_params, + batch_data_view, + orig_centroids_view, + batch_labels_view, + utils::mapping{}); + } vec_batches.prefetch_next_batch(); // User needs to make sure kernel finishes its work before we overwrite batch in the next // iteration if different streams are used for kernel and copy. @@ -232,23 +274,85 @@ void extend(raft::resources const& handle, // Calculate the centers and sizes on the new data, starting from the original values if (index->adaptive_centers()) { - auto centroids_view = raft::make_device_matrix_view( - index->centers().data_handle(), index->centers().extent(0), index->centers().extent(1)); auto list_sizes_view = raft::make_device_vector_view, IdxT>( list_sizes_ptr, n_lists); - for (const auto& batch : vec_batches) { - auto batch_data_view = - raft::make_device_matrix_view(batch.data(), batch.size(), index->dim()); - auto batch_labels_view = raft::make_device_vector_view( - new_labels.data_handle() + batch.offset(), batch.size()); - cuvs::cluster::kmeans_balanced::helpers::calc_centers_and_sizes(handle, - batch_data_view, - batch_labels_view, - centroids_view, - list_sizes_view, - false, - utils::mapping{}); + + if (index->binary_index()) { + if constexpr (std::is_same_v) { + // Accumulate exact per-bit one-counts, rather than rounded means or majority signs. + // Rounding a previous mean back to a sum can flip an exact majority tie. + vec_batches.reset(); + vec_batches.prefetch_next_batch(); + for (const auto& batch : vec_batches) { + const auto n_elements = uint64_t(batch.size()) * dim; + if (n_elements != 0) { + accumulate_binary_center_counts<<>>( + batch.data(), + new_labels.data_handle() + batch.offset(), + index->binary_center_counts().data_handle(), + n_elements, + dim); + RAFT_CUDA_TRY(cudaPeekAtLastError()); + } + vec_batches.prefetch_next_batch(); + if (enable_prefetch) { raft::resource::sync_stream(handle); } + } + raft::stats::histogram(raft::stats::HistTypeAuto, + reinterpret_cast(list_sizes_ptr), + IdxT(n_lists), + new_labels.data_handle(), + n_rows, + 1, + stream.get()); + raft::linalg::add( + handle, + raft::make_device_vector_view(list_sizes_ptr, n_lists), + raft::make_device_vector_view(old_list_sizes_dev.data_handle(), + n_lists), + raft::make_device_vector_view(list_sizes_ptr, n_lists)); + const auto* counts = index->binary_center_counts().data_handle(); + raft::linalg::map_offset( + handle, + index->binary_centers(), + [counts, list_sizes_ptr, dim] __device__(int64_t offset) { + uint8_t packed = 0; + for (int bit = 0; bit < 8; ++bit) { + if (uint64_t(counts[offset * 8 + bit]) * 2 > list_sizes_ptr[offset / dim]) { + packed |= uint8_t(1u << bit); + } + } + return packed; + }); + + } else { + // Error: BitwiseHamming with non-uint8_t type + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t data type, got %s", + typeid(T).name()); + } + } else { + auto centroids_view = raft::make_device_matrix_view( + index->centers().data_handle(), index->centers().extent(0), index->centers().extent(1)); + vec_batches.reset(); + vec_batches.prefetch_next_batch(); + for (const auto& batch : vec_batches) { + auto batch_data_view = + raft::make_device_matrix_view(batch.data(), batch.size(), index->dim()); + auto batch_labels_view = raft::make_device_vector_view( + new_labels.data_handle() + batch.offset(), batch.size()); + cuvs::cluster::kmeans_balanced::helpers::calc_centers_and_sizes(handle, + batch_data_view, + batch_labels_view, + centroids_view, + list_sizes_view, + false, + utils::mapping{}); + vec_batches.prefetch_next_batch(); + if (enable_prefetch) { raft::resource::sync_stream(handle); } + } } } else { raft::stats::histogram(raft::stats::HistTypeAuto, @@ -393,14 +497,22 @@ inline auto build(raft::resources const& handle, auto stream = raft::resource::get_cuda_stream(handle); cuvs::common::nvtx::range fun_scope( "ivf_flat::build(%zu, %u)", size_t(n_rows), dim); + + if (params.metric == cuvs::distance::DistanceType::BitwiseHamming && + !std::is_same_v) { + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t input type, got %s", + typeid(T).name()); + } static_assert(std::is_same_v || std::is_same_v || std::is_same_v || std::is_same_v, "unsupported data type"); + RAFT_EXPECTS(n_rows > 0 && dim > 0, "empty dataset"); RAFT_EXPECTS(n_rows >= params.n_lists, "number of rows can't be less than n_lists"); RAFT_EXPECTS(params.metric != cuvs::distance::DistanceType::CosineExpanded || dim > 1, "Cosine metric requires more than one dim"); index index(handle, params, dim); + utils::memzero( index.accum_sorted_sizes().data_handle(), index.accum_sorted_sizes().size(), stream); utils::memzero(index.list_sizes().data_handle(), index.list_sizes().size(), stream); @@ -412,6 +524,7 @@ inline auto build(raft::resources const& handle, auto trainset_ratio = std::max( 1, n_rows / std::max(params.kmeans_trainset_fraction * n_rows, index.n_lists())); auto n_rows_train = n_rows / trainset_ratio; + rmm::device_uvector trainset( n_rows_train * index.dim(), stream, raft::resource::get_large_workspace_resource_ref(handle)); // TODO: a proper sampling @@ -424,12 +537,40 @@ inline auto build(raft::resources const& handle, stream); auto trainset_const_view = raft::make_device_matrix_view(trainset.data(), n_rows_train, index.dim()); - auto centers_view = raft::make_device_matrix_view( - index.centers().data_handle(), index.n_lists(), index.dim()); + cuvs::cluster::kmeans::balanced_params kmeans_params; kmeans_params.n_iters = params.kmeans_n_iters; - kmeans_params.metric = index.metric(); - cuvs::cluster::kmeans::fit(handle, kmeans_params, trainset_const_view, centers_view); + kmeans_params.metric = + index.binary_index() ? cuvs::distance::DistanceType::L2Expanded : index.metric(); + kmeans_params.is_packed_binary = index.binary_index(); + if constexpr (std::is_same_v) { + if (index.binary_index()) { + rmm::device_uvector decoded_centers( + size_t(index.n_lists()) * index.dim() * 8, + stream, + raft::resource::get_workspace_resource_ref(handle)); + auto decoded_centers_view = raft::make_device_matrix_view( + decoded_centers.data(), index.n_lists(), IdxT(index.dim()) * 8); + + cuvs::cluster::kmeans_balanced::fit( + handle, kmeans_params, trainset_const_view, decoded_centers_view, raft::identity_op{}); + + // Convert decoded centers back to binary format + cuvs::preprocessing::quantize::binary::quantizer temp_quantizer(handle); + cuvs::preprocessing::quantize::binary::transform( + handle, temp_quantizer, decoded_centers_view, index.binary_centers()); + } else { + auto centers_view = raft::make_device_matrix_view( + index.centers().data_handle(), index.n_lists(), index.dim()); + cuvs::cluster::kmeans_balanced::fit( + handle, kmeans_params, trainset_const_view, centers_view, utils::mapping{}); + } + } else { + auto centers_view = raft::make_device_matrix_view( + index.centers().data_handle(), index.n_lists(), index.dim()); + cuvs::cluster::kmeans_balanced::fit( + handle, kmeans_params, trainset_const_view, centers_view, utils::mapping{}); + } } // add the data if necessary diff --git a/cpp/src/neighbors/ivf_flat/ivf_flat_interleaved_scan_jit.cuh b/cpp/src/neighbors/ivf_flat/ivf_flat_interleaved_scan_jit.cuh index 48aa62501c..2b98316851 100644 --- a/cpp/src/neighbors/ivf_flat/ivf_flat_interleaved_scan_jit.cuh +++ b/cpp/src/neighbors/ivf_flat/ivf_flat_interleaved_scan_jit.cuh @@ -313,6 +313,21 @@ void launch_with_fixed_consts(cuvs::distance::DistanceType metric, Args&&... arg tag_post_process_compose>( std::forward(args)...); // NB: update the description of `knn::ivf_flat::build` when // adding here a new metric. + case cuvs::distance::DistanceType::BitwiseHamming: + if constexpr (std::is_same_v) { + return launch_kernel(std::forward(args)...); + } else { + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t data type"); + } case cuvs::distance::DistanceType::CustomUDF: return launch_kernel& index, cuda::stream_ref stream, const std::optional& metric_udf) { + if (metric == cuvs::distance::DistanceType::BitwiseHamming && !std::is_same_v) { + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t data type, got %s", + typeid(T).name()); + } + const uint32_t n_probes_clamped = std::min(n_probes, index.n_lists()); const int capacity = raft::bound_by_power_of_two(k); diff --git a/cpp/src/neighbors/ivf_flat/ivf_flat_search.cuh b/cpp/src/neighbors/ivf_flat/ivf_flat_search.cuh index 6117935712..b99207775d 100644 --- a/cpp/src/neighbors/ivf_flat/ivf_flat_search.cuh +++ b/cpp/src/neighbors/ivf_flat/ivf_flat_search.cuh @@ -6,7 +6,6 @@ #pragma once #include "../../core/nvtx.hpp" -#include "../detail/ann_utils.cuh" #include "../ivf_common.cuh" // cuvs::neighbors::detail::ivf #include "ivf_flat_interleaved_scan_ext.cuh" // interleaved_scan #include // none_sample_filter @@ -26,6 +25,9 @@ #include +#include "../../distance/detail/distance_ops/bitwise_hamming.cuh" +#include "../../distance/detail/pairwise_matrix/dispatch.cuh" + #include namespace cuvs::neighbors::ivf_flat::detail { @@ -61,7 +63,8 @@ void search_impl(raft::resources const& handle, // The norm of query rmm::device_uvector query_norm_dev(n_queries, stream, search_mr); // The distance value of cluster(list) and queries - rmm::device_uvector distance_buffer_dev(n_queries * index.n_lists(), stream, search_mr); + rmm::device_uvector distance_buffer_dev( + size_t(n_queries) * index.n_lists(), stream, search_mr); // The topk distance value of cluster(list) and queries rmm::device_uvector coarse_distances_dev(n_queries_probes, stream, search_mr); // The topk index of cluster(list) and queries @@ -86,94 +89,133 @@ void search_impl(raft::resources const& handle, if constexpr (std::is_same_v) { float_query_size = 0; } else { - float_query_size = n_queries * index.dim(); + float_query_size = index.binary_index() ? 0 : size_t(n_queries) * index.dim(); } rmm::device_uvector converted_queries_dev(float_query_size, stream, search_mr); float* converted_queries_ptr = converted_queries_dev.data(); if constexpr (std::is_same_v) { converted_queries_ptr = const_cast(queries); - } else { + } else if (!index.binary_index()) { raft::linalg::map( handle, - raft::make_device_vector_view(converted_queries_ptr, n_queries * index.dim()), + raft::make_device_vector_view(converted_queries_ptr, float_query_size), utils::mapping{}, - raft::make_const_mdspan( - raft::make_device_vector_view(queries, n_queries * index.dim()))); + raft::make_const_mdspan(raft::make_device_vector_view(queries, float_query_size))); } - float alpha = 1.0f; - float beta = 0.0f; - - // todo(lsugy): raft distance? (if performance is similar/better than gemm) - switch (effective_metric) { - case cuvs::distance::DistanceType::L2Expanded: - case cuvs::distance::DistanceType::L2SqrtExpanded: { - alpha = -2.0f; - beta = 1.0f; - raft::linalg::norm( + // A custom scan metric still probes the coarse centers using the index's binary metric. + if (index.binary_index()) { + if constexpr (std::is_same_v) { + cuvs::distance::detail::ops::bitwise_hamming_distance_op distance_op{ + static_cast(index.dim())}; + + rmm::device_uvector uint32_distances( + size_t(n_queries) * index.n_lists(), stream, search_mr); + + cuvs::distance::detail::pairwise_matrix_dispatch(distance_op, + static_cast(n_queries), + static_cast(index.n_lists()), + static_cast(index.dim()), + queries, + index.binary_centers().data_handle(), + nullptr, + nullptr, + uint32_distances.data(), + raft::identity_op{}, + stream.get(), + true); + + // Cast uint32_t distances to float for the rest of the IVF-Flat pipeline. + raft::linalg::map( handle, - raft::make_device_matrix_view( - converted_queries_ptr, static_cast(n_queries), static_cast(index.dim())), - raft::make_device_vector_view(query_norm_dev.data(), - static_cast(n_queries))); - utils::outer_add(query_norm_dev.data(), - (IdxT)n_queries, - index.center_norms()->data_handle(), - (IdxT)index.n_lists(), - distance_buffer_dev.data(), - stream); - RAFT_LOG_TRACE_VEC(index.center_norms()->data_handle(), std::min(20, index.dim())); - RAFT_LOG_TRACE_VEC(distance_buffer_dev.data(), std::min(20, index.n_lists())); - break; + distance_buffer_dev_view, + raft::cast_op{}, + raft::make_const_mdspan(raft::make_device_matrix_view( + uint32_distances.data(), n_queries, index.n_lists()))); + } else { + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t data type"); } - case cuvs::distance::DistanceType::CosineExpanded: { - raft::linalg::norm( - handle, - raft::make_device_matrix_view( - converted_queries_ptr, static_cast(n_queries), static_cast(index.dim())), - raft::make_device_vector_view(query_norm_dev.data(), - static_cast(n_queries)), - raft::sqrt_op{}); - alpha = -1.0f; - beta = 0.0f; - break; - } - default: { - alpha = 1.0f; - beta = 0.0f; + } else { + float alpha = 1.0f; + float beta = 0.0f; + + // todo(lsugy): raft distance? (if performance is similar/better than gemm) + switch (effective_metric) { + case cuvs::distance::DistanceType::L2Expanded: + case cuvs::distance::DistanceType::L2SqrtExpanded: { + alpha = -2.0f; + beta = 1.0f; + raft::linalg::norm( + handle, + raft::make_device_matrix_view( + converted_queries_ptr, static_cast(n_queries), static_cast(index.dim())), + raft::make_device_vector_view(query_norm_dev.data(), + static_cast(n_queries))); + utils::outer_add(query_norm_dev.data(), + (IdxT)n_queries, + index.center_norms()->data_handle(), + (IdxT)index.n_lists(), + distance_buffer_dev.data(), + stream); + RAFT_LOG_TRACE_VEC(index.center_norms()->data_handle(), + std::min(20, index.dim())); + RAFT_LOG_TRACE_VEC(distance_buffer_dev.data(), std::min(20, index.n_lists())); + break; + } + case cuvs::distance::DistanceType::CosineExpanded: { + raft::linalg::norm( + handle, + raft::make_device_matrix_view( + converted_queries_ptr, static_cast(n_queries), static_cast(index.dim())), + raft::make_device_vector_view(query_norm_dev.data(), + static_cast(n_queries)), + raft::sqrt_op{}); + alpha = -1.0f; + beta = 0.0f; + break; + } + default: { + alpha = 1.0f; + beta = 0.0f; + } } - } - raft::linalg::gemm(handle, - true, - false, - index.n_lists(), - n_queries, - index.dim(), - &alpha, - index.centers().data_handle(), - index.dim(), - converted_queries_ptr, - index.dim(), - &beta, - distance_buffer_dev.data(), - index.n_lists(), - stream.get()); - - if (effective_metric == cuvs::distance::DistanceType::CosineExpanded) { - auto n_lists = index.n_lists(); - const auto* q_norm_ptr = query_norm_dev.data(); - const auto* index_center_norm_ptr = index.center_norms()->data_handle(); - raft::linalg::map_offset( - handle, - distance_buffer_dev_view, - [=] __device__(const uint32_t idx, const float dist) { - const auto query = idx / n_lists; - const auto cluster = idx % n_lists; - return dist / (q_norm_ptr[query] * index_center_norm_ptr[cluster]); - }, - raft::make_const_mdspan(distance_buffer_dev_view)); + raft::linalg::gemm(handle, + true, + false, + index.n_lists(), + n_queries, + index.dim(), + &alpha, + index.centers().data_handle(), + index.dim(), + converted_queries_ptr, + index.dim(), + &beta, + distance_buffer_dev.data(), + index.n_lists(), + stream.get()); + + if (effective_metric == cuvs::distance::DistanceType::CosineExpanded) { + auto n_lists = index.n_lists(); + const auto* q_norm_ptr = query_norm_dev.data(); + const auto* index_center_norm_ptr = index.center_norms()->data_handle(); + raft::linalg::map_offset( + handle, + distance_buffer_dev_view, + [=] __device__(const uint32_t idx, const float dist) { + const auto query = idx / n_lists; + const auto cluster = idx % n_lists; + return dist / (q_norm_ptr[query] * index_center_norm_ptr[cluster]); + }, + raft::make_const_mdspan(distance_buffer_dev_view)); + } } RAFT_LOG_TRACE_VEC(distance_buffer_dev.data(), std::min(20, index.n_lists())); @@ -292,7 +334,7 @@ void search_impl(raft::resources const& handle, cuvs::selection::SelectAlgo::kAuto, num_samples_vector); } - if (!manage_local_topk) { + if (!manage_local_topk && effective_metric != cuvs::distance::DistanceType::CustomUDF) { // post process distances && neighbor IDs ivf::detail::postprocess_distances( handle, distances, distances, effective_metric, n_queries, k, 1.0, false); @@ -345,9 +387,11 @@ inline void search_with_filtering(raft::resources const& handle, uint64_t max_ws_size = std::min(raft::resource::get_workspace_free_bytes(handle), kExpectedWsSize); - uint64_t ws_size_per_query = 4ull * (2 * n_probes + index.n_lists() + index.dim() + 1) + - (manage_local_topk ? ((sizeof(IdxT) + 4) * n_probes * k) - : (4ull * (max_samples + n_probes + 1))); + uint64_t ws_size_per_query = + 4ull * (2ull * n_probes + (index.binary_index() ? 2ull : 1ull) * index.n_lists() + + (index.binary_index() ? 0ull : index.dim()) + 1) + + (manage_local_topk ? ((sizeof(IdxT) + 4) * n_probes * k) + : (4ull * (max_samples + n_probes + 1))); const uint32_t max_queries = std::min(n_queries, raft::div_rounding_up_safe(max_ws_size, ws_size_per_query)); diff --git a/cpp/src/neighbors/ivf_flat/ivf_flat_serialize.cuh b/cpp/src/neighbors/ivf_flat/ivf_flat_serialize.cuh index 0dfcf5d552..e4bf4785ed 100644 --- a/cpp/src/neighbors/ivf_flat/ivf_flat_serialize.cuh +++ b/cpp/src/neighbors/ivf_flat/ivf_flat_serialize.cuh @@ -23,11 +23,12 @@ namespace cuvs::neighbors::ivf_flat::detail { // Serialization version -// No backward compatibility yet; that is, can't add additional fields without breaking -// backward compatibility. +// Version 6 combines packed binary centers and exact adaptive counts with the +// unpadded list lengths introduced in version 5. Versions 4 and 5 remain readable; +// version 4 binary indexes omitted their centers and cannot be recovered. // TODO(hcho3) Implement next-gen serializer for IVF that allows for expansion in a backward // compatible fashion. -constexpr int serialization_version = 5; +constexpr int serialization_version = 6; /** * Save the index to file. @@ -56,7 +57,14 @@ void serialize(raft::resources const& handle, Output& os, const index& serialize_scalar(handle, os, index_.metric()); serialize_scalar(handle, os, index_.adaptive_centers()); serialize_scalar(handle, os, index_.conservative_memory_allocation()); - cuvs::util::detail::serialize_mdspan(handle, os, index_.centers()); + if (index_.binary_index()) { + cuvs::util::detail::serialize_mdspan(handle, os, index_.binary_centers()); + if (index_.adaptive_centers()) { + cuvs::util::detail::serialize_mdspan(handle, os, index_.binary_center_counts()); + } + } else { + cuvs::util::detail::serialize_mdspan(handle, os, index_.centers()); + } if (index_.center_norms()) { bool has_norms = true; serialize_scalar(handle, os, has_norms); @@ -110,7 +118,7 @@ auto deserialize_impl(raft::resources const& handle, Input& input) -> index(handle, is); - if (ver != serialization_version) { + if (ver != serialization_version && ver != 5 && ver != 4) { RAFT_FAIL("serialization version mismatch, expected %d, got %d ", serialization_version, ver); } auto n_rows = raft::deserialize_scalar(handle, is); @@ -134,9 +142,19 @@ auto deserialize_impl(raft::resources const& handle, Input& input) -> index= 5 || metric != cuvs::distance::DistanceType::BitwiseHamming, + "ivf_flat::deserialize: version 4 binary indexes did not store their centers"); index index_ = index(handle, metric, n_lists, adaptive_centers, cma, dim); - cuvs::util::detail::deserialize_mdspan(handle, input, index_.centers()); + if (index_.binary_index()) { + cuvs::util::detail::deserialize_mdspan(handle, input, index_.binary_centers()); + if (index_.adaptive_centers()) { + cuvs::util::detail::deserialize_mdspan(handle, input, index_.binary_center_counts()); + } + } else { + cuvs::util::detail::deserialize_mdspan(handle, input, index_.centers()); + } bool has_norms = raft::deserialize_scalar(handle, is); if (has_norms) { index_.allocate_center_norms(handle); diff --git a/cpp/src/neighbors/ivf_flat_index.cpp b/cpp/src/neighbors/ivf_flat_index.cpp index 82e7f142bf..2b13c964c6 100644 --- a/cpp/src/neighbors/ivf_flat_index.cpp +++ b/cpp/src/neighbors/ivf_flat_index.cpp @@ -1,12 +1,27 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ +#include #include +#include #include +#include +#include +#include +#include +#include namespace cuvs::neighbors::ivf_flat { +namespace { +uint32_t expanded_binary_dim(uint32_t dim) +{ + RAFT_EXPECTS(dim <= std::numeric_limits::max() / 8, + "binary dimensionality is too large for expanded center statistics"); + return dim * 8; +} +} // namespace template index::index(raft::resources const& res) @@ -39,12 +54,33 @@ index::index(raft::resources const& res, conservative_memory_allocation_{conservative_memory_allocation}, lists_{n_lists}, list_sizes_{raft::make_device_vector(res, n_lists)}, - centers_(raft::make_device_matrix(res, n_lists, dim)), + centers_(metric != cuvs::distance::DistanceType::BitwiseHamming + ? raft::make_device_matrix(res, n_lists, dim) + : raft::make_device_matrix(res, 0, 0)), + binary_centers_(metric != cuvs::distance::DistanceType::BitwiseHamming + ? raft::make_device_matrix(res, 0, 0) + : raft::make_device_matrix(res, n_lists, dim)), + binary_center_counts_( + metric == cuvs::distance::DistanceType::BitwiseHamming && adaptive_centers + ? raft::make_device_matrix(res, n_lists, expanded_binary_dim(dim)) + : raft::make_device_matrix(res, 0, 0)), center_norms_(std::nullopt), + binary_index_(metric == cuvs::distance::DistanceType::BitwiseHamming), data_ptrs_{raft::make_device_vector(res, n_lists)}, inds_ptrs_{raft::make_device_vector(res, n_lists)}, accum_sorted_sizes_{raft::make_host_vector(n_lists + 1)} { + if (metric == cuvs::distance::DistanceType::BitwiseHamming && !std::is_same_v) { + RAFT_FAIL("BitwiseHamming distance is only supported with uint8_t data type, got %s", + typeid(T).name()); + } + + if (binary_center_counts_.size() != 0) { + RAFT_CUDA_TRY(cudaMemsetAsync(binary_center_counts_.data_handle(), + 0, + binary_center_counts_.size() * sizeof(uint32_t), + raft::resource::get_cuda_stream(res).get())); + } check_consistency(); accum_sorted_sizes_(n_lists) = 0; } @@ -92,6 +128,33 @@ raft::device_matrix_view index: return centers_.view(); } +template +raft::device_matrix_view +index::binary_centers() noexcept +{ + return binary_centers_.view(); +} + +template +raft::device_matrix_view index::binary_centers() + const noexcept +{ + return binary_centers_.view(); +} +template +raft::device_matrix_view +index::binary_center_counts() noexcept +{ + return binary_center_counts_.view(); +} + +template +raft::device_matrix_view +index::binary_center_counts() const noexcept +{ + return binary_center_counts_.view(); +} + template std::optional> index::center_norms() noexcept { @@ -136,7 +199,11 @@ IdxT index::size() const noexcept template uint32_t index::dim() const noexcept { - return centers_.extent(1); + if (binary_index_) { + return binary_centers_.extent(1); + } else { + return centers_.extent(1); + } } template @@ -210,10 +277,24 @@ void index::check_consistency() RAFT_EXPECTS(list_sizes_.extent(0) == n_lists, "inconsistent list size"); RAFT_EXPECTS(data_ptrs_.extent(0) == n_lists, "inconsistent list size"); RAFT_EXPECTS(inds_ptrs_.extent(0) == n_lists, "inconsistent list size"); - RAFT_EXPECTS( // - (centers_.extent(0) == list_sizes_.extent(0)) && // - (!center_norms_.has_value() || centers_.extent(0) == center_norms_->extent(0)), - "inconsistent number of lists (clusters)"); + if (binary_index_) { + RAFT_EXPECTS(binary_centers_.extent(0) == list_sizes_.extent(0), + "inconsistent number of lists (clusters)"); + RAFT_EXPECTS(!adaptive_centers_ || (binary_center_counts_.extent(0) == int64_t(n_lists) && + binary_center_counts_.extent(1) == int64_t(dim()) * 8), + "inconsistent binary center counts"); + } else { + RAFT_EXPECTS( // + (centers_.extent(0) == list_sizes_.extent(0)) && // + (!center_norms_.has_value() || centers_.extent(0) == center_norms_->extent(0)), + "inconsistent number of lists (clusters)"); + } +} + +template +bool index::binary_index() const noexcept +{ + return binary_index_; } template struct CUVS_EXPORT index; // Used for refine function diff --git a/cpp/tests/cluster/kmeans_balanced.cu b/cpp/tests/cluster/kmeans_balanced.cu index 777cb82b41..3d44444475 100644 --- a/cpp/tests/cluster/kmeans_balanced.cu +++ b/cpp/tests/cluster/kmeans_balanced.cu @@ -3,6 +3,7 @@ * SPDX-License-Identifier: Apache-2.0 */ +#include "../../src/cluster/detail/kmeans_balanced.cuh" #include "../test_utils.h" #include @@ -297,4 +298,199 @@ KB_TEST((KmeansBalancedTest, KmeansBalancedTestFI8I32I32_SEP, inputsf_i32); +TEST(KmeansBalancedBinary, PredictMatchesExpandedInput) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + constexpr int64_t n_rows = 97; + constexpr int64_t n_clusters = 7; + for (int64_t packed_dim : {int64_t{1}, int64_t{3}, int64_t{192}}) { + SCOPED_TRACE(packed_dim); + const int64_t dim = packed_dim * 8; + std::vector packed(n_rows * packed_dim); + std::vector expanded(n_rows * dim); + std::vector centers(n_clusters * dim); + for (size_t i = 0; i < packed.size(); ++i) { + packed[i] = uint8_t((i * 73 + i / 7 + 19) % 256); + } + for (size_t i = 0; i < expanded.size(); ++i) { + expanded[i] = ((packed[i / 8] >> (i % 8)) & 1) ? 1.0f : -1.0f; + } + for (size_t i = 0; i < centers.size(); ++i) { + centers[i] = (int((i * 19 + i / dim * 5) % 31) - 15) / 16.0f; + } + auto X = raft::make_device_matrix(handle, n_rows, packed_dim); + auto X_expanded = raft::make_device_matrix(handle, n_rows, dim); + auto C = raft::make_device_matrix(handle, n_clusters, dim); + auto labels = raft::make_device_vector(handle, n_rows); + auto expected = raft::make_device_vector(handle, n_rows); + raft::update_device(X.data_handle(), packed.data(), packed.size(), stream); + raft::update_device(X_expanded.data_handle(), expanded.data(), expanded.size(), stream); + raft::update_device(C.data_handle(), centers.data(), centers.size(), stream); + for (auto metric : {cuvs::distance::DistanceType::L2Expanded, + cuvs::distance::DistanceType::L2SqrtExpanded, + cuvs::distance::DistanceType::InnerProduct}) { + SCOPED_TRACE(int(metric)); + cuvs::cluster::kmeans::balanced_params params; + params.metric = metric; + params.is_packed_binary = true; + cuvs::cluster::kmeans::predict(handle, + params, + raft::make_const_mdspan(X.view()), + raft::make_const_mdspan(C.view()), + labels.view()); + params.is_packed_binary = false; + cuvs::cluster::kmeans::predict(handle, + params, + raft::make_const_mdspan(X_expanded.view()), + raft::make_const_mdspan(C.view()), + expected.view()); + std::vector actual_labels(n_rows), expected_labels(n_rows); + raft::update_host(actual_labels.data(), labels.data_handle(), n_rows, stream); + raft::update_host(expected_labels.data(), expected.data_handle(), n_rows, stream); + raft::resource::sync_stream(handle); + EXPECT_EQ(actual_labels, expected_labels); + } + } +} + +TEST(KmeansBalancedBinary, HierarchicalCentersPreserveEveryExpandedCoordinate) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + constexpr int64_t n_rows = 4096, packed_dim = 32, dim = 256, n_clusters = 256; + std::vector packed(n_rows * packed_dim); + for (int64_t row = 0; row < n_rows; ++row) { + // Initial mesocluster labels select the low four bits of the prototype ID. + // Within each mesocluster, fine-cluster labels select its high four bits. + // Every true centroid therefore has an exact {-1,+1} CPU oracle, with no + // dependence on centroid ordering or on private library symbols. + for (int64_t byte = 0; byte < packed_dim; ++byte) { + packed[row * packed_dim + byte] = (row & (int64_t{1} << (byte / 4))) ? 0xff : 0; + } + } + auto X = raft::make_device_matrix(handle, n_rows, packed_dim); + auto C = raft::make_device_matrix(handle, n_clusters, dim); + auto labels = raft::make_device_vector(handle, n_rows); + raft::update_device(X.data_handle(), packed.data(), packed.size(), stream); + raft::matrix::fill(handle, C.view(), 123.0f); + cuvs::cluster::kmeans::balanced_params params; + params.is_packed_binary = true; + params.n_iters = 1; + cuvs::cluster::kmeans::fit(handle, params, raft::make_const_mdspan(X.view()), C.view()); + cuvs::cluster::kmeans::predict(handle, + params, + raft::make_const_mdspan(X.view()), + raft::make_const_mdspan(C.view()), + labels.view()); + std::vector centers(n_clusters * dim); + std::vector actual_labels(n_rows); + raft::update_host(centers.data(), C.data_handle(), centers.size(), stream); + raft::update_host(actual_labels.data(), labels.data_handle(), actual_labels.size(), stream); + raft::resource::sync_stream(handle); + std::vector seen(n_clusters, false); + for (int64_t row = 0; row < n_rows; ++row) { + ASSERT_LT(actual_labels[row], n_clusters); + seen[actual_labels[row]] = true; + for (int64_t bit = 0; bit < dim; ++bit) { + const auto byte = packed[row * packed_dim + bit / 8]; + ASSERT_EQ(centers[actual_labels[row] * dim + bit], ((byte >> (bit % 8)) & 1) ? 1.0f : -1.0f) + << "row=" << row << ", bit=" << bit; + } + } + EXPECT_EQ(std::count(seen.begin(), seen.end(), true), n_clusters); +} + +TEST(KmeansBalancedBinary, FitAndPredictRecoverPackedClusters) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + constexpr int64_t n_rows = 128, packed_dim = 3, dim = 24, n_clusters = 4; + std::vector packed(n_rows * packed_dim); + for (int64_t row = 0; row < n_rows; ++row) { + packed[row * packed_dim] = (row & 1) ? 0xff : 0; + packed[row * packed_dim + 1] = (row & 2) ? 0xff : 0; + packed[row * packed_dim + 2] = (row & 2) ? 0xa5 : 0x5a; + } + auto X = raft::make_device_matrix(handle, n_rows, packed_dim); + auto C = raft::make_device_matrix(handle, n_clusters, dim); + auto labels = raft::make_device_vector(handle, n_rows); + raft::update_device(X.data_handle(), packed.data(), packed.size(), stream); + cuvs::cluster::kmeans::balanced_params params; + params.is_packed_binary = true; + params.n_iters = 2; + cuvs::cluster::kmeans::fit(handle, params, raft::make_const_mdspan(X.view()), C.view()); + cuvs::cluster::kmeans::predict(handle, + params, + raft::make_const_mdspan(X.view()), + raft::make_const_mdspan(C.view()), + labels.view()); + std::vector actual_labels(n_rows); + std::vector centers(n_clusters * dim); + raft::update_host(actual_labels.data(), labels.data_handle(), n_rows, stream); + raft::update_host(centers.data(), C.data_handle(), centers.size(), stream); + raft::resource::sync_stream(handle); + for (int64_t row = 0; row < n_rows; ++row) { + ASSERT_LT(actual_labels[row], n_clusters); + for (int64_t bit = 0; bit < dim; ++bit) { + const auto byte = packed[row * packed_dim + bit / 8]; + EXPECT_EQ(centers[actual_labels[row] * dim + bit], ((byte >> (bit % 8)) & 1) ? 1.0f : -1.0f); + } + } +} + +TEST(KmeansBalancedBinary, IncrementalCentersMatchSinglePass) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + constexpr int64_t n_rows = 9, packed_dim = 3, dim = 24, n_clusters = 3; + std::vector packed(n_rows * packed_dim); + std::vector labels(n_rows); + for (size_t i = 0; i < packed.size(); ++i) { + packed[i] = uint8_t(i * 73 + 19); + } + for (int64_t row = 0; row < n_rows; ++row) { + labels[row] = row % n_clusters; + } + auto X = raft::make_device_matrix(handle, n_rows, packed_dim); + auto L = raft::make_device_vector(handle, n_rows); + auto C = raft::make_device_matrix(handle, n_clusters, dim); + auto sizes = raft::make_device_vector(handle, n_clusters); + raft::update_device(X.data_handle(), packed.data(), packed.size(), stream); + raft::update_device(L.data_handle(), labels.data(), labels.size(), stream); + auto update = [&](int64_t offset, int64_t rows, bool reset) { + cuvs::cluster::kmeans::detail::calc_centers_and_sizes( + handle, + C.data_handle(), + sizes.data_handle(), + n_clusters, + packed_dim, + X.data_handle() + offset * packed_dim, + rows, + L.data_handle() + offset, + reset, + true, + raft::identity_op{}, + raft::resource::get_workspace_resource_ref(handle)); + }; + update(0, 4, true); + update(4, 5, false); + std::vector actual(n_clusters * dim); + std::vector actual_sizes(n_clusters); + raft::update_host(actual.data(), C.data_handle(), actual.size(), stream); + raft::update_host(actual_sizes.data(), sizes.data_handle(), actual_sizes.size(), stream); + raft::resource::sync_stream(handle); + for (int64_t cluster = 0; cluster < n_clusters; ++cluster) { + EXPECT_EQ(actual_sizes[cluster], 3); + for (int64_t bit = 0; bit < dim; ++bit) { + float expected = 0; + for (int64_t row = cluster; row < n_rows; row += n_clusters) { + const auto byte = packed[row * packed_dim + bit / 8]; + expected += ((byte >> (bit % 8)) & 1) ? 1.0f : -1.0f; + } + EXPECT_NEAR(actual[cluster * dim + bit], expected / 3.0f, 1e-6f); + } + } +} + } // namespace cuvs diff --git a/cpp/tests/neighbors/ann_ivf_flat.cuh b/cpp/tests/neighbors/ann_ivf_flat.cuh index 47a721bf2f..2b0f36cb2b 100644 --- a/cpp/tests/neighbors/ann_ivf_flat.cuh +++ b/cpp/tests/neighbors/ann_ivf_flat.cuh @@ -18,6 +18,8 @@ #include #include +#include +#include #include #include #include @@ -67,6 +69,11 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { void testIVFFlat() { + if ((ps.metric == cuvs::distance::DistanceType::BitwiseHamming) && + !(std::is_same_v)) { + GTEST_SKIP(); + } + size_t queries_size = ps.num_queries * ps.k; std::vector indices_ivfflat(queries_size); std::vector indices_naive(queries_size); @@ -190,6 +197,27 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { cuvs::neighbors::ivf_flat::index index_loaded(handle_); cuvs::neighbors::ivf_flat::deserialize(handle_, index_file.filename, &index_loaded); ASSERT_EQ(index_2.size(), index_loaded.size()); + if (index_2.binary_index()) { + ASSERT_TRUE(cuvs::devArrMatch(index_2.binary_centers().data_handle(), + index_loaded.binary_centers().data_handle(), + index_2.binary_centers().size(), + cuvs::Compare(), + stream_.get())); + } + if (index_2.binary_index() && index_2.adaptive_centers()) { + ASSERT_TRUE(cuvs::devArrMatch(index_2.binary_center_counts().data_handle(), + index_loaded.binary_center_counts().data_handle(), + index_2.binary_center_counts().size(), + cuvs::Compare(), + stream_.get())); + } + if (!index_2.binary_index()) { + ASSERT_TRUE(cuvs::devArrMatch(index_2.centers().data_handle(), + index_loaded.centers().data_handle(), + index_2.centers().size(), + cuvs::Compare(), + stream_.get())); + } cuvs::neighbors::ivf_flat::search(handle_, search_params, @@ -206,41 +234,53 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { // Test the centroid invariants if (index_2.adaptive_centers()) { - // The centers must be up-to-date with the corresponding data - std::vector list_sizes(index_2.n_lists()); - std::vector list_indices(index_2.n_lists()); - rmm::device_uvector centroid(ps.dim, stream_); - raft::copy( - list_sizes.data(), index_2.list_sizes().data_handle(), index_2.n_lists(), stream_); - raft::copy( - list_indices.data(), index_2.inds_ptrs().data_handle(), index_2.n_lists(), stream_); - raft::resource::sync_stream(handle_); - for (uint32_t l = 0; l < index_2.n_lists(); l++) { - if (list_sizes[l] == 0) continue; - rmm::device_uvector cluster_data(list_sizes[l] * ps.dim, stream_); - cuvs::spatial::knn::detail::utils::copy_selected((IdxT)list_sizes[l], - (IdxT)ps.dim, - database.data(), - list_indices[l], - (IdxT)ps.dim, - cluster_data.data(), - (IdxT)ps.dim, - stream_); - raft::stats::mean( - centroid.data(), cluster_data.data(), ps.dim, list_sizes[l], false, stream_.get()); - ASSERT_TRUE(cuvs::devArrMatch(index_2.centers().data_handle() + ps.dim * l, - centroid.data(), - ps.dim, - cuvs::CompareApprox(0.001), - stream_.get())); + // Skip centroid verification for BitwiseHamming metric + if (ps.metric != cuvs::distance::DistanceType::BitwiseHamming) { + // The centers must be up-to-date with the corresponding data + std::vector list_sizes(index_2.n_lists()); + std::vector list_indices(index_2.n_lists()); + rmm::device_uvector centroid(ps.dim, stream_); + raft::copy( + list_sizes.data(), index_2.list_sizes().data_handle(), index_2.n_lists(), stream_); + raft::copy( + list_indices.data(), index_2.inds_ptrs().data_handle(), index_2.n_lists(), stream_); + raft::resource::sync_stream(handle_); + for (uint32_t l = 0; l < index_2.n_lists(); l++) { + if (list_sizes[l] == 0) continue; + rmm::device_uvector cluster_data(list_sizes[l] * ps.dim, stream_); + cuvs::spatial::knn::detail::utils::copy_selected((IdxT)list_sizes[l], + (IdxT)ps.dim, + database.data(), + list_indices[l], + (IdxT)ps.dim, + cluster_data.data(), + (IdxT)ps.dim, + stream_); + raft::stats::mean( + centroid.data(), cluster_data.data(), ps.dim, list_sizes[l], false, stream_.get()); + ASSERT_TRUE(cuvs::devArrMatch(index_2.centers().data_handle() + ps.dim * l, + centroid.data(), + ps.dim, + cuvs::CompareApprox(0.001), + stream_.get())); + } } } else { // The centers must be immutable - ASSERT_TRUE(cuvs::devArrMatch(index_2.centers().data_handle(), - idx.centers().data_handle(), - index_2.centers().size(), - cuvs::Compare(), - stream_.get())); + if (ps.metric == cuvs::distance::DistanceType::BitwiseHamming) { + // For BitwiseHamming, compare binary centers + ASSERT_TRUE(cuvs::devArrMatch(index_2.binary_centers().data_handle(), + idx.binary_centers().data_handle(), + index_2.binary_centers().size(), + cuvs::Compare(), + stream_.get())); + } else { + ASSERT_TRUE(cuvs::devArrMatch(index_2.centers().data_handle(), + idx.centers().data_handle(), + index_2.centers().size(), + cuvs::Compare(), + stream_.get())); + } } } float eps = std::is_same_v ? 0.005 : 0.001; @@ -257,6 +297,11 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { void testPacker() { + if ((ps.metric == cuvs::distance::DistanceType::BitwiseHamming) && + !(std::is_same_v)) { + GTEST_SKIP(); + } + ivf_flat::index_params index_params; ivf_flat::search_params search_params; index_params.n_lists = ps.nlist; @@ -330,13 +375,11 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { [dim = idx.dim(), list_size, padded_list_size, - chunk_size = raft::util::FastIntDiv( - static_cast(idx.veclen()))] __device__(auto i) { + chunk_size = raft::util::FastIntDiv(idx.veclen())] __device__(auto i) { uint32_t max_group_offset = interleaved_group::roundDown(list_size); if (i < max_group_offset * dim) { return true; } - uint32_t surplus = (i - max_group_offset * dim); - uint32_t ingroup_id = - interleaved_group::mod(static_cast(surplus) / chunk_size); + uint32_t surplus = (i - max_group_offset * dim); + uint32_t ingroup_id = interleaved_group::mod(int64_t(surplus) / chunk_size); return ingroup_id < (list_size - max_group_offset); }); @@ -391,6 +434,11 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { void testFilter() { + if ((ps.metric == cuvs::distance::DistanceType::BitwiseHamming) && + !(std::is_same_v)) { + GTEST_SKIP(); + } + size_t queries_size = ps.num_queries * ps.k; std::vector indices_ivfflat(queries_size); std::vector indices_naive(queries_size); @@ -498,6 +546,12 @@ class AnnIVFFlatTest : public ::testing::TestWithParam> { handle_, r, database.data(), ps.num_db_vecs * ps.dim, DataT(0.1), DataT(2.0)); raft::random::uniform( handle_, r, search_queries.data(), ps.num_queries * ps.dim, DataT(0.1), DataT(2.0)); + } else if (ps.metric == cuvs::distance::DistanceType::BitwiseHamming && + std::is_same_v) { + raft::random::uniformInt( + handle_, r, database.data(), ps.num_db_vecs * ps.dim, DataT(0), DataT(255)); + raft::random::uniformInt( + handle_, r, search_queries.data(), ps.num_queries * ps.dim, DataT(0), DataT(255)); } else { raft::random::uniformInt( handle_, r, database.data(), ps.num_db_vecs * ps.dim, DataT(1), DataT(20)); @@ -527,14 +581,19 @@ const std::vector> inputs = { {1000, 10000, 1, 16, 40, 1024, cuvs::distance::DistanceType::L2Expanded, true}, {1000, 10000, 2, 16, 40, 1024, cuvs::distance::DistanceType::L2Expanded, false}, {1000, 10000, 2, 16, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, 10000, 2, 16, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 3, 16, 40, 1024, cuvs::distance::DistanceType::L2Expanded, true}, {1000, 10000, 3, 16, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, true}, + {1000, 10000, 3, 16, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 4, 16, 40, 1024, cuvs::distance::DistanceType::L2Expanded, false}, {1000, 10000, 4, 16, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, 10000, 4, 16, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 5, 16, 40, 1024, cuvs::distance::DistanceType::InnerProduct, false}, {1000, 10000, 5, 16, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, 10000, 5, 16, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 8, 16, 40, 1024, cuvs::distance::DistanceType::InnerProduct, true}, {1000, 10000, 8, 16, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, true}, + {1000, 10000, 8, 16, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 5, 16, 40, 1024, cuvs::distance::DistanceType::L2SqrtExpanded, false}, {1000, 10000, 5, 16, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, {1000, 10000, 8, 16, 40, 1024, cuvs::distance::DistanceType::L2SqrtExpanded, true}, @@ -561,50 +620,70 @@ const std::vector> inputs = { // various random combinations {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::L2Expanded, false}, {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::L2Expanded, false}, {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::L2Expanded, false}, {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::L2Expanded, false}, {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::CosineExpanded, false}, + {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::BitwiseHamming, false}, {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, true}, {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, true}, {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, true}, {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, true}, + {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false}, {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false}, + {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false}, // host input data {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::L2Expanded, false, true}, {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true}, {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::L2Expanded, false, true}, {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true}, {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::L2Expanded, false, true}, {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true}, {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::L2Expanded, false, true}, {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::BitwiseHamming, false, true}, {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false, true}, {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true}, {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false, true}, {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true}, {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false, true}, {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true}, + {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true}, // // host input data with prefetching for kernel copy overlapping {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {1000, 10000, 16, 10, 50, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {1000, 10000, 16, 10, 70, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {100, 10000, 16, 10, 20, 512, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {20, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {1000, 100000, 16, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::L2Expanded, false, true, true}, {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::CosineExpanded, false, true, true}, + {10000, 131072, 8, 10, 20, 1024, cuvs::distance::DistanceType::BitwiseHamming, false, true, true}, {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::InnerProduct, true}, {1000, 10000, 16, 10, 40, 1024, cuvs::distance::DistanceType::CosineExpanded, true}, @@ -627,10 +706,13 @@ const std::vector> inputs = { // test splitting the big query batches (> max gridDim.y) into smaller batches {100000, 1024, 32, 10, 64, 64, cuvs::distance::DistanceType::InnerProduct, false}, {100000, 1024, 32, 10, 64, 64, cuvs::distance::DistanceType::CosineExpanded, false}, + {100000, 1024, 32, 10, 64, 64, cuvs::distance::DistanceType::BitwiseHamming, false}, {1000000, 1024, 32, 10, 256, 256, cuvs::distance::DistanceType::InnerProduct, false}, {1000000, 1024, 32, 10, 256, 256, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000000, 1024, 32, 10, 256, 256, cuvs::distance::DistanceType::BitwiseHamming, false}, {98306, 1024, 32, 10, 64, 64, cuvs::distance::DistanceType::InnerProduct, true}, {98306, 1024, 32, 10, 64, 64, cuvs::distance::DistanceType::CosineExpanded, true}, + {98306, 1024, 32, 10, 64, 64, cuvs::distance::DistanceType::BitwiseHamming, false}, // test radix_sort for getting the cluster selection {1000, @@ -657,10 +739,24 @@ const std::vector> inputs = { raft::matrix::detail::select::warpsort::kMaxCapacity * 4, cuvs::distance::DistanceType::CosineExpanded, false}, + {1000, + 10000, + 16, + 10, + raft::matrix::detail::select::warpsort::kMaxCapacity * 4, + raft::matrix::detail::select::warpsort::kMaxCapacity * 4, + cuvs::distance::DistanceType::BitwiseHamming, + false}, // The following two test cases should show very similar recall. // num_queries, num_db_vecs, dim, k, nprobe, nlist, metric, adaptive_centers {20000, 8712, 3, 10, 51, 66, cuvs::distance::DistanceType::L2Expanded, false}, - {100000, 8712, 3, 10, 51, 66, cuvs::distance::DistanceType::L2Expanded, false}}; + {100000, 8712, 3, 10, 51, 66, cuvs::distance::DistanceType::L2Expanded, false}, + + // BitwiseHamming with adaptive centers + {1000, 10000, 32, 16, 20, 80, cuvs::distance::DistanceType::BitwiseHamming, true}, + {1000, 10000, 64, 16, 20, 80, cuvs::distance::DistanceType::BitwiseHamming, true}, + {1000, 10000, 128, 16, 20, 80, cuvs::distance::DistanceType::BitwiseHamming, true}, +}; } // namespace cuvs::neighbors::ivf_flat diff --git a/cpp/tests/neighbors/ann_ivf_flat/test_uint8_t_int64_t.cu b/cpp/tests/neighbors/ann_ivf_flat/test_uint8_t_int64_t.cu index 4a58d0f096..bd23c0ce13 100644 --- a/cpp/tests/neighbors/ann_ivf_flat/test_uint8_t_int64_t.cu +++ b/cpp/tests/neighbors/ann_ivf_flat/test_uint8_t_int64_t.cu @@ -1,14 +1,21 @@ /* - * SPDX-FileCopyrightText: Copyright (c) 2024, NVIDIA CORPORATION. + * SPDX-FileCopyrightText: Copyright (c) 2024-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. * SPDX-License-Identifier: Apache-2.0 */ #include +#include + +#include +#include +#include #include "../ann_ivf_flat.cuh" namespace cuvs::neighbors::ivf_flat { +CUVS_METRIC(binary_hamming, { acc += __popc(static_cast(x.raw() ^ y.raw())); }) + typedef AnnIVFFlatTest AnnIVFFlatTestF_uint8; TEST_P(AnnIVFFlatTestF_uint8, AnnIVFFlat) { @@ -19,4 +26,318 @@ TEST_P(AnnIVFFlatTestF_uint8, AnnIVFFlat) INSTANTIATE_TEST_CASE_P(AnnIVFFlatTest, AnnIVFFlatTestF_uint8, ::testing::ValuesIn(inputs)); +// Check exact distances rather than tied IDs: Hamming's kth neighbor often has many ties. +TEST(BinaryIvfFlatRegression, ExhaustiveScalarAndVectorizedScan) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + constexpr int64_t n_rows = 1536; + constexpr int64_t n_queries = 3; + std::mt19937 rng(42); + for (int64_t dim : {1, 3, 16, 192}) { + SCOPED_TRACE(dim); + std::vector data(n_rows * dim); + std::vector queries(n_queries * dim); + for (auto& x : data) { + x = static_cast(rng()); + } + for (auto& x : queries) { + x = static_cast(rng()); + } + auto data_dev = raft::make_device_matrix(handle, n_rows, dim); + auto queries_dev = raft::make_device_matrix(handle, n_queries, dim); + raft::update_device(data_dev.data_handle(), data.data(), data.size(), stream); + raft::update_device(queries_dev.data_handle(), queries.data(), queries.size(), stream); + index_params params; + params.metric = cuvs::distance::DistanceType::BitwiseHamming; + params.n_lists = 16; + params.kmeans_n_iters = 3; + params.kmeans_trainset_fraction = 1.0; + auto idx = build(handle, params, raft::make_const_mdspan(data_dev.view())); + ASSERT_EQ(idx.dim(), dim); + ASSERT_EQ(idx.size(), n_rows); + ASSERT_TRUE(idx.binary_index()); + ASSERT_EQ(idx.centers().size(), 0); + + // Verify packed-center persistence independently of the search itself. + std::stringstream bytes; + serialize(handle, bytes, idx); + index restored(handle); + deserialize(handle, bytes, &restored); + ASSERT_TRUE(cuvs::devArrMatch(idx.binary_centers().data_handle(), + restored.binary_centers().data_handle(), + idx.binary_centers().size(), + cuvs::Compare(), + stream.get())); + index moved(std::move(restored)); + restored = std::move(moved); + search_params search_params; + search_params.n_probes = params.n_lists; + for (int64_t k : {10, 100, 1025, 11, 1026}) { + SCOPED_TRACE(k); + search_params.metric_udf = + (k == 11 || k == 1026) ? std::make_optional(binary_hamming_udf()) : std::nullopt; + auto ids_dev = raft::make_device_matrix(handle, n_queries, k); + auto distances_dev = raft::make_device_matrix(handle, n_queries, k); + search(handle, + search_params, + restored, + raft::make_const_mdspan(queries_dev.view()), + ids_dev.view(), + distances_dev.view()); + std::vector ids(n_queries * k); + std::vector distances(n_queries * k); + raft::update_host(ids.data(), ids_dev.data_handle(), ids.size(), stream); + raft::update_host(distances.data(), distances_dev.data_handle(), distances.size(), stream); + raft::resource::sync_stream(handle); + for (int64_t q = 0; q < n_queries; ++q) { + std::vector expected(n_rows); + for (int64_t row = 0; row < n_rows; ++row) { + for (int64_t col = 0; col < dim; ++col) { + expected[row] += + __builtin_popcount(unsigned(data[row * dim + col] ^ queries[q * dim + col])); + } + } + auto sorted = expected; + std::sort(sorted.begin(), sorted.end()); + std::vector seen(n_rows); + auto actual = + std::vector(distances.begin() + q * k, distances.begin() + (q + 1) * k); + std::sort(actual.begin(), actual.end()); + for (int64_t j = 0; j < k; ++j) { + auto id = ids[q * k + j]; + ASSERT_GE(id, 0); + ASSERT_LT(id, n_rows); + ASSERT_FALSE(seen[id]); + seen[id] = true; + ASSERT_EQ(distances[q * k + j], expected[id]); + ASSERT_EQ(actual[j], sorted[j]); + } + } + } + } +} + +TEST(BinaryIvfFlatRegression, AdaptiveCountsSurviveCloneAndSerialization) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + const std::vector initial{1, 1, 1, 0}; + auto data = raft::make_device_matrix(handle, 4, 1); + raft::update_device(data.data_handle(), initial.data(), initial.size(), stream); + index_params params; + params.metric = cuvs::distance::DistanceType::BitwiseHamming; + params.n_lists = 1; + params.adaptive_centers = true; + params.kmeans_n_iters = 3; + params.kmeans_trainset_fraction = 1.0; + auto idx = build(handle, params, raft::make_const_mdspan(data.view())); + std::stringstream bytes; + serialize(handle, bytes, idx); + index restored(handle); + deserialize(handle, bytes, &restored); + ASSERT_EQ(restored.binary_center_counts().extent(1), 8); + ASSERT_TRUE(cuvs::devArrMatch(idx.binary_center_counts().data_handle(), + restored.binary_center_counts().data_handle(), + idx.binary_center_counts().size(), + cuvs::Compare(), + stream.get())); + auto more = raft::make_device_matrix(handle, 3, 1); + auto ids = raft::make_device_vector(handle, 3); + const std::vector more_host{0, 0, 0}; + const std::vector ids_host{4, 5, 6}; + raft::update_device(more.data_handle(), more_host.data(), more_host.size(), stream); + raft::update_device(ids.data_handle(), ids_host.data(), ids_host.size(), stream); + const std::optional> ids_view = + raft::make_const_mdspan(ids.view()); + // The returned-index overload exercises cloning of both center representations. + auto extended = extend(handle, raft::make_const_mdspan(more.view()), ids_view, restored); + uint8_t center; + uint8_t old_center; + uint32_t count; + raft::update_host(¢er, extended.binary_centers().data_handle(), 1, stream); + raft::update_host(&old_center, restored.binary_centers().data_handle(), 1, stream); + raft::update_host(&count, extended.binary_center_counts().data_handle(), 1, stream); + raft::resource::sync_stream(handle); + EXPECT_EQ(extended.size(), 7); + EXPECT_EQ(restored.size(), 4); + EXPECT_EQ(old_center, 1); + EXPECT_EQ(center, 0); // Three one bits out of seven, not a majority. + EXPECT_EQ(count, 3); +} + +TEST(BinaryIvfFlatRegression, RejectNonByteInputTypes) +{ + raft::resources handle; + index_params params; + params.metric = cuvs::distance::DistanceType::BitwiseHamming; + params.n_lists = 1; + EXPECT_THROW((index(handle, params, 16)), raft::logic_error); + EXPECT_THROW((index(handle, params, 16)), raft::logic_error); + EXPECT_THROW((index(handle, params, 16)), raft::logic_error); + params.adaptive_centers = true; + EXPECT_THROW((index(handle, params, std::numeric_limits::max())), + raft::logic_error); +} + +TEST(BinaryIvfFlatRegression, AdaptiveHostBatchesWithCopyStream) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + raft::resource::set_cuda_stream_pool(handle, std::make_shared(1)); + constexpr int64_t n_rows = 65539; // Cross the host staging boundary of 65536 rows. + auto data = raft::make_host_matrix(n_rows, 3); + for (int64_t row = 0; row < n_rows; ++row) { + data(row, 0) = row % 3 == 0 ? 0xff : 0; + data(row, 1) = row % 3 == 0 ? 0 : 0xff; + data(row, 2) = row % 2 == 0 ? 0x55 : 0xaa; + } + index_params params; + params.metric = cuvs::distance::DistanceType::BitwiseHamming; + params.n_lists = 1; + params.adaptive_centers = true; + params.kmeans_n_iters = 3; + auto idx = build(handle, params, raft::make_const_mdspan(data.view())); + std::vector center(3); + std::vector counts(24); + raft::update_host(center.data(), idx.binary_centers().data_handle(), center.size(), stream); + raft::update_host(counts.data(), idx.binary_center_counts().data_handle(), counts.size(), stream); + raft::resource::sync_stream(handle); + EXPECT_EQ(idx.size(), n_rows); + for (int64_t col = 0; col < 3; ++col) { + uint8_t expected = 0; + for (int bit = 0; bit < 8; ++bit) { + int sum = 0; + for (int64_t row = 0; row < n_rows; ++row) { + sum += ((data(row, col) >> bit) & 1) ? 1 : -1; + } + if (sum > 0) { expected |= uint8_t(1 << bit); } + EXPECT_EQ(counts[col * 8 + bit], (sum + n_rows) / 2); + } + EXPECT_EQ(center[col], expected); + } +} + +TEST(BinaryIvfFlatRegression, AdaptiveExactTieAfterTrainingOnly) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + // mean = 7 / 13 rounds upward in float; reconstructing its sum would make a tie positive. + std::vector initial(13); + std::fill(initial.begin(), initial.begin() + 10, 1); + auto data = raft::make_device_matrix(handle, 13, 1); + raft::update_device(data.data_handle(), initial.data(), initial.size(), stream); + index_params params; + params.metric = cuvs::distance::DistanceType::BitwiseHamming; + params.n_lists = 1; + params.adaptive_centers = true; + params.add_data_on_build = false; + params.kmeans_n_iters = 3; + params.kmeans_trainset_fraction = 1.0; + auto idx = build(handle, params, raft::make_const_mdspan(data.view())); + std::vector counts(8); + raft::update_host(counts.data(), idx.binary_center_counts().data_handle(), counts.size(), stream); + raft::resource::sync_stream(handle); + EXPECT_EQ(idx.size(), 0); + EXPECT_EQ(counts, std::vector(8, 0)); + const std::optional> no_ids = std::nullopt; + extend(handle, raft::make_const_mdspan(data.view()), no_ids, &idx); + auto more = raft::make_device_matrix(handle, 7, 1); + auto ids = raft::make_device_vector(handle, 7); + const std::vector more_host(7, 0); + const std::vector ids_host{13, 14, 15, 16, 17, 18, 19}; + raft::update_device(more.data_handle(), more_host.data(), more_host.size(), stream); + raft::update_device(ids.data_handle(), ids_host.data(), ids_host.size(), stream); + const std::optional> ids_view = + raft::make_const_mdspan(ids.view()); + extend(handle, raft::make_const_mdspan(more.view()), ids_view, &idx); + uint8_t center; + raft::update_host(¢er, idx.binary_centers().data_handle(), 1, stream); + raft::update_host(counts.data(), idx.binary_center_counts().data_handle(), counts.size(), stream); + raft::resource::sync_stream(handle); + EXPECT_EQ(idx.size(), 20); + EXPECT_EQ(counts[0], 10); + EXPECT_EQ(center, 0); // Exact ties use the same > 0 convention as binary quantization. +} + +TEST(BinaryIvfFlatRegression, LegacySerializationCompatibility) +{ + raft::resources handle; + auto stream = raft::resource::get_cuda_stream(handle); + const std::vector data_host{1, 2, 3, 4}; + auto data = raft::make_device_matrix(handle, 4, 1); + raft::update_device(data.data_handle(), data_host.data(), data_host.size(), stream); + for (auto metric : + {cuvs::distance::DistanceType::L2Expanded, cuvs::distance::DistanceType::BitwiseHamming}) { + index_params params; + params.metric = metric; + params.n_lists = 1; + params.kmeans_n_iters = 3; + params.kmeans_trainset_fraction = 1.0; + auto idx = build(handle, params, raft::make_const_mdspan(data.view())); + std::stringstream serialized; + serialize(handle, serialized, idx); + char dtype[4]; + serialized.read(dtype, 4); + EXPECT_EQ(raft::deserialize_scalar(handle, serialized), 6); + for (int version : {4, 5}) { + // Version 4 and the original binary version 5 stored padded list lengths and IDs. + // Upstream's nonbinary version 5 stores the actual list length with padded vector data. + const uint32_t stored_size = version == 4 || idx.binary_index() ? 32 : 4; + std::stringstream legacy; + legacy.write(dtype, 4); + raft::serialize_scalar(handle, legacy, version); + raft::serialize_scalar(handle, legacy, idx.size()); + raft::serialize_scalar(handle, legacy, idx.dim()); + raft::serialize_scalar(handle, legacy, idx.n_lists()); + raft::serialize_scalar(handle, legacy, idx.metric()); + raft::serialize_scalar(handle, legacy, idx.adaptive_centers()); + raft::serialize_scalar(handle, legacy, idx.conservative_memory_allocation()); + index restored(handle); + if (version == 4 && idx.binary_index()) { + EXPECT_THROW(deserialize(handle, legacy, &restored), raft::logic_error); + continue; + } + if (idx.binary_index()) { + raft::serialize_mdspan(handle, legacy, idx.binary_centers()); + } else { + raft::serialize_mdspan(handle, legacy, idx.centers()); + } + raft::serialize_scalar(handle, legacy, idx.center_norms().has_value()); + if (idx.center_norms()) { raft::serialize_mdspan(handle, legacy, *idx.center_norms()); } + raft::serialize_mdspan(handle, legacy, idx.list_sizes()); + raft::serialize_scalar(handle, legacy, stored_size); + raft::serialize_mdspan( + handle, + legacy, + raft::make_device_matrix_view(idx.lists()[0]->data_ptr(), 32, 1)); + raft::serialize_mdspan(handle, + legacy, + raft::make_device_vector_view( + idx.lists()[0]->indices_ptr(), stored_size)); + deserialize(handle, legacy, &restored); + ASSERT_EQ(restored.size(), idx.size()); + ASSERT_EQ(restored.metric(), idx.metric()); + if (idx.binary_index()) { + ASSERT_TRUE(cuvs::devArrMatch(idx.binary_centers().data_handle(), + restored.binary_centers().data_handle(), + idx.binary_centers().size(), + cuvs::Compare(), + stream.get())); + } else { + ASSERT_TRUE(cuvs::devArrMatch(idx.centers().data_handle(), + restored.centers().data_handle(), + idx.centers().size(), + cuvs::Compare(), + stream.get())); + } + ASSERT_TRUE(cuvs::devArrMatch(idx.lists()[0]->indices_ptr(), + restored.lists()[0]->indices_ptr(), + data_host.size(), + cuvs::Compare(), + stream.get())); + } + } +} + } // namespace cuvs::neighbors::ivf_flat diff --git a/cpp/tests/neighbors/ann_utils.cuh b/cpp/tests/neighbors/ann_utils.cuh index 346972d79d..35e940dc6f 100644 --- a/cpp/tests/neighbors/ann_utils.cuh +++ b/cpp/tests/neighbors/ann_utils.cuh @@ -111,6 +111,7 @@ inline auto operator<<(std::ostream& os, const print_metric& p) -> std::ostream& break; case cuvs::distance::DistanceType::DiceExpanded: os << "distance::DiceExpanded"; break; case cuvs::distance::DistanceType::Precomputed: os << "distance::Precomputed"; break; + case cuvs::distance::DistanceType::BitwiseHamming: os << "distance::BitwiseHamming"; break; default: RAFT_FAIL("unreachable code"); } return os; diff --git a/fern/pages/cpp_api/cpp-api-cluster-kmeans.md b/fern/pages/cpp_api/cpp-api-cluster-kmeans.md index 854db8c317..c05c2fa0dd 100644 --- a/fern/pages/cpp_api/cpp-api-cluster-kmeans.md +++ b/fern/pages/cpp_api/cpp-api-cluster-kmeans.md @@ -423,6 +423,8 @@ raft::device_matrix_view centroids, std::optional> inertia = std::nullopt); ``` +**Note:** When `params.is_packed_binary` is true, `X.extent(1)` counts packed bytes,
and centroids must have `8 * X.extent(1)` floating-point coordinates. Bits are
expanded least-significant bit first to \{-1, +1\}; the selected metric operates
on those expanded vectors. CosineExpanded is not supported in packed binary mode.
With the flag disabled, uint8_t values are numeric. + **Parameters** | Name | Direction | Type | Description | @@ -708,6 +710,8 @@ raft::device_matrix_view centroids, raft::device_vector_view labels); ``` +**Note:** When `params.is_packed_binary` is true, `X.extent(1)` counts packed bytes,
and centroids must have `8 * X.extent(1)` floating-point coordinates. Bits are
expanded least-significant bit first to \{-1, +1\}; the selected metric operates
on those expanded vectors. CosineExpanded is not supported in packed binary mode.
With the flag disabled, uint8_t values are numeric. + **Parameters** | Name | Direction | Type | Description | diff --git a/fern/pages/cpp_api/cpp-api-neighbors-ivf-flat.md b/fern/pages/cpp_api/cpp-api-neighbors-ivf-flat.md index 0cefe1a92d..25ef7b1dcd 100644 --- a/fern/pages/cpp_api/cpp-api-neighbors-ivf-flat.md +++ b/fern/pages/cpp_api/cpp-api-neighbors-ivf-flat.md @@ -184,7 +184,7 @@ NB: This may differ from the actual list size if the shared lists have been exte ### neighbors::ivf_flat::index::centers -k-means cluster centers corresponding to the lists [n_lists, dim] +Floating-point k-means centers [n_lists, dim]; empty for binary indexes. ```cpp raft::device_matrix_view centers() noexcept; @@ -194,6 +194,65 @@ raft::device_matrix_view centers() noexcept; `raft::device_matrix_view` + +### neighbors::ivf_flat::index::binary_centers + +Packed binary cluster centers, with `dim()` bytes per center. + +```cpp +raft::device_matrix_view binary_centers() noexcept; +``` + +**Returns** + +`raft::device_matrix_view` + +A mutable device view of shape [n_lists, dim], or an empty view for nonbinary indexes. + +**Additional overload:** `neighbors::ivf_flat::index::binary_centers` + +Packed binary cluster centers, with `dim()` bytes per center. + +```cpp +raft::device_matrix_view binary_centers() const noexcept; +``` + +**Returns** + +`raft::device_matrix_view` + +A read-only device view of shape [n_lists, dim], or an empty view for nonbinary indexes. + + +### neighbors::ivf_flat::index::binary_center_counts + +Exact per-bit one-counts for adaptive binary centers. Together with list_sizes(), these retain majority statistics across extensions. + +```cpp +raft::device_matrix_view binary_center_counts() noexcept; +``` + +**Returns** + +`raft::device_matrix_view` + +A mutable device view of shape [n_lists, dim * 8], or an empty view when unused. + +**Additional overload:** `neighbors::ivf_flat::index::binary_center_counts` + +Exact per-bit one-counts for adaptive binary centers. Together with list_sizes(), these retain majority statistics across extensions. + +```cpp +raft::device_matrix_view binary_center_counts() +const noexcept; +``` + +**Returns** + +`raft::device_matrix_view` + +A read-only device view of shape [n_lists, dim * 8], or an empty view when unused. + ### neighbors::ivf_flat::index::center_norms @@ -250,6 +309,8 @@ Dimensionality of the data. uint32_t dim() const noexcept; ``` +**Note:** For binary index, this returns the dimensionality of the byte dataset, which is the
number of bits / 8. + **Returns** `uint32_t` @@ -308,6 +369,19 @@ std::vector>>& lists() noexcept; `std::vector>>&` + +### neighbors::ivf_flat::index::binary_index + +Whether the index uses byte-packed vectors and BitwiseHamming distance. + +```cpp +bool binary_index() const noexcept; +``` + +**Returns** + +`bool` + ## IVF-Flat index build @@ -328,6 +402,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -360,6 +435,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -393,6 +469,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -425,6 +502,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -458,6 +536,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -490,6 +569,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -523,6 +603,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -555,6 +636,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Usage example: @@ -588,6 +670,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -622,6 +705,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -657,6 +741,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -691,6 +776,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -726,6 +812,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -760,6 +847,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -795,6 +883,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping. @@ -829,6 +918,7 @@ NB: Currently, the following distance metrics are supported: - L2Unexpanded - InnerProduct - CosineExpanded +- BitwiseHamming (uint8_t input only; dimensions are measured in packed bytes) Note, if index_params.add_data_on_build is set to true, the user can set a stream pool in the input raft::resource with at least one stream to enable kernel and copy overlapping.