diff --git a/cpp/src/neighbors/cagra.cuh b/cpp/src/neighbors/cagra.cuh index 0e1454d987..27e33d1b17 100644 --- a/cpp/src/neighbors/cagra.cuh +++ b/cpp/src/neighbors/cagra.cuh @@ -420,6 +420,28 @@ void search_with_filtering(raft::resources const& res, res, params, idx, queries, neighbors, distances, sample_filter); } +/** + * Fraction of rows removed by per-partition bitsets (one bitset per partition; an empty view means + * the partition is unfiltered), clamped to [0, 0.999] so the plan's itopk scaling stays finite. + */ +inline float bitset_filtering_rate( + raft::resources const& res, + const std::vector& partition_rows, + const std::vector>& partition_bitsets) +{ + int64_t total_rows = 0; + int64_t kept_rows = 0; + for (size_t i = 0; i < partition_rows.size(); i++) { + total_rows += partition_rows[i]; + const bool filtered = i < partition_bitsets.size() && partition_bitsets[i].data() != nullptr && + partition_bitsets[i].size() > 0; + kept_rows += + filtered ? static_cast(partition_bitsets[i].count(res)) : partition_rows[i]; + } + const float rate = static_cast(total_rows - kept_rows) / static_cast(total_rows); + return std::min(std::max(rate, 0.0f), 0.999f); +} + template (idx.dataset().n_rows())}, {sample_filter.bitset_view_}); } auto sample_filter_copy = sample_filter; return search_with_filtering( @@ -598,18 +616,27 @@ void search( } } + search_params params_copy = params; + if (params_copy.filtering_rate < 0.0f) { + std::vector partition_rows; + for (const auto* idx : indices) { + partition_rows.push_back(static_cast(idx->size())); + } + params_copy.filtering_rate = bitset_filtering_rate(res, partition_rows, partition_bitsets); + } + if (rep == nullptr) { cagra::detail::search_multi_partition( - res, params, indices, queries, partition_ids, neighbors, distances, partition_bitsets); + res, params_copy, indices, queries, partition_ids, neighbors, distances, partition_bitsets); } else { using bitset_filter_t = cuvs::neighbors::filtering::bitset_filter; cagra::detail::search_multi_partition( res, - params, + params_copy, indices, queries, partition_ids, diff --git a/cpp/src/neighbors/detail/cagra/search_plan.cuh b/cpp/src/neighbors/detail/cagra/search_plan.cuh index a6d1bd1204..d7f9b26fda 100644 --- a/cpp/src/neighbors/detail/cagra/search_plan.cuh +++ b/cpp/src/neighbors/detail/cagra/search_plan.cuh @@ -199,12 +199,34 @@ struct search_plan_impl : public search_plan_impl_base { void adjust_search_params() { + if (algo == search_algo::MULTI_CTA && (0.0 < filtering_rate && filtering_rate < 1.0)) { + size_t adjusted_itopk_size = + (size_t)((float)topk / (1.0 - filtering_rate) + + (float)(itopk_size - topk) / std::sqrt(1.0 - filtering_rate)); + if (adjusted_itopk_size % 32) { adjusted_itopk_size += 32 - (adjusted_itopk_size % 32); } + if (itopk_size < adjusted_itopk_size) { + RAFT_LOG_DEBUG( + "# internal_topk is increased from %lu to %lu, considering fintering rate %f.", + itopk_size, + adjusted_itopk_size, + filtering_rate); + itopk_size = adjusted_itopk_size; + } + } uint32_t _max_iterations = max_iterations; if (max_iterations == 0) { if (algo == search_algo::MULTI_CTA) { - constexpr uint32_t mc_itopk_size = 32; - constexpr uint32_t mc_search_width = 1; - _max_iterations = mc_itopk_size / mc_search_width; + constexpr size_t mc_itopk_size = 32; + const size_t minimum_depth = 16; + const auto effective_itopk_size = raft::ceildiv(itopk_size, mc_itopk_size) * mc_itopk_size; + // In multi-CTA algo, search_width and itopk are both knobs on num_ctas + const auto num_ctas = max(search_width, raft::ceildiv(effective_itopk_size, mc_itopk_size)); + + // Shrink max_iterations when num_ctas is large. In multi-CTA algo, larger num_ctas implies + // both more width and depth + _max_iterations = minimum_depth + raft::ceildiv(mc_itopk_size - minimum_depth, num_ctas); + // Compensate for the difficult case of large topk + _max_iterations += raft::ceildiv(static_cast(topk), mc_itopk_size) - 1; } else { _max_iterations = itopk_size / search_width; } @@ -220,20 +242,6 @@ struct search_plan_impl : public search_plan_impl_base { "# max_iterations is increased from %lu to %u.", max_iterations, _max_iterations); max_iterations = _max_iterations; } - if (algo == search_algo::MULTI_CTA && (0.0 < filtering_rate && filtering_rate < 1.0)) { - size_t adjusted_itopk_size = - (size_t)((float)topk / (1.0 - filtering_rate) + - (float)(itopk_size - topk) / std::sqrt(1.0 - filtering_rate)); - if (adjusted_itopk_size % 32) { adjusted_itopk_size += 32 - (adjusted_itopk_size % 32); } - if (itopk_size < adjusted_itopk_size) { - RAFT_LOG_DEBUG( - "# internal_topk is increased from %lu to %lu, considering fintering rate %f.", - itopk_size, - adjusted_itopk_size, - filtering_rate); - itopk_size = adjusted_itopk_size; - } - } if (itopk_size % 32) { uint32_t itopk32 = itopk_size; itopk32 += 32 - (itopk_size % 32);