Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
72 changes: 27 additions & 45 deletions diskann-benchmark/src/flat/search.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,13 @@

//! Backend for flat-index (brute-force kNN) benchmarks.
//!
//! This exercises [`diskann::flat::FlatIndex::knn_search`] over an in-memory
//! This exercises [`diskann::flat::knn_search`] over an in-memory
//! provider, measuring recall and latency.

use std::{io::Write, num::NonZeroUsize, sync::Arc};

use diskann::{
flat::{DistancesUnordered, FlatIndex, SearchStrategy},
flat::{knn_search, DistancesUnordered, SearchStrategy},
graph::{glue::CopyIds, SearchOutputBuffer},
provider::{DataProvider, DefaultContext, HasId, NoopGuard},
utils::VectorRepr,
Expand Down Expand Up @@ -132,10 +132,9 @@ where
);
writeln!(output, " Loaded {} vectors of dimension {}", nrows, ncols)?;

// Build the provider and wrap in FlatIndex
// Build the provider.
let data = Arc::new(data);
let provider = InMemProvider { data: data.clone() };
let index = FlatIndex::new(provider);

// Load queries and groundtruth
let queries: Matrix<T> =
Expand Down Expand Up @@ -172,7 +171,7 @@ where
let mut results = Vec::new();

let searcher = Arc::new(Searcher {
index,
provider,
queries,
strategy: Strategy::new(metric),
});
Expand Down Expand Up @@ -223,72 +222,56 @@ impl<T: VectorRepr> Strategy<T> {
}

/// The visitor that iterates over all vectors in the provider.
struct Visitor<'a, T> {
struct Visitor<'a, T: VectorRepr> {
data: &'a Matrix<T>,
computer: T::QueryDistance,
}

impl<T: VectorRepr> HasId for Visitor<'_, T> {
type Id = u32;
}

impl<T: VectorRepr> DistancesUnordered<T::QueryDistance> for Visitor<'_, T> {
type ElementRef<'a> = &'a [T];
impl<T: VectorRepr> DistancesUnordered for Visitor<'_, T> {
type Error = diskann::error::Infallible;

fn distances_unordered<F>(
&mut self,
computer: &T::QueryDistance,
mut f: F,
) -> impl SendFuture<Result<(), Self::Error>>
fn distances_unordered<F>(&mut self, mut f: F) -> impl SendFuture<Result<(), Self::Error>>
where
F: Send + FnMut(Self::Id, f32),
{
async move {
for (i, vector) in self.data.row_iter().enumerate() {
let dist = computer.evaluate_similarity(vector);
let dist = self.computer.evaluate_similarity(vector);
f(i as u32, dist);
}
Ok(())
}
}
}

impl<T: VectorRepr> SearchStrategy<InMemProvider<T>, &[T]> for Strategy<T> {
type ElementRef<'a> = &'a [T];
type QueryComputer = T::QueryDistance;
type QueryComputerError = diskann::error::Infallible;
type Visitor<'a>
= Visitor<'a, T>
where
Self: 'a,
InMemProvider<T>: 'a;
impl<'a, T: VectorRepr> SearchStrategy<'a, InMemProvider<T>, &'a [T]> for Strategy<T> {
type Visitor = Visitor<'a, T>;
type Error = diskann::error::Infallible;

fn create_visitor<'a>(
fn create_visitor(
&'a self,
provider: &'a InMemProvider<T>,
_context: &'a DefaultContext,
) -> Result<Self::Visitor<'a>, Self::Error> {
query: &'a [T],
) -> Result<Self::Visitor, Self::Error> {
Ok(Visitor {
data: &provider.data,
computer: T::query_distance(query, self.metric),
})
}

fn build_query_computer(
&self,
query: &[T],
) -> Result<Self::QueryComputer, Self::QueryComputerError> {
Ok(T::query_distance(query, self.metric))
}
}

//////////////////////////////////////////
// benchmark_core::search::Search impl //
//////////////////////////////////////////

/// Wraps a [`FlatIndex`] and queries to implement [`search::Search`].
/// Wraps a flat-search provider and queries to implement [`search::Search`].
struct Searcher<T: VectorRepr> {
index: FlatIndex<InMemProvider<T>>,
provider: InMemProvider<T>,
queries: Matrix<T>,
strategy: Strategy<T>,
}
Expand Down Expand Up @@ -334,17 +317,16 @@ where
let context = DefaultContext;
let query = self.queries.row(index);

