Skip to content
Open
Show file tree
Hide file tree
Changes from 2 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
63 changes: 56 additions & 7 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,52 @@ 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,
has_vector_filters: bool,
vector_filter: &'a HashSet<u32>,
Comment thread
dyhyfu marked this conversation as resolved.
Outdated
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,
has_vector_filters,
post_processor,
adaptive_l,
) {
(true, false, _, _) => SearchMode::flat(),
Comment thread
dyhyfu marked this conversation as resolved.
Outdated
(true, true, _, _) => {
SearchMode::flat_filtered(move |vid: &u32| vector_filter.contains(vid))
}
(false, false, Some(TopkPostProcessor::DeterminantDiversity(params)), _) => {
SearchMode::diverse_graph(*params)
}
(false, true, Some(TopkPostProcessor::DeterminantDiversity(params)), _) => {
SearchMode::diverse_graph_filtered(
move |vid: &u32| vector_filter.contains(vid),
*params,
)
}
(false, false, None, Some(adaptive_l)) => {
SearchMode::inline_filter(|_| true, Some(adaptive_l))
}
(false, true, None, Some(adaptive_l)) => SearchMode::inline_filter(
move |vid: &u32| vector_filter.contains(vid),
Some(adaptive_l),
),
(false, false, None, None) => SearchMode::graph(),
(false, true, 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,7 +233,7 @@ 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)?
Expand All @@ -199,7 +247,7 @@ where

// 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 @@ -271,11 +319,12 @@ where
// 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(
let has_filter = search_params.search_mode.vector_filters_file.is_some();
let mode: SearchMode<'_> = build_search_mode(
&search_params.search_mode,
has_filter,
vf,
Comment thread
dyhyfu marked this conversation as resolved.
Outdated
search_params.post_processor.as_ref(),
search_params.search_mode.post_processor.as_ref(),
);

match searcher.search(
Expand Down Expand Up @@ -351,7 +400,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