Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 9 additions & 8 deletions diskann-benchmark/example/disk-index-determinant-diversity.json
Original file line number Diff line number Diff line change
Expand Up @@ -27,14 +27,15 @@
"beam_width": 4,
"recall_at": 10,
"num_threads": 1,
"is_flat_search": false,
"distance": "squared_l2",
"vector_filters_file": null,
"post_processor": {
"type": "determinant-diversity",
"power": 2.0,
"eta": 1.0
}
"search_mode": {
"is_flat_search": false,
"post_processor": {
"type": "determinant-diversity",
"power": 2.0,
"eta": 1.0
}
},
"distance": "squared_l2"
}
}
}
Expand Down
16 changes: 10 additions & 6 deletions diskann-benchmark/example/disk-index-filter.json
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,11 @@
"beam_width": 4,
"recall_at": 10,
"num_threads": 1,
"is_flat_search": false,
"distance": "squared_l2",
"vector_filters_file": "disk_index_10pts_idx_uint32_range_res_r_100000.bin"
"search_mode": {
"is_flat_search": false,
"vector_filters_file": "disk_index_10pts_idx_uint32_range_res_r_100000.bin"
},
"distance": "squared_l2"
}
}
},
Expand Down Expand Up @@ -57,9 +59,11 @@
"beam_width": 4,
"recall_at": 10,
"num_threads": 1,
"is_flat_search": true,
"distance": "squared_l2",
"vector_filters_file": "disk_index_10pts_idx_uint32_range_res_r_100000.bin"
"search_mode": {
"is_flat_search": true,
"vector_filters_file": "disk_index_10pts_idx_uint32_range_res_r_100000.bin"
},
"distance": "squared_l2"
}
}
}
Expand Down
10 changes: 4 additions & 6 deletions diskann-benchmark/example/disk-index.json
Original file line number Diff line number Diff line change
Expand Up @@ -27,9 +27,8 @@
"beam_width": 4,
"recall_at": 10,
"num_threads": 1,
"is_flat_search": false,
"distance": "squared_l2",
"vector_filters_file": null
"search_mode": { "is_flat_search": false },
"distance": "squared_l2"
}
}
},
Expand All @@ -48,9 +47,8 @@
"beam_width": 4,
"recall_at": 10,
"num_threads": 1,
"is_flat_search": true,
"distance": "squared_l2",
"vector_filters_file": null
"search_mode": { "is_flat_search": true },
"distance": "squared_l2"
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,9 +29,8 @@
"beam_width": 4,
"recall_at": 100,
"num_threads": 4,
"is_flat_search": false,
"distance": "squared_l2",
"vector_filters_file": null
"search_mode": { "is_flat_search": false },
"distance": "squared_l2"
}
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -29,9 +29,8 @@
"beam_width": 4,
"recall_at": 100,
"num_threads": 4,
"is_flat_search": false,
"distance": "inner_product",
"vector_filters_file": null
"search_mode": { "is_flat_search": false },
"distance": "inner_product"
}
}
}
Expand Down
83 changes: 69 additions & 14 deletions diskann-benchmark/src/disk_index/search.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ use std::{collections::HashSet, fmt, sync::atomic::AtomicBool, time::Instant};
use opentelemetry::{global, trace::Span, trace::Tracer};
use opentelemetry_sdk::trace::SdkTracerProvider;

use diskann::graph;
use diskann::utils::VectorRepr;
use diskann_benchmark_runner::{files::InputFile, utils::MicroSeconds};
use diskann_disk::{
Expand Down Expand Up @@ -36,7 +37,8 @@ use serde::{Deserialize, Serialize};

use crate::{
disk_index::json_spancollector::JsonSpanCollector,
inputs::disk::{DiskIndexLoad, DiskSearchPhase},
inputs::disk::{DiskIndexLoad, DiskSearchMode, DiskSearchPhase},
inputs::post_processor::TopkPostProcessor,
utils::{datafiles, SimilarityMeasure},
};

Expand Down Expand Up @@ -158,6 +160,51 @@ impl DiskSearchResult {
}
}

/// Construct the disk [`SearchMode`] from the JSON-driven [`DiskSearchMode`]
/// config plus the per-query filter and post-processor supplied at search time.
fn build_search_mode<'a>(
mode: &'a DiskSearchMode,
vector_filter: Option<&'a HashSet<u32>>,
post_processor: Option<&TopkPostProcessor>,
) -> SearchMode<'a> {
let adaptive_l = mode.adaptive_l.as_ref().map(|adaptive_l| {
graph::search::AdaptiveL::new(adaptive_l.sample_count.into(), adaptive_l.scale_factor)
.expect("validated adaptive L must construct")
});