let stats = self
.index
.knn_search(
parameters.k,
&self.strategy,
CopyIds,
&context,
query,
buffer,
)
.await?;
let stats = knn_search(
&self.provider,
parameters.k,
&self.strategy,
CopyIds,
&context,
query,
buffer,
)
.await?;

Ok(Metrics {
comparisons: stats.cmps,
Expand Down
161 changes: 70 additions & 91 deletions diskann/src/flat/index.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,7 @@
* Licensed under the MIT license.
*/

//! [`FlatIndex`] — the index wrapper for a [`DataProvider`]
//! over which we do flat search.
//! Brute-force k-nearest-neighbor search over a [`DataProvider`].
use std::num::NonZeroUsize;

use diskann_utils::future::SendFuture;
Expand All @@ -28,74 +27,57 @@ pub struct SearchStats {
pub result_count: u32,
}

/// A thin wrapper around a [`DataProvider`] used for flat search.
#[derive(Debug)]
pub struct FlatIndex<P: DataProvider> {
/// The backing provider.
provider: P,
}

impl<P: DataProvider> FlatIndex<P> {
/// Construct a new [`FlatIndex`] around `provider`.
pub fn new(provider: P) -> Self {
Self { provider }
}

/// Borrow the underlying provider.
pub fn provider(&self) -> &P {
&self.provider
}

/// Brute-force k-nearest-neighbor flat search.
///
/// Streams every element produced by the strategy's visitor through the query
/// computer, keeps the best `k` candidates in a [`NeighborPriorityQueue`], then runs
/// `processor` over the survivors to populate `output`.
///
/// The post-processor [`SearchPostProcess::post_process`] outputs the number
/// of results that survive, which is returned as `SearchStats::result_count`.
pub fn knn_search<S, T, O, PP, OB>(
&self,
k: NonZeroUsize,
strategy: &S,
processor: PP,
context: &P::Context,
query: T,
output: &mut OB,
) -> impl SendFuture<ANNResult<SearchStats>>
where
S: SearchStrategy<P, T>,
T: Copy + Send + Sync,
O: Send,
PP: for<'a> SearchPostProcess<S::Visitor<'a>, T, O> + Send + Sync,
OB: SearchOutputBuffer<O> + Send + ?Sized,
{
async move {
let mut visitor = strategy
.create_visitor(&self.provider, context)
.into_ann_result()?;

let computer = strategy.build_query_computer(query).into_ann_result()?;

let k = k.get();
let mut queue = NeighborPriorityQueue::new(k);
let mut cmps: u32 = 0;

visitor
.distances_unordered(&computer, |id, dist| {
cmps += 1;
queue.insert(Neighbor::new(id, dist));
})
.await
.escalate("flat scan must complete to produce correct k-NN results")?;

let result_count = processor
.post_process(&mut visitor, query, queue.iter().take(k), output)
.await
.into_ann_result()? as u32;

Ok(SearchStats { cmps, result_count })
}
/// Brute-force k-nearest-neighbor search over a borrowed provider.
///
/// Streams every distance produced by the strategy's query-aware visitor, keeps the best `k`
/// candidates in a [`NeighborPriorityQueue`], then runs `processor` over the survivors
/// to populate `output`.
///
/// # Errors
///
/// Returns an error if visitor construction, distance scanning, or result
/// post-processing fails. Distance-scan errors are escalated because a
/// partial flat scan cannot produce correct k-nearest-neighbor results.
pub fn knn_search<'a, P, S, T, O, PP, OB>(
Comment thread
partychen marked this conversation as resolved.
Outdated
provider: &'a P,
k: NonZeroUsize,
strategy: &'a S,
processor: PP,
context: &'a P::Context,
query: T,
output: &mut OB,
) -> impl SendFuture<ANNResult<SearchStats>>
where
P: DataProvider,
S: SearchStrategy<'a, P, T>,
T: Copy + Send + Sync,
O: Send,
PP: SearchPostProcess<S::Visitor, T, O> + Send + Sync,
OB: SearchOutputBuffer<O> + Send + ?Sized,
{
async move {
let mut visitor = strategy
.create_visitor(provider, context, query)
.into_ann_result()?;

let k = k.get();
let mut queue = NeighborPriorityQueue::new(k);
let mut cmps: u32 = 0;

visitor
.distances_unordered(|id, dist| {
cmps += 1;
queue.insert(Neighbor::new(id, dist));
})
.await
.escalate("flat scan must complete to produce correct k-NN results")?;

let result_count = processor
.post_process(&mut visitor, query, queue.iter().take(k), output)
.await
.into_ann_result()? as u32;

Ok(SearchStats { cmps, result_count })
}
}

Expand All @@ -105,30 +87,27 @@ impl<P: DataProvider> FlatIndex<P> {

#[cfg(test)]
mod tests {
use crate::flat::{
FlatIndex,
test::{
harness::{CopyIdsOracle, EvenIdsOnlyOracle, KnnOracleRun, OracleProcessor},
provider::{self as flat_provider},
},
use crate::flat::test::{
harness::{CopyIdsOracle, EvenIdsOnlyOracle, KnnOracleRun, OracleProcessor},
provider::{self as flat_provider},
};
use crate::graph::test::synthetic::Grid;

fn fixture(grid: Grid, size: usize) -> (FlatIndex<flat_provider::Provider>, usize) {
fn fixture(grid: Grid, size: usize) -> (flat_provider::Provider, usize) {
let provider = flat_provider::Provider::grid(grid, size).unwrap();
let len = provider.len();
(FlatIndex::new(provider), len)
(provider, len)
}

/// `knn_search` returns a `Send` future, and a shared `&FlatIndex` can serve
/// `knn_search` returns a `Send` future, and a shared provider can serve
/// many concurrent searches on a multi-threaded runtime, each producing the
/// correct output independently.
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn multithreaded_knn_search() {
use std::sync::Arc;

let (index, len) = fixture(Grid::Two, 4);
let index = Arc::new(index);
let (provider, len) = fixture(Grid::Two, 4);
let provider = Arc::new(provider);

// Mix of corner, axis-aligned, and off-grid queries; k spans 1..=len.
let cases: &[(&[f32], usize)] = &[
Expand All @@ -145,20 +124,20 @@ mod tests {
/// Spawn every `(query, k)` case under `oracle` onto `set`.
fn spawn_cases<O>(
set: &mut tokio::task::JoinSet<(Vec<f32>, usize, KnnOracleRun)>,
index: &Arc<FlatIndex<flat_provider::Provider>>,
provider: &Arc<flat_provider::Provider>,
oracle: O,
cases: &[(&[f32], usize)],
) where
O: OracleProcessor + Copy + Send + Sync + 'static,
{
for (query, k) in cases {
let index = Arc::clone(index);
let provider = Arc::clone(provider);
let query: Vec<f32> = query.to_vec();
let k = *k;
set.spawn(async move {
let outcome = KnnOracleRun::run(
&index,
&flat_provider::Strategy::new(index.provider().dim()),
&provider,
&flat_provider::Strategy::new(provider.dim()),
&oracle,
&query,
k,
Expand All @@ -171,8 +150,8 @@ mod tests {
}

let mut set = tokio::task::JoinSet::new();
spawn_cases(&mut set, &index, CopyIdsOracle, cases);
spawn_cases(&mut set, &index, EvenIdsOnlyOracle, cases);
spawn_cases(&mut set, &provider, CopyIdsOracle, cases);
spawn_cases(&mut set, &provider, EvenIdsOnlyOracle, cases);

while let Some(joined) = set.join_next().await {
let (query, k, outcome) = joined.expect("task panicked");
Expand All @@ -199,8 +178,8 @@ mod tests {
for transient_ids in [&[0u32][..], &[3][..], &[1, 2, 5][..]] {
let strategy =
flat_provider::Strategy::with_transient(2, transient_ids.iter().copied());
let (index, _) = fixture(Grid::Two, 3);
let err = KnnOracleRun::run_sync(&index, &strategy, &CopyIdsOracle, &[1.0, 0.0], 4)
let (provider, _) = fixture(Grid::Two, 3);
let err = KnnOracleRun::run_sync(&provider, &strategy, &CopyIdsOracle, &[1.0, 0.0], 4)
.expect_err("transient error during full scan must escalate");

let msg = format!("{err}");
Expand All @@ -217,8 +196,8 @@ mod tests {
/// Run `knn_search` via the harness, assert it fails, and check the error
/// message contains `expected_msg`.
fn assert_search_error(strategy: &flat_provider::Strategy, query: &[f32], expected_msg: &str) {
let (index, _) = fixture(Grid::Two, 3);
let err = KnnOracleRun::run_sync(&index, strategy, &CopyIdsOracle, query, 4)
let (provider, _) = fixture(Grid::Two, 3);
let err = KnnOracleRun::run_sync(&provider, strategy, &CopyIdsOracle, query, 4)
.expect_err("expected knn_search to fail");

let msg = format!("{err}");
Expand Down
Loading
Loading