Skip to content
Merged
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
103 changes: 57 additions & 46 deletions diskann-benchmark-core/src/internal/buffer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,10 @@
* Licensed under the MIT license.
*/

use diskann::graph::{SearchOutputBuffer, search_output_buffer::BufferState};
use diskann::{
graph::{SearchOutputBuffer, search_output_buffer::BufferState},
neighbor::Neighbor,
};

/// A [`SearchOutputBuffer`] implementation that either references a slice in-place for
/// fixed sized outputs or references a growable vector.
Expand Down Expand Up @@ -58,7 +61,8 @@ impl<I, D> SearchOutputBuffer<I, D> for Buffer<'_, I> {
<Buffer<I>>::current_len(self)
}

fn push(&mut self, id: I, _distance: D) -> BufferState {
fn push(&mut self, neighbor: Neighbor<I, D>) -> BufferState {
let (id, _) = neighbor.as_tuple();
match &mut self.0 {
Inner::Slice { slice, written } => match slice.get_mut(*written) {
Some(slot) => {
Expand All @@ -81,14 +85,14 @@ impl<I, D> SearchOutputBuffer<I, D> for Buffer<'_, I> {

fn extend<Itr>(&mut self, itr: Itr) -> usize
where
Itr: IntoIterator<Item = (I, D)>,
Itr: IntoIterator<Item = Neighbor<I, D>>,
{
match &mut self.0 {
Inner::Slice { slice, written } => match slice.get_mut(*written..) {
Some(left) => {
let count = std::iter::zip(left.iter_mut(), itr)
.map(|(dst, src)| {
*dst = src.0;
*dst = src.as_tuple().0;
})
.count();
*written += count;
Expand All @@ -98,7 +102,7 @@ impl<I, D> SearchOutputBuffer<I, D> for Buffer<'_, I> {
},
Inner::Vec(vec) => {
let before = vec.len();
vec.extend(itr.into_iter().map(|i| i.0));
vec.extend(itr.into_iter().map(|i| i.as_tuple().0));
vec.len() - before
}
}
Expand All @@ -120,6 +124,10 @@ mod tests {
buffer.size_hint()
}

fn from_tuple<I, D>((id, distance): (I, D)) -> Neighbor<I, D> {
Neighbor::new(id, distance)
}

#[test]
fn test_slice_buffer_creation() {
let mut data = [0u32; 5];
Expand All @@ -144,7 +152,7 @@ mod tests {
let mut data = [0u32; 5];
let mut buffer = Buffer::slice(&mut data);

assert_eq!(buffer.push(42, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(42, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 1);
assert_eq!(data, [42, 0, 0, 0, 0]);
}
Expand All @@ -157,27 +165,27 @@ mod tests {
assert_eq!(buffer.current_len(), 0);
assert_eq!(size_hint(&buffer), Some(5));

assert_eq!(buffer.push(100, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(100, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 1);
assert_eq!(size_hint(&buffer), Some(4));

assert_eq!(buffer.push(200, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(200, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 2);
assert_eq!(size_hint(&buffer), Some(3));

assert_eq!(buffer.push(300, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(300, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 3);
assert_eq!(size_hint(&buffer), Some(2));

assert_eq!(buffer.push(400, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(400, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 4);
assert_eq!(size_hint(&buffer), Some(1));

assert_eq!(buffer.push(500, 0.0), BufferState::Full);
assert_eq!(buffer.push(Neighbor::new(500, 0.0)), BufferState::Full);
assert_eq!(buffer.current_len(), 5);
assert_eq!(size_hint(&buffer), Some(0));

assert_eq!(buffer.push(600, 0.0), BufferState::Full);
assert_eq!(buffer.push(Neighbor::new(600, 0.0)), BufferState::Full);
assert_eq!(buffer.current_len(), 5);
assert_eq!(size_hint(&buffer), Some(0));

Expand All @@ -190,7 +198,7 @@ mod tests {
let mut data = [0u32; 0];
let mut buffer = Buffer::slice(&mut data);

assert_eq!(buffer.push(42, 0.0), BufferState::Full);
assert_eq!(buffer.push(Neighbor::new(42, 0.0)), BufferState::Full);
}

#[test]
Expand All @@ -203,13 +211,13 @@ mod tests {
"vector-type buffers have no upper bound"
);

assert_eq!(buffer.push(42, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(42, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 1);

assert_eq!(buffer.push(50, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(50, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 2);

assert_eq!(buffer.push(3, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(3, 0.0)), BufferState::Available);
assert_eq!(buffer.current_len(), 3);

assert_eq!(&vec, &[42, 50, 3]);
Expand All @@ -220,7 +228,7 @@ mod tests {
let mut data = [0u32; 5];
let mut buffer = Buffer::slice(&mut data);

let items = vec![(10, 1.0), (20, 2.0), (30, 3.0)];
let items = [(10, 1.0), (20, 2.0), (30, 3.0)].map(from_tuple);
let count = buffer.extend(items);

assert_eq!(count, 3);
Expand All @@ -233,7 +241,7 @@ mod tests {
let mut data = [0u32; 3];
let mut buffer = Buffer::slice(&mut data);

let items = vec![(10, 1.0), (20, 2.0), (30, 3.0), (40, 4.0), (50, 5.0)];
let items = [(10, 1.0), (20, 2.0), (30, 3.0), (40, 4.0), (50, 5.0)].map(from_tuple);
let count = buffer.extend(items);

// Only first 3 items should be written
Expand All @@ -247,7 +255,7 @@ mod tests {
let mut vec = Vec::<u32>::new();
let mut buffer = Buffer::vector(&mut vec);

let items = vec![(100, 1.0), (200, 2.0)];
let items = [(100, 1.0), (200, 2.0)].map(from_tuple);
let count = buffer.extend(items);

assert_eq!(count, 2);
Expand All @@ -260,14 +268,14 @@ mod tests {
let mut vec = vec![1u32, 2, 3, 4, 5];
let mut buffer = Buffer::vector(&mut vec);

let items = vec![(10, 1.0), (20, 2.0)];
let items = [(10, 1.0), (20, 2.0)].map(from_tuple);
assert_eq!(buffer.extend(items), 2);

let items = vec![(21, 1.0), (22, 2.0)];
let items = [(21, 1.0), (22, 2.0)].map(from_tuple);
assert_eq!(buffer.extend(items), 2);

assert_eq!(
buffer.extend::<[(u32, f32); 0]>([]),
buffer.extend::<[Neighbor<u32, f32>; 0]>([]),
0,
"empty iterator should add nothing"
);
Expand All @@ -281,27 +289,30 @@ mod tests {
// Push then extend
let mut data = [0u32; 5];
let mut buffer = Buffer::slice(&mut data);
assert_eq!(buffer.push(1, 0.0), BufferState::Available);
assert_eq!(buffer.push(2, 0.0), BufferState::Available);
assert_eq!(buffer.extend(vec![(3, 0.0), (4, 0.0)]), 2);
assert_eq!(buffer.push(Neighbor::new(1, 0.0)), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(2, 0.0)), BufferState::Available);
assert_eq!(buffer.extend([(3, 0.0), (4, 0.0)].map(from_tuple)), 2);
assert_eq!(data, [1, 2, 3, 4, 0]);

// Extend then push to fill
let mut data = [0u32; 5];
let mut buffer = Buffer::slice(&mut data);
buffer.extend(vec![(10, 0.0), (20, 0.0)]);
assert_eq!(buffer.push(30, 0.0), BufferState::Available);
assert_eq!(buffer.push(40, 0.0), BufferState::Available);
assert_eq!(buffer.push(50, 0.0), BufferState::Full);
buffer.extend([(10, 0.0), (20, 0.0)].map(from_tuple));
assert_eq!(buffer.push(Neighbor::new(30, 0.0)), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(40, 0.0)), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(50, 0.0)), BufferState::Full);
assert_eq!(data, [10, 20, 30, 40, 50]);

// Interleaved operations
let mut data = [0u32; 6];
let mut buffer = Buffer::slice(&mut data);
assert_eq!(buffer.push(1, 0.0), BufferState::Available);
assert_eq!(buffer.extend(vec![(2, 0.0), (3, 0.0)]), 2);
assert_eq!(buffer.push(4, 0.0), BufferState::Available);
assert_eq!(buffer.extend(vec![(5, 0.0), (6, 0.0), (7, 0.0)]), 2);
assert_eq!(buffer.push(Neighbor::new(1, 0.0)), BufferState::Available);
assert_eq!(buffer.extend([(2, 0.0), (3, 0.0)].map(from_tuple)), 2);
assert_eq!(buffer.push(Neighbor::new(4, 0.0)), BufferState::Available);
assert_eq!(
buffer.extend([(5, 0.0), (6, 0.0), (7, 0.0)].map(from_tuple)),
2
);
assert_eq!(data, [1, 2, 3, 4, 5, 6]);
}

Expand All @@ -310,25 +321,25 @@ mod tests {
// Extend fills buffer, push returns Full
let mut data = [0u32; 2];
let mut buffer = Buffer::slice(&mut data);
buffer.extend(vec![(1, 0.0), (2, 0.0)]);
assert_eq!(buffer.push(99, 0.0), BufferState::Full);
buffer.extend([(1, 0.0), (2, 0.0)].map(from_tuple));
assert_eq!(buffer.push(Neighbor::new(99, 0.0)), BufferState::Full);
assert_eq!(data, [1, 2]);

// Push fills buffer, extend returns 0
let mut data = [0u32; 2];
let mut buffer = Buffer::slice(&mut data);
assert_eq!(buffer.push(1, 0.0), BufferState::Available);
assert_eq!(buffer.push(2, 0.0), BufferState::Full);
assert_eq!(buffer.extend(vec![(99, 0.0)]), 0);
assert_eq!(buffer.push(Neighbor::new(1, 0.0)), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(2, 0.0)), BufferState::Full);
assert_eq!(buffer.extend([Neighbor::new(99, 0.0)]), 0);
assert_eq!(data, [1, 2]);

// Extend truncates when exceeding remaining capacity after push
let mut data = [0u32; 4];
let mut buffer = Buffer::slice(&mut data);
assert_eq!(buffer.push(1, 0.0), BufferState::Available);
assert_eq!(buffer.push(2, 0.0), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(1, 0.0)), BufferState::Available);
assert_eq!(buffer.push(Neighbor::new(2, 0.0)), BufferState::Available);
assert_eq!(
buffer.extend(vec![(3, 0.0), (4, 0.0), (5, 0.0), (6, 0.0)]),
buffer.extend([(3, 0.0), (4, 0.0), (5, 0.0), (6, 0.0)].map(from_tuple)),
2
);
assert_eq!(data, [1, 2, 3, 4]);
Expand All @@ -340,11 +351,11 @@ mod tests {
let mut buffer = Buffer::vector(&mut vec);

// Interleave push and extend - vec has no capacity limit
assert_eq!(buffer.push(1, 0.0), BufferState::Available);
assert_eq!(buffer.extend(vec![(2, 0.0), (3, 0.0)]), 2);
assert_eq!(buffer.push(4, 0.0), BufferState::Available);
assert_eq!(buffer.extend::<[(u32, f32); 0]>([]), 0); // empty extend
assert_eq!(buffer.extend(vec![(5, 0.0)]), 1);
assert_eq!(buffer.push(Neighbor::new(1, 0.0)), BufferState::Available);
assert_eq!(buffer.extend([(2, 0.0), (3, 0.0)].map(from_tuple)), 2);
assert_eq!(buffer.push(Neighbor::new(4, 0.0)), BufferState::Available);
assert_eq!(buffer.extend::<[Neighbor<u32, f32>; 0]>([]), 0); // empty extend
assert_eq!(buffer.extend(vec![Neighbor::new(5, 0.0)]), 1);

assert_eq!(&vec, &[1, 2, 3, 4, 5]);
}
Expand Down
5 changes: 4 additions & 1 deletion diskann-benchmark-core/src/search/api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -584,6 +584,8 @@ mod tests {

use std::hash::{self, Hash, Hasher};

use diskann::neighbor::Neighbor;

// We intentionally do not derive `Clone` to ensure that it is not needed
// in the implementations.
#[derive(Debug)]
Expand Down Expand Up @@ -675,7 +677,8 @@ mod tests {
O: graph::SearchOutputBuffer<Self::Id> + Send,
{
let count = self.count(index, params);
let set = buffer.extend((0..count).map(|i| (self.format(index, i), i as f32)));
let set =
buffer.extend((0..count).map(|i| Neighbor::new(self.format(index, i), i as f32)));
assert_eq!(set, count);
Ok(count)
}
Expand Down
9 changes: 4 additions & 5 deletions diskann-bftree/src/provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@ use diskann::{
workingset::map,
AdjacencyList, SearchOutputBuffer,
},
neighbor::Neighbor,
neighbor::{self, Neighbor},
provider::{DataProvider, DefaultContext, Delete, ElementStatus, HasId, NoopGuard, SetElement},
utils::{IntoUsize, VectorRepr},
ANNError, ANNResult,
Expand Down Expand Up @@ -1536,15 +1536,15 @@ where
let provider = accessor.provider;
let f = T::distance(provider.metric, Some(provider.full_vectors.dim()));

let mut reranked: Vec<(u32, f32)> = Vec::new();
let mut reranked = Vec::new();
for n in candidates {
match provider
.full_vectors
.get_vector_sync(n.id().into_usize())
.allow_transient("stale candidate during rerank")
{
Ok(Some(vec)) => {
reranked.push((*n.id(), f.evaluate_similarity(query, &vec)));
reranked.push(Neighbor::new(*n.id(), f.evaluate_similarity(query, &vec)));
}
Ok(None) => {
// Transient (deleted/missing) — skip this candidate.
Expand All @@ -1553,8 +1553,7 @@ where
}
}

reranked
.sort_unstable_by(|a, b| (a.1).partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
reranked.sort_unstable_by(neighbor::ord::fast_distance);
std::future::ready(Ok(output.extend(reranked)))
}
}
Expand Down
18 changes: 9 additions & 9 deletions diskann-disk/src/search/provider/disk_provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ use diskann::{
search::{AdaptiveL, InlineFilterSearch, Knn},
search_output_buffer, DiskANNIndex,
},
neighbor::{Neighbor, NeighborPriorityQueue},
neighbor::{self, Neighbor, NeighborPriorityQueue},
provider::{DataProvider, DefaultContext, HasId, NoopGuard},
utils::{IntoUsize, VectorRepr},
ANNError, ANNResult,
Expand Down Expand Up @@ -349,10 +349,10 @@ where
let provider = accessor.provider;

let mut uncached_ids = Vec::new();
let mut reranked = {
let mut reranked: Vec<_> = {
let mut process = |n: u32| {
if let Some(entry) = accessor.scratch.distance_cache.get(&n) {
Some(Ok::<((u32, _), f32), ANNError>(((n, entry.1), entry.0)))
Some(Neighbor::new((n, entry.1), entry.0))
Comment thread
hildebrandmw marked this conversation as resolved.
} else {
uncached_ids.push(n);
None
Expand All @@ -362,12 +362,12 @@ where
PostprocessStrategy::AcceptAll => candidates
.map(|n| *n.id())
.filter_map(&mut process)
.collect::<Result<Vec<_>, _>>()?,
.collect(),
PostprocessStrategy::Apply(f) => candidates
.map(|n| *n.id())
.filter(|id| f(id))
.filter_map(&mut process)
.collect::<Result<Vec<_>, _>>()?,
.collect(),
}
};
if !uncached_ids.is_empty() {
Expand All @@ -376,13 +376,13 @@ where
let v = accessor.scratch.vertex_provider.get_vector(n)?;
let d = provider.distance_comparer.evaluate_similarity(query, v);
let a = accessor.scratch.vertex_provider.get_associated_data(n)?;
reranked.push(((*n, *a), d));
reranked.push(Neighbor::new((*n, *a), d));
}
}

// Sort the full precision distances.
reranked
.sort_unstable_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
reranked.sort_unstable_by(neighbor::ord::fast_distance);

// Store the reranked results.
Ok(output.extend(reranked))
}
Expand Down Expand Up @@ -462,7 +462,7 @@ where
Ok(output.extend(reranked.into_iter().map(|idx| {
let id = candidate_ids[idx];
let distance = candidate_distances[idx];
((id, associated_data[idx]), distance)
Neighbor::new((id, associated_data[idx]), distance)
})))
}
}
Expand Down
Loading
Loading