use std::hash::Hash;
use diskann::{
ANNError, ANNResult,
graph::{
AdjacencyList, SearchOutputBuffer,
glue::{self, HybridPredicate},
workingset,
},
neighbor::Neighbor,
provider,
};
use crate::{
counters::{Counters, LocalCounters},
ids::IdMap,
neighbors::Neighbors,
num::{IdLimit, MaxDegree},
repr,
};
pub trait Id: Send + Sync + Hash + Eq + Clone + 'static {}
impl<T> Id for T where T: Send + Sync + Hash + Eq + Clone + 'static {}
#[derive(Debug)]
pub struct Provider<R, M = u32>
where
M: Id,
{
representation: R,
mapping: IdMap<M>,
counters: Counters,
}
impl<R, M> Provider<R, M>
where
M: Id,
{
fn local_counters(&self) -> LocalCounters<'_> {
self.counters.local()
}
#[cfg(feature = "integration-test")]
pub fn counters(&self) -> crate::integration::counters::CounterSnapshot {
self.counters.snapshot()
}
}
impl<R, M> Provider<R, M>
where
R: repr::Representation,
M: Id,
{
pub fn new<C>(config: C) -> ANNResult<Self>
where
C: repr::RepresentationConfig<Representation = R>,
{
let representation = <_ as repr::RepresentationConfig>::build(config)?;
let mapping = IdMap::new(representation.capacity());
Ok(Self {
representation,
mapping,
counters: Counters::new(),
})
}
pub fn max_degree(&self) -> MaxDegree {
self.representation.max_degree()
}
}
#[derive(Debug, Clone, Default)]
pub struct Context;
impl diskann::provider::ExecutionContext for Context {}
impl<T, M> diskann::provider::DataProvider for Provider<T, M>
where
T: Send + Sync + 'static,
M: Id,
{
type Context = Context;
type InternalId = u32;
type ExternalId = M;
type Error = ANNError;
type Guard = diskann::provider::NoopGuard<u32>;
fn to_internal_id(
&self,
_context: &Self::Context,
gid: &M,
) -> Result<Self::InternalId, Self::Error> {
match self.mapping.to_internal(gid) {
Some(id) => Ok(id),
None => Err(ANNError::message("no mapping")),
}
}
fn to_external_id(
&self,
_context: &Self::Context,
id: Self::InternalId,
) -> Result<Self::ExternalId, Self::Error> {
match self.mapping.to_external(id) {
Some(gid) => Ok(gid),
None => Err(ANNError::message("no mapping")),
}
}
}
impl<R, M> diskann::provider::Delete for Provider<R, M>
where
R: repr::Representation,
M: Id,
{
async fn delete(&self, _context: &Context, gid: &M) -> ANNResult<()> {
let entry = match self.mapping.occupied_entry(gid.clone()) {
None => {
return Err(ANNError::message("id already deleted"));
}
Some(e) => e,
};
<R as repr::Representation>::retire(&self.representation, entry.internal())?;
entry.delete();
Ok(())
}
async fn release(&self, _context: &Context, _id: Self::InternalId) -> ANNResult<()> {
Ok(())
}
async fn status_by_internal_id(
&self,
_context: &Context,
id: u32,
) -> ANNResult<diskann::provider::ElementStatus> {
match <R as repr::Representation>::is_readable(&self.representation, id) {
Some(true) => Ok(diskann::provider::ElementStatus::Valid),
Some(false) => Ok(diskann::provider::ElementStatus::Deleted),
None => Err(ANNError::message("accessed invalid internal ID")),
}
}
async fn status_by_external_id(
&self,
_context: &Context,
gid: &M,
) -> ANNResult<diskann::provider::ElementStatus> {
if self.mapping.contains_external(gid) {
Ok(diskann::provider::ElementStatus::Valid)
} else {
Ok(diskann::provider::ElementStatus::Deleted)
}
}
}
fn ready<F, R>(f: F) -> std::future::Ready<R>
where
F: FnOnce() -> R,
{
std::future::ready(f())
}
impl<T, R, M> diskann::provider::SetElement<T> for Provider<R, M>
where
R: repr::Set<T>,
M: Id,
{
type SetError = ANNError;
fn set_element(
&self,
_context: &Self::Context,
id: &M,
element: T,
) -> impl std::future::Future<Output = Result<Self::Guard, Self::SetError>> + Send {
let work = move || {
let guard = <R as repr::Set<T>>::set(&self.representation, element)?;
let internal = <_ as repr::Guard>::id(&guard);
self.mapping.insert(id.clone(), internal)?;
<_ as repr::Guard>::publish(guard);
self.local_counters().set_vector(1);
Ok(diskann::provider::NoopGuard::new(internal))
};
ready(work)
}
}
#[derive(Debug)]
pub struct SearchAccessor<'a> {
neighbors: &'a Neighbors,
ids: AdjacencyList<u32>,
expand_beam: Box<dyn repr::ExpandBeam + 'a>,
id_limit: IdLimit,
buffer: Vec<Neighbor<u32>>,
post_process: Option<Box<dyn repr::PostProcess + 'a>>,
provider: &'a (dyn std::any::Any + Send + Sync),
start_points: std::ops::Range<u32>,
counters: LocalCounters<'a>,
}
impl<'a> SearchAccessor<'a> {
pub(crate) fn new(
neighbors: &'a Neighbors,
expand_beam: Box<dyn repr::ExpandBeam + 'a>,
post_process: Option<Box<dyn repr::PostProcess + 'a>>,
provider: &'a (dyn std::any::Any + Send + Sync),
start_points: std::ops::Range<u32>,
counters: LocalCounters<'a>,
) -> Self {
let id_limit = expand_beam.id_limit();
Self {
neighbors,
ids: AdjacencyList::with_capacity(neighbors.max_degree().value()),
expand_beam,
id_limit,
buffer: vec![Default::default(); neighbors.max_degree().value()],
post_process,
provider,
start_points,
counters,
}
}
#[cfg(all(test, feature = "quantization"))]
pub(crate) fn get_expand_beam(&self) -> &dyn repr::ExpandBeam {
&*self.expand_beam
}
#[cfg(all(test, feature = "quantization"))]
pub(crate) fn get_post_process(&mut self) -> Option<&mut (dyn repr::PostProcess + 'a)> {
self.post_process.as_deref_mut()
}
}
impl diskann::provider::HasId for SearchAccessor<'_> {
type Id = u32;
}
impl glue::SearchAccessor for SearchAccessor<'_> {
fn starting_points(
&self,
) -> impl std::future::Future<Output = ANNResult<Vec<Self::Id>>> + Send {
std::future::ready(Ok(self.start_points.clone().collect()))
}
fn start_point_distances<F>(
&mut self,
mut f: F,
) -> impl std::future::Future<Output = ANNResult<()>> + Send
where
F: FnMut(Self::Id, f32) + Send,
{
let work = move || {
for p in self.start_points.clone() {
match self.expand_beam.evaluate(p)? {
Some(distance) => {
self.counters.get_vector(1);
self.counters.query_distance(1);
f(p, distance);
}
None => {
return Err(ANNError::message("could not retrieve start point"));
}
}
}
Ok(())
};
ready(work)
}
fn expand_beam<Itr, P, F>(
&mut self,
ids: Itr,
mut pred: P,
mut on_neighbors: F,
) -> impl std::future::Future<Output = ANNResult<()>> + Send
where
Itr: Iterator<Item = Self::Id> + Send,
P: HybridPredicate<Self::Id> + Send + Sync,
F: FnMut(Self::Id, f32) + Send,
{
let work = move || -> ANNResult<()> {
for i in ids {
self.neighbors.get(i, &mut self.ids)?;
self.counters.get_neighbors(1);
self.ids
.retain(|i| pred.eval_mut(i) && self.id_limit.is_in_bounds(*i));
assert!(self.buffer.len() >= self.ids.len());
let processed =
unsafe { self.expand_beam.expand_beam(&self.ids, &mut self.buffer) }?;
self.counters.get_vector(processed as u64);
self.counters.query_distance(processed as u64);
self.buffer
.iter()
.take(processed)
.for_each(|neighbor| on_neighbors(*neighbor.id(), *neighbor.distance()));
}
Ok(())
};
ready(work)
}
}
#[derive(Debug)]
pub struct PruneAccessor<'a> {
prune: Box<dyn repr::Prune + 'a>,
keys: hashbrown::HashMap<u32, Option<repr::PruneKey>>,
neighbors: &'a Neighbors,
counters: LocalCounters<'a>,
}
impl<'a> PruneAccessor<'a> {
pub(crate) fn new(
prune: Box<dyn repr::Prune + 'a>,
neighbors: &'a Neighbors,
counters: LocalCounters<'a>,
) -> Self {
Self {
prune,
keys: hashbrown::HashMap::new(),
neighbors,
counters,
}
}
#[cfg(all(test, feature = "quantization"))]
pub(crate) fn get_prune(&mut self) -> &mut dyn repr::Prune {
&mut *self.prune
}
}
#[derive(Debug)]
pub struct Distance<'a> {
prune: &'a dyn repr::Prune,
counters: LocalCounters<'a>,
}
impl<'a> Distance<'a> {
fn new(prune: &'a dyn repr::Prune, counters: LocalCounters<'a>) -> Self {
Self { prune, counters }
}
}
#[derive(Debug, Clone, Copy)]
#[repr(transparent)]
pub struct ElementRef(repr::PruneKey);
impl<'a> diskann_utils::Reborrow<'a> for ElementRef {
type Target = ElementRef;
fn reborrow(&'a self) -> Self {
*self
}
}
impl diskann_vector::DistanceFunction<ElementRef, ElementRef, f32> for Distance<'_> {
#[inline]
fn evaluate_similarity(&self, x: ElementRef, y: ElementRef) -> f32 {
self.counters.distance_ref(1);
self.prune.evaluate(x.0, y.0)
}
}
impl diskann::provider::HasId for PruneAccessor<'_> {
type Id = u32;
}
impl glue::PruneAccessor for PruneAccessor<'_> {
type Neighbors<'a>
= provider::Neighbors<'a, Self>
where
Self: 'a;
type ElementRef<'a> = ElementRef;
type View<'a>
= &'a Self
where
Self: 'a;
type Distance<'a>
= Distance<'a>
where
Self: 'a;
fn neighbors(&mut self) -> Self::Neighbors<'_> {
provider::Neighbors(self)
}
async fn fill<'a, Itr>(
&'a mut self,
itr: Itr,
) -> ANNResult<(Self::View<'a>, Self::Distance<'a>)>
where
Itr: ExactSizeIterator<Item = Self::Id> + Clone + Send + Sync,
{
self.keys.clear();
self.keys.extend(itr.map(|i| (i, None)));
let count: usize = self.prune.prepare(self.keys.iter_mut())?;
self.counters.get_vector(count as u64);
Ok((self, Distance::new(&*self.prune, self.counters.fork())))
}
}
impl provider::NeighborAccessor for PruneAccessor<'_> {
fn get_neighbors(
&mut self,
id: Self::Id,
neighbors: &mut AdjacencyList<Self::Id>,
) -> impl std::future::Future<Output = ANNResult<()>> + Send {
let work = move || {
self.counters.get_neighbors(1);
Ok(self.neighbors.get(id, neighbors)?)
};
ready(work)
}
}
impl provider::NeighborAccessorMut for PruneAccessor<'_> {
fn set_neighbors(
&mut self,
id: Self::Id,
neighbors: &[Self::Id],
) -> impl std::future::Future<Output = ANNResult<()>> + Send {
let work = move || {
self.counters.set_neighbors(1);
Ok(self.neighbors.set(id, neighbors)?)
};
ready(work)
}
fn append_vector(
&mut self,
id: Self::Id,
neighbors: &[Self::Id],
) -> impl std::future::Future<Output = ANNResult<()>> + Send {
let work = move || -> ANNResult<()> {
self.counters.append_vector(1);
let lock = self.neighbors.lock(id)?;
if lock.len() + neighbors.len() > lock.capacity() {
let slack = lock.capacity() - lock.len();
lock.append(&neighbors[..slack])?;
} else {
lock.append(neighbors)?;
}
Ok(())
};
ready(work)
}
}
impl workingset::View<u32> for &PruneAccessor<'_> {
type ElementRef<'a> = ElementRef;
type Element<'a>
= ElementRef
where
Self: 'a;
fn get(&self, id: u32) -> Option<ElementRef> {
self.keys.get(&id)?.map(ElementRef)
}
}
#[derive(Debug, Clone, Copy)]
pub struct Strategy;
impl<'a, R, M> glue::SearchStrategy<'a, Provider<R, M>, R::Query<'a>> for Strategy
where
R: repr::Search,
M: Id,
{
type SearchAccessor = SearchAccessor<'a>;
type SearchAccessorError = ANNError;
fn search_accessor(
&'a self,
provider: &'a Provider<R, M>,
_context: &'a Context,
query: R::Query<'a>,
) -> ANNResult<SearchAccessor<'a>> {
<R as repr::Search>::search_accessor(
&provider.representation,
query,
provider,
provider.local_counters(),
)
}
}
pub fn test_function<'a>(
x: &'a Provider<repr::Full<u8>>,
strategy: &'a Strategy,
context: &'a Context,
query: &'a [u8],
) -> ANNResult<SearchAccessor<'a>> {
glue::SearchStrategy::search_accessor(strategy, x, context, query)
}
#[derive(Debug, Clone, Copy)]
pub struct Translate<R, M>(std::marker::PhantomData<(R, M)>);
impl<R, M> Default for Translate<R, M> {
fn default() -> Self {
Self(std::marker::PhantomData)
}
}
impl<'a, R, M> glue::SearchPostProcess<SearchAccessor<'a>, R::Query<'a>, M> for Translate<R, M>
where
R: repr::Search,
M: Id,
{
type Error = ANNError;
fn post_process<I, B>(
&self,
accessor: &mut SearchAccessor<'_>,
_query: R::Query<'a>,
candidates: I,
output: &mut B,
) -> impl std::future::Future<Output = Result<usize, Self::Error>> + Send
where
I: Iterator<Item = Neighbor<u32>> + Send,
B: SearchOutputBuffer<M> + Send + ?Sized,
{
let work = move || {
let provider = match accessor.provider.downcast_ref::<Provider<R, M>>() {
Some(provider) => provider,
None => return Err(ANNError::message("bad any cast")),
};
let mut count = 0;
let mut push = |neighbor: Neighbor<u32>| -> bool {
if let Some(ext) = provider.mapping.to_external(*neighbor.id()) {
if output
.push(Neighbor::new(ext, *neighbor.distance()))
.is_available()
{
count += 1;
true
} else {
false
}
} else {
true
}
};
if let Some(post_process) = &mut accessor.post_process {
accessor.buffer.clear();
Extend::extend(&mut accessor.buffer, candidates);
post_process.post_process(&mut accessor.buffer)?;
for neighbor in accessor.buffer.iter() {
if !push(*neighbor) {
break;
}
}
} else {
for neighbor in candidates {
if !push(neighbor) {
break;
}
}
}
Ok(count)
};
ready(work)
}
}
impl<'a, R, M> glue::DefaultPostProcessor<'a, Provider<R, M>, R::Query<'a>, M> for Strategy
where
R: repr::Search,
M: Id,
{
diskann::default_post_processor!(Translate<R, M>);
}
impl<R, M> glue::PruneStrategy<Provider<R, M>> for Strategy
where
R: repr::Insert,
M: Id,
{
type PruneAccessor<'a> = PruneAccessor<'a>;
type PruneAccessorError = ANNError;
fn prune_accessor<'a>(
&self,
provider: &'a Provider<R, M>,
_context: &'a Context,
_capacity: usize,
) -> ANNResult<PruneAccessor<'a>> {
<R as repr::Insert>::prune_accessor(&provider.representation, provider.local_counters())
}
}
impl<'a, R, M> glue::InsertStrategy<'a, Provider<R, M>, R::Query<'a>> for Strategy
where
R: repr::Insert,
M: Id,
{
type SearchAccessor = SearchAccessor<'a>;
type SearchAccessorError = ANNError;
type PruneStrategy = Self;
fn insert_search_accessor(
&'a self,
provider: &'a Provider<R, M>,
_context: &'a Context,
vector: R::Query<'a>,
) -> Result<Self::SearchAccessor, Self::SearchAccessorError> {
<R as repr::Insert>::insert_search_accessor(
&provider.representation,
vector,
provider,
provider.local_counters(),
)
}
fn prune_strategy(&self) -> Self::PruneStrategy {
*self
}
}
impl<T, M> glue::InplaceDeleteStrategy<Provider<repr::Full<T>, M>> for Strategy
where
M: Id,
T: repr::FullPrecision,
{
type DeleteElement<'a> = &'a [T];
type DeleteElementGuard = Box<[T]>;
type DeleteElementError = ANNError;
type PruneStrategy = Self;
type DeleteSearchAccessor<'a> = SearchAccessor<'a>;
type SearchPostProcessor = glue::CopyIds;
type SearchStrategy = Self;
fn prune_strategy(&self) -> Self {
*self
}
fn search_strategy(&self) -> Self {
*self
}
fn search_post_processor(&self) -> Self::SearchPostProcessor {
glue::CopyIds
}
fn get_delete_element<'a>(
&'a self,
provider: &'a Provider<repr::Full<T>, M>,
_context: &'a Context,
id: u32,
) -> impl Future<Output = Result<Self::DeleteElementGuard, Self::DeleteElementError>> + Send
{
let work = move || provider.representation.get(id);
ready(work)
}
}
#[cfg(test)]
mod tests {
use super::*;
use diskann::{
graph::{DiskANNIndex, InplaceDeleteMethod, search::Knn, test::synthetic::Grid},
neighbor::Neighbor,
provider::{DataProvider, Delete},
};
use diskann_utils::views::rowmajor::{self, Matrix};
use diskann_vector::distance::Metric;
use crate::num::Capacity;
#[tokio::test]
async fn smoke() {
let grid = Grid::Two;
let size = 5;
let data = grid.data(size);
let start = grid.start_point(size);
let degree = 6;
let config = repr::full::Config::new(
Capacity::new(grid.num_points(size)),
MaxDegree::new(degree),
Metric::L2,
rowmajor::Owned::row_vector(start.into()),
)
.unwrap();
let provider = Provider::<_, u64>::new(config).unwrap();
assert_eq!(provider.max_degree(), MaxDegree::new(degree));
let config = diskann::graph::config::Builder::new(
2 * (grid.dim() as usize),
diskann::graph::config::MaxDegree::new(provider.max_degree().value()),
10,
(Metric::L2).into(),
)
.build()
.unwrap();
let index = DiskANNIndex::new(config, provider, None);
for (i, data) in data.rows().enumerate() {
index
.insert(&Strategy, &Context, &((10 * i + 1) as u64), data)
.await
.unwrap();
}
for i in 0..data.nrows() {
let i = (10 * i + 1) as u64;
let internal = index.provider().to_internal_id(&Context, &i).unwrap();
assert_ne!(internal as u64, i);
assert_eq!(
index.provider().to_external_id(&Context, internal).unwrap(),
i
);
assert!(
!index
.provider()
.status_by_external_id(&Context, &i)
.await
.unwrap()
.is_deleted()
);
assert!(
!index
.provider()
.status_by_internal_id(&Context, internal)
.await
.unwrap()
.is_deleted()
);
}
assert!(index.provider().to_internal_id(&Context, &0).is_err());
assert!(index.provider().to_external_id(&Context, 26).is_err());
let knn = Knn::new(10, None).unwrap();
let mut neighbors = Vec::<Neighbor<u64>>::new();
index
.search(knn, &Strategy, &Context, &[0.0, 0.0], &mut neighbors)
.await
.unwrap();
assert_eq!(neighbors[0].as_tuple(), (1, 0.0));
assert_eq!(neighbors[1].as_tuple(), (11, 1.0)); assert_eq!(neighbors[2].as_tuple(), (51, 1.0));
assert_eq!(neighbors[3].as_tuple(), (61, 2.0));
index
.inplace_delete(
Strategy,
&Context,
&61,
3,
InplaceDeleteMethod::VisitedAndTopK {
k_value: 10,
l_value: 10,
},
)
.await
.unwrap();
assert!(
index
.provider()
.status_by_external_id(&Context, &61)
.await
.unwrap()
.is_deleted()
);
assert!(
index
.inplace_delete(
Strategy,
&Context,
&61,
3,
InplaceDeleteMethod::VisitedAndTopK {
k_value: 10,
l_value: 10
},
)
.await
.is_err()
);
let mut neighbors = Vec::<Neighbor<u64>>::new();
index
.search(knn, &Strategy, &Context, &[0.0, 0.0], &mut neighbors)
.await
.unwrap();
assert_eq!(neighbors[0].as_tuple(), (1, 0.0));
assert_eq!(neighbors[1].as_tuple(), (51, 1.0)); assert_eq!(neighbors[2].as_tuple(), (11, 1.0));
assert_eq!(neighbors[3].as_tuple(), (101, 4.0));
assert!(
index
.insert(&Strategy, &Context, &1, &[10.0, 10.0])
.await
.is_err()
);
assert!(
index
.insert(&Strategy, &Context, &2, &[10.0, 10.0, 10.0])
.await
.is_err()
);
index
.insert(&Strategy, &Context, &62, &[1.0, 1.0])
.await
.unwrap();
assert!(
index
.insert(&Strategy, &Context, &62, &[0.0, 0.0])
.await
.is_err()
);
let mut neighbors = Vec::<Neighbor<u64>>::new();
index
.search(knn, &Strategy, &Context, &[0.0, 0.0], &mut neighbors)
.await
.unwrap();
assert_eq!(neighbors[0].as_tuple(), (1, 0.0));
assert_eq!(neighbors[1].as_tuple(), (11, 1.0)); assert_eq!(neighbors[2].as_tuple(), (51, 1.0));
assert_eq!(neighbors[3].as_tuple(), (62, 2.0));
}
}