diff --git a/diskann-benchmark-core/src/search/graph/inline.rs b/diskann-benchmark-core/src/search/graph/inline.rs index 5e32e187d..9ce1306d0 100644 --- a/diskann-benchmark-core/src/search/graph/inline.rs +++ b/diskann-benchmark-core/src/search/graph/inline.rs @@ -12,7 +12,7 @@ use diskann::{ }; use diskann_utils::{future::AsyncFriendly, views::Matrix}; -use crate::search::{self, Search, graph::Strategy}; +use crate::search::{self, Search, graph::KnnParams, graph::Strategy}; /// A built-in helper for benchmarking filtered K-nearest neighbors search /// using the inline search method. @@ -22,7 +22,7 @@ use crate::search::{self, Search, graph::Strategy}; /// [`search::search_all`] is provided by the [`search::graph::knn::Aggregator`] type (same /// aggregator as [`search::graph::knn::KNN`]). /// -/// The provided implementation of [`Search`] accepts [`graph::search::Knn`] +/// The provided implementation of [`Search`] accepts [`KnnParams`] /// and returns [`search::graph::knn::Metrics`] as additional output. #[derive(Debug)] pub struct InlineFilterSearch @@ -100,7 +100,7 @@ where T: AsyncFriendly + Clone, { type Id = DP::ExternalId; - type Parameters = graph::search::Knn; + type Parameters = KnnParams; type Output = super::knn::Metrics; fn num_queries(&self) -> usize { @@ -121,8 +121,8 @@ where O: graph::SearchOutputBuffer + Send, { let context = DP::Context::default(); - let inline_search = - graph::search::InlineFilterSearch::new(*parameters, self.adaptive_l.clone()); + let knn = parameters.knn; + let inline_search = graph::search::InlineFilterSearch::new(knn, self.adaptive_l.clone()); let strategy = labeled::Filtered::new(self.strategy.get(index)?.clone(), &*self.labels[index]); @@ -199,7 +199,7 @@ mod tests { let rt = crate::tokio::runtime(2).unwrap(); let results = search::search( inline.clone(), - graph::search::Knn::new(nearest_neighbors, 10, None).unwrap(), + KnnParams::new(nearest_neighbors, 10).unwrap(), NonZeroUsize::new(2).unwrap(), &rt, ) @@ -227,11 +227,11 @@ mod tests { // Try the aggregated strategy. let parameters = [ search::Run::new( - graph::search::Knn::new(nearest_neighbors, 10, None).unwrap(), + KnnParams::new(nearest_neighbors, 10).unwrap(), setup.clone(), ), search::Run::new( - graph::search::Knn::new(nearest_neighbors, 15, None).unwrap(), + KnnParams::new(nearest_neighbors, 15).unwrap(), setup.clone(), ), ]; diff --git a/diskann-benchmark-core/src/search/graph/knn.rs b/diskann-benchmark-core/src/search/graph/knn.rs index f1c074d04..286be8871 100644 --- a/diskann-benchmark-core/src/search/graph/knn.rs +++ b/diskann-benchmark-core/src/search/graph/knn.rs @@ -5,10 +5,11 @@ //! A built-in helper for benchmarking K-nearest neighbors. -use std::sync::Arc; +use std::{num::NonZeroUsize, sync::Arc}; +use thiserror::Error; use diskann::{ - ANNResult, + ANNError, ANNResult, graph::{self, glue}, provider, }; @@ -30,7 +31,7 @@ use crate::{ /// the latter. Result aggregation for [`search::search_all`] is provided /// by the [`Aggregator`] type. /// -/// The provided implementation of [`Search`] accepts [`graph::search::Knn`] +/// The provided implementation of [`Search`] accepts [`KnnParams`] /// and returns [`Metrics`] as additional output. /// /// # Type Parameters @@ -194,7 +195,7 @@ where T: AsyncFriendly + Clone, { type Id = DP::ExternalId; - type Parameters = graph::search::Knn; + type Parameters = KnnParams; type Output = Metrics; fn num_queries(&self) -> usize { @@ -215,7 +216,7 @@ where O: graph::SearchOutputBuffer + Send, { let context = DP::Context::default(); - let knn_search = *parameters; + let knn_search = parameters.knn; let strategy = self.strategy.get(index)?; let processor = self.post_processor.as_post_processor(strategy); @@ -238,6 +239,51 @@ where } } +#[derive(Debug, Error)] +pub enum KnnParamsError { + #[error("k_value cannot be zero")] + KZero, + #[error("l_value ({l_value}) must be at least k_value ({k_value})")] + LLessThanK { l_value: usize, k_value: usize }, + #[error("invalid KNN parameters")] + InvalidKnnParameters, +} + +impl From for ANNError { + #[track_caller] + fn from(err: KnnParamsError) -> Self { + ANNError::opaque(err) + } +} + +/// A wrapper for the [`graph::search::Knn`] struct that also includes the `k` value. +#[derive(Debug, Copy, Clone)] +pub struct KnnParams { + k_value: NonZeroUsize, + pub knn: graph::search::Knn, +} + +impl KnnParams { + /// Construct a new [`KnnParams`]. + pub fn new(k_value: usize, l_value: usize) -> Result { + let k_value = NonZeroUsize::new(k_value).ok_or(KnnParamsError::KZero)?; + if l_value < k_value.get() { + return Err(KnnParamsError::LLessThanK { + l_value, + k_value: k_value.get(), + }); + } + + let knn = graph::search::Knn::new(l_value, None) + .map_err(|_| KnnParamsError::InvalidKnnParameters)?; + Ok(Self { k_value, knn }) + } + + pub fn k_value(&self) -> NonZeroUsize { + self.k_value + } +} + /// An [`search::Aggregate`]d summary of multiple [`KNN`] search runs /// returned by the provided [`Aggregator`]. /// @@ -249,7 +295,7 @@ pub struct Summary { pub setup: search::Setup, /// The [`Search::Parameters`] used for the batch of runs. - pub parameters: graph::search::Knn, + pub parameters: KnnParams, /// The end-to-end latency for each repetition in the batch. pub end_to_end_latencies: Vec, @@ -317,7 +363,7 @@ impl<'a, I> Aggregator<'a, I> { } } -impl search::Aggregate for Aggregator<'_, I> +impl search::Aggregate for Aggregator<'_, I> where I: crate::recall::RecallCompatible, { @@ -325,7 +371,7 @@ where fn aggregate( &mut self, - run: search::Run, + run: search::Run, mut results: Vec>, ) -> anyhow::Result { // Compute the recall using just the first result. @@ -422,7 +468,7 @@ mod tests { let rt = crate::tokio::runtime(2).unwrap(); let results = search::search( knn.clone(), - graph::search::Knn::new(nearest_neighbors, 10, None).unwrap(), + KnnParams::new(nearest_neighbors, 10).unwrap(), NonZeroUsize::new(2).unwrap(), &rt, ) @@ -446,11 +492,11 @@ mod tests { // Try the aggregated strategy. let parameters = [ search::Run::new( - graph::search::Knn::new(nearest_neighbors, 10, None).unwrap(), + KnnParams::new(nearest_neighbors, 10).unwrap(), setup.clone(), ), search::Run::new( - graph::search::Knn::new(nearest_neighbors, 15, None).unwrap(), + KnnParams::new(nearest_neighbors, 15).unwrap(), setup.clone(), ), ]; diff --git a/diskann-benchmark-core/src/search/graph/mod.rs b/diskann-benchmark-core/src/search/graph/mod.rs index 8063f0875..de2b69ba3 100644 --- a/diskann-benchmark-core/src/search/graph/mod.rs +++ b/diskann-benchmark-core/src/search/graph/mod.rs @@ -11,7 +11,7 @@ pub mod range; pub mod strategy; pub use inline::InlineFilterSearch; -pub use knn::KNN; +pub use knn::{KNN, KnnParams}; pub use multihop::MultiHop; pub use range::Range; diff --git a/diskann-benchmark-core/src/search/graph/multihop.rs b/diskann-benchmark-core/src/search/graph/multihop.rs index dba7c9925..53243b083 100644 --- a/diskann-benchmark-core/src/search/graph/multihop.rs +++ b/diskann-benchmark-core/src/search/graph/multihop.rs @@ -12,7 +12,7 @@ use diskann::{ }; use diskann_utils::{future::AsyncFriendly, views::Matrix}; -use crate::search::{self, Search, graph::Strategy}; +use crate::search::{self, Search, graph::KnnParams, graph::Strategy}; /// A built-in helper for benchmarking filtered K-nearest neighbors search /// using the multi-hop search method. @@ -22,7 +22,7 @@ use crate::search::{self, Search, graph::Strategy}; /// [`search::search_all`] is provided by the [`search::graph::knn::Aggregator`] type (same /// aggregator as [`search::graph::knn::KNN`]). /// -/// The provided implementation of [`Search`] accepts [`graph::search::Knn`] +/// The provided implementation of [`Search`] accepts [`KnnParams`] /// and returns [`search::graph::knn::Metrics`] as additional output. #[derive(Debug)] pub struct MultiHop @@ -97,7 +97,7 @@ where T: AsyncFriendly + Clone, { type Id = DP::ExternalId; - type Parameters = graph::search::Knn; + type Parameters = KnnParams; type Output = super::knn::Metrics; fn num_queries(&self) -> usize { @@ -118,7 +118,8 @@ where O: graph::SearchOutputBuffer + Send, { let context = DP::Context::default(); - let multihop_search = graph::search::MultihopFilterSearch::new(*parameters); + let knn = parameters.knn; + let multihop_search = graph::search::MultihopFilterSearch::new(knn); let strategy = labeled::Filtered::new(self.strategy.get(index)?.clone(), &*self.labels[index]); let stats = self @@ -191,7 +192,7 @@ mod tests { let rt = crate::tokio::runtime(2).unwrap(); let results = search::search( multihop.clone(), - graph::search::Knn::new(nearest_neighbors, 10, None).unwrap(), + KnnParams::new(nearest_neighbors, 10).unwrap(), NonZeroUsize::new(2).unwrap(), &rt, ) @@ -219,11 +220,11 @@ mod tests { // Try the aggregated strategy. let parameters = [ search::Run::new( - graph::search::Knn::new(nearest_neighbors, 10, None).unwrap(), + KnnParams::new(nearest_neighbors, 10).unwrap(), setup.clone(), ), search::Run::new( - graph::search::Knn::new(nearest_neighbors, 15, None).unwrap(), + KnnParams::new(nearest_neighbors, 15).unwrap(), setup.clone(), ), ]; diff --git a/diskann-benchmark/src/index/result.rs b/diskann-benchmark/src/index/result.rs index 8bd90b33e..149530b24 100644 --- a/diskann-benchmark/src/index/result.rs +++ b/diskann-benchmark/src/index/result.rs @@ -120,7 +120,7 @@ impl SearchResults { Self { num_tasks: setup.tasks.into(), search_n: parameters.k_value().get(), - search_l: parameters.l_value().get(), + search_l: parameters.knn.l_value().get(), qps, search_latencies: end_to_end_latencies, mean_latencies, diff --git a/diskann-benchmark/src/index/search/knn.rs b/diskann-benchmark/src/index/search/knn.rs index 14f426f46..afbf42017 100644 --- a/diskann-benchmark/src/index/search/knn.rs +++ b/diskann-benchmark/src/index/search/knn.rs @@ -5,8 +5,8 @@ use std::{num::NonZeroUsize, sync::Arc}; -use diskann_benchmark_core::recall::GroundTruthMode; use diskann_benchmark_core::{self as benchmark_core, search as core_search}; +use diskann_benchmark_core::{recall::GroundTruthMode, search::graph::KnnParams}; use crate::{index::result::SearchResults, inputs::graph_index::GraphSearch}; @@ -53,8 +53,7 @@ pub(crate) fn run( .search_l .iter() .map(|search_l| { - let search_params = - diskann::graph::search::Knn::new(run.search_n, *search_l, None).unwrap(); + let search_params = KnnParams::new(run.search_n, *search_l).unwrap(); core_search::Run::new(search_params, setup.clone()) }) @@ -73,7 +72,7 @@ pub(crate) fn run( Ok(all) } -type Run = core_search::Run; +type Run = core_search::Run; pub(crate) trait Knn { fn search_all( &self, @@ -94,13 +93,13 @@ where DP: diskann::provider::DataProvider, core_search::graph::KNN: core_search::Search< Id = DP::InternalId, - Parameters = diskann::graph::search::Knn, + Parameters = KnnParams, Output = core_search::graph::knn::Metrics, >, { fn search_all( &self, - parameters: Vec>, + parameters: Vec>, groundtruth: &dyn benchmark_core::recall::Rows, recall_k: usize, recall_n: usize, @@ -126,13 +125,13 @@ where DP: diskann::provider::DataProvider, core_search::graph::MultiHop: core_search::Search< Id = DP::InternalId, - Parameters = diskann::graph::search::Knn, + Parameters = KnnParams, Output = core_search::graph::knn::Metrics, >, { fn search_all( &self, - parameters: Vec>, + parameters: Vec>, groundtruth: &dyn benchmark_core::recall::Rows, recall_k: usize, recall_n: usize, @@ -158,13 +157,13 @@ where DP: diskann::provider::DataProvider, core_search::graph::InlineFilterSearch: core_search::Search< Id = DP::InternalId, - Parameters = diskann::graph::search::Knn, + Parameters = KnnParams, Output = core_search::graph::knn::Metrics, >, { fn search_all( &self, - parameters: Vec>, + parameters: Vec>, groundtruth: &dyn benchmark_core::recall::Rows, recall_k: usize, recall_n: usize, diff --git a/diskann-bftree/src/provider.rs b/diskann-bftree/src/provider.rs index 1a9f773ca..ff5022cac 100644 --- a/diskann-bftree/src/provider.rs +++ b/diskann-bftree/src/provider.rs @@ -2074,9 +2074,10 @@ mod tests { } let query = vec![3.0; 5]; - let params = Knn::new(5, 10, None).unwrap(); + let params = Knn::new(10, None).unwrap(); - let mut neighbors = vec![Neighbor::::default(); 5]; + let k = 5; + let mut neighbors = vec![Neighbor::::default(); k]; let res = index .search( params, @@ -2089,8 +2090,9 @@ mod tests { .unwrap(); assert_eq!( - res.result_count, 5, - "there are 15 points and we're asking for 5, we expect 5" + res.result_count, k as u32, + "there are 15 points and we're asking for {}, we expect {}", + 5, k ); assert_eq!(neighbors[0].id, 3); } @@ -2125,9 +2127,10 @@ mod tests { .unwrap(); let query = vec![3.0; 5]; - let params = Knn::new(5, 10, None).unwrap(); + let params = Knn::new(10, None).unwrap(); - let mut neighbors = vec![Neighbor::::default(); 5]; + let k = 5; + let mut neighbors = vec![Neighbor::::default(); k]; let res = index .search( params, @@ -2140,8 +2143,9 @@ mod tests { .unwrap(); assert_eq!( - res.result_count, 5, - "there are 15 points and we're asking for 5, we expect 5" + res.result_count, k as u32, + "there are 15 points and we're asking for {}, we expect {}", + 5, k ); let neighbor_ids: Vec = neighbors.iter().map(|n| n.id).collect(); for expected in 1u32..=5 { @@ -2175,9 +2179,10 @@ mod tests { .unwrap(); let query = vec![3.0; 5]; - let params = Knn::new(5, 10, None).unwrap(); + let params = Knn::new(10, None).unwrap(); - let mut neighbors = vec![Neighbor::::default(); 5]; + let k = 5; + let mut neighbors = vec![Neighbor::::default(); k]; let res = index .search( params, @@ -2189,7 +2194,7 @@ mod tests { .await .unwrap(); - assert_eq!(res.result_count, 5); + assert_eq!(res.result_count, k as u32); let neighbor_ids: Vec = neighbors.iter().map(|n| n.id).collect(); assert!(!neighbor_ids.contains(&2u32)); assert!(!neighbor_ids.contains(&4u32)); @@ -2246,9 +2251,10 @@ mod tests { } let query = vec![3.0; 5]; - let params = Knn::new(5, 10, None).unwrap(); + let params = Knn::new(10, None).unwrap(); - let mut neighbors = vec![Neighbor::::default(); 5]; + let k = 5; + let mut neighbors = vec![Neighbor::::default(); k]; let res = index .search( params, @@ -2261,8 +2267,9 @@ mod tests { .unwrap(); assert_eq!( - res.result_count, 5, - "there are 15 points and we're asking for 5, we expect 5" + res.result_count, k as u32, + "there are 15 points and we're asking for {}, we expect {}", + 5, k ); assert_eq!(neighbors[0].id, 3); } @@ -2302,9 +2309,10 @@ mod tests { .unwrap(); let query = vec![3.0; 5]; - let params = Knn::new(5, 10, None).unwrap(); + let params = Knn::new(10, None).unwrap(); - let mut neighbors = vec![Neighbor::::default(); 5]; + let k = 5; + let mut neighbors = vec![Neighbor::::default(); k]; let res = index .search( params, @@ -2316,7 +2324,7 @@ mod tests { .await .unwrap(); - assert_eq!(res.result_count, 5); + assert_eq!(res.result_count, k as u32); let neighbor_ids: Vec = neighbors.iter().map(|n| n.id).collect(); assert!(!neighbor_ids.contains(&2u32)); assert!(!neighbor_ids.contains(&4u32)); diff --git a/diskann-disk/src/search/provider/disk_provider.rs b/diskann-disk/src/search/provider/disk_provider.rs index c497414a6..c7ed61476 100644 --- a/diskann-disk/src/search/provider/disk_provider.rs +++ b/diskann-disk/src/search/provider/disk_provider.rs @@ -1058,6 +1058,13 @@ where let mut associated_data = vec![Data::AssociatedDataType::default(); return_list_size as usize]; + if search_list_size < return_list_size { + return Err(ANNError::message( + diskann::ANNErrorKind::IndexError, + "search list size must be at least as large as the number of results requested", + )); + } + let stats = self.search_internal( query, return_list_size as usize, @@ -1112,7 +1119,6 @@ where ); let timer = Instant::now(); - let k = k_value; let l = search_list_size as usize; let io_tracker = IOTracker::default(); @@ -1144,7 +1150,7 @@ where .as_deref() .map_or(PostprocessStrategy::AcceptAll, PostprocessStrategy::Apply), ); - let knn_search = Knn::new(k, l, beam_width)?; + let knn_search = Knn::new(l, beam_width)?; self.runtime.block_on(self.index.search( knn_search, &strategy, @@ -1158,7 +1164,7 @@ where // `labeled::Filtered` wrapper can own it; `io_tracker` keeps // its counters reachable from this scope. let strategy = self.search_strategy(&io_tracker, PostprocessStrategy::AcceptAll); - let knn_search = Knn::new(k, l, beam_width)?; + let knn_search = Knn::new(l, beam_width)?; self.runtime.block_on(self.filter_search( strategy, query, @@ -1176,7 +1182,7 @@ where .as_deref() .map_or(PostprocessStrategy::AcceptAll, PostprocessStrategy::Apply); let strategy = self.search_strategy(&io_tracker, postprocess_config); - let knn_search = Knn::new(k, l, beam_width)?; + let knn_search = Knn::new(l, beam_width)?; let processor = DiskSearchPostProcessor::DeterminantDiversity( DeterminantDiversityAndFilter::new(postprocess_config, *params), ); @@ -1230,10 +1236,7 @@ fn ensure_vertex_loaded>( mod disk_provider_tests { use crate::test_utils::{GraphDataF32VectorU32Data, GraphDataF32VectorUnitData}; use diskann::{ - graph::{ - search::{record::VisitedSearchRecord, Knn}, - KnnSearchError, - }, + graph::search::{record::VisitedSearchRecord, Knn}, utils::IntoUsize, ANNErrorKind, }; @@ -1695,15 +1698,8 @@ mod disk_provider_tests { "index_path is not correct" ); - // Test error case: l < k - let res = Knn::new_default(20, 10); - assert!(res.is_err()); - assert_eq!( - >::into(res.unwrap_err()).kind(), - ANNErrorKind::IndexError - ); // Test error case: beam_width = 0 - let res = Knn::new(10, 10, Some(0)); + let res = Knn::new(10, Some(0)); assert!(res.is_err()); let search_engine = @@ -1776,9 +1772,10 @@ mod disk_provider_tests { ); let query_vector: [f32; 128] = [1f32; 128]; - let mut indices = vec![0u32; 10]; - let mut distances = vec![0f32; 10]; - let mut associated_data = vec![(); 10]; + let k = 10; + let mut indices = vec![0u32; k]; + let mut distances = vec![0f32; k]; + let mut associated_data = vec![(); k]; let mut result_output_buffer = search_output_buffer::IdDistanceAssociatedData::new( &mut indices, @@ -1788,7 +1785,7 @@ mod disk_provider_tests { let io_tracker = IOTracker::default(); let strategy = search_engine.search_strategy(&io_tracker, PostprocessStrategy::AcceptAll); let mut search_record = VisitedSearchRecord::new(0); - let search_params = Knn::new(10, 10, Some(4)).unwrap(); + let search_params = Knn::new(10, Some(4)).unwrap(); let recorded_search = diskann::graph::search::RecordedKnn::new(search_params, &mut search_record); search_engine @@ -1814,11 +1811,10 @@ mod disk_provider_tests { assert_eq!(ids, &EXPECTED_NODES); - let return_list_size = 10; let search_list_size = 10; let result = search_engine.search( &query_vector, - return_list_size, + k as u32, search_list_size, Some(4), SearchMode::graph(), @@ -1826,8 +1822,8 @@ mod disk_provider_tests { assert!(result.is_ok(), "Expected search to succeed"); let search_result = result.unwrap(); assert_eq!( - search_result.results.len() as u32, - return_list_size, + search_result.results.len(), + k, "Expected result count to match" ); assert_eq!( @@ -1954,7 +1950,7 @@ mod disk_provider_tests { #[cfg(feature = "experimental_diversity_search")] #[test] fn test_disk_search_diversity_search() { - use diskann::graph::DiverseSearchParams; + use diskann::graph::search::DiverseSearchParams; use diskann::neighbor::AttributeValueProvider; use std::collections::HashMap; @@ -2012,9 +2008,10 @@ mod disk_provider_tests { // Wrap in Arc once to avoid cloning the HashMap later let attribute_provider = std::sync::Arc::new(attribute_provider); - let mut indices = vec![0u32; 10]; - let mut distances = vec![0f32; 10]; - let mut associated_data = vec![(); 10]; + let original_k = 10; + let mut indices = vec![0u32; original_k]; + let mut distances = vec![0f32; original_k]; + let mut associated_data = vec![(); original_k]; let mut result_output_buffer = search_output_buffer::IdDistanceAssociatedData::new( &mut indices, @@ -2028,12 +2025,15 @@ mod disk_provider_tests { let diverse_params = DiverseSearchParams::new( 0, // diverse_attribute_id 3, // diverse_results_k + original_k, attribute_provider.clone(), - ); + ) + .unwrap(); - let search_params = Knn::new(10, 20, None).unwrap(); + let search_params = Knn::new(20, None).unwrap(); - let diverse_search = diskann::graph::search::Diverse::new(search_params, diverse_params); + let diverse_search = + diskann::graph::search::Diverse::new(search_params, diverse_params).unwrap(); let stats = search_engine .runtime .block_on(search_engine.index.search( @@ -2051,19 +2051,20 @@ mod disk_provider_tests { "Expected to get some results during diversity search" ); - let return_list_size = 10; let search_list_size = 20; let diverse_results_k = 1; let diverse_params = DiverseSearchParams::new( 0, // diverse_attribute_id diverse_results_k, + original_k, attribute_provider.clone(), - ); + ) + .unwrap(); // Test diverse search using the search API - let mut indices2 = vec![0u32; return_list_size as usize]; - let mut distances2 = vec![0f32; return_list_size as usize]; - let mut associated_data2 = vec![(); return_list_size as usize]; + let mut indices2 = vec![0u32; original_k]; + let mut distances2 = vec![0f32; original_k]; + let mut associated_data2 = vec![(); original_k]; let mut result_output_buffer2 = search_output_buffer::IdDistanceAssociatedData::new( &mut indices2, &mut distances2, @@ -2071,10 +2072,10 @@ mod disk_provider_tests { ); let io_tracker2 = IOTracker::default(); let strategy2 = search_engine.search_strategy(&io_tracker2, PostprocessStrategy::AcceptAll); - let search_params2 = - Knn::new(return_list_size as usize, search_list_size as usize, None).unwrap(); + let search_params2 = Knn::new(search_list_size as usize, None).unwrap(); - let diverse_search2 = diskann::graph::search::Diverse::new(search_params2, diverse_params); + let diverse_search2 = + diskann::graph::search::Diverse::new(search_params2, diverse_params).unwrap(); let stats = search_engine .runtime .block_on(search_engine.index.search( @@ -2092,9 +2093,9 @@ mod disk_provider_tests { "Expected diversity search to return results" ); assert!( - stats.result_count <= return_list_size, + stats.result_count <= original_k as u32, "Expected result count to be <= {}", - return_list_size + original_k ); // Verify that we got some results @@ -2473,9 +2474,10 @@ mod disk_provider_tests { ); let query_vector: [f32; 128] = [1f32; 128]; - let mut indices = vec![0u32; 10]; - let mut distances = vec![0f32; 10]; - let mut associated_data = vec![(); 10]; + let k = 10; + let mut indices = vec![0u32; k]; + let mut distances = vec![0f32; k]; + let mut associated_data = vec![(); k]; let mut result_output_buffer = search_output_buffer::IdDistanceAssociatedData::new( &mut indices, @@ -2487,7 +2489,7 @@ mod disk_provider_tests { let strategy = search_engine.search_strategy(&io_tracker, PostprocessStrategy::AcceptAll); let mut search_record = VisitedSearchRecord::new(0); - let search_params = Knn::new(10, 10, Some(4)).unwrap(); + let search_params = Knn::new(10, Some(4)).unwrap(); let recorded_search = diskann::graph::search::RecordedKnn::new(search_params, &mut search_record); search_engine diff --git a/diskann-garnet/src/lib.rs b/diskann-garnet/src/lib.rs index b4d4977ca..5f18c8f57 100644 --- a/diskann-garnet/src/lib.rs +++ b/diskann-garnet/src/lib.rs @@ -635,7 +635,6 @@ pub unsafe extern "C" fn search_vector( ); let knn_params = match search::Knn::new( - output_distances_len, search_exploration_factor as usize, None, ) { @@ -704,7 +703,6 @@ pub unsafe extern "C" fn search_element( ); let knn_params = match search::Knn::new( - output_distances_len, search_exploration_factor as usize, None, ) { diff --git a/diskann-providers/src/index/diskann_async.rs b/diskann-providers/src/index/diskann_async.rs index 41389910e..1081119f8 100644 --- a/diskann-providers/src/index/diskann_async.rs +++ b/diskann-providers/src/index/diskann_async.rs @@ -343,8 +343,7 @@ pub(crate) mod tests { let mut distances = vec![0.0; parameters.search_k]; let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = - graph::search::Knn::new_default(parameters.search_k, parameters.search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(parameters.search_l).unwrap(); index .search( graph_search, @@ -1437,7 +1436,7 @@ pub(crate) mod tests { { let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); // Full Precision Search. index .search( @@ -1455,7 +1454,7 @@ pub(crate) mod tests { { let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); // Quantized Search index .search( @@ -1695,8 +1694,7 @@ pub(crate) mod tests { { let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = - graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); // Full Precision Search. index .search( @@ -1714,8 +1712,7 @@ pub(crate) mod tests { { let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = - graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); // Quantized Search index .search( @@ -1802,7 +1799,7 @@ pub(crate) mod tests { { let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = graph::search::Knn::new_default(top_k, top_k).unwrap(); + let graph_search = graph::search::Knn::new_default(top_k).unwrap(); // Quantized Search index .search( @@ -1916,7 +1913,7 @@ pub(crate) mod tests { // Full Precision Search. let mut output = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); index .search(graph_search, &FullPrecision, ctx, query, &mut output) .await @@ -1928,7 +1925,7 @@ pub(crate) mod tests { let strategy = inmem::spherical::Quantized::search( diskann_quantization::spherical::iface::QueryLayout::FourBitTransposed, ); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); index .search(graph_search, &strategy, ctx, query, &mut output) @@ -2032,7 +2029,7 @@ pub(crate) mod tests { let strategy = inmem::spherical::Quantized::search( diskann_quantization::spherical::iface::QueryLayout::FourBitTransposed, ); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); index .search(graph_search, &strategy, ctx, query, &mut output) @@ -2118,7 +2115,7 @@ pub(crate) mod tests { let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); // Full Precision Search. index .search( @@ -2697,7 +2694,7 @@ pub(crate) mod tests { let gt = groundtruth(queries.as_view(), query, |a, b| SquaredL2::evaluate(a, b)); let mut result_output_buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); - let graph_search = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let graph_search = graph::search::Knn::new_default(search_l).unwrap(); // Full Precision Search. index .search( @@ -2961,20 +2958,22 @@ pub(crate) mod tests { let mut result_output_buffer = diskann::graph::IdDistance::new(&mut indices, &mut distances); - let diverse_params = diskann::graph::DiverseSearchParams::new( + let diverse_params = diskann::graph::search::DiverseSearchParams::new( 0, // diverse_attribute_id diverse_results_k, + return_list_size, attribute_provider.clone(), - ); + ) + .unwrap(); let search_params = diskann::graph::search::Knn::new( - return_list_size, search_list_size, None, // beam_width ) .unwrap(); - let diverse_search = diskann::graph::search::Diverse::new(search_params, diverse_params); + let diverse_search = + diskann::graph::search::Diverse::new(search_params, diverse_params).unwrap(); let result = index .search( @@ -3125,7 +3124,7 @@ pub(crate) mod tests { let mut ids = vec![0; top_k]; let mut distances = vec![0.0; top_k]; let ctx = DefaultContext; - let search_params = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let search_params = graph::search::Knn::new_default(search_l).unwrap(); for i in 0..query_count { let query_vector = &queries[i * VECTORS_DIMENSION..(i + 1) * VECTORS_DIMENSION]; diff --git a/diskann-providers/src/index/wrapped_async.rs b/diskann-providers/src/index/wrapped_async.rs index 46084b6b1..a990ac2ba 100644 --- a/diskann-providers/src/index/wrapped_async.rs +++ b/diskann-providers/src/index/wrapped_async.rs @@ -775,7 +775,7 @@ mod tests { let mut output = search_output_buffer::IdDistance::new(&mut ids, &mut distances); let query = train_data.row(0); - let kind = graph::search::Knn::new_default(top_k, search_l).unwrap(); + let kind = graph::search::Knn::new_default(search_l).unwrap(); let stats = loaded .search(kind, &FullPrecision, &DefaultContext, query, &mut output) .unwrap(); diff --git a/diskann/src/graph/misc.rs b/diskann/src/graph/misc.rs index 067bd2d21..f039cef53 100644 --- a/diskann/src/graph/misc.rs +++ b/diskann/src/graph/misc.rs @@ -31,36 +31,6 @@ pub enum InplaceDeleteMethod { OneHop, } -// Parameters for diverse search -#[cfg(feature = "experimental_diversity_search")] -#[derive(Clone, Debug)] -pub struct DiverseSearchParams

