Skip to content
Open
4 changes: 2 additions & 2 deletions .github/workflows/disk-benchmarks.yml
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,7 @@ jobs:
working-directory: baseline
run: |
cargo run -p diskann-benchmark --features disk-index --release -- \
run --input-file ../diskann_rust/${{ env.PERF_INPUTS }}/${{ matrix.config }} \
run --input-file ${{ env.PERF_INPUTS }}/${{ matrix.config }} \
--output-file target/tmp/${{ matrix.dataset }}_baseline.json

- name: Run current branch benchmark
Expand Down Expand Up @@ -144,4 +144,4 @@ jobs:
path: |
diskann_rust/target/tmp/${{ matrix.dataset }}_target.json
baseline/target/tmp/${{ matrix.dataset }}_baseline.json
retention-days: 30
retention-days: 30
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": {
"mode": "graph",
"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": {
"mode": "graph",
"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": {
"mode": "flat",
"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": { "mode": "graph" },
"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": { "mode": "flat" },
"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": { "mode": "graph" },
"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": { "mode": "graph" },
"distance": "inner_product"
}
}
}
Expand Down
94 changes: 75 additions & 19 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,58 @@ 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>>,
) -> SearchMode<'a> {
match mode {
DiskSearchMode::Flat { .. } => match vector_filter {
None => SearchMode::flat(),
Some(vector_filter) => {
SearchMode::flat_filtered(move |vid: &u32| vector_filter.contains(vid))
}
},
DiskSearchMode::Graph {
adaptive_l,
post_processor,
..
} => {
let adaptive_l = 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 (post_processor, adaptive_l, vector_filter) {
(Some(TopkPostProcessor::DeterminantDiversity(params)), _, None) => {
SearchMode::diverse_graph(*params)
}
(Some(TopkPostProcessor::DeterminantDiversity(params)), _, Some(vector_filter)) => {
SearchMode::diverse_graph_filtered(
move |vid: &u32| vector_filter.contains(vid),
*params,
)
}
(None, Some(adaptive_l), None) => {
SearchMode::inline_filter(|_| true, Some(adaptive_l))
}
(None, Some(adaptive_l), Some(vector_filter)) => SearchMode::inline_filter(
move |vid: &u32| vector_filter.contains(vid),
Some(adaptive_l),
),
(None, None, None) => SearchMode::graph(),
(None, None, Some(vector_filter)) => {
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 +239,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,24 +319,20 @@ 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)| {
// 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(),
);
|(((((query_index, q), id_chunk), dist_chunk), stats), rc)| {
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);

match searcher.search(
q,
Expand Down Expand Up @@ -349,9 +405,9 @@ where
num_threads: search_params.num_threads,
beam_width: search_params.beam_width,
recall_at: search_params.recall_at,
is_flat_search: search_params.search_mode.is_flat_search,
is_flat_search: matches!(search_params.search_mode, DiskSearchMode::Flat { .. }),
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