use std::{
fmt::Debug,
future::Future,
io::{Read, Write},
num::NonZeroUsize,
str::FromStr,
};
use diskann_quantization::{
alloc::{GlobalAllocator, Poly},
spherical::iface::{self as spherical_iface, try_deserialize, Opaque, Quantizer},
};
use serde::{Deserialize, Serialize};
use bf_tree::{BfTree, Config};
use diskann::{
default_post_processor,
error::{ErrorExt, Infallible, RankedError},
graph::{
glue::{
self, Batch, CopyIds, DefaultPostProcessor, InplaceDeleteStrategy, InsertStrategy,
MultiInsertStrategy, PruneStrategy, SearchStrategy,
},
strategy::{FullPrecision, Quantized},
workingset::map,
AdjacencyList, SearchOutputBuffer,
},
neighbor::{self, Neighbor},
provider::{DataProvider, DefaultContext, Delete, ElementStatus, HasId, NoopGuard, SetElement},
utils::VectorRepr,
ANNError, ANNResult,
};
use diskann_utils::{
future::{AsyncFriendly, SendFuture},
lazy_format,
views::MatrixView,
};
use diskann_vector::{distance::Metric, DistanceFunction, PreprocessedDistanceFunction};
use super::{
neighbors::{NeighborAccessor, NeighborProvider},
quant::QuantVectorProvider,
vectors::VectorProvider,
AccessError, BfTreeId, NoStore,
};
use crate::locks::StripedLocks;
use diskann_providers::model::graph::provider::async_::distances::UnwrapErr;
use diskann_providers::storage::{LoadWith, SaveWith, StorageReadProvider, StorageWriteProvider};
pub struct BfTreeProvider<T, Q = QuantVectorProvider, I = u32>
where
T: VectorRepr,
I: BfTreeId,
{
pub(super) quant_vectors: Q,
pub(super) full_vectors: VectorProvider<T>,
pub(crate) neighbor_provider: NeighborProvider<I>,
pub(super) metric: Metric,
pub(crate) graph_params: Option<GraphParams>,
pub(crate) use_snapshot: bool,
pub(crate) locks: StripedLocks,
}
#[derive(Debug, Clone)]
pub struct BfTreeProviderParameters {
pub max_points: usize,
pub num_start_points: NonZeroUsize,
pub dim: usize,
pub metric: Metric,
pub max_degree: u32,
pub vector_provider_config: Config,
pub quant_vector_provider_config: Config,
pub neighbor_list_provider_config: Config,
pub graph_params: Option<GraphParams>,
pub use_snapshot: bool,
}
impl<T, Q, I> BfTreeProvider<T, Q, I>
where
T: VectorRepr,
I: BfTreeId,
{
fn new_empty<TQ>(mut params: BfTreeProviderParameters, quant_precursor: TQ) -> ANNResult<Self>
where
Self: StartPoint<T>,
TQ: CreateQuantProvider<Target = Q>,
{
params
.vector_provider_config
.use_snapshot(params.use_snapshot);
params
.neighbor_list_provider_config
.use_snapshot(params.use_snapshot);
params
.quant_vector_provider_config
.use_snapshot(params.use_snapshot);
let id_capacity = params
.max_points
.checked_add(params.num_start_points.get())
.ok_or_else(|| {
ANNError::message("max_points + num_start_points overflows usize".to_string())
})?;
crate::id::validate_id_capacity::<I>(id_capacity)?;
Ok(Self {
quant_vectors: quant_precursor.create(params.quant_vector_provider_config)?,
full_vectors: VectorProvider::new_with_config(
params.max_points,
params.dim,
params.num_start_points.get(),
params.vector_provider_config,
)?,
neighbor_provider: NeighborProvider::new_with_config(
params.max_degree,
params.neighbor_list_provider_config,
)?,
metric: params.metric,
graph_params: params.graph_params,
use_snapshot: params.use_snapshot,
locks: StripedLocks::new(),
})
}
pub fn new<TQ>(
params: BfTreeProviderParameters,
start_points: MatrixView<'_, T>,
quant_precursor: TQ,
) -> ANNResult<Self>
where
Self: StartPoint<T>,
TQ: CreateQuantProvider<Target = Q>,
{
if start_points.nrows() != params.num_start_points.get() {
return Err(ANNError::message(format!(
"start_points matrix has {} rows, but params.num_start_points is {}",
start_points.nrows(),
params.num_start_points.get(),
)));
}
let provider = Self::new_empty(params.clone(), quant_precursor)?;
provider.set_start_points(Hidden(()), start_points)?;
Ok(provider)
}
pub fn starting_points(&self) -> ANNResult<Vec<I>> {
self.full_vectors.starting_points()
}
pub fn iter(&self) -> crate::id::IdRange<I> {
I::id_range(self.full_vectors.total())
}
pub fn num_start_points(&self) -> usize {
self.full_vectors.num_start_points
}
pub fn max_points(&self) -> usize {
self.full_vectors.max_vectors
}
pub fn dim(&self) -> usize {
self.full_vectors.dim()
}
pub fn metric(&self) -> Metric {
self.metric
}
pub fn max_degree(&self) -> u32 {
self.neighbor_provider.max_degree()
}
}
impl<T, I> BfTreeProvider<T, QuantVectorProvider, I>
where
T: VectorRepr,
I: BfTreeId,
{
pub fn counts_for_get_vector(&self) -> (usize, usize) {
(
self.full_vectors.num_get_calls.get(),
self.quant_vectors.num_get_calls.get(),
)
}
}
impl<T, I> BfTreeProvider<T, NoStore, I>
where
T: VectorRepr,
I: BfTreeId,
{
pub fn counts_for_get_vector(&self) -> (usize, usize) {
(self.full_vectors.num_get_calls.get(), 0)
}
}
pub(crate) trait DeleteQuant {
fn delete_vector(&self, id: usize);
}
impl DeleteQuant for QuantVectorProvider {
fn delete_vector(&self, id: usize) {
QuantVectorProvider::delete_vector(self, id);
}
}
impl DeleteQuant for NoStore {
fn delete_vector(&self, _id: usize) {}
}
impl<T, Q, I> Delete for BfTreeProvider<T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly + DeleteQuant,
I: BfTreeId,
{
fn release(
&self,
_context: &Self::Context,
_id: Self::InternalId,
) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
std::future::ready(Ok(()))
}
fn delete(
&self,
_context: &Self::Context,
gid: &Self::ExternalId,
) -> impl std::future::Future<Output = Result<(), Self::Error>> + Send {
let id = *gid;
let _guard = self.locks.lock(id.as_index());
self.full_vectors.delete_vector(id.as_index());
self.quant_vectors.delete_vector(id.as_index());
std::future::ready(Ok(()))
}
fn status_by_external_id(
&self,
context: &Self::Context,
gid: &Self::ExternalId,
) -> impl std::future::Future<Output = Result<diskann::provider::ElementStatus, Self::Error>> + Send
{
self.status_by_internal_id(context, *gid)
}
fn status_by_internal_id(
&self,
_context: &Self::Context,
id: Self::InternalId,
) -> impl std::future::Future<Output = Result<diskann::provider::ElementStatus, Self::Error>> + Send
{
let status = match self.full_vectors.get_vector_sync(id.as_index()) {
Ok(_) => Ok(ElementStatus::Valid),
Err(RankedError::Transient(_)) => Ok(ElementStatus::Deleted),
Err(RankedError::Error(e)) => Err(e),
};
std::future::ready(status)
}
}
impl<T, Q, I> IntoIterator for &BfTreeProvider<T, Q, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Item = I;
type IntoIter = crate::id::IdRange<I>;
fn into_iter(self) -> Self::IntoIter {
self.iter()
}
}
pub trait CreateQuantProvider {
type Target;
fn create(self, bf_tree_config: Config) -> ANNResult<Self::Target>;
}
impl CreateQuantProvider for NoStore {
type Target = NoStore;
fn create(self, _bf_tree_config: Config) -> ANNResult<Self::Target> {
Ok(self)
}
}
impl CreateQuantProvider for Poly<dyn Quantizer> {
type Target = QuantVectorProvider;
fn create(self, bf_tree_config: Config) -> ANNResult<Self::Target> {
QuantVectorProvider::new_with_config(self, bf_tree_config)
}
}
impl<T, Q, I> BfTreeProvider<T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
pub fn neighbors(&self) -> &NeighborProvider<I> {
&self.neighbor_provider
}
}
impl<T, Q, I> DataProvider for BfTreeProvider<T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type Context = DefaultContext;
type InternalId = I;
type ExternalId = I;
type Error = ANNError;
type Guard = NoopGuard<I>;
fn to_internal_id(
&self,
_context: &DefaultContext,
gid: &Self::ExternalId,
) -> Result<Self::InternalId, Self::Error> {
Ok(*gid)
}
fn to_external_id(
&self,
_context: &DefaultContext,
id: Self::InternalId,
) -> Result<Self::ExternalId, Self::Error> {
Ok(id)
}
}
impl<T, Q, I> HasId for BfTreeProvider<T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type Id = I;
}
impl<T, I> SetElement<&[T]> for BfTreeProvider<T, QuantVectorProvider, I>
where
T: VectorRepr,
I: BfTreeId,
{
type SetError = ANNError;
fn set_element(
&self,
_context: &Self::Context,
id: &I,
element: &[T],
) -> impl Future<Output = Result<Self::Guard, Self::SetError>> + Send {
let _guard = self.locks.lock(id.as_index());
if let Err(err) = self.full_vectors.set_vector_sync(id.as_index(), element) {
return std::future::ready(Err(err));
}
if let Err(err) = self.quant_vectors.set_vector_sync(id.as_index(), element) {
debug_assert!(
false,
"quant write failed after full-precision success: {err}"
);
return std::future::ready(Err(err));
}
std::future::ready(Ok(NoopGuard::new(*id)))
}
}
impl<T, I> SetElement<&[T]> for BfTreeProvider<T, NoStore, I>
where
T: VectorRepr,
I: BfTreeId,
{
type SetError = ANNError;
fn set_element(
&self,
_context: &Self::Context,
id: &I,
element: &[T],
) -> impl Future<Output = Result<Self::Guard, Self::SetError>> + Send {
let _guard = self.locks.lock(id.as_index());
if let Err(err) = self.full_vectors.set_vector_sync(id.as_index(), element) {
return std::future::ready(Err(err));
}
std::future::ready(Ok(NoopGuard::new(*id)))
}
}
pub struct Hidden(());
pub trait StartPoint<T> {
#[doc(hidden)]
fn set_start_points(&self, hidden: Hidden, start_points: MatrixView<'_, T>) -> ANNResult<()>;
}
impl<T, I> StartPoint<T> for BfTreeProvider<T, QuantVectorProvider, I>
where
T: VectorRepr,
I: BfTreeId,
{
fn set_start_points(&self, _hidden: Hidden, start_points: MatrixView<'_, T>) -> ANNResult<()> {
let start_point_ids: Vec<I> = self.full_vectors.starting_points()?;
if start_points.nrows() != start_point_ids.len() {
return Err(ANNError::message(format!(
"expected start_points to contain `{}` rows, instead it has {}",
start_point_ids.len(),
start_points.nrows(),
)));
}
let mut scratch = self.neighbor_provider.scratch(&self.locks);
for (id, v) in std::iter::zip(start_point_ids, start_points.row_iter()) {
self.full_vectors.set_vector_sync(id.as_index(), v)?;
self.quant_vectors.set_vector_sync(id.as_index(), v)?;
scratch.write_neighbors(id, &[])?;
}
Ok(())
}
}
impl<T, I> StartPoint<T> for BfTreeProvider<T, NoStore, I>
where
T: VectorRepr,
I: BfTreeId,
{
fn set_start_points(&self, _hidden: Hidden, start_points: MatrixView<'_, T>) -> ANNResult<()> {
let start_point_ids: Vec<I> = self.full_vectors.starting_points()?;
if start_points.nrows() != start_point_ids.len() {
return Err(ANNError::message(format!(
"expected start_points to contain `{}` rows, instead it has {}",
start_point_ids.len(),
start_points.nrows(),
)));
}
let mut scratch = self.neighbor_provider.scratch(&self.locks);
for (id, v) in std::iter::zip(start_point_ids, start_points.row_iter()) {
self.full_vectors.set_vector_sync(id.as_index(), v)?;
scratch.write_neighbors(id, &[])?;
}
Ok(())
}
}
pub struct FullAccessor<'a, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
provider: &'a BfTreeProvider<T, Q, I>,
computer: T::QueryDistance,
element: Box<[T]>,
}
impl<'a, T, Q, I> FullAccessor<'a, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
pub(crate) fn new(provider: &'a BfTreeProvider<T, Q, I>, query: &[T]) -> Self {
Self {
provider,
computer: T::query_distance(query, provider.metric),
element: (0..provider.full_vectors.dim())
.map(|_| T::default())
.collect(),
}
}
fn get_distance(&mut self, id: I) -> Result<f32, AccessError> {
self.provider
.full_vectors
.get_vector_into(id.as_index(), &mut self.element)
.map(|_: ()| self.computer.evaluate_similarity(&self.element))
}
}
impl<T, Q, I> HasId for FullAccessor<'_, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type Id = I;
}
impl<T, Q, I> glue::SearchAccessor for FullAccessor<'_, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
fn starting_points(&self) -> impl Future<Output = ANNResult<Vec<I>>> {
std::future::ready(self.provider.starting_points())
}
async fn start_point_distances<F>(&mut self, mut f: F) -> ANNResult<()>
where
F: FnMut(Self::Id, f32) + Send,
{
for i in self.provider.starting_points()? {
f(
i,
self.get_distance(i)
.escalate("starting point retrieval must succeed")?,
)
}
Ok(())
}
async fn expand_beam<Itr, P, F>(
&mut self,
ids: Itr,
mut pred: P,
mut on_neighbors: F,
) -> ANNResult<()>
where
Itr: Iterator<Item = Self::Id> + Send,
P: glue::HybridPredicate<Self::Id> + Send + Sync,
F: FnMut(Self::Id, f32) + Send,
{
let mut neighbors = AdjacencyList::new();
for n in ids {
self.provider.neighbors().get_neighbors(n, &mut neighbors)?;
for &id in neighbors.iter().filter(|i| pred.eval_mut(i)) {
if let Some(distance) = self
.get_distance(id)
.allow_transient("skipping deleted vectors")?
{
on_neighbors(id, distance)
}
}
}
Ok(())
}
}
pub struct QuantAccessor<'a, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
computer: super::quant::QuantQueryComputer,
element: Box<[u8]>,
}
impl<'a, T, I> QuantAccessor<'a, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
pub(crate) fn new(
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
query: &[T],
) -> ANNResult<Self> {
let computer = provider.quant_vectors.query_computer(query)?;
Ok(Self {
provider,
computer,
element: (0..provider.quant_vectors.quantizer.bytes())
.map(|_| u8::default())
.collect(),
})
}
fn get_distance(&mut self, id: I) -> Result<f32, AccessError> {
match self
.provider
.quant_vectors
.get_vector_into(id.as_index(), &mut self.element)
{
Ok(()) => self
.computer
.evaluate(&self.element)
.map_err(RankedError::Error),
Err(err) => Err(err),
}
}
}
impl<T, I> HasId for QuantAccessor<'_, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Id = I;
}
impl<T, I> glue::SearchAccessor for QuantAccessor<'_, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
fn starting_points(&self) -> impl Future<Output = ANNResult<Vec<I>>> {
std::future::ready(self.provider.starting_points())
}
async fn start_point_distances<F>(&mut self, mut f: F) -> ANNResult<()>
where
F: FnMut(Self::Id, f32) + Send,
{
for i in self.provider.starting_points()? {
f(
i,
self.get_distance(i)
.escalate("starting point retrieval must succeed")?,
)
}
Ok(())
}
async fn expand_beam<Itr, P, F>(
&mut self,
ids: Itr,
mut pred: P,
mut on_neighbors: F,
) -> ANNResult<()>
where
Itr: Iterator<Item = Self::Id> + Send,
P: glue::HybridPredicate<Self::Id> + Send + Sync,
F: FnMut(Self::Id, f32) + Send,
{
let mut neighbors = AdjacencyList::new();
for n in ids {
self.provider.neighbors().get_neighbors(n, &mut neighbors)?;
for &id in neighbors.iter().filter(|i| pred.eval_mut(i)) {
if let Some(distance) = self
.get_distance(id)
.allow_transient("skipping deleted vectors")?
{
on_neighbors(id, distance)
}
}
}
Ok(())
}
}
pub struct FullPruneAccessor<'a, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
provider: &'a BfTreeProvider<T, Q, I>,
neighbors: NeighborAccessor<'a, I>,
set: map::Map<I, Box<[T]>, map::Ref<[T]>>,
distance: T::Distance,
}
impl<'a, T, Q, I> FullPruneAccessor<'a, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
fn new(
provider: &'a BfTreeProvider<T, Q, I>,
set: map::Map<I, Box<[T]>, map::Ref<[T]>>,
) -> Self {
Self {
provider,
neighbors: provider.neighbor_provider.scratch(&provider.locks),
set,
distance: T::distance(provider.metric, Some(provider.full_vectors.dim())),
}
}
}
impl<T, Q, I> HasId for FullPruneAccessor<'_, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type Id = I;
}
impl<'q, T, Q, I> glue::PruneAccessor for FullPruneAccessor<'q, T, Q, I>
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type ElementRef<'a> = &'a [T];
type View<'a>
= map::View<'a, I, Box<[T]>, map::Ref<[T]>>
where
Self: 'a;
type Distance<'a>
= &'a T::Distance
where
Self: 'a;
type Neighbors<'a>
= diskann::provider::Neighbors<'a, NeighborAccessor<'q, I>>
where
Self: 'a;
fn fill<Itr>(
&mut self,
itr: Itr,
) -> impl SendFuture<ANNResult<(Self::View<'_>, Self::Distance<'_>)>>
where
Itr: ExactSizeIterator<Item = Self::Id> + Clone + Send + Sync,
{
let mut buf: Option<Box<[T]>> = None;
let view = self.set.fill(itr, |i: I| -> ANNResult<_> {
let mut b = match buf.take() {
Some(b) => b,
None => std::iter::repeat_n(T::default(), self.provider.dim()).collect(),
};
match self
.provider
.full_vectors
.get_vector_into(i.as_index(), &mut b)
.allow_transient("transient errors allowed during fill")?
{
Some(()) => Ok(Some(b)),
None => {
buf = Some(b);
Ok(None)
}
}
});
let result = view.map(|v| (v, &self.distance));
std::future::ready(result)
}
fn neighbors(&mut self) -> Self::Neighbors<'_> {
diskann::provider::Neighbors(&mut self.neighbors)
}
}
pub struct QuantPruneAccessor<'a, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
neighbors: NeighborAccessor<'a, I>,
set: map::Map<I, Owned>,
distance: UnwrapErr<spherical_iface::DistanceComputer, spherical_iface::DistanceError>,
}
impl<'a, T, I> QuantPruneAccessor<'a, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
fn new(
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
capacity: usize,
) -> ANNResult<Self> {
let distance = provider
.quant_vectors
.distance_computer()
.map(UnwrapErr::new)?;
let set = map::Builder::new(map::Capacity::Default).build(capacity);
Ok(Self {
provider,
neighbors: provider.neighbor_provider.scratch(&provider.locks),
set,
distance,
})
}
}
impl<T, I> HasId for QuantPruneAccessor<'_, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Id = I;
}
impl<'q, T, I> glue::PruneAccessor for QuantPruneAccessor<'q, T, I>
where
T: VectorRepr,
I: BfTreeId,
{
type ElementRef<'a> = Opaque<'a>;
type View<'a>
= map::View<'a, I, Owned>
where
Self: 'a;
type Distance<'a>
= &'a UnwrapErr<spherical_iface::DistanceComputer, spherical_iface::DistanceError>
where
Self: 'a;
type Neighbors<'a>
= diskann::provider::Neighbors<'a, NeighborAccessor<'q, I>>
where
Self: 'a;
fn fill<Itr>(
&mut self,
itr: Itr,
) -> impl SendFuture<ANNResult<(Self::View<'_>, Self::Distance<'_>)>>
where
Itr: ExactSizeIterator<Item = Self::Id> + Clone + Send + Sync,
{
let mut buf: Option<Box<[u8]>> = None;
let bytes = self.provider.quant_vectors.quantizer.bytes();
let view = self.set.fill(itr, |i: I| -> ANNResult<_> {
let mut b = match buf.take() {
Some(b) => b,
None => std::iter::repeat_n(0, bytes).collect(),
};
match self
.provider
.quant_vectors
.get_vector_into(i.as_index(), &mut b)
.allow_transient("transient errors allowed during fill")?
{
Some(()) => Ok(Some(Owned(b))),
None => {
buf = Some(b);
Ok(None)
}
}
});
let result = view.map(|v| (v, &self.distance));
std::future::ready(result)
}
fn neighbors(&mut self) -> Self::Neighbors<'_> {
diskann::provider::Neighbors(&mut self.neighbors)
}
}
pub struct Owned(Box<[u8]>);
impl<'short> diskann_utils::Reborrow<'short> for Owned {
type Target = Opaque<'short>;
fn reborrow(&'short self) -> Self::Target {
Opaque::new(&self.0)
}
}
impl<'a, T, Q, I> SearchStrategy<'a, BfTreeProvider<T, Q, I>, &'a [T]> for FullPrecision
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type SearchAccessor = FullAccessor<'a, T, Q, I>;
type SearchAccessorError = Infallible;
fn search_accessor(
&'a self,
provider: &'a BfTreeProvider<T, Q, I>,
_context: &'a DefaultContext,
query: &'a [T],
) -> Result<Self::SearchAccessor, Self::SearchAccessorError> {
Ok(FullAccessor::new(provider, query))
}
}
impl<'a, T, Q, I> DefaultPostProcessor<'a, BfTreeProvider<T, Q, I>, &'a [T]> for FullPrecision
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
default_post_processor!(glue::Pipeline<glue::FilterStartPoints, CopyIds>);
}
impl<T, Q, I> PruneStrategy<BfTreeProvider<T, Q, I>> for FullPrecision
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type PruneAccessor<'a> = FullPruneAccessor<'a, T, Q, I>;
type PruneAccessorError = diskann::error::Infallible;
fn prune_accessor<'a>(
&'a self,
provider: &'a BfTreeProvider<T, Q, I>,
_context: &'a DefaultContext,
capacity: usize,
) -> Result<Self::PruneAccessor<'a>, Self::PruneAccessorError> {
let set = map::Builder::new(map::Capacity::Default).build(capacity);
Ok(FullPruneAccessor::new(provider, set))
}
}
impl<'a, T, Q, I> InsertStrategy<'a, BfTreeProvider<T, Q, I>, &'a [T]> for FullPrecision
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type PruneStrategy = Self;
fn prune_strategy(&self) -> Self::PruneStrategy {
*self
}
}
impl<T, Q, B, I> MultiInsertStrategy<BfTreeProvider<T, Q, I>, B> for FullPrecision
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
B: for<'a> Batch<Element<'a> = &'a [T]> + Debug,
{
type Seed = map::Builder<I, map::Ref<[T]>>;
type FinishError = diskann::error::Infallible;
type PruneStrategy = Self;
type InsertStrategy = Self;
fn insert_strategy(&self) -> Self::InsertStrategy {
*self
}
fn finish<Itr>(
&self,
_provider: &BfTreeProvider<T, Q, I>,
_ctx: &DefaultContext,
batch: &std::sync::Arc<B>,
ids: Itr,
) -> impl std::future::Future<Output = Result<Self::Seed, Self::FinishError>> + Send
where
Itr: ExactSizeIterator<Item = I> + Send,
{
let overlay = map::Overlay::from_batch(batch.clone(), ids);
let builder = map::Builder::new(map::Capacity::Default).with_overlay(overlay);
std::future::ready(Ok(builder))
}
fn seeded_prune_accessor<'a>(
&'a self,
provider: &'a BfTreeProvider<T, Q, I>,
_context: &'a DefaultContext,
seed: &'a Self::Seed,
capacity: usize,
) -> ANNResult<FullPruneAccessor<'a, T, Q, I>> {
let set = seed.clone().build(capacity);
Ok(FullPruneAccessor::new(provider, set))
}
}
impl<T, Q, I> InplaceDeleteStrategy<BfTreeProvider<T, Q, I>> for FullPrecision
where
T: VectorRepr,
Q: AsyncFriendly,
I: BfTreeId,
{
type DeleteElementError = ANNError;
type DeleteElement<'a> = &'a [T];
type DeleteElementGuard = Box<[T]>;
type PruneStrategy = Self;
type DeleteSearchAccessor<'a> = FullAccessor<'a, T, Q, I>;
type SearchPostProcessor = CopyIds;
type SearchStrategy = Self;
fn search_strategy(&self) -> Self::SearchStrategy {
Self
}
fn prune_strategy(&self) -> Self::PruneStrategy {
Self
}
fn search_post_processor(&self) -> Self::SearchPostProcessor {
CopyIds
}
async fn get_delete_element<'a>(
&'a self,
provider: &'a BfTreeProvider<T, Q, I>,
_context: &'a DefaultContext,
id: I,
) -> Result<Self::DeleteElementGuard, Self::DeleteElementError> {
use diskann::error::ErrorExt;
let elt = provider
.full_vectors
.get_vector_sync(id.as_index())
.escalate("get_delete_element: failed to read vector for inplace delete")?
.into();
Ok(elt)
}
}
impl<'a, T, I> SearchStrategy<'a, BfTreeProvider<T, QuantVectorProvider, I>, &'a [T]> for Quantized
where
T: VectorRepr,
I: BfTreeId,
{
type SearchAccessor = QuantAccessor<'a, T, I>;
type SearchAccessorError = ANNError;
fn search_accessor(
&'a self,
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
_context: &'a DefaultContext,
query: &'a [T],
) -> Result<Self::SearchAccessor, Self::SearchAccessorError> {
QuantAccessor::new(provider, query)
}
}
impl<'a, T, I> DefaultPostProcessor<'a, BfTreeProvider<T, QuantVectorProvider, I>, &'a [T]>
for Quantized
where
T: VectorRepr,
I: BfTreeId,
{
default_post_processor!(glue::Pipeline<glue::FilterStartPoints, Rerank>);
}
impl<'a, T, I> InsertStrategy<'a, BfTreeProvider<T, QuantVectorProvider, I>, &'a [T]> for Quantized
where
T: VectorRepr,
I: BfTreeId,
{
type PruneStrategy = Self;
fn prune_strategy(&self) -> Self::PruneStrategy {
*self
}
}
impl<T, B, I> MultiInsertStrategy<BfTreeProvider<T, QuantVectorProvider, I>, B> for Quantized
where
T: VectorRepr,
I: BfTreeId,
B: glue::Batch,
B: for<'a> Batch<Element<'a> = &'a [T]> + Debug,
{
type Seed = ();
type FinishError = diskann::error::Infallible;
type PruneStrategy = Self;
type InsertStrategy = Self;
fn insert_strategy(&self) -> Self::InsertStrategy {
*self
}
fn finish<Itr>(
&self,
_provider: &BfTreeProvider<T, QuantVectorProvider, I>,
_ctx: &DefaultContext,
_batch: &std::sync::Arc<B>,
_ids: Itr,
) -> impl std::future::Future<Output = Result<Self::Seed, Self::FinishError>> + Send
where
Itr: ExactSizeIterator<Item = I> + Send,
{
std::future::ready(Ok(()))
}
fn seeded_prune_accessor<'a>(
&'a self,
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
_context: &'a DefaultContext,
_seed: &'a (),
capacity: usize,
) -> ANNResult<QuantPruneAccessor<'a, T, I>> {
QuantPruneAccessor::new(provider, capacity)
}
}
impl<T, I> InplaceDeleteStrategy<BfTreeProvider<T, QuantVectorProvider, I>> for Quantized
where
T: VectorRepr,
I: BfTreeId,
{
type DeleteElementError = ANNError;
type DeleteElement<'a> = &'a [T];
type DeleteElementGuard = Box<[T]>;
type PruneStrategy = Self;
type DeleteSearchAccessor<'a> = QuantAccessor<'a, T, I>;
type SearchPostProcessor = Rerank;
type SearchStrategy = Self;
fn search_strategy(&self) -> Self::SearchStrategy {
*self
}
fn prune_strategy(&self) -> Self::PruneStrategy {
*self
}
fn search_post_processor(&self) -> Self::SearchPostProcessor {
Rerank
}
async fn get_delete_element<'a>(
&'a self,
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
_context: &'a DefaultContext,
id: I,
) -> Result<Self::DeleteElementGuard, Self::DeleteElementError> {
use diskann::error::ErrorExt;
provider
.full_vectors
.get_vector_sync(id.as_index())
.escalate("get_delete_element: failed to read vector for inplace delete")
.map(Into::into)
}
}
impl<T, I> PruneStrategy<BfTreeProvider<T, QuantVectorProvider, I>> for Quantized
where
T: VectorRepr,
I: BfTreeId,
{
type PruneAccessor<'a> = QuantPruneAccessor<'a, T, I>;
type PruneAccessorError = ANNError;
fn prune_accessor<'a>(
&'a self,
provider: &'a BfTreeProvider<T, QuantVectorProvider, I>,
_context: &'a DefaultContext,
capacity: usize,
) -> Result<Self::PruneAccessor<'a>, Self::PruneAccessorError> {
QuantPruneAccessor::new(provider, capacity)
}
}
#[derive(Debug, Default, Clone, Copy)]
pub struct Rerank;
impl<'a, T, I> glue::SearchPostProcess<QuantAccessor<'a, T, I>, &[T]> for Rerank
where
T: VectorRepr,
I: BfTreeId,
{
type Error = ANNError;
fn post_process<Itr, B>(
&self,
accessor: &mut QuantAccessor<'a, T, I>,
query: &[T],
candidates: Itr,
output: &mut B,
) -> impl Future<Output = Result<usize, Self::Error>> + Send
where
Itr: Iterator<Item = Neighbor<I>> + Send,
B: SearchOutputBuffer<I> + Send + ?Sized,
{
use diskann::error::ErrorExt;
let provider = accessor.provider;
let f = T::distance(provider.metric, Some(provider.full_vectors.dim()));
let mut reranked = Vec::new();
for n in candidates {
match provider
.full_vectors
.get_vector_sync(n.id().as_index())
.allow_transient("stale candidate during rerank")
{
Ok(Some(vec)) => {
reranked.push(Neighbor::new(*n.id(), f.evaluate_similarity(query, &vec)));
}
Ok(None) => {
}
Err(e) => return std::future::ready(Err(e)),
}
}
reranked.sort_unstable_by(neighbor::ord::fast_distance);
std::future::ready(Ok(output.extend(reranked)))
}
}
#[derive(Serialize, Deserialize, Clone)]
pub struct BfTreeParams {
pub bytes: usize,
pub max_record_size: usize,
pub leaf_page_size: usize,
}
#[derive(Serialize, Deserialize, Clone)]
pub struct QuantParams {
pub params_quant: BfTreeParams,
}
#[derive(Serialize, Deserialize, Clone)]
pub struct SavedParams {
pub max_points: usize,
pub frozen_points: NonZeroUsize,
pub dim: usize,
pub metric: String,
pub max_degree: u32,
pub prefix: String,
pub params_vector: BfTreeParams,
pub params_neighbor: BfTreeParams,
pub quant_params: Option<QuantParams>,
pub graph_params: Option<GraphParams>,
pub is_memory: bool,
#[serde(default)]
pub use_snapshot: bool,
#[serde(default = "default_id_width")]
pub id_width: usize,
}
fn default_id_width() -> usize {
std::mem::size_of::<u32>()
}
fn validate_loaded_id_params<I: BfTreeId>(saved_params: &SavedParams) -> ANNResult<()> {
let expected = std::mem::size_of::<I>();
if saved_params.id_width != expected {
return Err(ANNError::message(format!(
"index was built with {}-byte vertex ids but is being loaded with a {expected}-byte \
id type; load it with the matching id type",
saved_params.id_width,
)));
}
let id_capacity = saved_params
.max_points
.checked_add(saved_params.frozen_points.get())
.ok_or_else(|| {
ANNError::message("max_points + frozen_points overflows usize".to_string())
})?;
crate::id::validate_id_capacity::<I>(id_capacity)
}
#[derive(Serialize, Deserialize, Clone, Debug, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum VectorDtype {
F32,
F16,
U8,
I8,
}
pub trait AsVectorDtype {
const DATA_TYPE: VectorDtype;
}
impl AsVectorDtype for f32 {
const DATA_TYPE: VectorDtype = VectorDtype::F32;
}
impl AsVectorDtype for half::f16 {
const DATA_TYPE: VectorDtype = VectorDtype::F16;
}
impl AsVectorDtype for i8 {
const DATA_TYPE: VectorDtype = VectorDtype::I8;
}
impl AsVectorDtype for u8 {
const DATA_TYPE: VectorDtype = VectorDtype::U8;
}
#[derive(Serialize, Deserialize, Clone, Debug)]
pub struct GraphParams {
pub l_build: usize,
pub alpha: f32,
pub backedge_ratio: f32,
pub vector_dtype: VectorDtype,
}
pub struct BfTreePaths;
impl BfTreePaths {
pub fn params_json(prefix: &str) -> String {
format!("{}_params.json", prefix)
}
pub fn vectors_bftree(prefix: &str) -> std::path::PathBuf {
std::path::PathBuf::from(format!("{}_vectors.bftree", prefix))
}
pub fn neighbors_bftree(prefix: &str) -> std::path::PathBuf {
std::path::PathBuf::from(format!("{}_neighbors.bftree", prefix))
}
pub fn quant_bftree(prefix: &str) -> std::path::PathBuf {
std::path::PathBuf::from(format!("{}_quant.bftree", prefix))
}
pub fn delete_bin(prefix: &str) -> String {
format!("{}_delete.bin", prefix)
}
pub fn quant_data_bin(prefix: &str) -> String {
format!("{}_quant_data.bin", prefix)
}
}
fn save_bftree(
tree: &BfTree,
target_path: std::path::PathBuf,
use_snapshot: bool,
) -> ANNResult<()> {
if !use_snapshot {
return Err(ANNError::message(
"cannot snapshot a BfTree that was not configured with use_snapshot(true)",
));
}
tree.cpr_snapshot(&target_path);
Ok(())
}
fn load_bftree(snapshot_path: std::path::PathBuf, use_snapshot: bool) -> Result<BfTree, ANNError> {
BfTree::new_from_cpr_snapshot(snapshot_path, use_snapshot, None, None, None)
.map_err(|e| ANNError::from(super::ConfigError(e)))
}
impl<T, I> SaveWith<String> for BfTreeProvider<T, NoStore, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Ok = usize;
type Error = ANNError;
async fn save_with<P>(&self, storage: &P, prefix: &String) -> Result<Self::Ok, Self::Error>
where
P: StorageWriteProvider,
{
let saved_params = SavedParams {
max_points: self.max_points(),
frozen_points: NonZeroUsize::new(self.num_start_points())
.ok_or_else(|| ANNError::message("num_start_points is zero"))?,
dim: self.dim(),
metric: self.metric().as_str().to_string(),
max_degree: self.max_degree(),
prefix: prefix.clone(),
params_vector: BfTreeParams {
bytes: self.full_vectors.config().get_cb_size_byte(),
max_record_size: self.full_vectors.config().get_cb_max_record_size(),
leaf_page_size: self.full_vectors.config().get_leaf_page_size(),
},
params_neighbor: BfTreeParams {
bytes: self.neighbor_provider.config().get_cb_size_byte(),
max_record_size: self.neighbor_provider.config().get_cb_max_record_size(),
leaf_page_size: self.neighbor_provider.config().get_leaf_page_size(),
},
quant_params: None,
graph_params: self.graph_params.clone(),
is_memory: self.full_vectors.config().is_memory_backend(),
use_snapshot: self.use_snapshot,
id_width: std::mem::size_of::<I>(),
};
debug_assert_eq!(
self.full_vectors.config().is_memory_backend(),
self.neighbor_provider.config().is_memory_backend(),
"Vector and neighbor stores have mismatched storage backends"
);
{
let params_filename = BfTreePaths::params_json(&saved_params.prefix);
let params_json = serde_json::to_string(&saved_params).map_err(|e| {
ANNError::message(lazy_format!(move, "Failed to serialize params: {}", e))
})?;
let mut params_writer = storage.create_for_write(¶ms_filename)?;
params_writer.write_all(params_json.as_bytes())?;
}
save_bftree(
self.full_vectors.bftree(),
BfTreePaths::vectors_bftree(&saved_params.prefix),
self.use_snapshot,
)?;
save_bftree(
self.neighbor_provider.bftree(),
BfTreePaths::neighbors_bftree(&saved_params.prefix),
self.use_snapshot,
)?;
Ok(0)
}
}
impl<T, I> LoadWith<String> for BfTreeProvider<T, NoStore, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Error = ANNError;
async fn load_with<P>(storage: &P, prefix: &String) -> Result<Self, Self::Error>
where
P: StorageReadProvider,
{
let saved_params: SavedParams = {
let params_filename = BfTreePaths::params_json(prefix);
let mut params_reader = storage.open_reader(¶ms_filename)?;
let mut params_json = String::new();
params_reader.read_to_string(&mut params_json)?;
serde_json::from_str(¶ms_json).map_err(|e| {
ANNError::message(lazy_format!(move, "Failed to deserialize params: {}", e))
})?
};
validate_loaded_id_params::<I>(&saved_params)?;
let metric = Metric::from_str(&saved_params.metric)
.map_err(|e| ANNError::message(lazy_format!(move, "Failed to parse metric: {}", e)))?;
let vector_index = load_bftree(
BfTreePaths::vectors_bftree(&saved_params.prefix),
saved_params.use_snapshot,
)?;
let full_vectors = VectorProvider::<T>::new_from_bftree(
saved_params.max_points,
saved_params.dim,
saved_params.frozen_points.get(),
vector_index,
);
let adjacency_list_index = load_bftree(
BfTreePaths::neighbors_bftree(&saved_params.prefix),
saved_params.use_snapshot,
)?;
let neighbor_provider =
NeighborProvider::<I>::new_from_bftree(saved_params.max_degree, adjacency_list_index)?;
Ok(Self {
quant_vectors: NoStore,
full_vectors,
neighbor_provider,
metric,
graph_params: saved_params.graph_params,
use_snapshot: saved_params.use_snapshot,
locks: StripedLocks::new(),
})
}
}
impl<T, I> SaveWith<String> for BfTreeProvider<T, QuantVectorProvider, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Ok = usize;
type Error = ANNError;
async fn save_with<P>(&self, storage: &P, prefix: &String) -> Result<Self::Ok, Self::Error>
where
P: StorageWriteProvider,
{
let saved_params = SavedParams {
max_points: self.max_points(),
frozen_points: NonZeroUsize::new(self.num_start_points())
.ok_or_else(|| ANNError::message("num_start_points is zero"))?,
dim: self.dim(),
metric: self.metric().as_str().to_string(),
max_degree: self.max_degree(),
prefix: prefix.clone(),
params_vector: BfTreeParams {
bytes: self.full_vectors.config().get_cb_size_byte(),
max_record_size: self.full_vectors.config().get_cb_max_record_size(),
leaf_page_size: self.full_vectors.config().get_leaf_page_size(),
},
params_neighbor: BfTreeParams {
bytes: self.neighbor_provider.config().get_cb_size_byte(),
max_record_size: self.neighbor_provider.config().get_cb_max_record_size(),
leaf_page_size: self.neighbor_provider.config().get_leaf_page_size(),
},
quant_params: Some(QuantParams {
params_quant: BfTreeParams {
bytes: self.quant_vectors.config().get_cb_size_byte(),
max_record_size: self.quant_vectors.config().get_cb_max_record_size(),
leaf_page_size: self.quant_vectors.config().get_leaf_page_size(),
},
}),
graph_params: self.graph_params.clone(),
is_memory: self.full_vectors.config().is_memory_backend(),
use_snapshot: self.use_snapshot,
id_width: std::mem::size_of::<I>(),
};
debug_assert_eq!(
self.full_vectors.config().is_memory_backend(),
self.neighbor_provider.config().is_memory_backend(),
"Vector and neighbor stores have mismatched storage backends"
);
debug_assert_eq!(
self.full_vectors.config().is_memory_backend(),
self.quant_vectors.config().is_memory_backend(),
"Vector and quant stores have mismatched storage backends"
);
{
let params_filename = BfTreePaths::params_json(&saved_params.prefix);
let params_json = serde_json::to_string(&saved_params).map_err(|e| {
ANNError::message(lazy_format!(move, "Failed to serialize params: {}", e))
})?;
let mut params_writer = storage.create_for_write(¶ms_filename)?;
params_writer.write_all(params_json.as_bytes())?;
}
save_bftree(
self.full_vectors.bftree(),
BfTreePaths::vectors_bftree(&saved_params.prefix),
self.use_snapshot,
)?;
save_bftree(
self.neighbor_provider.bftree(),
BfTreePaths::neighbors_bftree(&saved_params.prefix),
self.use_snapshot,
)?;
save_bftree(
self.quant_vectors.bftree(),
BfTreePaths::quant_bftree(&saved_params.prefix),
self.use_snapshot,
)?;
let filename = BfTreePaths::quant_data_bin(&saved_params.prefix);
let serialized = self
.quant_vectors
.quantizer
.serialize(GlobalAllocator)
.map_err(|e| ANNError::message(lazy_format!(move, "{e}")))?;
let mut writer = storage.create_for_write(&filename)?;
writer.write_all(&serialized)?;
Ok(0)
}
}
impl<T, I> LoadWith<String> for BfTreeProvider<T, QuantVectorProvider, I>
where
T: VectorRepr,
I: BfTreeId,
{
type Error = ANNError;
async fn load_with<P>(storage: &P, prefix: &String) -> Result<Self, Self::Error>
where
P: StorageReadProvider,
{
let saved_params: SavedParams = {
let params_filename = BfTreePaths::params_json(prefix);
let mut params_reader = storage.open_reader(¶ms_filename)?;
let mut params_json = String::new();
params_reader.read_to_string(&mut params_json)?;
serde_json::from_str(¶ms_json).map_err(|e| {
ANNError::message(lazy_format!(move, "Failed to deserialize params: {}", e))
})?
};
validate_loaded_id_params::<I>(&saved_params)?;
let _quant_params = saved_params.quant_params.ok_or_else(|| {
ANNError::message("Missing quant_params in saved params for quantized provider")
})?;
let metric = Metric::from_str(&saved_params.metric)
.map_err(|e| ANNError::message(lazy_format!(move, "Failed to parse metric: {}", e)))?;
let vector_index = load_bftree(
BfTreePaths::vectors_bftree(&saved_params.prefix),
saved_params.use_snapshot,
)?;
let full_vectors = VectorProvider::<T>::new_from_bftree(
saved_params.max_points,
saved_params.dim,
saved_params.frozen_points.get(),
vector_index,
);
let adjacency_list_index = load_bftree(
BfTreePaths::neighbors_bftree(&saved_params.prefix),
saved_params.use_snapshot,
)?;
let neighbor_provider =
NeighborProvider::<I>::new_from_bftree(saved_params.max_degree, adjacency_list_index)?;
let filename = BfTreePaths::quant_data_bin(&saved_params.prefix);
let mut reader = storage.open_reader(&filename)?;
let mut bytes = Vec::new();
reader.read_to_end(&mut bytes)?;
let quantizer: Poly<dyn Quantizer> = try_deserialize(&bytes, GlobalAllocator)
.map_err(|e| ANNError::message(lazy_format!(move, "{e}")))?;
let quant_vector_index = load_bftree(
BfTreePaths::quant_bftree(&saved_params.prefix),
saved_params.use_snapshot,
)?;
let quant_vectors = QuantVectorProvider::new_from_bftree(quantizer, quant_vector_index);
Ok(Self {
quant_vectors,
full_vectors,
neighbor_provider,
metric,
graph_params: saved_params.graph_params,
use_snapshot: saved_params.use_snapshot,
locks: StripedLocks::new(),
})
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::neighbors::NeighborProvider;
use crate::quant::create_test_quantizer;
use crate::vectors::VectorProvider;
use diskann::{
graph::DiskANNIndex,
graph::{self, search::Knn},
neighbor::BackInserter,
};
use diskann_providers::storage::FileStorageProvider;
use diskann_utils::views::{Init, Matrix};
fn create_quant_index() -> Arc<DiskANNIndex<BfTreeProvider<f32, QuantVectorProvider>>> {
let start_point = Matrix::new(Init(|| 0.0f32), 1, 5);
let dim = 5;
let logical_max_degree = 6;
let physical_max_degree = (logical_max_degree as f32 * 1.3) as u32;
let metric = Metric::L2;
let provider = BfTreeProvider::new(
BfTreeProviderParameters {
max_points: 20,
num_start_points: NonZeroUsize::new(1).unwrap(),
dim,
metric,
max_degree: physical_max_degree,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_point.as_view(),
create_test_quantizer(5),
)
.unwrap();
let index_config = graph::config::Builder::new_with(
logical_max_degree as usize,
graph::config::MaxDegree::new(physical_max_degree as usize),
10,
metric.into(),
|_| {},
)
.build()
.unwrap();
Arc::new(DiskANNIndex::new(index_config, provider, None))
}
#[tokio::test]
async fn test_quantized_index_search() {
let index = create_quant_index();
let ctx = &DefaultContext;
for i in 0..15 {
let point = vec![i as f32; 5];
index
.insert(&Quantized, ctx, &i, point.as_slice())
.await
.unwrap();
}
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let k = 5;
let mut neighbors = vec![Neighbor::<u32>::default(); k];
let res = index
.search(
params,
&Quantized,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(
res.result_count, k as u32,
"there are 15 points and we're asking for {}, we expect {}",
5, k
);
assert_eq!(*neighbors[0].id(), 3);
}
#[tokio::test]
async fn test_quantized_index_search_u64_ids() {
let start_point = Matrix::new(Init(|| 0.0f32), 1, 5);
let dim = 5;
let logical_max_degree = 6;
let physical_max_degree = (logical_max_degree as f32 * 1.3) as u32;
let metric = Metric::L2;
let provider: BfTreeProvider<f32, QuantVectorProvider, u64> = BfTreeProvider::new(
BfTreeProviderParameters {
max_points: 20,
num_start_points: NonZeroUsize::new(1).unwrap(),
dim,
metric,
max_degree: physical_max_degree,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_point.as_view(),
create_test_quantizer(5),
)
.unwrap();
let index_config = graph::config::Builder::new_with(
logical_max_degree as usize,
graph::config::MaxDegree::new(physical_max_degree as usize),
10,
metric.into(),
|_| {},
)
.build()
.unwrap();
let index = Arc::new(DiskANNIndex::new(index_config, provider, None));
let ctx = &DefaultContext;
for i in 0u64..15 {
let point = vec![i as f32; 5];
index
.insert(&Quantized, ctx, &i, point.as_slice())
.await
.unwrap();
}
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let mut neighbors = vec![Neighbor::<u64>::default(); 5];
let res = index
.search(
params,
&Quantized,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(
res.result_count, 5,
"there are 15 points and we're asking for 5, we expect 5"
);
assert_eq!(*neighbors[0].id(), 3u64);
}
#[tokio::test]
async fn test_quantized_index_search_u64_high_ids() {
let start_point = Matrix::new(Init(|| 0.0f32), 1, 5);
let dim = 5;
let logical_max_degree = 6;
let physical_max_degree = (logical_max_degree as f32 * 1.3) as u32;
let metric = Metric::L2;
let base: u64 = (u32::MAX as u64) + 1;
let num_points: u64 = 15;
let provider: BfTreeProvider<f32, QuantVectorProvider, u64> = BfTreeProvider::new(
BfTreeProviderParameters {
max_points: base as usize + num_points as usize + 1,
num_start_points: NonZeroUsize::new(1).unwrap(),
dim,
metric,
max_degree: physical_max_degree,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_point.as_view(),
create_test_quantizer(5),
)
.unwrap();
for id in provider.starting_points().unwrap() {
assert!(
id > u32::MAX as u64,
"start point id {id} should be in the u64-only range"
);
}
let index_config = graph::config::Builder::new_with(
logical_max_degree as usize,
graph::config::MaxDegree::new(physical_max_degree as usize),
10,
metric.into(),
|_| {},
)
.build()
.unwrap();
let index = Arc::new(DiskANNIndex::new(index_config, provider, None));
let ctx = &DefaultContext;
for k in 0u64..num_points {
let id = base + k;
let point = vec![k as f32; 5];
index
.insert(&Quantized, ctx, &id, point.as_slice())
.await
.unwrap();
}
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let mut neighbors = vec![Neighbor::<u64>::default(); 5];
let res = index
.search(
params,
&Quantized,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(
res.result_count, 5,
"there are 15 points and we're asking for 5, we expect 5"
);
assert_eq!(*neighbors[0].id(), base + 3);
assert!(
*neighbors[0].id() > u32::MAX as u64,
"returned id {} must be in the u64-only range",
neighbors[0].id()
);
}
#[tokio::test]
async fn test_quantized_index_multi_insert_search() {
let index = create_quant_index();
let ctx = &DefaultContext;
let data = Matrix::new(
Init({
let mut row = 0usize;
let mut col = 0usize;
move || {
let val = row as f32;
col += 1;
if col == 5 {
col = 0;
row += 1;
}
val
}
}),
15,
5,
);
let ids: Arc<[u32]> = (0u32..15).collect::<Vec<_>>().into();
let batch: Arc<Matrix<f32>> = Arc::new(data);
index
.multi_insert::<Quantized, Matrix<f32>>(Quantized, ctx, batch, ids)
.await
.unwrap();
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let k = 5;
let mut neighbors = vec![Neighbor::<u32>::default(); k];
let res = index
.search(
params,
&Quantized,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(
res.result_count, k as u32,
"there are 15 points and we're asking for {}, we expect {}",
5, k
);
let neighbor_ids: Vec<u32> = neighbors.iter().map(|n| *n.id()).collect();
for expected in 1u32..=5 {
assert!(
neighbor_ids.contains(&expected),
"expected id {expected} in results, got {neighbor_ids:?}"
);
}
}
#[tokio::test]
async fn test_quantized_delete_and_search() {
let index = create_quant_index();
let ctx = &DefaultContext;
for i in 0..15 {
let point = vec![i as f32; 5];
index
.insert(&Quantized, ctx, &i, point.as_slice())
.await
.unwrap();
}
index
.inplace_delete(Quantized, ctx, &2u32, 2, graph::InplaceDeleteMethod::OneHop)
.await
.unwrap();
index
.inplace_delete(Quantized, ctx, &4u32, 2, graph::InplaceDeleteMethod::OneHop)
.await
.unwrap();
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let k = 5;
let mut neighbors = vec![Neighbor::<u32>::default(); k];
let res = index
.search(
params,
&Quantized,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(res.result_count, k as u32);
let neighbor_ids: Vec<u32> = neighbors.iter().map(|n| *n.id()).collect();
assert!(!neighbor_ids.contains(&2u32));
assert!(!neighbor_ids.contains(&4u32));
}
fn create_full_precision_index() -> Arc<DiskANNIndex<BfTreeProvider<f32, NoStore>>> {
let start_point = Matrix::new(Init(|| 0.0f32), 1, 5);
let logical_max_degree = 6;
let physical_max_degree = (logical_max_degree as f32 * 1.3) as u32;
let metric = Metric::L2;
let provider = BfTreeProvider::new(
BfTreeProviderParameters {
max_points: 20,
num_start_points: NonZeroUsize::new(1).unwrap(),
dim: 5,
metric,
max_degree: physical_max_degree,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_point.as_view(),
NoStore,
)
.unwrap();
let index_config = graph::config::Builder::new_with(
logical_max_degree as usize,
graph::config::MaxDegree::new(physical_max_degree as usize),
10,
metric.into(),
|_| {},
)
.build()
.unwrap();
Arc::new(DiskANNIndex::new(index_config, provider, None))
}
#[tokio::test]
async fn test_full_precision_index_search() {
let index = create_full_precision_index();
let ctx = &DefaultContext;
for i in 0u32..15 {
let point = vec![i as f32; 5];
index
.insert(&FullPrecision, ctx, &i, point.as_slice())
.await
.unwrap();
}
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let k = 5;
let mut neighbors = vec![Neighbor::<u32>::default(); k];
let res = index
.search(
params,
&FullPrecision,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(
res.result_count, k as u32,
"there are 15 points and we're asking for {}, we expect {}",
5, k
);
assert_eq!(*neighbors[0].id(), 3);
}
#[tokio::test]
async fn test_u64_id_recall_parity_on_transposed_dataset() {
const N: usize = 1500;
const DIM: usize = 16;
const NUM_QUERIES: usize = 30;
const K: usize = 10;
const L: usize = 64;
struct SplitMix64(u64);
impl SplitMix64 {
fn next_u64(&mut self) -> u64 {
self.0 = self.0.wrapping_add(0x9E37_79B9_7F4A_7C15);
let mut z = self.0;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^ (z >> 31)
}
fn next_f32(&mut self) -> f32 {
(self.next_u64() >> 40) as f32 / (1u64 << 24) as f32
}
}
let mut rng = SplitMix64(0x0DDB_1A5E_5EED_1234);
let gen_vec = |rng: &mut SplitMix64| (0..DIM).map(|_| rng.next_f32()).collect::<Vec<f32>>();
let data: Vec<Vec<f32>> = (0..N).map(|_| gen_vec(&mut rng)).collect();
let queries: Vec<Vec<f32>> = (0..NUM_QUERIES).map(|_| gen_vec(&mut rng)).collect();
let l2 = |a: &[f32], b: &[f32]| -> f32 {
a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum::<f32>()
};
let ground_truth: Vec<Vec<usize>> = queries
.iter()
.map(|q| {
let mut scored: Vec<(f32, usize)> = data
.iter()
.enumerate()
.map(|(i, v)| (l2(q, v), i))
.collect();
scored.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap().then(a.1.cmp(&b.1)));
scored.into_iter().take(K).map(|(_, i)| i).collect()
})
.collect();
async fn build_and_search<I, F>(
data: &[Vec<f32>],
queries: &[Vec<f32>],
max_points: usize,
id_of: F,
) -> Vec<Vec<I>>
where
I: BfTreeId,
F: Fn(usize) -> I,
{
let logical_max_degree = 32usize;
let physical_max_degree = (logical_max_degree as f32 * 1.3) as u32;
let metric = Metric::L2;
let start_point = Matrix::new(Init(|| 0.0f32), 1, DIM);
let provider: BfTreeProvider<f32, NoStore, I> = BfTreeProvider::new(
BfTreeProviderParameters {
max_points,
num_start_points: NonZeroUsize::new(1).unwrap(),
dim: DIM,
metric,
max_degree: physical_max_degree,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_point.as_view(),
NoStore,
)
.unwrap();
let index_config = graph::config::Builder::new_with(
logical_max_degree,
graph::config::MaxDegree::new(physical_max_degree as usize),
L,
metric.into(),
|_| {},
)
.build()
.unwrap();
let index = Arc::new(DiskANNIndex::new(index_config, provider, None));
let ctx = &DefaultContext;
for (row, vector) in data.iter().enumerate() {
index
.insert(&FullPrecision, ctx, &id_of(row), vector.as_slice())
.await
.unwrap();
}
let mut results = Vec::with_capacity(queries.len());
for query in queries {
let params = Knn::new(L, None).unwrap();
let mut neighbors = vec![Neighbor::<I>::default(); K];
let res = index
.search(
params,
&FullPrecision,
ctx,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
results.push(
neighbors[..res.result_count as usize]
.iter()
.map(|n| *n.id())
.collect(),
);
}
results
}
let u32_results = build_and_search::<u32, _>(&data, &queries, N, |row| row as u32).await;
let base: u64 = (u32::MAX as u64) + 1;
let u64_results =
build_and_search::<u64, _>(&data, &queries, base as usize + N, |row| base + row as u64)
.await;
let mut recall_hits_u32 = 0usize;
let mut recall_hits_u64 = 0usize;
let total = NUM_QUERIES * K;
for qi in 0..NUM_QUERIES {
let mut rows_u32: Vec<usize> = u32_results[qi].iter().map(|&id| id as usize).collect();
let mut rows_u64: Vec<usize> = Vec::with_capacity(u64_results[qi].len());
for &id in &u64_results[qi] {
assert!(
id > u32::MAX as u64,
"query {qi}: transposed id {id} must live in the u64-only range"
);
rows_u64.push((id - base) as usize);
}
let mut sorted_u32 = rows_u32.clone();
let mut sorted_u64 = rows_u64.clone();
sorted_u32.sort_unstable();
sorted_u64.sort_unstable();
assert_eq!(
sorted_u32, sorted_u64,
"query {qi}: u64-transposed run returned a different neighbor set than the u32 baseline"
);
let gt: std::collections::HashSet<usize> = ground_truth[qi].iter().copied().collect();
rows_u32.retain(|r| gt.contains(r));
rows_u64.retain(|r| gt.contains(r));
recall_hits_u32 += rows_u32.len();
recall_hits_u64 += rows_u64.len();
}
let recall_u32 = recall_hits_u32 as f64 / total as f64;
let recall_u64 = recall_hits_u64 as f64 / total as f64;
assert_eq!(
recall_hits_u32, recall_hits_u64,
"recall must be identical between the u32 baseline and the u64-transposed run"
);
assert!(
recall_u64 > 0.8,
"recall floor not met (u32={recall_u32:.4}, u64={recall_u64:.4}); \
parity is only meaningful between two working indexes"
);
}
#[tokio::test]
async fn test_full_precision_delete_and_search() {
let index = create_full_precision_index();
let ctx = &DefaultContext;
for i in 0u32..15 {
let point = vec![i as f32; 5];
index
.insert(&FullPrecision, ctx, &i, point.as_slice())
.await
.unwrap();
}
index
.inplace_delete(
FullPrecision,
ctx,
&2u32,
2,
graph::InplaceDeleteMethod::OneHop,
)
.await
.unwrap();
index
.inplace_delete(
FullPrecision,
ctx,
&4u32,
2,
graph::InplaceDeleteMethod::OneHop,
)
.await
.unwrap();
let query = vec![3.0; 5];
let params = Knn::new(10, None).unwrap();
let k = 5;
let mut neighbors = vec![Neighbor::<u32>::default(); k];
let res = index
.search(
params,
&FullPrecision,
&DefaultContext,
query.as_slice(),
&mut BackInserter::new(neighbors.as_mut_slice()),
)
.await
.unwrap();
assert_eq!(res.result_count, k as u32);
let neighbor_ids: Vec<u32> = neighbors.iter().map(|n| *n.id()).collect();
assert!(!neighbor_ids.contains(&2u32));
assert!(!neighbor_ids.contains(&4u32));
}
#[tokio::test]
async fn test_data_provider_and_delete_interface() {
let ctx = &DefaultContext;
let num_start_points = 2;
let dim = 5;
let start_points = Matrix::try_from(
vec![0.0f32; dim]
.into_iter()
.chain(vec![0.5f32; dim])
.collect::<Box<[_]>>(),
num_start_points,
dim,
)
.unwrap();
let provider = BfTreeProvider::new(
BfTreeProviderParameters {
max_points: 10,
num_start_points: NonZeroUsize::new(num_start_points).unwrap(),
dim,
metric: Metric::L2,
max_degree: 64,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_points.as_view(),
NoStore,
)
.unwrap();
assert!((&provider).into_iter().eq(0u32..(10 + 2)));
let iter = provider.iter();
for i in iter.clone() {
let vector: Vec<f32> = (0..5).map(|j| (i * 5 + j) as f32).collect();
provider.set_element(ctx, &i, &vector).await.unwrap();
}
for i in iter.clone() {
assert_eq!(provider.to_external_id(ctx, i).unwrap(), i);
assert_eq!(provider.to_internal_id(ctx, &i).unwrap(), i);
assert_eq!(
provider.status_by_internal_id(ctx, i).await.unwrap(),
ElementStatus::Valid
);
assert_eq!(
provider.status_by_external_id(ctx, &i).await.unwrap(),
ElementStatus::Valid
);
provider.delete(ctx, &i).await.unwrap();
assert_eq!(
provider.status_by_internal_id(ctx, i).await.unwrap(),
ElementStatus::Deleted
);
assert_eq!(
provider.status_by_external_id(ctx, &i).await.unwrap(),
ElementStatus::Deleted
);
}
for i in iter.clone() {
provider.release(ctx, i).await.unwrap();
assert_eq!(
provider.status_by_internal_id(ctx, i).await.unwrap(),
ElementStatus::Deleted
);
}
assert!(provider
.set_element(ctx, &100, &[1.0, 2.0, 3.0, 4.0])
.await
.is_err());
}
#[tokio::test]
async fn test_empty_neighbor_list() {
let num_points = 100u32;
let ctx = &DefaultContext;
let num_start_points = 2;
let dim = 3;
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points, dim);
let provider = BfTreeProvider::<f32, _>::new(
BfTreeProviderParameters {
max_points: num_points as usize,
num_start_points: NonZeroUsize::new(num_start_points).unwrap(),
dim,
metric: Metric::L2,
max_degree: 64,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
},
start_points.as_view(),
NoStore,
)
.unwrap();
let mut scratch = provider.neighbor_provider.scratch(&provider.locks);
for i in 0..num_points {
let vector = vec![i as f32, (i + 1) as f32, (i + 2) as f32];
provider.set_element(ctx, &i, &vector).await.unwrap();
let mut out = AdjacencyList::new();
provider
.neighbor_provider
.get_neighbors(i, &mut out)
.unwrap();
assert!(out.is_empty());
scratch.write_neighbors(i, &[]).unwrap();
provider
.neighbor_provider
.get_neighbors(i, &mut out)
.unwrap();
assert!(out.is_empty());
}
for i in 0..num_points {
let mut out = AdjacencyList::new();
let neighbors = vec![10, 20, 30, 40, 50, 60, 70, 80, 90, 100];
scratch.write_neighbors(i, &neighbors).unwrap();
provider
.neighbor_provider
.get_neighbors(i, &mut out)
.unwrap();
assert_eq!(&*out, &[10, 20, 30, 40, 50, 60, 70, 80, 90, 100]);
scratch.write_neighbors(i, &[]).unwrap();
provider
.neighbor_provider
.get_neighbors(i, &mut out)
.unwrap();
assert!(out.is_empty());
}
let mut out = AdjacencyList::from_iter_untrusted([10, 20, 30, 40, 50, 60, 70, 80, 90, 100]);
provider
.neighbor_provider
.get_neighbors(200, &mut out)
.unwrap();
assert!(out.is_empty());
}
use tempfile::tempdir;
#[tokio::test]
async fn test_bf_tree_provider_save_load_no_quant() {
let num_points = 50usize;
let dim = 4usize;
let max_degree = 32u32;
let num_start_points = NonZeroUsize::new(2).unwrap();
let ctx = &DefaultContext;
let temp_dir = tempdir().unwrap();
let temp_path = temp_dir.path();
let prefix = temp_path
.join("test_bf_tree_provider")
.to_string_lossy()
.to_string();
let vector_path = BfTreePaths::vectors_bftree(&prefix);
let neighbor_path = BfTreePaths::neighbors_bftree(&prefix);
let bytes_vector = 1024 * 1024;
let mut vector_config = Config::new(&vector_path, bytes_vector);
vector_config.leaf_page_size(8192);
vector_config.cb_max_record_size(1024);
vector_config.storage_backend(bf_tree::StorageBackend::Std);
vector_config.use_snapshot(true);
let bytes_neighbor = 1024 * 1024;
let mut neighbor_config = Config::new(&neighbor_path, bytes_neighbor);
neighbor_config.storage_backend(bf_tree::StorageBackend::Std);
neighbor_config.use_snapshot(true);
let params = BfTreeProviderParameters {
max_points: num_points,
num_start_points,
dim,
metric: Metric::L2,
max_degree,
vector_provider_config: vector_config.clone(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: neighbor_config.clone(),
graph_params: None,
use_snapshot: true,
};
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points.into(), dim);
let provider =
BfTreeProvider::<f32, NoStore>::new(params.clone(), start_points.as_view(), NoStore)
.unwrap();
for i in 0..num_points {
let vector: Vec<f32> = (0..dim).map(|j| (i * dim + j) as f32 * 0.1).collect();
provider
.set_element(ctx, &(i as u32), &vector)
.await
.unwrap();
}
let mut scratch = provider.neighbor_provider.scratch(&provider.locks);
for i in 0..num_points as u32 {
let neighbors: Vec<u32> = (0..std::cmp::min(i, max_degree))
.map(|j| (i + j) % num_points as u32)
.collect();
scratch.write_neighbors(i, &neighbors).unwrap();
}
assert_eq!(vector_config.get_leaf_page_size(), 8192);
assert_eq!(vector_config.get_cb_max_record_size(), 1024);
let storage = FileStorageProvider;
let save_dir = tempdir().unwrap();
let save_prefix = save_dir
.path()
.join("saved_bf_tree_provider")
.to_string_lossy()
.to_string();
provider.save_with(&storage, &save_prefix).await.unwrap();
let loaded_provider = BfTreeProvider::<f32, NoStore>::load_with(&storage, &save_prefix)
.await
.unwrap();
for i in 0..num_points as u32 {
let original = provider.full_vectors.get_vector_sync(i as usize).unwrap();
let loaded = loaded_provider
.full_vectors
.get_vector_sync(i as usize)
.unwrap();
assert_eq!(original, loaded, "Vector mismatch at index {}", i);
}
for i in 0..num_points as u32 {
let mut original_list = AdjacencyList::new();
let mut loaded_list = AdjacencyList::new();
provider
.neighbor_provider
.get_neighbors(i, &mut original_list)
.unwrap();
loaded_provider
.neighbor_provider
.get_neighbors(i, &mut loaded_list)
.unwrap();
assert_eq!(
&*original_list, &*loaded_list,
"Neighbor list mismatch at index {}",
i
);
}
}
#[tokio::test]
async fn test_bf_tree_provider_save_load_quant() {
let num_points = 50usize;
let dim = 8usize;
let max_degree = 32u32;
let num_start_points = NonZeroUsize::new(2).unwrap();
let ctx = &DefaultContext;
let temp_dir = tempdir().unwrap();
let temp_path = temp_dir.path();
let prefix = temp_path
.join("test_bf_tree_provider_quant")
.to_string_lossy()
.to_string();
let vector_path = BfTreePaths::vectors_bftree(&prefix);
let neighbor_path = BfTreePaths::neighbors_bftree(&prefix);
let quant_path = BfTreePaths::quant_bftree(&prefix);
let bytes_vector = 1024 * 1024;
let mut vector_config = Config::new(&vector_path, bytes_vector);
vector_config.storage_backend(bf_tree::StorageBackend::Std);
vector_config.use_snapshot(true);
let bytes_neighbor = 1024 * 1024;
let mut neighbor_config = Config::new(&neighbor_path, bytes_neighbor);
neighbor_config.storage_backend(bf_tree::StorageBackend::Std);
neighbor_config.use_snapshot(true);
let bytes_quant = 1024 * 1024;
let mut quant_config = Config::new(&quant_path, bytes_quant);
quant_config.storage_backend(bf_tree::StorageBackend::Std);
quant_config.use_snapshot(true);
let quantizer = create_test_quantizer(dim);
let params = BfTreeProviderParameters {
max_points: num_points,
num_start_points,
dim,
metric: Metric::L2,
max_degree,
vector_provider_config: vector_config.clone(),
quant_vector_provider_config: quant_config.clone(),
neighbor_list_provider_config: neighbor_config.clone(),
graph_params: None,
use_snapshot: true,
};
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points.into(), dim);
let provider = BfTreeProvider::<f32, QuantVectorProvider>::new(
params.clone(),
start_points.as_view(),
quantizer,
)
.unwrap();
for i in 0..num_points {
let vector: Vec<f32> = (0..dim).map(|j| (i * dim + j) as f32 * 0.1).collect();
provider
.set_element(ctx, &(i as u32), &vector)
.await
.unwrap();
}
let mut scratch = provider.neighbor_provider.scratch(&provider.locks);
for i in 0..num_points as u32 {
let neighbors: Vec<u32> = (0..std::cmp::min(i, max_degree))
.map(|j| (i + j) % num_points as u32)
.collect();
scratch.write_neighbors(i, &neighbors).unwrap();
}
let storage = FileStorageProvider;
let save_dir = tempdir().unwrap();
let save_prefix = save_dir
.path()
.join("saved_bf_tree_provider_quant")
.to_string_lossy()
.to_string();
provider.save_with(&storage, &save_prefix).await.unwrap();
let loaded_provider =
BfTreeProvider::<f32, QuantVectorProvider>::load_with(&storage, &save_prefix)
.await
.unwrap();
assert_eq!(
provider.quant_vectors.quantizer.full_dim(),
loaded_provider.quant_vectors.quantizer.full_dim(),
"Quantizer full_dim mismatch"
);
assert_eq!(
provider.quant_vectors.quantizer.bytes(),
loaded_provider.quant_vectors.quantizer.bytes(),
"Quantizer bytes mismatch"
);
assert_eq!(
provider.quant_vectors.quantizer.nbits(),
loaded_provider.quant_vectors.quantizer.nbits(),
"Quantizer nbits mismatch"
);
for i in 0..num_points as u32 {
let original = provider.full_vectors.get_vector_sync(i as usize).unwrap();
let loaded = loaded_provider
.full_vectors
.get_vector_sync(i as usize)
.unwrap();
assert_eq!(original, loaded, "Vector mismatch at index {}", i);
}
for i in 0..num_points as u32 {
let original = provider.quant_vectors.get_vector_sync(i as usize).unwrap();
let loaded = loaded_provider
.quant_vectors
.get_vector_sync(i as usize)
.unwrap();
assert_eq!(original, loaded, "Quant vector mismatch at index {}", i);
}
for i in 0..num_points as u32 {
let mut original_list = AdjacencyList::new();
let mut loaded_list = AdjacencyList::new();
provider
.neighbor_provider
.get_neighbors(i, &mut original_list)
.unwrap();
loaded_provider
.neighbor_provider
.get_neighbors(i, &mut loaded_list)
.unwrap();
assert_eq!(
&*original_list, &*loaded_list,
"Neighbor list mismatch at index {}",
i
);
}
}
#[tokio::test]
async fn test_bf_tree_provider_memory_save_load_no_quant() {
let num_points = 20usize;
let dim = 4usize;
let max_degree = 16u32;
let num_start_points = NonZeroUsize::new(1).unwrap();
let ctx = &DefaultContext;
let mut vector_config = Config::default();
vector_config.use_snapshot(true);
let mut neighbor_config = Config::default();
neighbor_config.use_snapshot(true);
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points.into(), dim);
let provider = BfTreeProvider::<f32, NoStore>::new(
BfTreeProviderParameters {
max_points: num_points,
num_start_points,
dim,
metric: Metric::L2,
max_degree,
vector_provider_config: vector_config,
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: neighbor_config,
graph_params: None,
use_snapshot: true,
},
start_points.as_view(),
NoStore,
)
.unwrap();
for i in 0..num_points {
let vector: Vec<f32> = (0..dim).map(|j| (i * dim + j) as f32 * 0.1).collect();
provider
.set_element(ctx, &(i as u32), &vector)
.await
.unwrap();
}
let mut scratch = provider.neighbor_provider.scratch(&provider.locks);
for i in 0..num_points as u32 {
let neighbors: Vec<u32> = (0..std::cmp::min(i, max_degree))
.map(|j| (i + j) % num_points as u32)
.collect();
scratch.write_neighbors(i, &neighbors).unwrap();
}
provider.delete(ctx, &3u32).await.unwrap();
provider.delete(ctx, &7u32).await.unwrap();
let save_dir = tempdir().unwrap();
let save_prefix = save_dir
.path()
.join("mem_no_quant")
.to_string_lossy()
.to_string();
let storage = FileStorageProvider;
provider.save_with(&storage, &save_prefix).await.unwrap();
let loaded = BfTreeProvider::<f32, NoStore>::load_with(&storage, &save_prefix)
.await
.unwrap();
for i in 0..num_points as u32 {
if i == 3 || i == 7 {
continue;
}
assert_eq!(
provider.full_vectors.get_vector_sync(i as usize).unwrap(),
loaded.full_vectors.get_vector_sync(i as usize).unwrap(),
"Vector mismatch at {}",
i
);
}
for i in 0..num_points as u32 {
let mut orig = AdjacencyList::new();
let mut load = AdjacencyList::new();
provider
.neighbor_provider
.get_neighbors(i, &mut orig)
.unwrap();
loaded
.neighbor_provider
.get_neighbors(i, &mut load)
.unwrap();
assert_eq!(&*orig, &*load, "Neighbor mismatch at {}", i);
}
assert_eq!(
loaded.status_by_internal_id(ctx, 3).await.unwrap(),
ElementStatus::Deleted
);
assert_eq!(
loaded.status_by_internal_id(ctx, 7).await.unwrap(),
ElementStatus::Deleted
);
assert_eq!(
loaded.status_by_internal_id(ctx, 0).await.unwrap(),
ElementStatus::Valid
);
}
#[tokio::test]
async fn test_bf_tree_provider_memory_save_load_quant() {
let num_points = 20usize;
let dim = 8usize;
let max_degree = 16u32;
let num_start_points = NonZeroUsize::new(1).unwrap();
let ctx = &DefaultContext;
let quantizer = create_test_quantizer(dim);
let mut vector_config = Config::default();
vector_config.use_snapshot(true);
let mut neighbor_config = Config::default();
neighbor_config.use_snapshot(true);
let mut quant_config = Config::default();
quant_config.use_snapshot(true);
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points.into(), dim);
let provider = BfTreeProvider::<f32, QuantVectorProvider>::new(
BfTreeProviderParameters {
max_points: num_points,
num_start_points,
dim,
metric: Metric::L2,
max_degree,
vector_provider_config: vector_config,
quant_vector_provider_config: quant_config,
neighbor_list_provider_config: neighbor_config,
graph_params: None,
use_snapshot: true,
},
start_points.as_view(),
quantizer,
)
.unwrap();
for i in 0..num_points {
let vector: Vec<f32> = (0..dim).map(|j| (i * dim + j) as f32 * 0.1).collect();
provider
.set_element(ctx, &(i as u32), &vector)
.await
.unwrap();
}
let mut scratch = provider.neighbor_provider.scratch(&provider.locks);
for i in 0..num_points as u32 {
let neighbors: Vec<u32> = (0..std::cmp::min(i, max_degree))
.map(|j| (i + j) % num_points as u32)
.collect();
scratch.write_neighbors(i, &neighbors).unwrap();
}
provider.delete(ctx, &2u32).await.unwrap();
let save_dir = tempdir().unwrap();
let save_prefix = save_dir
.path()
.join("mem_quant")
.to_string_lossy()
.to_string();
let storage = FileStorageProvider;
provider.save_with(&storage, &save_prefix).await.unwrap();
let loaded = BfTreeProvider::<f32, QuantVectorProvider>::load_with(&storage, &save_prefix)
.await
.unwrap();
for i in 0..num_points as u32 {
if i == 2 {
continue;
}
assert_eq!(
provider.full_vectors.get_vector_sync(i as usize).unwrap(),
loaded.full_vectors.get_vector_sync(i as usize).unwrap(),
"Vector mismatch at {}",
i
);
}
for i in 0..num_points as u32 {
if i == 2 {
continue;
}
assert_eq!(
provider.quant_vectors.get_vector_sync(i as usize).unwrap(),
loaded.quant_vectors.get_vector_sync(i as usize).unwrap(),
"Quant vector mismatch at {}",
i
);
}
for i in 0..num_points as u32 {
if i == 2 {
continue;
}
let mut orig = AdjacencyList::new();
let mut load = AdjacencyList::new();
provider
.neighbor_provider
.get_neighbors(i, &mut orig)
.unwrap();
loaded
.neighbor_provider
.get_neighbors(i, &mut load)
.unwrap();
assert_eq!(&*orig, &*load, "Neighbor mismatch at {}", i);
}
assert_eq!(
loaded.status_by_internal_id(ctx, 2).await.unwrap(),
ElementStatus::Deleted
);
assert_eq!(
loaded.status_by_internal_id(ctx, 0).await.unwrap(),
ElementStatus::Valid
);
}
#[test]
fn test_validate_rejects_undersized_vector_config() {
let result = VectorProvider::<f32>::new_with_config(
100,
1536,
1,
Config::default(), );
let err = result.err().expect("should fail").to_string();
assert!(
err.contains("vector_provider"),
"should name the failing config; got: {err}"
);
assert!(
err.contains("6152"),
"should state the required size; got: {err}"
);
}
#[test]
fn test_validate_rejects_undersized_neighbor_config() {
let mut neighbor_config = Config::default();
neighbor_config.cb_max_record_size(1952);
let result = NeighborProvider::<u32>::new_with_config(500, neighbor_config);
let err = result.err().expect("should fail").to_string();
assert!(
err.contains("neighbor_provider"),
"should name the failing config; got: {err}"
);
}
#[test]
fn test_validate_accepts_valid_config() {
if let Err(e) = VectorProvider::<f32>::new_with_config(100, 128, 1, Config::default()) {
panic!("VectorProvider should succeed: {e}");
}
if let Err(e) = NeighborProvider::<u32>::new_with_config(64, Config::default()) {
panic!("NeighborProvider should succeed: {e}");
}
}
#[tokio::test]
async fn test_bf_tree_provider_save_load_u64_ids() {
let num_points = 16usize;
let dim = 4usize;
let max_degree = 16u32;
let num_start_points = NonZeroUsize::new(2).unwrap();
let ctx = &DefaultContext;
let temp_dir = tempdir().unwrap();
let prefix = temp_dir
.path()
.join("u64_provider")
.to_string_lossy()
.to_string();
let mut vector_config = Config::new(BfTreePaths::vectors_bftree(&prefix), 1024 * 1024);
vector_config.storage_backend(bf_tree::StorageBackend::Std);
vector_config.use_snapshot(true);
let mut neighbor_config = Config::new(BfTreePaths::neighbors_bftree(&prefix), 1024 * 1024);
neighbor_config.storage_backend(bf_tree::StorageBackend::Std);
neighbor_config.use_snapshot(true);
let params = BfTreeProviderParameters {
max_points: num_points,
num_start_points,
dim,
metric: Metric::L2,
max_degree,
vector_provider_config: vector_config,
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: neighbor_config,
graph_params: None,
use_snapshot: true,
};
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points.into(), dim);
let provider =
BfTreeProvider::<f32, NoStore, u64>::new(params, start_points.as_view(), NoStore)
.unwrap();
for i in 0..num_points {
let vector: Vec<f32> = (0..dim).map(|j| (i * dim + j) as f32 * 0.1).collect();
provider
.set_element(ctx, &(i as u64), &vector)
.await
.unwrap();
}
let mut scratch = provider.neighbor_provider.scratch(&provider.locks);
for i in 0..num_points as u64 {
let neighbors: Vec<u64> = (0..std::cmp::min(i, max_degree as u64))
.map(|j| (i + j) % num_points as u64)
.collect();
scratch.write_neighbors(i, &neighbors).unwrap();
}
drop(scratch);
let storage = FileStorageProvider;
let save_dir = tempdir().unwrap();
let save_prefix = save_dir
.path()
.join("saved_u64_provider")
.to_string_lossy()
.to_string();
provider.save_with(&storage, &save_prefix).await.unwrap();
let loaded = BfTreeProvider::<f32, NoStore, u64>::load_with(&storage, &save_prefix)
.await
.unwrap();
for i in 0..num_points as u64 {
let mut original_list = AdjacencyList::new();
let mut loaded_list = AdjacencyList::new();
provider
.neighbor_provider
.get_neighbors(i, &mut original_list)
.unwrap();
loaded
.neighbor_provider
.get_neighbors(i, &mut loaded_list)
.unwrap();
assert_eq!(&*original_list, &*loaded_list, "neighbor mismatch at {i}");
}
let mismatch = BfTreeProvider::<f32, NoStore, u32>::load_with(&storage, &save_prefix).await;
assert!(
mismatch.is_err(),
"loading a u64 index with a u32 id type must fail"
);
}
#[tokio::test]
async fn test_new_rejects_capacity_exceeding_id_type() {
let dim = 4usize;
let num_start_points = NonZeroUsize::new(1).unwrap();
let start_points = Matrix::new(Init(|| 0.0f32), num_start_points.into(), dim);
let params = BfTreeProviderParameters {
max_points: u32::MAX as usize + 2,
num_start_points,
dim,
metric: Metric::L2,
max_degree: 8,
vector_provider_config: Config::default(),
quant_vector_provider_config: Config::default(),
neighbor_list_provider_config: Config::default(),
graph_params: None,
use_snapshot: false,
};
let result =
BfTreeProvider::<f32, NoStore, u32>::new(params, start_points.as_view(), NoStore);
assert!(
result.is_err(),
"u32 provider must reject a capacity exceeding u32::MAX"
);
}
}