-where - P: crate::neighbor::AttributeValueProvider, -{ - pub diverse_attribute_id: usize, - pub diverse_results_k: usize, - pub attribute_provider: std::sync::Arc

, -} - -#[cfg(feature = "experimental_diversity_search")] -impl

DiverseSearchParams

-where - P: crate::neighbor::AttributeValueProvider, -{ - pub fn new( - diverse_attribute_id: usize, - diverse_results_k: usize, - attribute_provider: std::sync::Arc

, - ) -> Self { - Self { - diverse_attribute_id, - diverse_results_k, - attribute_provider, - } - } -} - /////////// // Tests // /////////// diff --git a/diskann/src/graph/mod.rs b/diskann/src/graph/mod.rs index 6bf4be7dd..374efc644 100644 --- a/diskann/src/graph/mod.rs +++ b/diskann/src/graph/mod.rs @@ -23,9 +23,6 @@ pub use start_point::{SampleableForStart, StartPointStrategy}; mod misc; pub use misc::{ConsolidateKind, InplaceDeleteMethod}; -#[cfg(feature = "experimental_diversity_search")] -pub use misc::DiverseSearchParams; - pub mod glue; pub mod search; pub mod workingset; diff --git a/diskann/src/graph/search/diverse_search.rs b/diskann/src/graph/search/diverse_search.rs index bc5b7de47..383b36eab 100644 --- a/diskann/src/graph/search/diverse_search.rs +++ b/diskann/src/graph/search/diverse_search.rs @@ -7,13 +7,14 @@ use diskann_utils::future::SendFuture; use hashbrown::HashSet; +use std::num::NonZeroUsize; +use thiserror::Error; use super::{Knn, Search, record::NoopSearchRecord, scratch::SearchScratch}; use crate::{ - ANNResult, + ANNError, ANNErrorKind, ANNResult, error::IntoANNResult, graph::{ - DiverseSearchParams, glue::{SearchAccessor, SearchPostProcess, SearchStrategy}, index::{DiskANNIndex, SearchStats}, search_output_buffer::SearchOutputBuffer, @@ -22,6 +23,75 @@ use crate::{ provider::DataProvider, }; +/// Error type for [`DiverseSearchParams`] parameter validation. +#[derive(Debug, Error)] +pub enum DiverseSearchError { + #[error("total k_value cannot be zero")] + TotalKZero, + #[error("diverse k_value cannot be zero")] + DiverseKZero, +} + +impl From for ANNError { + #[track_caller] + fn from(err: DiverseSearchError) -> Self { + Self::new(ANNErrorKind::IndexError, err) + } +} + +/// Error type for [`Diverse`] parameter validation. +#[derive(Debug, Error)] +pub enum DiverseError { + #[error("l_value ({l_value}) must be greater than or equal to total_k_value ({total_k_value})")] + LValueTooSmall { + l_value: usize, + total_k_value: usize, + }, +} + +impl From for ANNError { + #[track_caller] + fn from(err: DiverseError) -> Self { + Self::new(ANNErrorKind::IndexError, err) + } +} + +// Parameters for diverse search +#[derive(Clone, Debug)] +pub struct DiverseSearchParams