match (
mode.is_flat_search,
vector_filter,
post_processor,
adaptive_l,
) {
(true, None, _, _) => SearchMode::flat(),
(true, Some(vector_filter), _, _) => {
SearchMode::flat_filtered(move |vid: &u32| vector_filter.contains(vid))
}
(false, None, Some(TopkPostProcessor::DeterminantDiversity(params)), _) => {
SearchMode::diverse_graph(*params)
}
(false, Some(vector_filter), Some(TopkPostProcessor::DeterminantDiversity(params)), _) => {
SearchMode::diverse_graph_filtered(
move |vid: &u32| vector_filter.contains(vid),
*params,
)
}
(false, None, None, Some(adaptive_l)) => {
SearchMode::inline_filter(|_| true, Some(adaptive_l))
}
(false, Some(vector_filter), None, Some(adaptive_l)) => SearchMode::inline_filter(
move |vid: &u32| vector_filter.contains(vid),
Some(adaptive_l),
),
(false, None, None, None) => SearchMode::graph(),
(false, Some(vector_filter), None, None) => {
SearchMode::graph_filtered(move |vid: &u32| vector_filter.contains(vid))
}
}
}

pub(super) fn search_disk_index<T, StorageType>(
index_load: &DiskIndexLoad,
search_params: &DiskSearchPhase,
Expand Down Expand Up @@ -185,21 +232,27 @@ where
let num_queries = queries.nrows();

// Load the vector filters
let vector_filters = match &search_params.vector_filters_file {
let vector_filters = match &search_params.search_mode.vector_filters_file {
Some(vector_filters_file) => {
let vector_filters_file = vector_filters_file.to_string_lossy().to_string();
search_index_utils::load_vector_filters(storage_provider, &vector_filters_file)?
Some(search_index_utils::load_vector_filters(
storage_provider,
&vector_filters_file,
)?)
}
None => vec![HashSet::<u32>::new(); num_queries],
None => None,
};

if vector_filters.len() != num_queries {
if vector_filters
.as_ref()
.is_some_and(|filters| filters.len() != num_queries)
{
anyhow::bail!("Mismatch in query and vector filter sizes");
}

// Prepare ground truth context
let gt_context = prepare_ground_truth_context(
search_params.vector_filters_file.is_some(),
search_params.search_mode.vector_filters_file.is_some(),
&search_params.groundtruth,
search_params.recall_at,
storage_provider,
Expand Down Expand Up @@ -259,23 +312,25 @@ where

let zipped = queries
.par_row_iter()
.zip(vector_filters.par_iter())
.enumerate()
.zip(result_ids.par_chunks_mut(search_params.recall_at as usize))
.zip(result_dists.par_chunks_mut(search_params.recall_at as usize))
.zip(statistics_vec.par_iter_mut())
.zip(result_counts.par_iter_mut());

zipped.for_each_in_pool(
pool.as_ref(),
|(((((q, vf), id_chunk), dist_chunk), stats), rc)| {
|(((((query_index, q), id_chunk), dist_chunk), stats), rc)| {
// Construct the SearchMode from the JSON-driven
// `adaptive_l` is now encapsulated in `DiskSearchMode`, so the
// benchmark only supplies the per-query filter and post-processor.
let has_filter = search_params.vector_filters_file.is_some();
let mode: SearchMode<'_> = search_params.search_mode.search_mode(
has_filter,
vf,
search_params.post_processor.as_ref(),
let vector_filter = vector_filters
.as_ref()
.and_then(|filters| filters.get(query_index));
let mode: SearchMode<'_> = build_search_mode(
&search_params.search_mode,
vector_filter,
search_params.search_mode.post_processor.as_ref(),
);

match searcher.search(
Expand Down Expand Up @@ -351,7 +406,7 @@ where
recall_at: search_params.recall_at,
is_flat_search: search_params.search_mode.is_flat_search,
distance: search_params.distance,
uses_vector_filters: search_params.vector_filters_file.is_some(),
uses_vector_filters: search_params.search_mode.vector_filters_file.is_some(),
num_nodes_to_cache: search_params.num_nodes_to_cache,
search_results_per_l,
span_metrics,
Expand Down
Loading
Loading