From 250c536e667882cc9092f957413467d740a83aeb Mon Sep 17 00:00:00 2001 From: Junkui Chen Date: Sat, 1 Aug 2026 23:08:06 +0800 Subject: [PATCH 1/2] Refine Vamana index builders Replace the trait-object in-memory builder facade with an explicit Vamana build index and split strategy, one-shot, merged, and test responsibilities into private modules. This refinement is important because it aligns the code with the algorithm it implements, makes FP/SQ/PQ dispatch explicit, and isolates the two build modes without changing index behavior. These boundaries reduce accidental coupling and make future Vamana changes safer to review and extend. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- diskann-disk/src/build/builder/build.rs | 249 +------- .../src/build/builder/inmem_builder.rs | 290 --------- diskann-disk/src/build/builder/mod.rs | 6 +- diskann-disk/src/build/builder/tests.rs | 2 +- .../src/build/builder/vamana/index.rs | 204 ++++++ .../src/build/builder/vamana/merged.rs | 475 ++++++++++++++ diskann-disk/src/build/builder/vamana/mod.rs | 16 + .../src/build/builder/vamana/one_shot.rs | 253 ++++++++ .../src/build/builder/vamana/strategy.rs | 121 ++++ .../builder/{core.rs => vamana/tests.rs} | 585 +----------------- .../src/search/provider/disk_provider.rs | 2 +- .../src/utils/instrumentation/perf_logger.rs | 6 +- 12 files changed, 1113 insertions(+), 1096 deletions(-) delete mode 100644 diskann-disk/src/build/builder/inmem_builder.rs create mode 100644 diskann-disk/src/build/builder/vamana/index.rs create mode 100644 diskann-disk/src/build/builder/vamana/merged.rs create mode 100644 diskann-disk/src/build/builder/vamana/mod.rs create mode 100644 diskann-disk/src/build/builder/vamana/one_shot.rs create mode 100644 diskann-disk/src/build/builder/vamana/strategy.rs rename diskann-disk/src/build/builder/{core.rs => vamana/tests.rs} (52%) diff --git a/diskann-disk/src/build/builder/build.rs b/diskann-disk/src/build/builder/build.rs index c0df5d6b7..b5258eac6 100644 --- a/diskann-disk/src/build/builder/build.rs +++ b/diskann-disk/src/build/builder/build.rs @@ -4,40 +4,26 @@ */ //! Async disk index builder implementation. -use std::{ - marker::PhantomData, - num::NonZeroUsize, - sync::{Arc, Mutex}, -}; +use std::marker::PhantomData; use crate::data_model::GraphDataType; -use diskann::{ - utils::{async_tools, VectorRepr, ONE}, - ANNResult, -}; +use diskann::{utils::VectorRepr, ANNResult}; use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider}; use diskann_providers::{ - model::{ - graph::provider::async_::inmem::DefaultProviderParameters, IndexConfiguration, - MAX_PQ_TRAINING_SET_SIZE, NUM_KMEANS_REPS_PQ, NUM_PQ_CENTROIDS, - }, - storage::{DiskGraphOnly, PQStorage}, - utils::{ - create_thread_pool, find_medoid_with_sampling, RayonThreadPoolRef, VectorDataIterator, - MAX_MEDOID_SAMPLE_SIZE, - }, + model::{IndexConfiguration, MAX_PQ_TRAINING_SET_SIZE, NUM_KMEANS_REPS_PQ, NUM_PQ_CENTROIDS}, + storage::PQStorage, + utils::{create_thread_pool, RayonThreadPoolRef}, }; -use tokio::task::JoinSet; -use tracing::{debug, info}; +use tracing::info; use crate::{ build::builder::{ - core::{determine_build_strategy, IndexBuildStrategy, MergedVamanaIndexBuilder}, - inmem_builder::{new_inmem_index_builder, InmemIndexBuilder}, quantizer::BuildQuantizer, tokio::create_runtime, + vamana::{ + determine_build_strategy, IndexBuildStrategy, MergedVamanaBuilder, OneShotVamanaBuilder, + }, }, - error::{diskann_error, ErrorKind}, storage::{ quant::{PQGeneration, PQGenerationContext, QuantDataGenerator}, DiskIndexWriter, @@ -123,8 +109,8 @@ where self.generate_compressed_data(pool.as_ref())?; logger.log_checkpoint(DiskIndexBuildCheckpoint::PqConstruction); - self.build_inmem_index(pool.as_ref()).await?; - logger.log_checkpoint(DiskIndexBuildCheckpoint::InmemIndexBuild); + self.build_vamana_index(pool.as_ref()).await?; + logger.log_checkpoint(DiskIndexBuildCheckpoint::VamanaIndexBuild); // Use physical file to pass the memory index to the disk writer self.create_disk_layout()?; @@ -172,31 +158,34 @@ where ) } - async fn build_inmem_index(&mut self, pool: RayonThreadPoolRef<'_>) -> ANNResult<()> { - match determine_build_strategy::( + async fn build_vamana_index(&mut self, pool: RayonThreadPoolRef<'_>) -> ANNResult<()> { + let strategy = determine_build_strategy::( &self.index_configuration, self.disk_build_param.build_memory_limit().in_bytes() as f64, self.disk_build_param.build_quantization(), - ) { - IndexBuildStrategy::Merged => { - MergedVamanaIndexBuilder::::new( + ); + + match strategy { + IndexBuildStrategy::OneShot => { + OneShotVamanaBuilder::::new( &self.index_configuration, - &self.disk_build_param, - &self.index_writer, &self.build_quantizer, + self.index_writer.get_dataset_file(), + self.index_writer.get_mem_index_file(), self.storage_provider, ) - .build(pool) + .build() .await } - IndexBuildStrategy::OneShot => { - build_inmem_index::( - self.index_configuration.clone(), + IndexBuildStrategy::Merged => { + MergedVamanaBuilder::::new( + &self.index_configuration, + &self.disk_build_param, + &self.index_writer, &self.build_quantizer, - &self.index_writer.get_dataset_file(), - &self.index_writer.get_mem_index_file(), self.storage_provider, ) + .build(pool) .await } } @@ -211,187 +200,3 @@ where Ok(()) } } - -pub(super) async fn build_inmem_index( - config: IndexConfiguration, - quantizer: &BuildQuantizer, - data_path: &str, - save_path: &str, - storage_provider: &StorageProvider, -) -> ANNResult<()> -where - T: VectorRepr, - StorageProvider: StorageReadProvider + StorageWriteProvider + 'static, - ::Reader: std::marker::Send, -{ - // use either user-specified number of threads or default to available parallelism - let num_tasks = NonZeroUsize::new(config.num_threads) - .or_else(|| std::thread::available_parallelism().ok()) - .ok_or_else(|| { - diskann_error!( - ErrorKind::IndexError, - "Failed to determine number of threads" - ) - })?; - - // Associated data will only be used in the write_disk_layout function which only requires the none-partitioned associated data stream. - let dataset_iter = Arc::new(Mutex::new({ - let iter = VectorDataIterator::<_, T>::new(data_path, Option::None, storage_provider)?; - iter.enumerate() - })); - - let index_config = config.config.clone(); - let provider_parameters = DefaultProviderParameters { - max_points: config.max_points, - frozen_points: ONE, - metric: config.dist_metric, - dim: config.dim, - max_degree: index_config.max_degree_u32().get(), - prefetch_lookahead: config.prefetch_lookahead.map(|x| x.get()), - prefetch_cache_line_level: config.prefetch_cache_line_level, - }; - let index = new_inmem_index_builder::(index_config, provider_parameters, quantizer)?; - let medoid_id = - set_start_point_to_medoid::(&index, data_path, config.random_seed, storage_provider)?; - let start_point = u32_try_from(medoid_id)?; - - run_build(&index, dataset_iter, num_tasks).await?; - - #[cfg(debug_assertions)] - log_build_stats::<_>(&index).await?; - - run_final_prune(&index, num_tasks).await?; - index - .save_graph( - storage_provider, - &(start_point, DiskGraphOnly::new(save_path)), - ) - .await?; - - Ok(()) -} - -#[cfg(debug_assertions)] -/// Log statistics about the build process -async fn log_build_stats(index: &Arc>) -> ANNResult<()> { - debug!( - "Number of points reachable in the graph: {}", - index.count_reachable_nodes().await? - ); - - let (full_vector, quant_vector) = index.counts_for_get_vector(); - let capacity = index.capacity(); - debug!( - "Number of get vector calls per insert: {}", - full_vector as f32 / capacity as f32 - ); - debug!( - "Number of get quantized vector calls per insert: {}", - quant_vector as f32 / capacity as f32 - ); - - Ok(()) -} - -/// Convert a `usize` index into the `u32` internal id type, erroring if it does not fit. -/// -/// The async index uses `u32` internal ids, so positions in the dataset must not exceed -/// `u32::MAX`. -fn u32_try_from(value: usize) -> ANNResult { - u32::try_from(value) - .map_err(|_| diskann_error!(ErrorKind::IndexError, "id {value} exceeds u32::MAX")) -} - -fn set_start_point_to_medoid( - index: &Arc>, - path: &str, - random_seed: Option, - reader: &StorageReader, -) -> ANNResult -where - T: VectorRepr, - StorageReader: StorageReadProvider, -{ - let mut rng = diskann_providers::utils::create_rnd_from_optional_seed(random_seed); - let (medoid, medoid_id) = - find_medoid_with_sampling::(path, reader, MAX_MEDOID_SAMPLE_SIZE, &mut rng)?; - - index.set_start_point(medoid.as_slice())?; - - debug!("Set start point to medoid ID: {}", medoid_id); - - Ok(medoid_id) -} - -async fn run_build( - index: &Arc>, - iterator: Arc>, - num_tasks: NonZeroUsize, -) -> ANNResult<()> -where - T: VectorRepr, - I: Iterator, ()))> + Send + 'static, -{ - let total_points = index.capacity(); - let partitions = async_tools::PartitionIter::new(total_points, num_tasks); - - let mut tasks = JoinSet::new(); - - for partition in partitions { - let index_clone = index.clone(); - let iterator_clone = iterator.clone(); - tasks.spawn(async move { - for _ in partition { - let vector_data = { - let mut guard = iterator_clone.lock().map_err(|_| { - diskann_error!(ErrorKind::IndexError, "Poisoned mutex during construction") - })?; - guard.next() - }; - - match vector_data { - Some((i, (vector, _))) => { - let id = u32_try_from(i)?; - index_clone.insert_vector(id, vector.as_ref()).await?; - } - None => break, - } - } - ANNResult::Ok(()) - }); - } - - // Wait for all tasks to complete. - while let Some(res) = tasks.join_next().await { - res.map_err(|_| diskann_error!(ErrorKind::IndexError, "A spawned insert task failed"))??; - } - - info!("Linked all points. Num points: #{}", total_points); - Ok(()) -} - -async fn run_final_prune( - index: &Arc>, - num_tasks: NonZeroUsize, -) -> ANNResult<()> { - let partitions = async_tools::PartitionIter::new(index.total_points(), num_tasks); - - let mut tasks = JoinSet::new(); - - for partition in partitions { - let index_clone = index.clone(); - tasks.spawn(async move { - let range = u32_try_from(partition.start)?..u32_try_from(partition.end)?; - index_clone.final_prune(range).await - }); - } - - // Wait for all final prune tasks to complete - while let Some(res) = tasks.join_next().await { - res.map_err(|_| { - diskann_error!(ErrorKind::IndexError, "A spawned final prune task failed") - })??; - } - - Ok(()) -} diff --git a/diskann-disk/src/build/builder/inmem_builder.rs b/diskann-disk/src/build/builder/inmem_builder.rs deleted file mode 100644 index 2eeab9780..000000000 --- a/diskann-disk/src/build/builder/inmem_builder.rs +++ /dev/null @@ -1,290 +0,0 @@ -/* - * Copyright (c) Microsoft Corporation. - * Licensed under the MIT license. - */ - -use std::{marker::PhantomData, pin::Pin, sync::Arc}; - -use diskann::{ - graph::{ - glue::{InsertStrategy, PruneStrategy}, - Config, DiskANNIndex, - }, - provider::DefaultContext, - utils::VectorRepr, - ANNError, ANNResult, -}; -use diskann_providers::storage::{DynWriteProvider, WriteProviderWrapper}; -use diskann_providers::{ - index::diskann_async, - model::graph::provider::async_::{ - common::{FullPrecision, NoDeletes, NoStore, Quantized, SetElementHelper, VectorStore}, - inmem::{ - DefaultProvider, DefaultProviderParameters, FullPrecisionProvider, SetStartPoints, - }, - }, - storage::{DiskGraphOnly, SaveWith}, -}; -use diskann_utils::future::{AsyncFriendly, SendFuture}; - -use super::quantizer::BuildQuantizer; - -/// Builder facade for in memory index construction and persistence. -/// -/// Thread safety: -/// Implementors must be `Send` and `Sync`. Methods can be called from many tasks. -pub(super) trait InmemIndexBuilder: Send + Sync { - /// Return the total capacity of the provider, **excluding** start points. - fn capacity(&self) -> usize; - - /// Return the total capacity of the provider, **including** start points. - fn total_points(&self) -> usize; - - /// Set a single start point to search. - /// - /// The slice must match the underlying vector type, else `WrongDataType` is returned. - fn set_start_point(&self, start_point: &[T]) -> ANNResult<()>; - - /// Insert a vector with a `id`. - /// - /// The slice must match the underlying vector type. - fn insert_vector<'a>( - &'a self, - id: u32, - vector: &'a [T], - ) -> Pin> + 'a>>; - - /// Prune the built graph over `[range.start, range.end)`. - fn final_prune( - &self, - range: core::ops::Range, - ) -> Pin> + '_>>; - - /// Persist only the graph file set. - fn save_graph<'a>( - &'a self, - storage_provider: &'a dyn DynWriteProvider, - start_point_and_path: &'a (u32, DiskGraphOnly), - ) -> Pin> + 'a>>; - - /// Return the number of vector reads for full_precision and quantized stores respectively. - #[cfg(debug_assertions)] - fn counts_for_get_vector(&self) -> (usize, usize); - - /// Count the number of nodes in the graph reachable from the given `start_points`. - /// - /// This function has a large memory footprint for large graphs and should not be called - /// frequently. This is mainly for analysis and sanity tests. - #[cfg(debug_assertions)] - fn count_reachable_nodes(&self) -> Pin> + '_>>; -} - -////////////////////////////////// -// FullPrecision Implementation // -////////////////////////////////// - -impl InmemIndexBuilder for DiskANNIndex> -where - T: VectorRepr, -{ - fn capacity(&self) -> usize { - self.provider().capacity() - } - - fn total_points(&self) -> usize { - self.provider().total_points() - } - - fn set_start_point(&self, start_point: &[T]) -> ANNResult<()> { - self.provider() - .set_start_points(std::iter::once(start_point)) - } - - fn insert_vector<'a>( - &'a self, - id: u32, - vector: &'a [T], - ) -> Pin> + 'a>> { - Box::pin(async move { - self.insert(&FullPrecision, &DefaultContext, &id, vector) - .await - }) - } - - fn final_prune( - &self, - range: core::ops::Range, - ) -> Pin> + '_>> { - Box::pin(async move { - self.prune_range(&FullPrecision, &DefaultContext, range) - .await - }) - } - - fn save_graph<'a>( - &'a self, - storage_provider: &'a dyn DynWriteProvider, - start_point_and_path: &'a (u32, DiskGraphOnly), - ) -> Pin> + 'a>> { - Box::pin(async move { - let wrapper = WriteProviderWrapper::new(storage_provider); - self.save_with(&wrapper, start_point_and_path).await - }) - } - - #[cfg(debug_assertions)] - fn counts_for_get_vector(&self) -> (usize, usize) { - self.provider().counts_for_get_vector() - } - - #[cfg(debug_assertions)] - fn count_reachable_nodes(&self) -> Pin> + '_>> { - Box::pin(async move { - let provider = self.provider(); - let start_points = provider.starting_points()?; - let mut neighbor_accessor = provider.neighbors(); - self.count_reachable_nodes(&start_points, &mut neighbor_accessor) - .await - }) - } -} - -////////////////////////// -// Quant Implementation // -////////////////////////// - -pub(super) struct QuantInMemBuilder -where - Q: AsyncFriendly, -{ - index: DiskANNIndex>, - _vector_data_type: PhantomData, -} - -impl QuantInMemBuilder -where - Q: AsyncFriendly, -{ - pub fn new(index: DiskANNIndex>) -> Self { - Self { - index, - _vector_data_type: PhantomData, - } - } - - fn index(&self) -> &DiskANNIndex> { - &self.index - } -} - -impl InmemIndexBuilder for QuantInMemBuilder -where - T: VectorRepr, - Q: AsyncFriendly + VectorStore + SetElementHelper, - Quantized: for<'a> InsertStrategy<'a, DefaultProvider, &'a [T]> - + PruneStrategy>, - DefaultProvider: SaveWith<(u32, u32, DiskGraphOnly), Error = ANNError>, -{ - fn capacity(&self) -> usize { - self.index().provider().capacity() - } - - fn total_points(&self) -> usize { - self.index().provider().total_points() - } - - fn set_start_point(&self, start_point: &[T]) -> ANNResult<()> { - self.index() - .provider() - .set_start_points(std::iter::once(start_point)) - } - - fn insert_vector<'a>( - &'a self, - id: u32, - vector: &'a [T], - ) -> Pin> + 'a>> { - Box::pin(async move { - self.index() - .insert(&Quantized, &DefaultContext, &id, vector) - .await - }) - } - - fn final_prune( - &self, - range: core::ops::Range, - ) -> Pin> + '_>> { - Box::pin(async move { - self.index() - .prune_range(&Quantized, &DefaultContext, range) - .await - }) - } - - fn save_graph<'a>( - &'a self, - storage_provider: &'a dyn DynWriteProvider, - start_point_and_path: &'a (u32, DiskGraphOnly), - ) -> Pin> + 'a>> { - Box::pin(async move { - let wrapper = WriteProviderWrapper::new(storage_provider); - self.index().save_with(&wrapper, start_point_and_path).await - }) - } - - #[cfg(debug_assertions)] - fn counts_for_get_vector(&self) -> (usize, usize) { - self.index().provider().counts_for_get_vector() - } - - #[cfg(debug_assertions)] - fn count_reachable_nodes(&self) -> Pin> + '_>> { - Box::pin(async move { - let provider = self.index().provider(); - let start_points = provider.starting_points()?; - let mut neighbor_accessor = provider.neighbors(); - self.index() - .count_reachable_nodes(&start_points, &mut neighbor_accessor) - .await - }) - } -} - -/// Create a new in-memory index builder for vectors of type `T`. -/// -/// Chooses the builder implementation based on the given `BuildQuantizer`. -/// - `NoQuant` uses a plain index with no quantization. -/// - `Scalar1Bit` and `PQ` create quantized only indexes backed by `QuantInMemBuilder`. -/// -/// # Parameters -/// * `config` – Index configuration. -/// * `build_quantizer` – Quantization strategy to apply. -/// -/// # Returns -/// An `Arc` wrapped in `ANNResult`. -/// -/// # Errors -/// Returns an error if the underlying index creation fails. -pub(super) fn new_inmem_index_builder( - config: Config, - params: DefaultProviderParameters, - build_quantizer: &BuildQuantizer, -) -> ANNResult>> -where - T: VectorRepr, -{ - match &build_quantizer { - BuildQuantizer::NoQuant(_) => diskann_async::new_index::(config, params, NoDeletes) - .map(|index| index as Arc>), - BuildQuantizer::Scalar1Bit(q) => { - let index = diskann_async::new_quant_only_index(config, params, q.clone(), NoDeletes)?; - Ok(Arc::new(QuantInMemBuilder::::new(index))) - } - BuildQuantizer::PQ(table) => { - let index = - diskann_async::new_quant_only_index(config, params, table.clone(), NoDeletes)?; - Ok(Arc::new(QuantInMemBuilder::::new(index))) - } - } -} diff --git a/diskann-disk/src/build/builder/mod.rs b/diskann-disk/src/build/builder/mod.rs index 9c22a0766..b7b5e15ef 100644 --- a/diskann-disk/src/build/builder/mod.rs +++ b/diskann-disk/src/build/builder/mod.rs @@ -5,11 +5,13 @@ //! Disk index builders and related functionality. pub mod build; -pub mod core; pub mod quantizer; -pub mod inmem_builder; pub mod tokio; +mod vamana; + +#[cfg(test)] +pub(crate) use vamana::tests::disk_index_builder_tests; #[cfg(test)] mod tests; diff --git a/diskann-disk/src/build/builder/tests.rs b/diskann-disk/src/build/builder/tests.rs index 8ce51a3c5..71925194e 100644 --- a/diskann-disk/src/build/builder/tests.rs +++ b/diskann-disk/src/build/builder/tests.rs @@ -10,7 +10,7 @@ mod disk_index_build_tests { use rstest::rstest; use crate::{ - build::builder::core::disk_index_builder_tests::{ + build::builder::disk_index_builder_tests::{ new_vfs, verify_search_result_with_ground_truth, IndexBuildFixture, TestParams, }, QuantizationType, diff --git a/diskann-disk/src/build/builder/vamana/index.rs b/diskann-disk/src/build/builder/vamana/index.rs new file mode 100644 index 000000000..20302863f --- /dev/null +++ b/diskann-disk/src/build/builder/vamana/index.rs @@ -0,0 +1,204 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::sync::Arc; + +use diskann::{ + graph::{Config, DiskANNIndex}, + provider::DefaultContext, + utils::VectorRepr, + ANNResult, +}; +use diskann_providers::{ + index::diskann_async, + model::graph::provider::async_::{ + common::{FullPrecision, NoDeletes, NoStore, Quantized}, + inmem::{ + DefaultProvider, DefaultProviderParameters, DefaultQuant, FullPrecisionProvider, + SQStore, SetStartPoints, + }, + }, + storage::{DiskGraphOnly, DynWriteProvider, SaveWith, WriteProviderWrapper}, +}; + +use crate::build::builder::quantizer::BuildQuantizer; + +type FullPrecisionIndex = DiskANNIndex>; +type ScalarQuantizedIndex = DiskANNIndex>>; +type ProductQuantizedIndex = DiskANNIndex>; + +/// Index implementation used while constructing a Vamana graph. +pub(super) enum VamanaBuildIndex +where + T: VectorRepr, +{ + FullPrecision(Arc>), + ScalarQuantized(Arc), + ProductQuantized(Arc), +} + +/// Manual implementation: `#[derive(Clone)]` would incorrectly require `T: Clone`, +/// even though `T` only appears behind `Arc`. +impl Clone for VamanaBuildIndex +where + T: VectorRepr, +{ + fn clone(&self) -> Self { + match self { + Self::FullPrecision(index) => Self::FullPrecision(Arc::clone(index)), + Self::ScalarQuantized(index) => Self::ScalarQuantized(Arc::clone(index)), + Self::ProductQuantized(index) => Self::ProductQuantized(Arc::clone(index)), + } + } +} + +impl VamanaBuildIndex +where + T: VectorRepr, +{ + pub(super) fn new( + config: Config, + params: DefaultProviderParameters, + build_quantizer: &BuildQuantizer, + ) -> ANNResult { + match build_quantizer { + BuildQuantizer::NoQuant(_) => { + diskann_async::new_index::(config, params, NoDeletes).map(Self::FullPrecision) + } + BuildQuantizer::Scalar1Bit(quantizer) => { + let index = diskann_async::new_quant_only_index( + config, + params, + quantizer.clone(), + NoDeletes, + )?; + Ok(Self::ScalarQuantized(Arc::new(index))) + } + BuildQuantizer::PQ(quantizer) => { + let index = diskann_async::new_quant_only_index( + config, + params, + quantizer.clone(), + NoDeletes, + )?; + Ok(Self::ProductQuantized(Arc::new(index))) + } + } + } + + pub(super) fn capacity(&self) -> usize { + match self { + Self::FullPrecision(index) => index.provider().capacity(), + Self::ScalarQuantized(index) => index.provider().capacity(), + Self::ProductQuantized(index) => index.provider().capacity(), + } + } + + pub(super) fn total_points(&self) -> usize { + match self { + Self::FullPrecision(index) => index.provider().total_points(), + Self::ScalarQuantized(index) => index.provider().total_points(), + Self::ProductQuantized(index) => index.provider().total_points(), + } + } + + pub(super) fn set_start_point(&self, start_point: &[T]) -> ANNResult<()> { + match self { + Self::FullPrecision(index) => index + .provider() + .set_start_points(std::iter::once(start_point)), + Self::ScalarQuantized(index) => index + .provider() + .set_start_points(std::iter::once(start_point)), + Self::ProductQuantized(index) => index + .provider() + .set_start_points(std::iter::once(start_point)), + } + } + + pub(super) async fn insert_vector(&self, id: u32, vector: &[T]) -> ANNResult<()> { + match self { + Self::FullPrecision(index) => { + index + .insert(&FullPrecision, &DefaultContext, &id, vector) + .await + } + Self::ScalarQuantized(index) => { + index.insert(&Quantized, &DefaultContext, &id, vector).await + } + Self::ProductQuantized(index) => { + index.insert(&Quantized, &DefaultContext, &id, vector).await + } + } + } + + pub(super) async fn final_prune(&self, range: core::ops::Range) -> ANNResult<()> { + match self { + Self::FullPrecision(index) => { + index + .prune_range(&FullPrecision, &DefaultContext, range) + .await + } + Self::ScalarQuantized(index) => { + index.prune_range(&Quantized, &DefaultContext, range).await + } + Self::ProductQuantized(index) => { + index.prune_range(&Quantized, &DefaultContext, range).await + } + } + } + + pub(super) async fn save_graph( + &self, + storage_provider: &dyn DynWriteProvider, + start_point_and_path: &(u32, DiskGraphOnly), + ) -> ANNResult<()> { + let wrapper = WriteProviderWrapper::new(storage_provider); + match self { + Self::FullPrecision(index) => index.save_with(&wrapper, start_point_and_path).await, + Self::ScalarQuantized(index) => index.save_with(&wrapper, start_point_and_path).await, + Self::ProductQuantized(index) => index.save_with(&wrapper, start_point_and_path).await, + } + } + + #[cfg(debug_assertions)] + pub(super) fn counts_for_get_vector(&self) -> (usize, usize) { + match self { + Self::FullPrecision(index) => index.provider().counts_for_get_vector(), + Self::ScalarQuantized(index) => index.provider().counts_for_get_vector(), + Self::ProductQuantized(index) => index.provider().counts_for_get_vector(), + } + } + + #[cfg(debug_assertions)] + pub(super) async fn count_reachable_nodes(&self) -> ANNResult { + match self { + Self::FullPrecision(index) => { + let provider = index.provider(); + let start_points = provider.starting_points()?; + let mut neighbor_accessor = provider.neighbors(); + index + .count_reachable_nodes(&start_points, &mut neighbor_accessor) + .await + } + Self::ScalarQuantized(index) => { + let provider = index.provider(); + let start_points = provider.starting_points()?; + let mut neighbor_accessor = provider.neighbors(); + index + .count_reachable_nodes(&start_points, &mut neighbor_accessor) + .await + } + Self::ProductQuantized(index) => { + let provider = index.provider(); + let start_points = provider.starting_points()?; + let mut neighbor_accessor = provider.neighbors(); + index + .count_reachable_nodes(&start_points, &mut neighbor_accessor) + .await + } + } + } +} diff --git a/diskann-disk/src/build/builder/vamana/merged.rs b/diskann-disk/src/build/builder/vamana/merged.rs new file mode 100644 index 000000000..f68740991 --- /dev/null +++ b/diskann-disk/src/build/builder/vamana/merged.rs @@ -0,0 +1,475 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::{ + marker::PhantomData, + mem::{self, size_of}, +}; + +use diskann::{utils::VectorRepr, ANNResult}; +use diskann_providers::{ + model::{IndexConfiguration, MAX_PQ_TRAINING_SET_SIZE}, + storage::{StorageReadProvider, StorageWriteProvider}, + utils::{ + load_metadata_from_file, RayonThreadPoolRef, SampleVectorReader, SamplingDensity, + READ_WRITE_BLOCK_SIZE, + }, +}; +use diskann_utils::io::read_bin; +use rand::seq::SliceRandom; +use tracing::info; + +use crate::{ + build::builder::quantizer::BuildQuantizer, + data_model::GraphDataType, + storage::{CachedReader, CachedWriter, DiskIndexWriter}, + utils::{ + instrumentation::{BuildMergedVamanaIndexCheckpoint, PerfLogger}, + partition_with_ram_budget, + }, + DiskIndexBuildParameters, +}; + +use super::{one_shot::OneShotVamanaBuilder, strategy::estimate_build_index_ram_usage}; + +/// Number of nearest shards each vector is assigned to during partitioning. +const PARTITION_ASSIGNMENTS_PER_VECTOR: usize = 2; +/// Builds a merged Vamana index from overlapping dataset shards. +pub(in crate::build::builder) struct MergedVamanaBuilder<'a, Data, StorageProvider> +where + Data: GraphDataType, + StorageProvider: StorageReadProvider + StorageWriteProvider, +{ + index_configuration: &'a IndexConfiguration, + disk_build_param: &'a DiskIndexBuildParameters, + index_writer: &'a DiskIndexWriter, + build_quantizer: &'a BuildQuantizer, + storage_provider: &'a StorageProvider, + rng: diskann_providers::utils::StandardRng, + _phantom: PhantomData, +} + +impl<'a, Data, StorageProvider> MergedVamanaBuilder<'a, Data, StorageProvider> +where + Data: GraphDataType, + Data::VectorDataType: VectorRepr, + StorageProvider: StorageReadProvider + StorageWriteProvider + 'static, + ::Reader: Send, +{ + pub(in crate::build::builder) fn new( + index_configuration: &'a IndexConfiguration, + disk_build_param: &'a DiskIndexBuildParameters, + index_writer: &'a DiskIndexWriter, + build_quantizer: &'a BuildQuantizer, + storage_provider: &'a StorageProvider, + ) -> Self { + Self { + index_configuration, + disk_build_param, + index_writer, + build_quantizer, + storage_provider, + rng: diskann_providers::utils::create_rnd_from_optional_seed( + index_configuration.random_seed, + ), + _phantom: PhantomData, + } + } + + pub(in crate::build::builder) async fn build( + mut self, + pool: RayonThreadPoolRef<'_>, + ) -> ANNResult<()> { + let mut logger = PerfLogger::new_disk_index_build_logger(); + let dataset_file = self.index_writer.get_dataset_file(); + let merged_index_prefix = self.index_writer.get_merged_index_prefix(); + let output_vamana = self.index_writer.get_mem_index_file(); + let max_degree = self.index_configuration.config.pruned_degree_u32().get(); + + let num_parts = + self.partition_data(&dataset_file, &merged_index_prefix, max_degree, pool)?; + logger.log_checkpoint(BuildMergedVamanaIndexCheckpoint::PartitionData); + + for shard_id in 0..num_parts { + self.build_shard_index(&dataset_file, &merged_index_prefix, shard_id) + .await?; + } + logger.log_checkpoint(BuildMergedVamanaIndexCheckpoint::BuildIndicesOnShards); + + self.merge_and_cleanup(&merged_index_prefix, num_parts, max_degree, output_vamana)?; + logger.log_checkpoint(BuildMergedVamanaIndexCheckpoint::MergeIndices); + + Ok(()) + } + + fn create_shard_index_config(&self, shard_base_file: &str) -> ANNResult { + let base_config = self.index_configuration; + let storage_provider = self.storage_provider; + + let search_list_size = base_config.config.l_build().get(); + let pruned_degree = base_config.config.pruned_degree().get(); + + let low_degree_params = diskann::graph::config::Builder::new( + 2 * pruned_degree / 3, + diskann::graph::config::MaxDegree::default_slack(), + search_list_size, + base_config.dist_metric.into(), + ) + .build()?; + + let metadata = load_metadata_from_file(storage_provider, shard_base_file)?; + + let mut index_config = (*base_config).clone(); + index_config.max_points = metadata.npoints(); + index_config.config = low_degree_params; + + Ok(index_config) + } + + fn retrieve_shard_data_from_ids( + &self, + dataset_file: &str, + shard_ids_file: &str, + shard_base_file: &str, + ) -> ANNResult<()> + where + T: Default + bytemuck::Pod, + { + let storage_provider = self.storage_provider; + let shard_ids = read_bin::(&mut storage_provider.open_reader(shard_ids_file)?)?; + let shard_size = shard_ids.nrows(); + info!("Loaded {} shard ids from {}", shard_size, shard_ids_file); + let max_id = shard_ids.as_slice().iter().max().copied().unwrap_or(0); + let sampling_rate = shard_ids.as_slice().len() as f64 / (max_id + 1) as f64; + + let mut dataset_reader: SampleVectorReader = SampleVectorReader::new( + dataset_file, + SamplingDensity::from_sample_rate(sampling_rate), + storage_provider, + )?; + + let (_npts, dim) = dataset_reader.get_dataset_headers(); + + let mut shard_base_cached_writer = CachedWriter::::new( + shard_base_file, + READ_WRITE_BLOCK_SIZE, + storage_provider.create_for_write(shard_base_file)?, + )?; + + let dummy_size: u32 = 0; + shard_base_cached_writer.write(&dummy_size.to_le_bytes())?; + shard_base_cached_writer.write(&dim.to_le_bytes())?; + + let mut num_written: u32 = 0; + dataset_reader.read_vectors(shard_ids.as_slice().iter().copied(), |vector_t| { + // Casting Pod type to bytes always succeeds (u8 has alignment of 1) + let vector_bytes: &[u8] = bytemuck::must_cast_slice(vector_t); + shard_base_cached_writer.write(vector_bytes)?; + num_written += 1; + Ok(()) + })?; + + info!( + "Written file: {} with {} points", + shard_base_file, num_written + ); + + shard_base_cached_writer.flush()?; + shard_base_cached_writer.reset()?; + shard_base_cached_writer.write(&num_written.to_le_bytes())?; + + Ok(()) + } + + async fn build_shard_index( + &self, + dataset_file: &str, + merged_index_prefix: &str, + shard_id: usize, + ) -> ANNResult<()> { + let shard_base_file = + DiskIndexWriter::get_merged_index_subshard_data_file(merged_index_prefix, shard_id); + let shard_ids_file = + DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard_id); + self.retrieve_shard_data_from_ids::( + dataset_file, + &shard_ids_file, + &shard_base_file, + )?; + info!("Generated data for shard {}", shard_id); + + let index_config = self.create_shard_index_config(&shard_base_file)?; + let shard_index_file = DiskIndexWriter::get_merged_index_subshard_mem_index_file( + merged_index_prefix, + shard_id, + ); + + OneShotVamanaBuilder::::new( + &index_config, + self.build_quantizer, + shard_base_file, + shard_index_file, + self.storage_provider, + ) + .build() + .await + } + + fn merge_shards( + &mut self, + merged_index_prefix: &str, + num_parts: usize, + max_degree: u32, + output_vamana: String, + ) -> ANNResult<()> { + // Read ID maps + let mut vamana_names = vec![String::new(); num_parts]; + let mut id_maps: Vec> = vec![Vec::new(); num_parts]; + for shard in 0..num_parts { + vamana_names[shard] = DiskIndexWriter::get_merged_index_subshard_mem_index_file( + merged_index_prefix, + shard, + ); + + let id_maps_file = + DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard); + id_maps[shard] = self.read_idmap(id_maps_file)?; + } + + // find max node id + let num_nodes: u32 = *id_maps.iter().flatten().max().unwrap_or(&0) + 1; + let num_elements: u32 = id_maps.iter().map(|idmap| idmap.len() as u32).sum(); + info!("# nodes: {}, max degree: {}", num_nodes, max_degree); + + // compute inverse map: node -> shards + let mut node_shard: Vec<(u32, u32)> = Vec::with_capacity(num_elements as usize); + for (shard, id_map) in id_maps.iter().enumerate() { + info!("Creating inverse map -- shard #{}", shard); + node_shard.extend(id_map.iter().map(|node_id| (*node_id, shard as u32))); + } + node_shard.sort_unstable_by(|left, right| { + left.0.cmp(&right.0).then_with(|| left.1.cmp(&right.1)) + }); + + info!("Finished computing node -> shards map"); + + // create cached vamana readers + let mut vamana_readers = Vec::new(); + for name in &vamana_names { + let reader = CachedReader::::new( + name, + READ_WRITE_BLOCK_SIZE, + self.storage_provider, + )?; + vamana_readers.push(reader); + } + + // create cached vamana writers + let mut merged_vamana_cached_writer = CachedWriter::::new( + &output_vamana, + READ_WRITE_BLOCK_SIZE, + self.storage_provider.create_for_write(&output_vamana)?, + )?; + + // expected file size + max degree + medoid_id + frozen_point info + let vamana_metadata_size = + size_of::() + size_of::() + size_of::() + size_of::(); + + // we initialize the size of the merged index to the metadata size + // we will overwrite the index size at the end + let mut merged_index_size: u64 = vamana_metadata_size as u64; + merged_vamana_cached_writer.write(&merged_index_size.to_le_bytes())?; + + let mut read_buf_8_bytes = [0u8; 8]; + + // get max input width + let mut max_input_width = 0; + // read width from each vamana to advance buffer by sizeof(uint32_t) bytes + for reader in &mut vamana_readers { + reader.read(&mut read_buf_8_bytes)?; + let _expected_file_size: u64 = u64::from_le_bytes(read_buf_8_bytes); + let input_width = reader.read_u32()?; + max_input_width = input_width.max(max_input_width); + } + + // write max_degree to merged_vamana_index + let output_width: u32 = max_degree; + info!( + "Max input width: {}, output width: {}", + max_input_width, output_width + ); + + merged_vamana_cached_writer.write(&output_width.to_le_bytes())?; + + // write medoid to merged_vamana_index + for shard in 0..num_parts { + // read medoid + let mut medoid: u32 = vamana_readers[shard].read_u32()?; + vamana_readers[shard].read(&mut read_buf_8_bytes)?; + let vamana_index_frozen: u64 = u64::from_le_bytes(read_buf_8_bytes); + debug_assert_eq!(vamana_index_frozen, 0); + + // rename medoid + medoid = id_maps[shard][medoid as usize]; + + // write renamed medoid + if shard == (num_parts - 1) { + // uncomment if running hierarchical + merged_vamana_cached_writer.write(&medoid.to_le_bytes())?; + } + } + + let vamana_index_frozen: u64 = 0; // as of now the functionality to merge many overlapping vamana + // indices is supported only for bulk indices without frozen point. + // Hence the final index will also not have any frozen points. + merged_vamana_cached_writer.write(&vamana_index_frozen.to_le_bytes())?; + + info!("Starting merge"); + + let mut nbr_set = vec![false; num_nodes as usize]; + let mut final_nbrs: Vec = Vec::new(); + let mut cur_id = 0; + for pair in &node_shard { + let (node_id, shard_id) = *pair; + if cur_id < node_id { + final_nbrs.shuffle(&mut self.rng); + + let nnbrs: u32 = std::cmp::min(final_nbrs.len() as u32, max_degree); + merged_vamana_cached_writer.write(&nnbrs.to_le_bytes())?; + + let bytes = final_nbrs + .iter() + .take(nnbrs as usize) + .flat_map(|x| x.to_le_bytes()) + .collect::>(); + merged_vamana_cached_writer.write(&bytes)?; + + merged_index_size += (size_of::() + nnbrs as usize * size_of::()) as u64; + if cur_id % 499999 == 1 { + print!("."); + } + cur_id = node_id; + + final_nbrs.iter().for_each(|p| nbr_set[*p as usize] = false); + final_nbrs.clear(); + } + + // read num of neighbors from vamana index + let num_nbrs = vamana_readers[shard_id as usize].read_u32()?; + + if num_nbrs == 0 { + info!( + "WARNING: shard #{}, node_id {} has 0 nbrs", + shard_id, node_id + ); + } else { + let mut nbrs_bytes = vec![0u8; num_nbrs as usize * mem::size_of::()]; + vamana_readers[shard_id as usize].read(&mut nbrs_bytes)?; + let nbrs: &[u32] = bytemuck::cast_slice(&nbrs_bytes); + + // rename nodes + for j in 0..num_nbrs { + let nbr = nbrs[j as usize]; + let renamed_node = id_maps[shard_id as usize][nbr as usize]; + if !nbr_set[renamed_node as usize] { + nbr_set[renamed_node as usize] = true; + final_nbrs.push(renamed_node); + } + } + } + } + + // write the last node, to be refactored... + final_nbrs.shuffle(&mut self.rng); + + let nnbrs: u32 = std::cmp::min(final_nbrs.len() as u32, max_degree); + merged_vamana_cached_writer.write(&nnbrs.to_le_bytes())?; + + let bytes = final_nbrs + .iter() + .take(nnbrs as usize) + .flat_map(|x| x.to_le_bytes()) + .collect::>(); + merged_vamana_cached_writer.write(&bytes)?; + + merged_index_size += (size_of::() + nnbrs as usize * size_of::()) as u64; + + nbr_set.clear(); + final_nbrs.clear(); + + info!("Expected size: {}", merged_index_size); + merged_vamana_cached_writer.reset()?; + merged_vamana_cached_writer.write(&merged_index_size.to_le_bytes())?; + + info!("Finished merge"); + Ok(()) + } + + fn read_idmap(&self, idmaps_path: String) -> Result, diskann_utils::io::ReadBinError> { + let data = read_bin::(&mut self.storage_provider.open_reader(&idmaps_path)?)?; + Ok(data.into_inner().into_vec()) + } + + fn partition_data( + &mut self, + dataset_file: &str, + merged_index_prefix: &str, + max_degree: u32, + pool: RayonThreadPoolRef<'_>, + ) -> ANNResult { + let sampling_rate = MAX_PQ_TRAINING_SET_SIZE / self.index_configuration.max_points as f64; + let ram_budget_in_bytes = self.disk_build_param.build_memory_limit().in_bytes() as f64; + + partition_with_ram_budget::( + dataset_file, + self.index_configuration.dim, + sampling_rate, + ram_budget_in_bytes, + PARTITION_ASSIGNMENTS_PER_VECTOR, + merged_index_prefix, + self.storage_provider, + &mut self.rng, + pool, + |num_points, dim| { + let datasize = std::mem::size_of::() as u64; + let graph_degree = 2 * max_degree / 3; + estimate_build_index_ram_usage( + num_points, + dim, + datasize, + graph_degree as u64, + self.disk_build_param.build_quantization(), + ) + }, + ) + } + + fn merge_and_cleanup( + &mut self, + merged_index_prefix: &str, + num_parts: usize, + max_degree: u32, + output_vamana: String, + ) -> ANNResult<()> { + // merge all in-memory indices into one + self.merge_shards(merged_index_prefix, num_parts, max_degree, output_vamana)?; + + // delete tempFiles + for p in 0..num_parts { + let shard_base_file = + DiskIndexWriter::get_merged_index_subshard_data_file(merged_index_prefix, p); + let shard_ids_file = + DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, p); + let shard_index_file = + DiskIndexWriter::get_merged_index_subshard_mem_index_file(merged_index_prefix, p); + + self.storage_provider.delete(&shard_base_file)?; + self.storage_provider.delete(&shard_ids_file)?; + self.storage_provider.delete(&shard_index_file)?; + } + + Ok(()) + } +} diff --git a/diskann-disk/src/build/builder/vamana/mod.rs b/diskann-disk/src/build/builder/vamana/mod.rs new file mode 100644 index 000000000..b8edecf59 --- /dev/null +++ b/diskann-disk/src/build/builder/vamana/mod.rs @@ -0,0 +1,16 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +mod index; +mod merged; +mod one_shot; +mod strategy; + +#[cfg(test)] +pub(super) mod tests; + +pub(super) use merged::MergedVamanaBuilder; +pub(super) use one_shot::OneShotVamanaBuilder; +pub(super) use strategy::{determine_build_strategy, IndexBuildStrategy}; diff --git a/diskann-disk/src/build/builder/vamana/one_shot.rs b/diskann-disk/src/build/builder/vamana/one_shot.rs new file mode 100644 index 000000000..2f0e2f447 --- /dev/null +++ b/diskann-disk/src/build/builder/vamana/one_shot.rs @@ -0,0 +1,253 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::{ + marker::PhantomData, + num::NonZeroUsize, + sync::{Arc, Mutex}, +}; + +use diskann::{ + utils::{async_tools, VectorRepr, ONE}, + ANNResult, +}; +use diskann_providers::{ + model::{graph::provider::async_::inmem::DefaultProviderParameters, IndexConfiguration}, + storage::{DiskGraphOnly, StorageReadProvider, StorageWriteProvider}, + utils::{find_medoid_with_sampling, VectorDataIterator, MAX_MEDOID_SAMPLE_SIZE}, +}; +use tokio::task::JoinSet; +use tracing::{debug, info}; + +use crate::{ + build::builder::quantizer::BuildQuantizer, + error::{diskann_error, ErrorKind}, +}; + +use super::index::VamanaBuildIndex; +/// Builds a complete Vamana graph from a dataset in one pass. +pub(in crate::build::builder) struct OneShotVamanaBuilder<'a, T, StorageProvider> +where + T: VectorRepr, + StorageProvider: StorageReadProvider + StorageWriteProvider, +{ + config: &'a IndexConfiguration, + quantizer: &'a BuildQuantizer, + data_path: String, + save_path: String, + storage_provider: &'a StorageProvider, + _phantom: PhantomData, +} + +impl<'a, T, StorageProvider> OneShotVamanaBuilder<'a, T, StorageProvider> +where + T: VectorRepr, + StorageProvider: StorageReadProvider + StorageWriteProvider + 'static, + ::Reader: Send, +{ + pub(in crate::build::builder) fn new( + config: &'a IndexConfiguration, + quantizer: &'a BuildQuantizer, + data_path: String, + save_path: String, + storage_provider: &'a StorageProvider, + ) -> Self { + Self { + config, + quantizer, + data_path, + save_path, + storage_provider, + _phantom: PhantomData, + } + } + + pub(in crate::build::builder) async fn build(self) -> ANNResult<()> { + let Self { + config, + quantizer, + data_path, + save_path, + storage_provider, + .. + } = self; + + // use either user-specified number of threads or default to available parallelism + let num_tasks = NonZeroUsize::new(config.num_threads) + .or_else(|| std::thread::available_parallelism().ok()) + .ok_or_else(|| { + diskann_error!( + ErrorKind::IndexError, + "Failed to determine number of threads" + ) + })?; + + // Associated data will only be used in the write_disk_layout function which only requires the none-partitioned associated data stream. + let dataset_iter = Arc::new(Mutex::new({ + let iter = VectorDataIterator::<_, T>::new(&data_path, None, storage_provider)?; + iter.enumerate() + })); + + let index_config = config.config.clone(); + let provider_parameters = DefaultProviderParameters { + max_points: config.max_points, + frozen_points: ONE, + metric: config.dist_metric, + dim: config.dim, + max_degree: index_config.max_degree_u32().get(), + prefetch_lookahead: config.prefetch_lookahead.map(|x| x.get()), + prefetch_cache_line_level: config.prefetch_cache_line_level, + }; + let index = VamanaBuildIndex::::new(index_config, provider_parameters, quantizer)?; + let medoid_id = Self::set_start_point_to_medoid( + &index, + &data_path, + config.random_seed, + storage_provider, + )?; + let start_point = Self::u32_try_from(medoid_id)?; + + Self::run_build(&index, dataset_iter, num_tasks).await?; + + #[cfg(debug_assertions)] + Self::log_build_stats(&index).await?; + + Self::run_final_prune(&index, num_tasks).await?; + index + .save_graph( + storage_provider, + &(start_point, DiskGraphOnly::new(&save_path)), + ) + .await?; + + Ok(()) + } + + /// Log statistics about the build process + #[cfg(debug_assertions)] + async fn log_build_stats(index: &VamanaBuildIndex) -> ANNResult<()> { + debug!( + "Number of points reachable in the graph: {}", + index.count_reachable_nodes().await? + ); + + let (full_vector, quant_vector) = index.counts_for_get_vector(); + let capacity = index.capacity(); + debug!( + "Number of get vector calls per insert: {}", + full_vector as f32 / capacity as f32 + ); + debug!( + "Number of get quantized vector calls per insert: {}", + quant_vector as f32 / capacity as f32 + ); + + Ok(()) + } + + /// Convert a `usize` index into the `u32` internal id type, erroring if it does not fit. + /// + /// The async index uses `u32` internal ids, so positions in the dataset must not exceed + /// `u32::MAX`. + fn u32_try_from(value: usize) -> ANNResult { + u32::try_from(value) + .map_err(|_| diskann_error!(ErrorKind::IndexError, "id {value} exceeds u32::MAX")) + } + + fn set_start_point_to_medoid( + index: &VamanaBuildIndex, + path: &str, + random_seed: Option, + reader: &StorageProvider, + ) -> ANNResult { + let mut rng = diskann_providers::utils::create_rnd_from_optional_seed(random_seed); + let (medoid, medoid_id) = + find_medoid_with_sampling::(path, reader, MAX_MEDOID_SAMPLE_SIZE, &mut rng)?; + + index.set_start_point(medoid.as_slice())?; + + debug!("Set start point to medoid ID: {}", medoid_id); + + Ok(medoid_id) + } + + async fn run_build( + index: &VamanaBuildIndex, + iterator: Arc>, + num_tasks: NonZeroUsize, + ) -> ANNResult<()> + where + I: Iterator, ()))> + Send + 'static, + { + let total_points = index.capacity(); + let partitions = async_tools::PartitionIter::new(total_points, num_tasks); + + let mut tasks = JoinSet::new(); + + for partition in partitions { + let index_clone = index.clone(); + let iterator_clone = iterator.clone(); + tasks.spawn(async move { + for _ in partition { + let vector_data = { + let mut guard = iterator_clone.lock().map_err(|_| { + diskann_error!( + ErrorKind::IndexError, + "Poisoned mutex during construction" + ) + })?; + guard.next() + }; + + match vector_data { + Some((i, (vector, _))) => { + let id = Self::u32_try_from(i)?; + index_clone.insert_vector(id, vector.as_ref()).await?; + } + None => break, + } + } + ANNResult::Ok(()) + }); + } + + // Wait for all tasks to complete. + while let Some(res) = tasks.join_next().await { + res.map_err(|_| { + diskann_error!(ErrorKind::IndexError, "A spawned insert task failed") + })??; + } + + info!("Linked all points. Num points: #{}", total_points); + Ok(()) + } + + async fn run_final_prune( + index: &VamanaBuildIndex, + num_tasks: NonZeroUsize, + ) -> ANNResult<()> { + let partitions = async_tools::PartitionIter::new(index.total_points(), num_tasks); + + let mut tasks = JoinSet::new(); + + for partition in partitions { + let index_clone = index.clone(); + tasks.spawn(async move { + let range = + Self::u32_try_from(partition.start)?..Self::u32_try_from(partition.end)?; + index_clone.final_prune(range).await + }); + } + + // Wait for all final prune tasks to complete + while let Some(res) = tasks.join_next().await { + res.map_err(|_| { + diskann_error!(ErrorKind::IndexError, "A spawned final prune task failed") + })??; + } + + Ok(()) + } +} diff --git a/diskann-disk/src/build/builder/vamana/strategy.rs b/diskann-disk/src/build/builder/vamana/strategy.rs new file mode 100644 index 000000000..8014c13a3 --- /dev/null +++ b/diskann-disk/src/build/builder/vamana/strategy.rs @@ -0,0 +1,121 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::mem; + +use diskann_providers::model::{IndexConfiguration, GRAPH_SLACK_FACTOR}; +use tracing::info; + +use crate::{data_model::GraphDataType, disk_index_build_parameter::BYTES_IN_GB, QuantizationType}; +/// Overhead factor for RAM estimation during index build (10% buffer). +const OVERHEAD_FACTOR: f64 = 1.1f64; + +/// Estimate RAM usage in bytes for building an index. +#[inline] +pub(super) fn estimate_build_index_ram_usage( + num_points: u64, + dim: u64, + datasize: u64, + graph_degree: u64, + build_quantization_type: &QuantizationType, +) -> f64 { + let graph_size = + (num_points * graph_degree * mem::size_of::() as u64) as f64 * GRAPH_SLACK_FACTOR; + + let single_vec_size = match *build_quantization_type { + QuantizationType::FP => dim.next_multiple_of(8u64) * datasize, + // We can skip PQ pivots data as it is very small(~3MB) for even large datasets like OAI-3072. + QuantizationType::PQ { num_chunks } => num_chunks as u64, + // `+ std::mem::size_of::()` for f32 compensation metadata for the scalar quantizer. + QuantizationType::SQ { nbits, .. } => { + (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::() as u64 + } + }; + + OVERHEAD_FACTOR * (graph_size + (single_vec_size * num_points) as f64) +} + +pub(in crate::build::builder) enum IndexBuildStrategy { + OneShot, + Merged, +} + +pub(in crate::build::builder) fn determine_build_strategy( + index_configuration: &IndexConfiguration, + index_build_ram_limit_in_bytes: f64, + build_quantization_type: &QuantizationType, +) -> IndexBuildStrategy { + let estimated_index_ram_in_bytes = estimate_build_index_ram_usage( + index_configuration.max_points as u64, + index_configuration.dim as u64, + mem::size_of::() as u64, + index_configuration.config.max_degree().get() as u64, + build_quantization_type, + ); + + info!( + "Estimated index RAM usage: {} GB, index_build_ram_limit={} GB", + estimated_index_ram_in_bytes / BYTES_IN_GB, + index_build_ram_limit_in_bytes / BYTES_IN_GB + ); + + if estimated_index_ram_in_bytes >= index_build_ram_limit_in_bytes { + info!( + "Insufficient memory budget for index build in one shot, index_build_ram_limit={} GB estimated_index_ram={} GB", + index_build_ram_limit_in_bytes / BYTES_IN_GB, + estimated_index_ram_in_bytes / BYTES_IN_GB, + ); + IndexBuildStrategy::Merged + } else { + info!( + "Full index fits in RAM budget, should consume at most {} GBs, so building in one shot", + estimated_index_ram_in_bytes / BYTES_IN_GB + ); + IndexBuildStrategy::OneShot + } +} + +#[cfg(test)] +mod ram_estimation_tests { + use rstest::rstest; + + use super::*; + use crate::QuantizationType; + + #[rstest] + #[case(QuantizationType::FP)] + #[case(QuantizationType::PQ { num_chunks: 15 })] + #[case(QuantizationType::SQ { nbits: 1, standard_deviation: None })] + fn test_estimate_build_index_ram_usage(#[case] build_quantization_type: QuantizationType) { + let num_points = 1000; + let dim = 128; + let size_of_t = std::mem::size_of::() as u64; + let graph_degree = 50; + + let single_vec_size = match build_quantization_type { + QuantizationType::FP => dim * size_of_t, + QuantizationType::PQ { num_chunks } => num_chunks as u64, + QuantizationType::SQ { nbits, .. } => { + (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::() as u64 + } + }; + let mut expected_ram_usage = (num_points as f64) + * (graph_degree as f64) + * (std::mem::size_of::() as f64) + * GRAPH_SLACK_FACTOR + + (num_points * single_vec_size) as f64; + expected_ram_usage *= OVERHEAD_FACTOR; + + let actual_ram_usage = estimate_build_index_ram_usage( + num_points, + dim, + size_of_t, + graph_degree, + &build_quantization_type, + ); + + assert_eq!(actual_ram_usage, expected_ram_usage); + } +} diff --git a/diskann-disk/src/build/builder/core.rs b/diskann-disk/src/build/builder/vamana/tests.rs similarity index 52% rename from diskann-disk/src/build/builder/core.rs rename to diskann-disk/src/build/builder/vamana/tests.rs index 6f7e157d2..336b9efed 100644 --- a/diskann-disk/src/build/builder/core.rs +++ b/diskann-disk/src/build/builder/vamana/tests.rs @@ -2,538 +2,6 @@ * Copyright (c) Microsoft Corporation. * Licensed under the MIT license. */ -use std::{ - marker::PhantomData, - mem::{self, size_of}, -}; - -use crate::data_model::GraphDataType; -use diskann::{utils::VectorRepr, ANNResult}; -use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider}; -use diskann_providers::{ - model::{IndexConfiguration, GRAPH_SLACK_FACTOR, MAX_PQ_TRAINING_SET_SIZE}, - utils::{ - load_metadata_from_file, RayonThreadPoolRef, SampleVectorReader, SamplingDensity, - READ_WRITE_BLOCK_SIZE, - }, -}; -use diskann_utils::io::read_bin; -use rand::seq::SliceRandom; -use tracing::info; - -use crate::{ - build::builder::{build::build_inmem_index, quantizer::BuildQuantizer}, - disk_index_build_parameter::BYTES_IN_GB, - storage::{CachedReader, CachedWriter, DiskIndexWriter}, - utils::instrumentation::{BuildMergedVamanaIndexCheckpoint, PerfLogger}, - utils::partition_with_ram_budget, - DiskIndexBuildParameters, QuantizationType, -}; - -/// Overhead factor for RAM estimation during index build (10% buffer). -const OVERHEAD_FACTOR: f64 = 1.1f64; - -/// Number of nearest shards each vector is assigned to during partitioning. -const PARTITION_ASSIGNMENTS_PER_VECTOR: usize = 2; - -/// Estimate RAM usage in bytes for building an index. -#[inline] -fn estimate_build_index_ram_usage( - num_points: u64, - dim: u64, - datasize: u64, - graph_degree: u64, - build_quantization_type: &QuantizationType, -) -> f64 { - let graph_size = - (num_points * graph_degree * mem::size_of::() as u64) as f64 * GRAPH_SLACK_FACTOR; - - let single_vec_size = match *build_quantization_type { - QuantizationType::FP => dim.next_multiple_of(8u64) * datasize, - // We can skip PQ pivots data as it is very small(~3MB) for even large datasets like OAI-3072. - QuantizationType::PQ { num_chunks } => num_chunks as u64, - // `+ std::mem::size_of::()` for f32 compensation metadata for the scalar quantizer. - QuantizationType::SQ { nbits, .. } => { - (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::() as u64 - } - }; - - OVERHEAD_FACTOR * (graph_size + (single_vec_size * num_points) as f64) -} - -/// Builds a merged Vamana index from overlapping dataset shards. -pub(super) struct MergedVamanaIndexBuilder<'a, Data, StorageProvider> -where - Data: GraphDataType, - StorageProvider: StorageReadProvider + StorageWriteProvider, -{ - index_configuration: &'a IndexConfiguration, - disk_build_param: &'a DiskIndexBuildParameters, - index_writer: &'a DiskIndexWriter, - build_quantizer: &'a BuildQuantizer, - storage_provider: &'a StorageProvider, - rng: diskann_providers::utils::StandardRng, - _phantom: PhantomData, -} - -impl<'a, Data, StorageProvider> MergedVamanaIndexBuilder<'a, Data, StorageProvider> -where - Data: GraphDataType, - Data::VectorDataType: VectorRepr, - StorageProvider: StorageReadProvider + StorageWriteProvider + 'static, - ::Reader: Send, -{ - pub(super) fn new( - index_configuration: &'a IndexConfiguration, - disk_build_param: &'a DiskIndexBuildParameters, - index_writer: &'a DiskIndexWriter, - build_quantizer: &'a BuildQuantizer, - storage_provider: &'a StorageProvider, - ) -> Self { - Self { - index_configuration, - disk_build_param, - index_writer, - build_quantizer, - storage_provider, - rng: diskann_providers::utils::create_rnd_from_optional_seed( - index_configuration.random_seed, - ), - _phantom: PhantomData, - } - } - - pub(super) async fn build(mut self, pool: RayonThreadPoolRef<'_>) -> ANNResult<()> { - let mut logger = PerfLogger::new_disk_index_build_logger(); - let dataset_file = self.index_writer.get_dataset_file(); - let merged_index_prefix = self.index_writer.get_merged_index_prefix(); - let output_vamana = self.index_writer.get_mem_index_file(); - let max_degree = self.index_configuration.config.pruned_degree_u32().get(); - - let num_parts = - self.partition_data(&dataset_file, &merged_index_prefix, max_degree, pool)?; - logger.log_checkpoint(BuildMergedVamanaIndexCheckpoint::PartitionData); - - for shard_id in 0..num_parts { - self.build_shard_index(&dataset_file, &merged_index_prefix, shard_id) - .await?; - } - logger.log_checkpoint(BuildMergedVamanaIndexCheckpoint::BuildIndicesOnShards); - - self.merge_and_cleanup(&merged_index_prefix, num_parts, max_degree, output_vamana)?; - logger.log_checkpoint(BuildMergedVamanaIndexCheckpoint::MergeIndices); - - Ok(()) - } - - fn create_shard_index_config(&self, shard_base_file: &str) -> ANNResult { - let base_config = self.index_configuration; - let storage_provider = self.storage_provider; - - let search_list_size = base_config.config.l_build().get(); - let pruned_degree = base_config.config.pruned_degree().get(); - - let low_degree_params = diskann::graph::config::Builder::new( - 2 * pruned_degree / 3, - diskann::graph::config::MaxDegree::default_slack(), - search_list_size, - base_config.dist_metric.into(), - ) - .build()?; - - let metadata = load_metadata_from_file(storage_provider, shard_base_file)?; - - let mut index_config = (*base_config).clone(); - index_config.max_points = metadata.npoints(); - index_config.config = low_degree_params; - - Ok(index_config) - } - - fn retrieve_shard_data_from_ids( - &self, - dataset_file: &str, - shard_ids_file: &str, - shard_base_file: &str, - ) -> ANNResult<()> - where - T: Default + bytemuck::Pod, - { - let storage_provider = self.storage_provider; - let shard_ids = read_bin::(&mut storage_provider.open_reader(shard_ids_file)?)?; - let shard_size = shard_ids.nrows(); - info!("Loaded {} shard ids from {}", shard_size, shard_ids_file); - let max_id = shard_ids.as_slice().iter().max().copied().unwrap_or(0); - let sampling_rate = shard_ids.as_slice().len() as f64 / (max_id + 1) as f64; - - let mut dataset_reader: SampleVectorReader = SampleVectorReader::new( - dataset_file, - SamplingDensity::from_sample_rate(sampling_rate), - storage_provider, - )?; - - let (_npts, dim) = dataset_reader.get_dataset_headers(); - - let mut shard_base_cached_writer = CachedWriter::::new( - shard_base_file, - READ_WRITE_BLOCK_SIZE, - storage_provider.create_for_write(shard_base_file)?, - )?; - - let dummy_size: u32 = 0; - shard_base_cached_writer.write(&dummy_size.to_le_bytes())?; - shard_base_cached_writer.write(&dim.to_le_bytes())?; - - let mut num_written: u32 = 0; - dataset_reader.read_vectors(shard_ids.as_slice().iter().copied(), |vector_t| { - // Casting Pod type to bytes always succeeds (u8 has alignment of 1) - let vector_bytes: &[u8] = bytemuck::must_cast_slice(vector_t); - shard_base_cached_writer.write(vector_bytes)?; - num_written += 1; - Ok(()) - })?; - - info!( - "Written file: {} with {} points", - shard_base_file, num_written - ); - - shard_base_cached_writer.flush()?; - shard_base_cached_writer.reset()?; - shard_base_cached_writer.write(&num_written.to_le_bytes())?; - - Ok(()) - } - - async fn build_shard_index( - &self, - dataset_file: &str, - merged_index_prefix: &str, - shard_id: usize, - ) -> ANNResult<()> { - let shard_base_file = - DiskIndexWriter::get_merged_index_subshard_data_file(merged_index_prefix, shard_id); - let shard_ids_file = - DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard_id); - self.retrieve_shard_data_from_ids::( - dataset_file, - &shard_ids_file, - &shard_base_file, - )?; - info!("Generated data for shard {}", shard_id); - - let index_config = self.create_shard_index_config(&shard_base_file)?; - let shard_index_file = DiskIndexWriter::get_merged_index_subshard_mem_index_file( - merged_index_prefix, - shard_id, - ); - - build_inmem_index::( - index_config, - self.build_quantizer, - &shard_base_file, - &shard_index_file, - self.storage_provider, - ) - .await - } - - fn merge_shards( - &mut self, - merged_index_prefix: &str, - num_parts: usize, - max_degree: u32, - output_vamana: String, - ) -> ANNResult<()> { - // Read ID maps - let mut vamana_names = vec![String::new(); num_parts]; - let mut id_maps: Vec> = vec![Vec::new(); num_parts]; - for shard in 0..num_parts { - vamana_names[shard] = DiskIndexWriter::get_merged_index_subshard_mem_index_file( - merged_index_prefix, - shard, - ); - - let id_maps_file = - DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard); - id_maps[shard] = self.read_idmap(id_maps_file)?; - } - - // find max node id - let num_nodes: u32 = *id_maps.iter().flatten().max().unwrap_or(&0) + 1; - let num_elements: u32 = id_maps.iter().map(|idmap| idmap.len() as u32).sum(); - info!("# nodes: {}, max degree: {}", num_nodes, max_degree); - - // compute inverse map: node -> shards - let mut node_shard: Vec<(u32, u32)> = Vec::with_capacity(num_elements as usize); - for (shard, id_map) in id_maps.iter().enumerate() { - info!("Creating inverse map -- shard #{}", shard); - node_shard.extend(id_map.iter().map(|node_id| (*node_id, shard as u32))); - } - node_shard.sort_unstable_by(|left, right| { - left.0.cmp(&right.0).then_with(|| left.1.cmp(&right.1)) - }); - - info!("Finished computing node -> shards map"); - - // create cached vamana readers - let mut vamana_readers = Vec::new(); - for name in &vamana_names { - let reader = CachedReader::::new( - name, - READ_WRITE_BLOCK_SIZE, - self.storage_provider, - )?; - vamana_readers.push(reader); - } - - // create cached vamana writers - let mut merged_vamana_cached_writer = CachedWriter::::new( - &output_vamana, - READ_WRITE_BLOCK_SIZE, - self.storage_provider.create_for_write(&output_vamana)?, - )?; - - // expected file size + max degree + medoid_id + frozen_point info - let vamana_metadata_size = - size_of::() + size_of::() + size_of::() + size_of::(); - - // we initialize the size of the merged index to the metadata size - // we will overwrite the index size at the end - let mut merged_index_size: u64 = vamana_metadata_size as u64; - merged_vamana_cached_writer.write(&merged_index_size.to_le_bytes())?; - - let mut read_buf_8_bytes = [0u8; 8]; - - // get max input width - let mut max_input_width = 0; - // read width from each vamana to advance buffer by sizeof(uint32_t) bytes - for reader in &mut vamana_readers { - reader.read(&mut read_buf_8_bytes)?; - let _expected_file_size: u64 = u64::from_le_bytes(read_buf_8_bytes); - let input_width = reader.read_u32()?; - max_input_width = input_width.max(max_input_width); - } - - // write max_degree to merged_vamana_index - let output_width: u32 = max_degree; - info!( - "Max input width: {}, output width: {}", - max_input_width, output_width - ); - - merged_vamana_cached_writer.write(&output_width.to_le_bytes())?; - - // write medoid to merged_vamana_index - for shard in 0..num_parts { - // read medoid - let mut medoid: u32 = vamana_readers[shard].read_u32()?; - vamana_readers[shard].read(&mut read_buf_8_bytes)?; - let vamana_index_frozen: u64 = u64::from_le_bytes(read_buf_8_bytes); - debug_assert_eq!(vamana_index_frozen, 0); - - // rename medoid - medoid = id_maps[shard][medoid as usize]; - - // write renamed medoid - if shard == (num_parts - 1) { - // uncomment if running hierarchical - merged_vamana_cached_writer.write(&medoid.to_le_bytes())?; - } - } - - let vamana_index_frozen: u64 = 0; // as of now the functionality to merge many overlapping vamana - // indices is supported only for bulk indices without frozen point. - // Hence the final index will also not have any frozen points. - merged_vamana_cached_writer.write(&vamana_index_frozen.to_le_bytes())?; - - info!("Starting merge"); - - let mut nbr_set = vec![false; num_nodes as usize]; - let mut final_nbrs: Vec = Vec::new(); - let mut cur_id = 0; - for pair in &node_shard { - let (node_id, shard_id) = *pair; - if cur_id < node_id { - final_nbrs.shuffle(&mut self.rng); - - let nnbrs: u32 = std::cmp::min(final_nbrs.len() as u32, max_degree); - merged_vamana_cached_writer.write(&nnbrs.to_le_bytes())?; - - let bytes = final_nbrs - .iter() - .take(nnbrs as usize) - .flat_map(|x| x.to_le_bytes()) - .collect::>(); - merged_vamana_cached_writer.write(&bytes)?; - - merged_index_size += (size_of::() + nnbrs as usize * size_of::()) as u64; - if cur_id % 499999 == 1 { - print!("."); - } - cur_id = node_id; - - final_nbrs.iter().for_each(|p| nbr_set[*p as usize] = false); - final_nbrs.clear(); - } - - // read num of neighbors from vamana index - let num_nbrs = vamana_readers[shard_id as usize].read_u32()?; - - if num_nbrs == 0 { - info!( - "WARNING: shard #{}, node_id {} has 0 nbrs", - shard_id, node_id - ); - } else { - let mut nbrs_bytes = vec![0u8; num_nbrs as usize * mem::size_of::()]; - vamana_readers[shard_id as usize].read(&mut nbrs_bytes)?; - let nbrs: &[u32] = bytemuck::cast_slice(&nbrs_bytes); - - // rename nodes - for j in 0..num_nbrs { - let nbr = nbrs[j as usize]; - let renamed_node = id_maps[shard_id as usize][nbr as usize]; - if !nbr_set[renamed_node as usize] { - nbr_set[renamed_node as usize] = true; - final_nbrs.push(renamed_node); - } - } - } - } - - // write the last node, to be refactored... - final_nbrs.shuffle(&mut self.rng); - - let nnbrs: u32 = std::cmp::min(final_nbrs.len() as u32, max_degree); - merged_vamana_cached_writer.write(&nnbrs.to_le_bytes())?; - - let bytes = final_nbrs - .iter() - .take(nnbrs as usize) - .flat_map(|x| x.to_le_bytes()) - .collect::>(); - merged_vamana_cached_writer.write(&bytes)?; - - merged_index_size += (size_of::() + nnbrs as usize * size_of::()) as u64; - - nbr_set.clear(); - final_nbrs.clear(); - - info!("Expected size: {}", merged_index_size); - merged_vamana_cached_writer.reset()?; - merged_vamana_cached_writer.write(&merged_index_size.to_le_bytes())?; - - info!("Finished merge"); - Ok(()) - } - - fn read_idmap(&self, idmaps_path: String) -> Result, diskann_utils::io::ReadBinError> { - let data = read_bin::(&mut self.storage_provider.open_reader(&idmaps_path)?)?; - Ok(data.into_inner().into_vec()) - } - - fn partition_data( - &mut self, - dataset_file: &str, - merged_index_prefix: &str, - max_degree: u32, - pool: RayonThreadPoolRef<'_>, - ) -> ANNResult { - let sampling_rate = MAX_PQ_TRAINING_SET_SIZE / self.index_configuration.max_points as f64; - let ram_budget_in_bytes = self.disk_build_param.build_memory_limit().in_bytes() as f64; - - partition_with_ram_budget::( - dataset_file, - self.index_configuration.dim, - sampling_rate, - ram_budget_in_bytes, - PARTITION_ASSIGNMENTS_PER_VECTOR, - merged_index_prefix, - self.storage_provider, - &mut self.rng, - pool, - |num_points, dim| { - let datasize = std::mem::size_of::() as u64; - let graph_degree = 2 * max_degree / 3; - estimate_build_index_ram_usage( - num_points, - dim, - datasize, - graph_degree as u64, - self.disk_build_param.build_quantization(), - ) - }, - ) - } - - fn merge_and_cleanup( - &mut self, - merged_index_prefix: &str, - num_parts: usize, - max_degree: u32, - output_vamana: String, - ) -> ANNResult<()> { - // merge all in-memory indices into one - self.merge_shards(merged_index_prefix, num_parts, max_degree, output_vamana)?; - - // delete tempFiles - for p in 0..num_parts { - let shard_base_file = - DiskIndexWriter::get_merged_index_subshard_data_file(merged_index_prefix, p); - let shard_ids_file = - DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, p); - let shard_index_file = - DiskIndexWriter::get_merged_index_subshard_mem_index_file(merged_index_prefix, p); - - self.storage_provider.delete(&shard_base_file)?; - self.storage_provider.delete(&shard_ids_file)?; - self.storage_provider.delete(&shard_index_file)?; - } - - Ok(()) - } -} - -pub(crate) enum IndexBuildStrategy { - OneShot, - Merged, -} - -pub(crate) fn determine_build_strategy( - index_configuration: &IndexConfiguration, - index_build_ram_limit_in_bytes: f64, - build_quantization_type: &QuantizationType, -) -> IndexBuildStrategy { - let estimated_index_ram_in_bytes = estimate_build_index_ram_usage( - index_configuration.max_points as u64, - index_configuration.dim as u64, - mem::size_of::() as u64, - index_configuration.config.max_degree().get() as u64, - build_quantization_type, - ); - - info!( - "Estimated index RAM usage: {} GB, index_build_ram_limit={} GB", - estimated_index_ram_in_bytes / BYTES_IN_GB, - index_build_ram_limit_in_bytes / BYTES_IN_GB - ); - - if estimated_index_ram_in_bytes >= index_build_ram_limit_in_bytes { - info!( - "Insufficient memory budget for index build in one shot, index_build_ram_limit={} GB estimated_index_ram={} GB", - index_build_ram_limit_in_bytes / BYTES_IN_GB, - estimated_index_ram_in_bytes / BYTES_IN_GB, - ); - IndexBuildStrategy::Merged - } else { - info!( - "Full index fits in RAM budget, should consume at most {} GBs, so building in one shot", - estimated_index_ram_in_bytes / BYTES_IN_GB - ); - IndexBuildStrategy::OneShot - } -} #[cfg(test)] pub(crate) mod disk_index_builder_tests { @@ -558,7 +26,6 @@ pub(crate) mod disk_index_builder_tests { use rstest::rstest; use vfs::OverlayFS; - use super::*; use crate::{ build::builder::build::DiskIndexBuilder, data_model::{CachingStrategy, GraphHeader}, @@ -567,9 +34,16 @@ pub(crate) mod disk_index_builder_tests { aligned_file_reader::VirtualAlignedReaderFactory, disk_provider::DiskIndexSearcher, disk_vertex_provider_factory::DiskVertexProviderFactory, }, - storage::disk_index_reader::DiskIndexReader, + storage::{disk_index_reader::DiskIndexReader, DiskIndexWriter}, utils::QueryStatistics, }; + use crate::{data_model::GraphDataType, QuantizationType}; + use diskann_providers::{ + model::IndexConfiguration, + storage::{StorageReadProvider, StorageWriteProvider}, + utils::load_metadata_from_file, + }; + use diskann_utils::io::read_bin; const DEFAULT_DISK_SECTOR_LEN: usize = 4096; pub const TEST_DATA_FILE: &str = "/sift/siftsmall_learn_256pts.fbin"; /// We can use the same index prefix for all tests since we use virtual storage provider @@ -1155,46 +629,3 @@ pub(crate) mod disk_index_builder_tests { ) } } - -#[cfg(test)] -mod ram_estimation_tests { - use rstest::rstest; - - use super::*; - use crate::QuantizationType; - - #[rstest] - #[case(QuantizationType::FP)] - #[case(QuantizationType::PQ { num_chunks: 15 })] - #[case(QuantizationType::SQ { nbits: 1, standard_deviation: None })] - fn test_estimate_build_index_ram_usage(#[case] build_quantization_type: QuantizationType) { - let num_points = 1000; - let dim = 128; - let size_of_t = std::mem::size_of::() as u64; - let graph_degree = 50; - - let single_vec_size = match build_quantization_type { - QuantizationType::FP => dim * size_of_t, - QuantizationType::PQ { num_chunks } => num_chunks as u64, - QuantizationType::SQ { nbits, .. } => { - (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::() as u64 - } - }; - let mut expected_ram_usage = (num_points as f64) - * (graph_degree as f64) - * (std::mem::size_of::() as f64) - * GRAPH_SLACK_FACTOR - + (num_points * single_vec_size) as f64; - expected_ram_usage *= OVERHEAD_FACTOR; - - let actual_ram_usage = estimate_build_index_ram_usage( - num_points, - dim, - size_of_t, - graph_degree, - &build_quantization_type, - ); - - assert_eq!(actual_ram_usage, expected_ram_usage); - } -} diff --git a/diskann-disk/src/search/provider/disk_provider.rs b/diskann-disk/src/search/provider/disk_provider.rs index a4986729d..ac789978b 100644 --- a/diskann-disk/src/search/provider/disk_provider.rs +++ b/diskann-disk/src/search/provider/disk_provider.rs @@ -1261,7 +1261,7 @@ mod disk_provider_tests { use super::*; use crate::{ - build::builder::core::disk_index_builder_tests::{IndexBuildFixture, TestParams}, + build::builder::disk_index_builder_tests::{IndexBuildFixture, TestParams}, error::{error_kind, ErrorKind}, search::provider::aligned_file_reader::VirtualAlignedReaderFactory, utils::QueryStatistics, diff --git a/diskann-disk/src/utils/instrumentation/perf_logger.rs b/diskann-disk/src/utils/instrumentation/perf_logger.rs index bdd19a9a5..03cea4bd8 100644 --- a/diskann-disk/src/utils/instrumentation/perf_logger.rs +++ b/diskann-disk/src/utils/instrumentation/perf_logger.rs @@ -24,7 +24,7 @@ mod scenario { #[derive(Debug)] pub enum DiskIndexBuildCheckpoint { PqConstruction, - InmemIndexBuild, + VamanaIndexBuild, DiskLayout, } @@ -185,7 +185,7 @@ mod perf_logger_tests { assert!(logger.log_enabled()); logger.log_checkpoint(DiskIndexBuildCheckpoint::PqConstruction); logger.start(); - logger.log_checkpoint(DiskIndexBuildCheckpoint::InmemIndexBuild); + logger.log_checkpoint(DiskIndexBuildCheckpoint::VamanaIndexBuild); } #[test] @@ -195,6 +195,6 @@ mod perf_logger_tests { assert!(!logger.log_enabled()); logger.log_checkpoint(DiskIndexBuildCheckpoint::PqConstruction); logger.start(); - logger.log_checkpoint(DiskIndexBuildCheckpoint::InmemIndexBuild); + logger.log_checkpoint(DiskIndexBuildCheckpoint::VamanaIndexBuild); } } From 94feca0f091b74bc5913b85af78806084d8ee24e Mon Sep 17 00:00:00 2001 From: Junkui Chen Date: Sat, 1 Aug 2026 23:48:39 +0800 Subject: [PATCH 2/2] Address Vamana builder review feedback Move the graph output path into DiskGraphOnly, make the borrowed save tuple explicit across the await, and cover both Vamana build strategy decisions at the RAM budget boundary. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- .../src/build/builder/vamana/one_shot.rs | 8 +-- .../src/build/builder/vamana/strategy.rs | 50 ++++++++++++++++++- 2 files changed, 51 insertions(+), 7 deletions(-) diff --git a/diskann-disk/src/build/builder/vamana/one_shot.rs b/diskann-disk/src/build/builder/vamana/one_shot.rs index 2f0e2f447..996ae8cae 100644 --- a/diskann-disk/src/build/builder/vamana/one_shot.rs +++ b/diskann-disk/src/build/builder/vamana/one_shot.rs @@ -115,12 +115,8 @@ where Self::log_build_stats(&index).await?; Self::run_final_prune(&index, num_tasks).await?; - index - .save_graph( - storage_provider, - &(start_point, DiskGraphOnly::new(&save_path)), - ) - .await?; + let graph_output = (start_point, DiskGraphOnly::new(save_path)); + index.save_graph(storage_provider, &graph_output).await?; Ok(()) } diff --git a/diskann-disk/src/build/builder/vamana/strategy.rs b/diskann-disk/src/build/builder/vamana/strategy.rs index 8014c13a3..cb4a20f5a 100644 --- a/diskann-disk/src/build/builder/vamana/strategy.rs +++ b/diskann-disk/src/build/builder/vamana/strategy.rs @@ -79,10 +79,12 @@ pub(in crate::build::builder) fn determine_build_strategy( #[cfg(test)] mod ram_estimation_tests { + use diskann::{graph::config, utils::ONE}; + use diskann_vector::distance::Metric::L2; use rstest::rstest; use super::*; - use crate::QuantizationType; + use crate::{test_utils::GraphDataF32VectorUnitData, QuantizationType}; #[rstest] #[case(QuantizationType::FP)] @@ -118,4 +120,50 @@ mod ram_estimation_tests { assert_eq!(actual_ram_usage, expected_ram_usage); } + + #[test] + fn selects_one_shot_when_index_fits_memory_budget() { + let index_configuration = index_configuration(); + + let strategy = determine_build_strategy::( + &index_configuration, + f64::INFINITY, + &QuantizationType::FP, + ); + + assert!(matches!(strategy, IndexBuildStrategy::OneShot)); + } + + #[test] + fn selects_merged_when_index_meets_memory_budget() { + let index_configuration = index_configuration(); + let estimated_usage = estimate_build_index_ram_usage( + index_configuration.max_points as u64, + index_configuration.dim as u64, + std::mem::size_of::() as u64, + index_configuration.config.max_degree().get() as u64, + &QuantizationType::FP, + ); + + let strategy = determine_build_strategy::( + &index_configuration, + estimated_usage, + &QuantizationType::FP, + ); + + assert!(matches!(strategy, IndexBuildStrategy::Merged)); + } + + fn index_configuration() -> IndexConfiguration { + IndexConfiguration::new( + L2, + 128, + 1000, + ONE, + 1, + config::Builder::new(16, config::MaxDegree::default_slack(), 64, L2.into()) + .build() + .unwrap(), + ) + } }