+where + P: crate::neighbor::AttributeValueProvider, +{ + pub diverse_attribute_id: usize, + pub diverse_results_k: NonZeroUsize, + pub total_k_value: NonZeroUsize, + pub attribute_provider: std::sync::Arc

, +} + +impl

DiverseSearchParams

+where + P: crate::neighbor::AttributeValueProvider, +{ + pub fn new( + diverse_attribute_id: usize, + diverse_results_k: usize, + total_k_value: usize, + attribute_provider: std::sync::Arc

, + ) -> Result { + let diverse_results_k = + NonZeroUsize::new(diverse_results_k).ok_or(DiverseSearchError::DiverseKZero)?; + let total_k_value = + NonZeroUsize::new(total_k_value).ok_or(DiverseSearchError::TotalKZero)?; + + Ok(Self { + diverse_attribute_id, + diverse_results_k, + total_k_value, + attribute_provider, + }) + } +} + /// Parameters for diversity-aware search. /// /// Returns results that are diverse across a specified attribute. @@ -41,11 +111,21 @@ where P: AttributeValueProvider, { /// Create new diverse search parameters. - pub fn new(inner: Knn, diverse_params: DiverseSearchParams

) -> Self { - Self { + pub fn new(inner: Knn, diverse_params: DiverseSearchParams

) -> Result { + let l_value = inner.l_value().get(); + let total_k_value = diverse_params.total_k_value.get(); + + if l_value < total_k_value { + return Err(DiverseError::LValueTooSmall { + l_value, + total_k_value, + }); + } + + Ok(Self { inner, diverse_params, - } + }) } /// Returns a reference to the inner k-NN search parameters. @@ -72,8 +152,8 @@ where let attribute_provider = self.diverse_params.attribute_provider.clone(); let diverse_queue = DiverseNeighborQueue::new( self.inner.l_value().get(), - self.inner.k_value(), - self.diverse_params.diverse_results_k, + self.diverse_params.total_k_value, + self.diverse_params.diverse_results_k.get(), attribute_provider, ); diff --git a/diskann/src/graph/search/knn_search.rs b/diskann/src/graph/search/knn_search.rs index 9f61f1a10..59a9e871b 100644 --- a/diskann/src/graph/search/knn_search.rs +++ b/diskann/src/graph/search/knn_search.rs @@ -26,12 +26,8 @@ use crate::{ /// Error type for [`Knn`] parameter validation. #[derive(Debug, Error)] pub enum KnnSearchError { - #[error("l_value ({l_value}) cannot be less than k_value ({k_value})")] - LLessThanK { l_value: usize, k_value: usize }, #[error("beam width cannot be zero")] BeamWidthZero, - #[error("k_value cannot be zero")] - KZero, #[error("l_value cannot be zero")] LZero, } @@ -60,7 +56,6 @@ impl From for ANNError { /// /// # Parameters /// -/// - `k_value`: Number of nearest neighbors to return /// - `l_value`: Search list size (larger values improve recall at cost of latency) /// - `beam_width`: Optional parallel exploration width /// @@ -69,13 +64,11 @@ impl From for ANNError { /// ```ignore /// use diskann::graph::{search::Knn, Search}; /// -/// let params = Knn::new(10, 100, None)?; +/// let params = Knn::new(100, None)?; /// let stats = index.search(params, &strategy, &context, &query, &mut output).await?; /// ``` #[derive(Debug, Clone, Copy)] pub struct Knn { - /// Number of results to return (k in k-NN). - k_value: NonZeroUsize, /// Search list size - controls accuracy vs speed tradeoff. l_value: NonZeroUsize, /// Beam width for parallel graph exploration (defaults to 1). @@ -89,21 +82,10 @@ impl Knn { /// /// # Errors /// - /// Returns an error if `k_value` is zero, `l_value` is zero, - /// `l_value < k_value`, or if `beam_width` is `Some(0)`. - pub fn new( - k_value: usize, - l_value: usize, - beam_width: Option, - ) -> Result { - let k_value = NonZeroUsize::new(k_value).ok_or(KnnSearchError::KZero)?; + /// Returns an error if `l_value` is zero, + /// or if `beam_width` is `Some(0)`. + pub fn new(l_value: usize, beam_width: Option) -> Result { let l_value = NonZeroUsize::new(l_value).ok_or(KnnSearchError::LZero)?; - if k_value > l_value { - return Err(KnnSearchError::LLessThanK { - l_value: l_value.get(), - k_value: k_value.get(), - }); - } const ONE: NonZeroUsize = NonZeroUsize::new(1).unwrap(); let beam_width = match beam_width { @@ -112,21 +94,14 @@ impl Knn { }; Ok(Self { - k_value, l_value, beam_width, }) } /// Create parameters with default beam width. - pub fn new_default(k_value: usize, l_value: usize) -> Result { - Self::new(k_value, l_value, None) - } - - /// Returns the number of results to return (k in k-NN). - #[inline] - pub fn k_value(&self) -> NonZeroUsize { - self.k_value + pub fn new_default(l_value: usize) -> Result { + Self::new(l_value, None) } /// Returns the search list size. @@ -299,25 +274,15 @@ mod tests { #[test] fn test_knn_search_validation() { // Valid - assert!(Knn::new(10, 100, None).is_ok()); - assert!(Knn::new(10, 100, Some(4)).is_ok()); - assert!(Knn::new(10, 10, None).is_ok()); // k == l is valid - - // Invalid: k = 0 - assert!(matches!(Knn::new(0, 100, None), Err(KnnSearchError::KZero))); + assert!(Knn::new(100, None).is_ok()); + assert!(Knn::new(100, Some(4)).is_ok()); // Invalid: l = 0 - assert!(matches!(Knn::new(10, 0, None), Err(KnnSearchError::LZero))); - - // Invalid: l < k - assert!(matches!( - Knn::new(100, 10, None), - Err(KnnSearchError::LLessThanK { .. }) - )); + assert!(matches!(Knn::new(0, None), Err(KnnSearchError::LZero))); // Invalid: zero beam_width assert!(matches!( - Knn::new(10, 100, Some(0)), + Knn::new(100, Some(0)), Err(KnnSearchError::BeamWidthZero) )); } diff --git a/diskann/src/graph/search/mod.rs b/diskann/src/graph/search/mod.rs index 4e2aedc61..0f27c7db3 100644 --- a/diskann/src/graph/search/mod.rs +++ b/diskann/src/graph/search/mod.rs @@ -123,3 +123,6 @@ mod diverse_search; #[cfg(feature = "experimental_diversity_search")] pub use diverse_search::Diverse; + +#[cfg(feature = "experimental_diversity_search")] +pub use diverse_search::DiverseSearchParams; diff --git a/diskann/src/graph/test/cases/grid_insert.rs b/diskann/src/graph/test/cases/grid_insert.rs index 0c0e24724..ed26f970a 100644 --- a/diskann/src/graph/test/cases/grid_insert.rs +++ b/diskann/src/graph/test/cases/grid_insert.rs @@ -201,10 +201,11 @@ fn run_searches( let mut results = Vec::new(); for (query, desc) in queries { - let params = Knn::new(10, 10, None).unwrap(); + let k_value = 10; + let params = Knn::new(10, None).unwrap(); let search_ctx = test_provider::Context::new(); - let mut neighbors = vec![Neighbor::::default(); params.k_value().get()]; + let mut neighbors = vec![Neighbor::::default(); k_value]; let graph::index::SearchStats { cmps, hops, diff --git a/diskann/src/graph/test/cases/grid_search.rs b/diskann/src/graph/test/cases/grid_search.rs index 7afbb2214..dab9001cb 100644 --- a/diskann/src/graph/test/cases/grid_search.rs +++ b/diskann/src/graph/test/cases/grid_search.rs @@ -127,10 +127,11 @@ fn _grid_search(grid: Grid, size: usize, mut parent: TestPath<'_>) { // are correct. let index = setup_grid_search(grid, size); - let params = Knn::new(10, 10, Some(beam_width)).unwrap(); + let k_value = 10; + let params = Knn::new(10, Some(beam_width)).unwrap(); let context = test_provider::Context::new(); - let mut neighbors = vec![Neighbor::::default(); params.k_value().get()]; + let mut neighbors = vec![Neighbor::::default(); k_value]; let graph::index::SearchStats { cmps, hops, @@ -147,7 +148,7 @@ fn _grid_search(grid: Grid, size: usize, mut parent: TestPath<'_>) { .unwrap(); assert!( - result_count.into_usize() <= params.k_value().get(), + result_count.into_usize() <= k_value, "grid search should not return more than the requested number of neighbors", ); diff --git a/diskann/src/graph/test/cases/inline.rs b/diskann/src/graph/test/cases/inline.rs index cf8c42360..33e0d566c 100644 --- a/diskann/src/graph/test/cases/inline.rs +++ b/diskann/src/graph/test/cases/inline.rs @@ -348,7 +348,7 @@ fn run_inline_on_grid( adaptive_l: Option, ) -> InlineFilterBaseline { let rt = current_thread_runtime(); - let inline = InlineFilterSearch::new(Knn::new_default(k, l).unwrap(), adaptive_l); + let inline = InlineFilterSearch::new(Knn::new_default(l).unwrap(), adaptive_l); let mut ids = vec![0u32; k]; let mut distances = vec![0.0f32; k]; @@ -391,7 +391,7 @@ fn inline_search_returns_only_final_level_matches() { let filter = LevelLabelProvider::new(); let k = 8; let l = 32; - let inline = InlineFilterSearch::new(Knn::new_default(k, l).unwrap(), None); + let inline = InlineFilterSearch::new(Knn::new_default(l).unwrap(), None); let mut ids = vec![0u32; k]; let mut distances = vec![0.0f32; k]; @@ -448,7 +448,7 @@ fn inline_search_three_level_no_adaptive_l_with_l1_finds_no_matches() { let filter = LevelLabelProvider::new(); let k = 1; let l = 1; - let inline = InlineFilterSearch::new(Knn::new_default(k, l).unwrap(), None); + let inline = InlineFilterSearch::new(Knn::new_default(l).unwrap(), None); let mut ids = vec![0u32; k]; let mut distances = vec![0.0f32; k]; @@ -497,11 +497,11 @@ fn inline_search_three_level_adaptive_l_with_l1_finds_matches() { let index = build_three_level_index(); let filter = LevelLabelProvider::new(); - let k = 1; let l = 1; let adaptive_l = AdaptiveL::new(1, 16.0).unwrap(); - let inline = InlineFilterSearch::new(Knn::new_default(k, l).unwrap(), Some(adaptive_l)); + let inline = InlineFilterSearch::new(Knn::new_default(l).unwrap(), Some(adaptive_l)); + let k = 1; let mut ids = vec![0u32; k]; let mut distances = vec![0.0f32; k]; let mut buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); @@ -624,11 +624,11 @@ fn inline_search_reaches_matches_through_non_matching_nodes() { let filter = EvenFilter; - let k = 5; let l = 20; - let search_params = Knn::new_default(k, l).unwrap(); + let search_params = Knn::new_default(l).unwrap(); let inline = InlineFilterSearch::new(search_params, None); + let k = 5; let mut ids = vec![0u32; k]; let mut distances = vec![0.0f32; k]; let mut buffer = search_output_buffer::IdDistance::new(&mut ids, &mut distances); @@ -645,8 +645,8 @@ fn inline_search_reaches_matches_through_non_matching_nodes() { let result_count = stats.result_count as usize; let baseline = InlineBaseline { - query: vec![2.0f32], k, + query: vec![2.0f32], l, result_count, results: ids[..result_count] diff --git a/diskann/src/graph/test/cases/multihop.rs b/diskann/src/graph/test/cases/multihop.rs index 1020964ff..be834722f 100644 --- a/diskann/src/graph/test/cases/multihop.rs +++ b/diskann/src/graph/test/cases/multihop.rs @@ -113,13 +113,12 @@ pub(super) fn build_1d_index( fn run( index: &DiskANNIndex, query: &[f32], - k: usize, l: usize, filter: &dyn labeled::QueryLabelProvider, ) -> (graph::index::SearchStats, Vec>) { let rt = current_thread_runtime(); rt.block_on(async { - let multihop = MultihopFilterSearch::new(Knn::new_default(k, l).unwrap()); + let multihop = MultihopFilterSearch::new(Knn::new_default(l).unwrap()); let mut neighbors = Vec::>::new(); let stats = index @@ -163,7 +162,7 @@ fn accept_all_finds_all_nodes() { 3, ); - let (stats, results) = run(&index, &[1.5], 3, 10, &AcceptAll); + let (stats, results) = run(&index, &[1.5], 10, &AcceptAll); let ids: Vec = results.iter().map(|n| n.id).collect(); assert!(ids.contains(&0), "node 0 should be found"); @@ -204,7 +203,7 @@ fn reject_triggers_two_hop_expansion() { ); let filter = EvenFilter; - let (stats, results) = run(&index, &[2.0], 5, 20, &filter); + let (stats, results) = run(&index, &[2.0], 20, &filter); let ids: Vec = results.iter().map(|n| n.id).collect(); @@ -252,7 +251,7 @@ fn reject_all_yields_nothing() { 2, ); - let (_stats, results) = run(&index, &[0.5], 5, 10, &RejectAll); + let (_stats, results) = run(&index, &[0.5], 10, &RejectAll); // Nothing should be present in the result, not even the start point since it does not // satisfy the predicate. @@ -349,7 +348,7 @@ fn two_hop_reaches_through_non_matching() { let k = 5; let l = 20; - let search_params = Knn::new_default(k, l).unwrap(); + let search_params = Knn::new_default(l).unwrap(); let multihop = MultihopFilterSearch::new(search_params); let mut ids = vec![0u32; k]; @@ -418,7 +417,7 @@ fn even_filtering_grid() { let k = 20; let l = 40; - let search_params = Knn::new_default(k, l).unwrap(); + let search_params = Knn::new_default(l).unwrap(); let multihop = MultihopFilterSearch::new(search_params); let mut ids = vec![0u32; k];