Skip to main content

diskann_disk/build/builder/
core.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5use std::mem::{self, size_of};
6
7use crate::data_model::GraphDataType;
8use diskann::ANNResult;
9use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
10use diskann_providers::{
11    model::{IndexConfiguration, GRAPH_SLACK_FACTOR, MAX_PQ_TRAINING_SET_SIZE},
12    storage::PQStorage,
13    utils::{
14        load_metadata_from_file, RayonThreadPoolRef, SampleVectorReader, SamplingDensity,
15        READ_WRITE_BLOCK_SIZE,
16    },
17};
18use diskann_utils::io::read_bin;
19use rand::{seq::SliceRandom, Rng};
20use tracing::info;
21
22use crate::{
23    build::chunking::{
24        checkpoint::{
25            CheckpointContext, CheckpointManager, CheckpointManagerExt, Progress, WorkStage,
26        },
27        continuation::ChunkingConfig,
28    },
29    disk_index_build_parameter::BYTES_IN_GB,
30    storage::{CachedReader, CachedWriter, DiskIndexWriter},
31    utils::partition_with_ram_budget,
32    DiskIndexBuildParameters, QuantizationType,
33};
34
35/// Overhead factor for RAM estimation during index build (10% buffer).
36const OVERHEAD_FACTOR: f64 = 1.1f64;
37
38/// Estimate RAM usage in bytes for building an index.
39#[inline]
40fn estimate_build_index_ram_usage(
41    num_points: u64,
42    dim: u64,
43    datasize: u64,
44    graph_degree: u64,
45    build_quantization_type: &QuantizationType,
46) -> f64 {
47    let graph_size =
48        (num_points * graph_degree * mem::size_of::<u32>() as u64) as f64 * GRAPH_SLACK_FACTOR;
49
50    let single_vec_size = match *build_quantization_type {
51        QuantizationType::FP => dim.next_multiple_of(8u64) * datasize,
52        // We can skip PQ pivots data as it is very small(~3MB) for even large datasets like OAI-3072.
53        QuantizationType::PQ { num_chunks } => num_chunks as u64,
54        // `+ std::mem::size_of::<f32>()` for f32 compensation metadata for the scalar quantizer.
55        QuantizationType::SQ { nbits, .. } => {
56            (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::<f32>() as u64
57        }
58    };
59
60    OVERHEAD_FACTOR * (graph_size + (single_vec_size * num_points) as f64)
61}
62
63/// Core shared functionality between sync and async disk index builders.
64/// Contains only fields and methods that are truly needed by both builder types.
65pub struct DiskIndexBuilderCore<'a, Data, StorageProvider>
66where
67    Data: GraphDataType<VectorIdType = u32>,
68    StorageProvider: StorageReadProvider + StorageWriteProvider,
69{
70    pub index_writer: DiskIndexWriter,
71
72    pub pq_storage: PQStorage,
73
74    pub disk_build_param: DiskIndexBuildParameters,
75
76    pub index_configuration: IndexConfiguration,
77
78    pub chunking_config: ChunkingConfig,
79
80    pub checkpoint_record_manager: Box<dyn CheckpointManager>,
81
82    pub storage_provider: &'a StorageProvider,
83
84    pub _phantom: std::marker::PhantomData<Data>,
85}
86
87impl<'a, Data, StorageProvider> DiskIndexBuilderCore<'a, Data, StorageProvider>
88where
89    Data: GraphDataType<VectorIdType = u32>,
90    StorageProvider: StorageReadProvider + StorageWriteProvider,
91{
92    pub(crate) fn create_disk_layout(&mut self) -> ANNResult<()> {
93        self.checkpoint_record_manager.execute_stage(
94            WorkStage::WriteDiskLayout,
95            WorkStage::End,
96            || {
97                self.index_writer
98                    .create_disk_layout::<Data, StorageProvider>(self.storage_provider)?;
99                Ok(())
100            },
101            || Ok(()),
102        )?;
103
104        self.index_writer
105            .index_build_cleanup(self.storage_provider)?;
106
107        Ok(())
108    }
109
110    pub(crate) fn create_shard_index_config(
111        &self,
112        shard_base_file: &str,
113    ) -> ANNResult<IndexConfiguration> {
114        let base_config = &self.index_configuration;
115        let storage_provider = self.storage_provider;
116
117        let search_list_size = base_config.config.l_build().get();
118        let pruned_degree = base_config.config.pruned_degree().get();
119
120        let low_degree_params = diskann::graph::config::Builder::new(
121            2 * pruned_degree / 3,
122            diskann::graph::config::MaxDegree::default_slack(),
123            search_list_size,
124            base_config.dist_metric.into(),
125        )
126        .build()?;
127
128        let metadata = load_metadata_from_file(storage_provider, shard_base_file)?;
129
130        let mut index_config = base_config.clone();
131        index_config.max_points = metadata.npoints();
132        index_config.config = low_degree_params;
133
134        Ok(index_config)
135    }
136
137    pub(crate) fn retrieve_shard_data_from_ids<T>(
138        &self,
139        dataset_file: &str,
140        shard_ids_file: &str,
141        shard_base_file: &str,
142    ) -> ANNResult<()>
143    where
144        T: Default + bytemuck::Pod,
145    {
146        let storage_provider = self.storage_provider;
147        let shard_ids = read_bin::<u32>(&mut storage_provider.open_reader(shard_ids_file)?)?;
148        let shard_size = shard_ids.nrows();
149        info!("Loaded {} shard ids from {}", shard_size, shard_ids_file);
150        let max_id = shard_ids.as_slice().iter().max().copied().unwrap_or(0);
151        let sampling_rate = shard_ids.as_slice().len() as f64 / (max_id + 1) as f64;
152
153        let mut dataset_reader: SampleVectorReader<T, _> = SampleVectorReader::new(
154            dataset_file,
155            SamplingDensity::from_sample_rate(sampling_rate),
156            storage_provider,
157        )?;
158
159        let (_npts, dim) = dataset_reader.get_dataset_headers();
160
161        let mut shard_base_cached_writer = CachedWriter::<StorageProvider>::new(
162            shard_base_file,
163            READ_WRITE_BLOCK_SIZE,
164            storage_provider.create_for_write(shard_base_file)?,
165        )?;
166
167        let dummy_size: u32 = 0;
168        shard_base_cached_writer.write(&dummy_size.to_le_bytes())?;
169        shard_base_cached_writer.write(&dim.to_le_bytes())?;
170
171        let mut num_written: u32 = 0;
172        dataset_reader.read_vectors(shard_ids.as_slice().iter().copied(), |vector_t| {
173            // Casting Pod type to bytes always succeeds (u8 has alignment of 1)
174            let vector_bytes: &[u8] = bytemuck::must_cast_slice(vector_t);
175            shard_base_cached_writer.write(vector_bytes)?;
176            num_written += 1;
177            Ok(())
178        })?;
179
180        info!(
181            "Written file: {} with {} points",
182            shard_base_file, num_written
183        );
184
185        shard_base_cached_writer.flush()?;
186        shard_base_cached_writer.reset()?;
187        shard_base_cached_writer.write(&num_written.to_le_bytes())?;
188
189        Ok(())
190    }
191
192    #[allow(clippy::too_many_arguments)]
193    fn merge_shards(
194        &self,
195        merged_index_prefix: &str,
196        num_parts: usize,
197        max_degree: u32,
198        output_vamana: String,
199        rng: &mut impl Rng,
200    ) -> ANNResult<()> {
201        // Read ID maps
202        let mut vamana_names = vec![String::new(); num_parts];
203        let mut id_maps: Vec<Vec<u32>> = vec![Vec::new(); num_parts];
204        for shard in 0..num_parts {
205            vamana_names[shard] = DiskIndexWriter::get_merged_index_subshard_mem_index_file(
206                merged_index_prefix,
207                shard,
208            );
209
210            let id_maps_file =
211                DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard);
212            id_maps[shard] = self.read_idmap(id_maps_file)?;
213        }
214
215        // find max node id
216        let num_nodes: u32 = *id_maps.iter().flatten().max().unwrap_or(&0) + 1;
217        let num_elements: u32 = id_maps.iter().map(|idmap| idmap.len() as u32).sum();
218        info!("# nodes: {}, max degree: {}", num_nodes, max_degree);
219
220        // compute inverse map: node -> shards
221        let mut node_shard: Vec<(u32, u32)> = Vec::with_capacity(num_elements as usize);
222        for (shard, id_map) in id_maps.iter().enumerate() {
223            info!("Creating inverse map -- shard #{}", shard);
224            node_shard.extend(id_map.iter().map(|node_id| (*node_id, shard as u32)));
225        }
226        node_shard.sort_unstable_by(|left, right| {
227            left.0.cmp(&right.0).then_with(|| left.1.cmp(&right.1))
228        });
229
230        info!("Finished computing node -> shards map");
231
232        // create cached vamana readers
233        let mut vamana_readers = Vec::new();
234        for name in &vamana_names {
235            let reader = CachedReader::<StorageProvider>::new(
236                name,
237                READ_WRITE_BLOCK_SIZE,
238                self.storage_provider,
239            )?;
240            vamana_readers.push(reader);
241        }
242
243        // create cached vamana writers
244        let mut merged_vamana_cached_writer = CachedWriter::<StorageProvider>::new(
245            &output_vamana,
246            READ_WRITE_BLOCK_SIZE,
247            self.storage_provider.create_for_write(&output_vamana)?,
248        )?;
249
250        // expected file size + max degree + medoid_id + frozen_point info
251        let vamana_metadata_size =
252            size_of::<u64>() + size_of::<u32>() + size_of::<u32>() + size_of::<u64>();
253
254        // we initialize the size of the merged index to the metadata size
255        // we will overwrite the index size at the end
256        let mut merged_index_size: u64 = vamana_metadata_size as u64;
257        merged_vamana_cached_writer.write(&merged_index_size.to_le_bytes())?;
258
259        let mut read_buf_8_bytes = [0u8; 8];
260
261        // get max input width
262        let mut max_input_width = 0;
263        // read width from each vamana to advance buffer by sizeof(uint32_t) bytes
264        for reader in &mut vamana_readers {
265            reader.read(&mut read_buf_8_bytes)?;
266            let _expected_file_size: u64 = u64::from_le_bytes(read_buf_8_bytes);
267            let input_width = reader.read_u32()?;
268            max_input_width = input_width.max(max_input_width);
269        }
270
271        // write max_degree to merged_vamana_index
272        let output_width: u32 = max_degree;
273        info!(
274            "Max input width: {}, output width: {}",
275            max_input_width, output_width
276        );
277
278        merged_vamana_cached_writer.write(&output_width.to_le_bytes())?;
279
280        // write medoid to merged_vamana_index
281        for shard in 0..num_parts {
282            // read medoid
283            let mut medoid: u32 = vamana_readers[shard].read_u32()?;
284            vamana_readers[shard].read(&mut read_buf_8_bytes)?;
285            let vamana_index_frozen: u64 = u64::from_le_bytes(read_buf_8_bytes);
286            debug_assert_eq!(vamana_index_frozen, 0);
287
288            // rename medoid
289            medoid = id_maps[shard][medoid as usize];
290
291            // write renamed medoid
292            if shard == (num_parts - 1) {
293                // uncomment if running hierarchical
294                merged_vamana_cached_writer.write(&medoid.to_le_bytes())?;
295            }
296        }
297
298        let vamana_index_frozen: u64 = 0; // as of now the functionality to merge many overlapping vamana
299                                          // indices is supported only for bulk indices without frozen point.
300                                          // Hence the final index will also not have any frozen points.
301        merged_vamana_cached_writer.write(&vamana_index_frozen.to_le_bytes())?;
302
303        info!("Starting merge");
304
305        let mut nbr_set = vec![false; num_nodes as usize];
306        let mut final_nbrs: Vec<u32> = Vec::new();
307        let mut cur_id = 0;
308        for pair in &node_shard {
309            let (node_id, shard_id) = *pair;
310            if cur_id < node_id {
311                final_nbrs.shuffle(rng);
312
313                let nnbrs: u32 = std::cmp::min(final_nbrs.len() as u32, max_degree);
314                merged_vamana_cached_writer.write(&nnbrs.to_le_bytes())?;
315
316                let bytes = final_nbrs
317                    .iter()
318                    .take(nnbrs as usize)
319                    .flat_map(|x| x.to_le_bytes())
320                    .collect::<Vec<u8>>();
321                merged_vamana_cached_writer.write(&bytes)?;
322
323                merged_index_size += (size_of::<u32>() + nnbrs as usize * size_of::<u32>()) as u64;
324                if cur_id % 499999 == 1 {
325                    print!(".");
326                }
327                cur_id = node_id;
328
329                final_nbrs.iter().for_each(|p| nbr_set[*p as usize] = false);
330                final_nbrs.clear();
331            }
332
333            // read num of neighbors from vamana index
334            let num_nbrs = vamana_readers[shard_id as usize].read_u32()?;
335
336            if num_nbrs == 0 {
337                info!(
338                    "WARNING: shard #{}, node_id {} has 0 nbrs",
339                    shard_id, node_id
340                );
341            } else {
342                let mut nbrs_bytes = vec![0u8; num_nbrs as usize * mem::size_of::<u32>()];
343                vamana_readers[shard_id as usize].read(&mut nbrs_bytes)?;
344                let nbrs: &[u32] = bytemuck::cast_slice(&nbrs_bytes);
345
346                // rename nodes
347                for j in 0..num_nbrs {
348                    let nbr = nbrs[j as usize];
349                    let renamed_node = id_maps[shard_id as usize][nbr as usize];
350                    if !nbr_set[renamed_node as usize] {
351                        nbr_set[renamed_node as usize] = true;
352                        final_nbrs.push(renamed_node);
353                    }
354                }
355            }
356        }
357
358        // write the last node, to be refactored...
359        final_nbrs.shuffle(rng);
360
361        let nnbrs: u32 = std::cmp::min(final_nbrs.len() as u32, max_degree);
362        merged_vamana_cached_writer.write(&nnbrs.to_le_bytes())?;
363
364        let bytes = final_nbrs
365            .iter()
366            .take(nnbrs as usize)
367            .flat_map(|x| x.to_le_bytes())
368            .collect::<Vec<u8>>();
369        merged_vamana_cached_writer.write(&bytes)?;
370
371        merged_index_size += (size_of::<u32>() + nnbrs as usize * size_of::<u32>()) as u64;
372
373        nbr_set.clear();
374        final_nbrs.clear();
375
376        info!("Expected size: {}", merged_index_size);
377        merged_vamana_cached_writer.reset()?;
378        merged_vamana_cached_writer.write(&merged_index_size.to_le_bytes())?;
379
380        info!("Finished merge");
381        Ok(())
382    }
383
384    fn read_idmap(&self, idmaps_path: String) -> Result<Vec<u32>, diskann_utils::io::ReadBinError> {
385        let data = read_bin::<u32>(&mut self.storage_provider.open_reader(&idmaps_path)?)?;
386        Ok(data.into_inner().into_vec())
387    }
388
389    fn merge_shards_and_cleanup(
390        &self,
391        merged_index_prefix: &str,
392        num_parts: usize,
393        max_degree: u32,
394        rng: &mut impl Rng,
395    ) -> ANNResult<()> {
396        // merge all in-memory indices into one
397        self.merge_shards(
398            merged_index_prefix,
399            num_parts,
400            max_degree,
401            self.index_writer.get_mem_index_file(),
402            rng,
403        )?;
404
405        // delete tempFiles
406        for p in 0..num_parts {
407            let shard_base_file =
408                DiskIndexWriter::get_merged_index_subshard_data_file(merged_index_prefix, p);
409            let shard_ids_file =
410                DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, p);
411            let shard_index_file =
412                DiskIndexWriter::get_merged_index_subshard_mem_index_file(merged_index_prefix, p);
413            let shard_index_file_data =
414                DiskIndexWriter::get_merged_index_subshard_mem_dataset_file(&shard_index_file);
415
416            self.storage_provider.delete(&shard_base_file)?;
417            self.storage_provider.delete(&shard_ids_file)?;
418            self.storage_provider.delete(&shard_index_file)?;
419            // Check if shard dataset file exists before deleting it.
420            // Async build path doesn't always create this file.
421            if self.storage_provider.exists(&shard_index_file_data) {
422                self.storage_provider.delete(&shard_index_file_data)?;
423            }
424        }
425
426        Ok(())
427    }
428}
429
430pub(crate) enum IndexBuildStrategy {
431    OneShot,
432    Merged,
433}
434
435pub(crate) fn determine_build_strategy<Data: GraphDataType>(
436    index_configuration: &IndexConfiguration,
437    index_build_ram_limit_in_bytes: f64,
438    build_quantization_type: &QuantizationType,
439) -> IndexBuildStrategy {
440    let estimated_index_ram_in_bytes = estimate_build_index_ram_usage(
441        index_configuration.max_points as u64,
442        index_configuration.dim as u64,
443        mem::size_of::<Data::VectorDataType>() as u64,
444        index_configuration.config.max_degree().get() as u64,
445        build_quantization_type,
446    );
447
448    info!(
449        "Estimated index RAM usage: {} GB, index_build_ram_limit={} GB",
450        estimated_index_ram_in_bytes / BYTES_IN_GB,
451        index_build_ram_limit_in_bytes / BYTES_IN_GB
452    );
453
454    if estimated_index_ram_in_bytes >= index_build_ram_limit_in_bytes {
455        info!(
456            "Insufficient memory budget for index build in one shot, index_build_ram_limit={} GB estimated_index_ram={} GB",
457            index_build_ram_limit_in_bytes / BYTES_IN_GB,
458            estimated_index_ram_in_bytes / BYTES_IN_GB,
459        );
460        IndexBuildStrategy::Merged
461    } else {
462        info!(
463            "Full index fits in RAM budget, should consume at most {} GBs, so building in one shot",
464            estimated_index_ram_in_bytes / BYTES_IN_GB
465        );
466        IndexBuildStrategy::OneShot
467    }
468}
469
470pub(crate) struct MergedVamanaIndexWorkflow<'a> {
471    pool: RayonThreadPoolRef<'a>,
472    rng: diskann_providers::utils::StandardRng,
473    dataset_file: String,
474    max_degree: u32,
475    pub merged_index_prefix: String,
476}
477
478impl<'a> MergedVamanaIndexWorkflow<'a> {
479    pub(crate) fn new<Data, StorageProvider>(
480        builder: &mut DiskIndexBuilderCore<'_, Data, StorageProvider>,
481        pool: RayonThreadPoolRef<'a>,
482    ) -> Self
483    where
484        Data: GraphDataType<VectorIdType = u32>,
485        StorageProvider: StorageReadProvider + StorageWriteProvider,
486    {
487        let rng = diskann_providers::utils::create_rnd_from_optional_seed(
488            builder.index_configuration.random_seed,
489        );
490        let dataset_file = builder.index_writer.get_dataset_file();
491        let merged_index_prefix = builder.index_writer.get_merged_index_prefix();
492        let max_degree = builder.index_configuration.config.pruned_degree_u32().get();
493
494        Self {
495            pool,
496            rng,
497            dataset_file,
498            merged_index_prefix,
499            max_degree,
500        }
501    }
502
503    pub(crate) fn partition_data<Data, StorageProvider>(
504        &mut self,
505        builder: &mut DiskIndexBuilderCore<'_, Data, StorageProvider>,
506    ) -> ANNResult<usize>
507    where
508        Data: GraphDataType<VectorIdType = u32>,
509        StorageProvider: StorageReadProvider + StorageWriteProvider,
510    {
511        // Advance to PartitionData stage if current stage is InMemIndexBuild
512        builder.checkpoint_record_manager.execute_stage(
513            WorkStage::InMemIndexBuild,
514            WorkStage::PartitionData,
515            || Ok(()),
516            || Ok(()),
517        )?;
518
519        // Partition data stage
520        builder.checkpoint_record_manager.execute_stage(
521            WorkStage::PartitionData,
522            WorkStage::BuildIndicesOnShards(0),
523            || {
524                let num_points = builder.index_configuration.max_points;
525                let sampling_rate = MAX_PQ_TRAINING_SET_SIZE / num_points as f64;
526
527                let ram_budget_in_bytes =
528                    builder.disk_build_param.build_memory_limit().in_bytes() as f64;
529                // calculate how many partitions we need, in order to fit in RAM budget
530                // save id_map for each partition to disk
531                partition_with_ram_budget::<Data::VectorDataType, _, _>(
532                    &self.dataset_file,
533                    builder.index_configuration.dim,
534                    sampling_rate,
535                    ram_budget_in_bytes,
536                    2, // k_base
537                    &self.merged_index_prefix,
538                    builder.storage_provider,
539                    &mut self.rng,
540                    self.pool,
541                    |num_points, dim| {
542                        let datasize = std::mem::size_of::<Data::VectorDataType>() as u64;
543                        let graph_degree = 2 * self.max_degree / 3;
544                        estimate_build_index_ram_usage(
545                            num_points,
546                            dim,
547                            datasize,
548                            graph_degree as u64,
549                            builder.disk_build_param.build_quantization(),
550                        )
551                    },
552                )
553            },
554            || {
555                // load num_parts based on file names
556                let mut p = 0;
557                while builder.storage_provider.exists(
558                    &DiskIndexWriter::get_merged_index_subshard_id_map_file(
559                        &self.merged_index_prefix,
560                        p,
561                    ),
562                ) {
563                    p += 1;
564                }
565                info!("Found {} existing partitions from previous run", p);
566                Ok(p)
567            },
568        )
569    }
570
571    pub(crate) fn merge_and_cleanup<Data, StorageProvider>(
572        &mut self,
573        builder: &mut DiskIndexBuilderCore<'_, Data, StorageProvider>,
574        num_parts: usize,
575    ) -> ANNResult<()>
576    where
577        Data: GraphDataType<VectorIdType = u32>,
578        StorageProvider: StorageReadProvider + StorageWriteProvider,
579    {
580        if builder
581            .checkpoint_record_manager
582            .get_resumption_point(WorkStage::MergeIndices)?
583            .is_some()
584        {
585            builder.merge_shards_and_cleanup(
586                &self.merged_index_prefix,
587                num_parts,
588                self.max_degree,
589                &mut self.rng,
590            )?;
591            builder
592                .checkpoint_record_manager
593                .update(Progress::Completed, WorkStage::WriteDiskLayout)?;
594        }
595
596        Ok(())
597    }
598
599    pub(crate) fn get_shard_context<'b, Data, StorageProvider>(
600        &self,
601        builder: &'b DiskIndexBuilderCore<'_, Data, StorageProvider>,
602        p: usize,
603        num_parts: usize,
604    ) -> CheckpointContext<'b>
605    where
606        Data: GraphDataType<VectorIdType = u32>,
607        StorageProvider: StorageReadProvider + StorageWriteProvider,
608    {
609        let current_stage = WorkStage::BuildIndicesOnShards(p);
610        let next_stage = if p == num_parts - 1 {
611            // If this is the last shard, next stage is MergeIndices
612            WorkStage::MergeIndices
613        } else {
614            // Otherwise, continue with the next shard
615            WorkStage::BuildIndicesOnShards(p + 1)
616        };
617        CheckpointContext::new(
618            builder.checkpoint_record_manager.as_ref(),
619            current_stage,
620            next_stage,
621        )
622    }
623}
624
625#[cfg(test)]
626pub(crate) mod disk_index_builder_tests {
627    use std::{io::Read, sync::Arc};
628
629    use crate::test_utils::{GraphDataF32VectorU32Data, GraphDataF32VectorUnitData};
630    use diskann::{
631        graph::config,
632        utils::{IntoUsize, VectorRepr, ONE},
633        ANNResult,
634    };
635    use diskann_providers::storage::VirtualStorageProvider;
636    use diskann_providers::{
637        storage::{get_compressed_pq_file, get_disk_index_file, get_pq_pivot_file},
638        utils::Timer,
639    };
640    use diskann_utils::test_data_root;
641    use diskann_vector::{
642        distance::Metric::{self, L2},
643        DistanceFunction,
644    };
645    use rstest::rstest;
646    use vfs::OverlayFS;
647
648    use super::*;
649    use crate::{
650        build::builder::build::DiskIndexBuilder,
651        data_model::{CachingStrategy, GraphHeader},
652        disk_index_build_parameter::{DiskIndexBuildParameters, MemoryBudget, NumPQChunks},
653        search::provider::{
654            disk_provider::DiskIndexSearcher,
655            disk_vertex_provider_factory::DiskVertexProviderFactory,
656        },
657        storage::disk_index_reader::DiskIndexReader,
658        utils::{QueryStatistics, VirtualAlignedReaderFactory},
659    };
660    const DEFAULT_DISK_SECTOR_LEN: usize = 4096;
661    pub const TEST_DATA_FILE: &str = "/sift/siftsmall_learn_256pts.fbin";
662    /// We can use the same index prefix for all tests since we use virtual storage provider
663    const INDEX_PATH_PREFIX: &str = "/disk_index_build/sift_learn_test_disk_index_build";
664    const TRUTH_INDEX_PATH_PREFIX_R4_L50: &str = "/disk_index_build/truth_sift_learn_R4_L50";
665
666    pub struct CheckpointParams {
667        pub chunking_config: ChunkingConfig,
668        pub checkpoint_record_manager: Box<dyn CheckpointManager>,
669    }
670
671    pub struct TestParams {
672        pub dim: usize,
673        pub full_dim: usize,
674        pub max_degree: u32,
675        pub num_pq_chunks: usize,
676        pub build_quantization_type: QuantizationType,
677        pub l_build: u32,
678        pub data_path: String,
679        pub index_path_prefix: String,
680        pub associated_data_path: Option<String>,
681        pub index_build_ram_gb: f64,
682        pub checkpoint_params: Option<CheckpointParams>,
683        pub num_threads: usize,
684        pub metric: Metric,
685    }
686
687    impl Default for TestParams {
688        fn default() -> Self {
689            Self {
690                dim: 128, // D
691                full_dim: 128,
692                max_degree: 4, // R
693                num_pq_chunks: 128,
694                build_quantization_type: QuantizationType::FP, // No quantization, i.e. QuantizationType::FP
695                l_build: 50,
696                data_path: TEST_DATA_FILE.to_string(),
697                index_path_prefix: INDEX_PATH_PREFIX.to_string(),
698                associated_data_path: None,
699                index_build_ram_gb: 1.0,
700                checkpoint_params: None,
701                num_threads: 1,
702                metric: L2,
703            }
704        }
705    }
706
707    impl TestParams {
708        /// Returns the appropriate truth index path prefix for build comparison.
709        fn truth_index_path_prefix(&self) -> &str {
710            match (self.max_degree, self.l_build, self.index_build_ram_gb) {
711                (4, 50, 1.0) => TRUTH_INDEX_PATH_PREFIX_R4_L50,
712                (max_degree, l_build, index_build_ram_gb) => panic!(
713                    "Truth index path not found for max_degree={}, l_build={}, index_build_ram_gb={}",
714                    max_degree, l_build, index_build_ram_gb
715                ),
716            }
717        }
718        pub fn truth_pq_compressed_path(&self) -> String {
719            let prefix = match self.num_pq_chunks {
720                128 => TRUTH_INDEX_PATH_PREFIX_R4_L50,
721                num_pq_chunks => panic!(
722                    "Truth pq compressed path not found for num_pq_chunks={}",
723                    num_pq_chunks,
724                ),
725            };
726            get_compressed_pq_file(prefix)
727        }
728
729        pub fn pq_compressed_path(&self) -> String {
730            get_compressed_pq_file(&self.index_path_prefix)
731        }
732    }
733
734    pub fn new_vfs() -> VirtualStorageProvider<OverlayFS> {
735        VirtualStorageProvider::new_overlay(test_data_root())
736    }
737
738    pub struct IndexBuildFixture<StorageProvider: StorageReadProvider + StorageWriteProvider> {
739        pub storage_provider: Arc<StorageProvider>,
740        pub params: TestParams,
741    }
742
743    impl<StorageProvider: StorageReadProvider + StorageWriteProvider + 'static>
744        IndexBuildFixture<StorageProvider>
745    {
746        pub fn new(storage_provider: StorageProvider, params: TestParams) -> ANNResult<Self> {
747            Ok(Self {
748                storage_provider: Arc::new(storage_provider),
749                params,
750            })
751        }
752
753        pub fn build<T>(&self) -> ANNResult<()>
754        where
755            T: GraphDataType<VectorIdType = u32>,
756            StorageProvider::Reader: std::marker::Send + Read,
757        {
758            // Create disk index build parameters
759            let disk_index_build_parameters = DiskIndexBuildParameters::new(
760                MemoryBudget::try_from_gb(self.params.index_build_ram_gb)?,
761                self.params.build_quantization_type,
762                NumPQChunks::new_with(self.params.num_pq_chunks, self.params.full_dim)?,
763            );
764
765            let config = config::Builder::new_with(
766                self.params.max_degree.into_usize(),
767                config::MaxDegree::default_slack(),
768                self.params.l_build.into_usize(),
769                self.params.metric.into(),
770                |b| {
771                    b.saturate_after_prune(true);
772                },
773            )
774            .build()?;
775
776            let metadata =
777                load_metadata_from_file(self.storage_provider.as_ref(), &self.params.data_path)
778                    .unwrap();
779
780            assert_eq!(
781                self.params.dim,
782                metadata.ndims(),
783                "Parameters dimension {} and data dimension {} are not equal",
784                self.params.dim,
785                metadata.ndims(),
786            );
787
788            let config = IndexConfiguration::new(
789                self.params.metric,
790                self.params.dim,
791                metadata.npoints(),
792                ONE,
793                self.params.num_threads,
794                config,
795            )
796            .with_pseudo_rng_from_seed(100);
797
798            let disk_index_writer = DiskIndexWriter::new(
799                self.params.data_path.clone(),
800                self.params.index_path_prefix.clone(),
801                self.params.associated_data_path.clone(),
802                DEFAULT_DISK_SECTOR_LEN,
803            )?;
804
805            let mut disk_index = match self.params.checkpoint_params {
806                Some(ref checkpoint_params) => {
807                    let checkpoint_record_manager =
808                        checkpoint_params.checkpoint_record_manager.clone_box();
809                    let chunking_config = checkpoint_params.chunking_config.clone();
810                    DiskIndexBuilder::<T, _>::new_with_chunking_config(
811                        self.storage_provider.as_ref(),
812                        disk_index_build_parameters,
813                        config,
814                        disk_index_writer,
815                        chunking_config,
816                        checkpoint_record_manager,
817                    )
818                }
819                None => DiskIndexBuilder::<T, _>::new(
820                    self.storage_provider.as_ref(),
821                    disk_index_build_parameters,
822                    config,
823                    disk_index_writer,
824                ),
825            }?;
826
827            let timer = Timer::new();
828            disk_index.build()?;
829            println!("Indexing time: {} seconds", timer.elapsed().as_secs_f64());
830
831            Ok(())
832        }
833
834        pub fn compare_pq_compressed_files(&self) {
835            self.compare_files(
836                &self.params.pq_compressed_path(),
837                &self.params.truth_pq_compressed_path(),
838            );
839        }
840
841        pub fn assert_index_max_degree<T: GraphDataType>(&self) -> ANNResult<()> {
842            let index_file_path = get_disk_index_file(&self.params.index_path_prefix);
843            let file_data = load_file_to_vec(self.storage_provider.as_ref(), &index_file_path);
844            let graph_header = GraphHeader::try_from(&file_data[8..])?;
845            let max_degree = graph_header.max_degree::<T::VectorDataType>()?;
846            assert_eq!(
847                max_degree, self.params.max_degree as usize,
848                "Max degree mismatch: expected {}, got {}",
849                self.params.max_degree, max_degree
850            );
851
852            Ok(())
853        }
854
855        fn compare_disk_index_with_associated_data(
856            &self,
857            pivot_file_prefix_test: &str,
858            pivot_file_prefix_expected: &str,
859            index_file_suffix: &str,
860        ) {
861            let pq_pivot_path = pivot_file_prefix_test.to_string() + index_file_suffix;
862            let pq_pivot_path_truth = pivot_file_prefix_expected.to_string() + index_file_suffix;
863            let file1 = load_file_to_vec(self.storage_provider.as_ref(), &pq_pivot_path);
864            let file2 = load_file_to_vec(self.storage_provider.as_ref(), &pq_pivot_path_truth);
865            compare_disk_index_graphs(&file1, &file2)
866        }
867
868        pub fn compare_files(&self, file_path1: &str, file_path2: &str) {
869            let file1 = load_file_to_vec(self.storage_provider.as_ref(), file_path1);
870            let file2 = load_file_to_vec(self.storage_provider.as_ref(), file_path2);
871
872            assert_eq!(file1.len(), file2.len());
873            assert_eq!(file1, file2)
874        }
875    }
876
877    /// Common helper function for one-shot async index build tests
878    fn run_one_shot_test<F>(index_path_prefix: String, params_customizer: F)
879    where
880        F: FnOnce(TestParams) -> TestParams,
881    {
882        let l_build = 64;
883        let max_degree = 16;
884        let top_k = 10;
885        let search_l = 32;
886
887        let base_params = TestParams {
888            l_build,
889            max_degree,
890            index_path_prefix,
891            ..TestParams::default()
892        };
893
894        let params = params_customizer(base_params);
895
896        let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
897        fixture.build::<GraphDataF32VectorUnitData>().unwrap();
898
899        // Validate search recall against ground truth for async tests
900        verify_search_result_with_ground_truth::<GraphDataF32VectorUnitData>(
901            &fixture.params,
902            top_k,
903            search_l,
904            &fixture.storage_provider,
905        )
906        .unwrap();
907
908        fixture
909            .assert_index_max_degree::<GraphDataF32VectorUnitData>()
910            .unwrap();
911
912        // Assert that all data was kept in memory and no files were written to the disk.
913        let mem_index_file_path = format!("{}_mem.index.data", fixture.params.index_path_prefix);
914        assert!(!fixture.storage_provider.exists(&mem_index_file_path));
915    }
916
917    #[rstest]
918    fn test_build_from_iter_one_shot_with_metric(
919        #[values(Metric::L2, Metric::InnerProduct, Metric::Cosine)] metric: Metric,
920    ) {
921        let index_path_prefix = format!("{}_metric_{:?}", INDEX_PATH_PREFIX, metric);
922
923        run_one_shot_test(index_path_prefix, |params| TestParams { metric, ..params });
924    }
925
926    #[test]
927    fn test_build_from_iter_one_shot_with_associated_data() {
928        // Set up test data
929        let params = TestParams {
930            associated_data_path: Some(
931                "/sift/siftsmall_learn_256pts_u32_associated_data.fbin".to_string(),
932            ),
933            ..TestParams::default()
934        };
935
936        // Create fixture with virtual storage provider
937        let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
938
939        // Build the index with the associated data
940        fixture.build::<GraphDataF32VectorU32Data>().unwrap();
941
942        // Assert that all data was kept in memory and no files were written to the disk.
943        let mem_index_file_path = format!("{}_mem.index.data", fixture.params.index_path_prefix);
944        let mem_index_associated_data_path = format!(
945            "{}_mem.index.associated_data",
946            fixture.params.index_path_prefix
947        );
948        assert!(!fixture.storage_provider.exists(&mem_index_file_path));
949        assert!(!fixture
950            .storage_provider
951            .exists(&mem_index_associated_data_path));
952
953        // assert index files are expected.
954        fixture.compare_disk_index_with_associated_data(
955            &fixture.params.index_path_prefix,
956            fixture.params.truth_index_path_prefix(),
957            "_disk.index",
958        );
959    }
960
961    #[test]
962    fn test_build_from_iter_merged_index() {
963        // Use the same parameters from [test_sift_build_and_search] in diskann_index
964        let l_build = 64;
965        let max_degree = 16;
966        let top_k = 10;
967        let search_l = 32;
968
969        let index_path_prefix =
970            "/disk_index_build/disk_index_sift_learn_test_disk_index_build_merged".to_string();
971        let params = TestParams {
972            l_build,
973            max_degree,
974            index_path_prefix,
975            index_build_ram_gb: 0.0001, // small enough to trigger merged index build
976            ..TestParams::default()
977        };
978
979        let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
980
981        fixture.build::<GraphDataF32VectorUnitData>().unwrap();
982
983        verify_search_result_with_ground_truth::<GraphDataF32VectorUnitData>(
984            &fixture.params,
985            top_k,
986            search_l,
987            &fixture.storage_provider,
988        )
989        .unwrap();
990
991        fixture
992            .assert_index_max_degree::<GraphDataF32VectorUnitData>()
993            .unwrap();
994    }
995
996    #[rstest]
997    #[case(QuantizationType::SQ { nbits: 2, standard_deviation: None }, "SQ quantization is only supported for 1 bit")]
998    fn test_build_quantization_type_failure_cases(
999        #[case] build_quantization_type: QuantizationType,
1000        #[case] error_message: &str,
1001    ) {
1002        let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
1003        let disk_index_builder = create_disk_index_builder(
1004            1000, // num_points
1005            128,  // dim
1006            128,  // num_pq_chunks
1007            &storage_provider,
1008            build_quantization_type,
1009        );
1010
1011        let err = disk_index_builder.err().unwrap();
1012        assert!(err.to_string().contains(error_message));
1013    }
1014
1015    fn load_file_to_vec<StorageType: StorageReadProvider>(
1016        storage_provider: &StorageType,
1017        file_path: &str,
1018    ) -> Vec<u8> {
1019        let mut file = storage_provider.open_reader(file_path).unwrap();
1020        let mut buffer = vec![];
1021        file.read_to_end(&mut buffer).unwrap();
1022        buffer
1023    }
1024
1025    /// Verifies that search results exactly match the ground truth of nearest neighbors
1026    ///
1027    /// This function performs validation of search results by:
1028    /// 1. Running searches on the index using actual data points from the dataset as queries
1029    /// 2. Computing the exact ground truth results using direct distance calculations
1030    /// 3. Verifying that the search engine returns precisely the same results as the ground truth
1031    pub(crate) fn verify_search_result_with_ground_truth<
1032        G: GraphDataType<VectorIdType = u32, AssociatedDataType = ()>,
1033    >(
1034        params: &TestParams,
1035        top_k: usize,
1036        search_l: u32,
1037        storage_provider: &Arc<VirtualStorageProvider<OverlayFS>>,
1038    ) -> ANNResult<()> {
1039        let pq_pivot_path = get_pq_pivot_file(&params.index_path_prefix);
1040        let pq_compressed_path = get_compressed_pq_file(&params.index_path_prefix);
1041        let index_file_path = get_disk_index_file(&params.index_path_prefix);
1042
1043        let index_reader = DiskIndexReader::<G::VectorDataType>::new(
1044            pq_pivot_path,
1045            pq_compressed_path,
1046            storage_provider.as_ref(),
1047        )?;
1048
1049        let vertex_provider_factory = DiskVertexProviderFactory::new(
1050            VirtualAlignedReaderFactory::new(index_file_path, Arc::clone(storage_provider)),
1051            CachingStrategy::None,
1052        )?;
1053
1054        let search_engine = DiskIndexSearcher::<G, DiskVertexProviderFactory<G, _>>::new(
1055            1,
1056            u32::MAX as usize,
1057            &index_reader,
1058            vertex_provider_factory,
1059            params.metric,
1060            None,
1061        )?;
1062
1063        let data =
1064            read_bin::<G::VectorDataType>(&mut storage_provider.open_reader(&params.data_path)?)?;
1065        let dim = data.ncols();
1066        let distance = <G::VectorDataType>::distance(params.metric, Some(dim));
1067
1068        // Here, we use elements of the dataset to search the dataset itself.
1069        //
1070        // We do this for each query, computing the expected ground truth and verifying
1071        // that our simple graph search matches.
1072        //
1073        // Because this dataset is small, we can expect exact equality.
1074        for (q, query_data) in data.row_iter().enumerate() {
1075            let gt =
1076                diskann_providers::test_utils::groundtruth(data.as_view(), query_data, |a, b| {
1077                    distance.evaluate_similarity(a, b)
1078                });
1079
1080            let mut query_stats = QueryStatistics::default();
1081
1082            let mut indices = vec![0u32; top_k];
1083            let mut distances = vec![0f32; top_k];
1084            let mut associated_data = vec![(); top_k];
1085
1086            _ = search_engine.search_internal(
1087                query_data,
1088                top_k,
1089                search_l,
1090                None, // beam_width
1091                &mut query_stats,
1092                &mut indices,
1093                &mut distances,
1094                &mut associated_data,
1095                &|_| true,
1096                false,
1097            );
1098
1099            diskann_providers::test_utils::assert_top_k_exactly_match(
1100                q, &gt, &indices, &distances, top_k,
1101            );
1102        }
1103
1104        Ok(())
1105    }
1106
1107    // Compare that the index built in test is the same as the truth index. The truth index doesn't have associated data, we are only comparing the vector and neighbor data.
1108    pub fn compare_disk_index_graphs(graph_data: &[u8], truth_graph_data: &[u8]) {
1109        let graph_header = GraphHeader::try_from(&graph_data[8..]).unwrap();
1110        let truth_graph_header = GraphHeader::try_from(&truth_graph_data[8..]).unwrap();
1111
1112        let test_node_per_block = graph_header.metadata().num_nodes_per_block;
1113        let test_max_node_length = graph_header.metadata().node_len;
1114
1115        let truth_node_per_block = truth_graph_header.metadata().num_nodes_per_block;
1116        let truth_max_node_length = truth_graph_header.metadata().node_len;
1117
1118        assert_eq!(
1119            graph_header.metadata().num_pts,
1120            truth_graph_header.metadata().num_pts
1121        );
1122
1123        assert_eq!(
1124            graph_header.metadata().dims,
1125            truth_graph_header.metadata().dims
1126        );
1127
1128        let num_pts = graph_header.metadata().num_pts as usize;
1129        let dim = graph_header.metadata().dims;
1130
1131        for idx in 0..num_pts {
1132            let test_node_id_offset = node_data_offset(
1133                idx,
1134                test_max_node_length as usize,
1135                test_node_per_block as usize,
1136                DEFAULT_DISK_SECTOR_LEN,
1137            );
1138
1139            let truth_node_id_offset = node_data_offset(
1140                idx,
1141                truth_max_node_length as usize,
1142                truth_node_per_block as usize,
1143                DEFAULT_DISK_SECTOR_LEN,
1144            );
1145
1146            // Assert that the vector data is the same between the test and truth graphs for this node.
1147            assert_eq!(
1148                &graph_data
1149                    [test_node_id_offset..test_node_id_offset + dim * std::mem::size_of::<f32>()],
1150                &truth_graph_data
1151                    [truth_node_id_offset..truth_node_id_offset + dim * std::mem::size_of::<f32>()]
1152            );
1153
1154            // Assert that the neighbor count is the same between the test and truth graphs for this node.
1155            let test_nbr_cnt_offset = test_node_id_offset + dim * std::mem::size_of::<f32>();
1156            let truth_nbr_cnt_offset = truth_node_id_offset + dim * std::mem::size_of::<f32>();
1157
1158            let test_nbr_count = u32::from_le_bytes([
1159                graph_data[test_nbr_cnt_offset],
1160                graph_data[test_nbr_cnt_offset + 1],
1161                graph_data[test_nbr_cnt_offset + 2],
1162                graph_data[test_nbr_cnt_offset + 3],
1163            ]);
1164
1165            let truth_nbr_count = u32::from_le_bytes([
1166                truth_graph_data[truth_nbr_cnt_offset],
1167                truth_graph_data[truth_nbr_cnt_offset + 1],
1168                truth_graph_data[truth_nbr_cnt_offset + 2],
1169                truth_graph_data[truth_nbr_cnt_offset + 3],
1170            ]);
1171
1172            assert_eq!(test_nbr_count, truth_nbr_count);
1173
1174            // Assert the neighbors (u32) are the same between the test and truth graphs for this node.
1175            let test_nbr_offset = test_nbr_cnt_offset + 4;
1176            let truth_nbr_offset = truth_nbr_cnt_offset + 4;
1177            assert_eq!(
1178                graph_data[test_nbr_offset..test_nbr_offset + test_nbr_count as usize * 4],
1179                truth_graph_data[truth_nbr_offset..truth_nbr_offset + truth_nbr_count as usize * 4]
1180            );
1181        }
1182    }
1183
1184    pub fn node_data_offset(
1185        node_id: usize,
1186        node_length: usize,
1187        nodes_per_block: usize,
1188        block_size: usize,
1189    ) -> usize {
1190        let block_id = node_id / nodes_per_block;
1191        let node_id_in_block = node_id % nodes_per_block;
1192        let offset = block_id * block_size + node_id_in_block * node_length;
1193        offset + block_size
1194    }
1195
1196    fn create_disk_index_builder(
1197        num_points: usize,
1198        dim: usize,
1199        num_of_pq_chunks: usize,
1200        storage_provider: &VirtualStorageProvider<OverlayFS>,
1201        build_quantization_type: QuantizationType,
1202    ) -> ANNResult<
1203        DiskIndexBuilder<'_, GraphDataF32VectorUnitData, VirtualStorageProvider<OverlayFS>>,
1204    > {
1205        let memory_budget = MemoryBudget::try_from_gb(1.0)?;
1206        let num_pq_chunks = NumPQChunks::new_with(num_of_pq_chunks, dim)?;
1207
1208        let build_parameters =
1209            DiskIndexBuildParameters::new(memory_budget, build_quantization_type, num_pq_chunks);
1210
1211        let index_configuration = IndexConfiguration::new(
1212            L2,
1213            dim,
1214            num_points,
1215            ONE,
1216            1,
1217            config::Builder::new_with(4, config::MaxDegree::default_slack(), 50, L2.into(), |b| {
1218                b.saturate_after_prune(true);
1219            })
1220            .build()?,
1221        );
1222
1223        let disk_index_writer = DiskIndexWriter::new(
1224            "data_path".to_string(),
1225            "index_path_prefix".to_string(),
1226            None,
1227            DEFAULT_DISK_SECTOR_LEN,
1228        )?;
1229
1230        DiskIndexBuilder::<GraphDataF32VectorUnitData, VirtualStorageProvider<OverlayFS>>::new(
1231            storage_provider,
1232            build_parameters,
1233            index_configuration,
1234            disk_index_writer,
1235        )
1236    }
1237}
1238
1239#[cfg(test)]
1240mod ram_estimation_tests {
1241    use rstest::rstest;
1242
1243    use super::*;
1244    use crate::QuantizationType;
1245
1246    #[rstest]
1247    #[case(QuantizationType::FP)]
1248    #[case(QuantizationType::PQ { num_chunks: 15 })]
1249    #[case(QuantizationType::SQ { nbits: 1, standard_deviation: None })]
1250    fn test_estimate_build_index_ram_usage(#[case] build_quantization_type: QuantizationType) {
1251        let num_points = 1000;
1252        let dim = 128;
1253        let size_of_t = std::mem::size_of::<f32>() as u64;
1254        let graph_degree = 50;
1255
1256        let single_vec_size = match build_quantization_type {
1257            QuantizationType::FP => dim * size_of_t,
1258            QuantizationType::PQ { num_chunks } => num_chunks as u64,
1259            QuantizationType::SQ { nbits, .. } => {
1260                (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::<f32>() as u64
1261            }
1262        };
1263        let mut expected_ram_usage = (num_points as f64)
1264            * (graph_degree as f64)
1265            * (std::mem::size_of::<u32>() as f64)
1266            * GRAPH_SLACK_FACTOR
1267            + (num_points * single_vec_size) as f64;
1268        expected_ram_usage *= OVERHEAD_FACTOR;
1269
1270        let actual_ram_usage = estimate_build_index_ram_usage(
1271            num_points,
1272            dim,
1273            size_of_t,
1274            graph_degree,
1275            &build_quantization_type,
1276        );
1277
1278        assert_eq!(actual_ram_usage, expected_ram_usage);
1279    }
1280}