1use 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
35const OVERHEAD_FACTOR: f64 = 1.1f64;
37
38#[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 QuantizationType::PQ { num_chunks } => num_chunks as u64,
54 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
63pub 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 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 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 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 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 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 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 let vamana_metadata_size =
252 size_of::<u64>() + size_of::<u32>() + size_of::<u32>() + size_of::<u64>();
253
254 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 let mut max_input_width = 0;
263 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 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 for shard in 0..num_parts {
282 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 medoid = id_maps[shard][medoid as usize];
290
291 if shard == (num_parts - 1) {
293 merged_vamana_cached_writer.write(&medoid.to_le_bytes())?;
295 }
296 }
297
298 let vamana_index_frozen: u64 = 0; 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 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 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 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 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 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 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 builder.checkpoint_record_manager.execute_stage(
513 WorkStage::InMemIndexBuild,
514 WorkStage::PartitionData,
515 || Ok(()),
516 || Ok(()),
517 )?;
518
519 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 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, &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 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 WorkStage::MergeIndices
613 } else {
614 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, time::Instant};
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::storage::{
637 get_compressed_pq_file, get_disk_index_file, get_pq_pivot_file,
638 };
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 aligned_file_reader::VirtualAlignedReaderFactory, disk_provider::DiskIndexSearcher,
655 disk_vertex_provider_factory::DiskVertexProviderFactory,
656 },
657 storage::disk_index_reader::DiskIndexReader,
658 utils::QueryStatistics,
659 };
660 const DEFAULT_DISK_SECTOR_LEN: usize = 4096;
661 pub const TEST_DATA_FILE: &str = "/sift/siftsmall_learn_256pts.fbin";
662 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 pub block_size: usize,
686 }
687
688 impl Default for TestParams {
689 fn default() -> Self {
690 Self {
691 dim: 128, full_dim: 128,
693 max_degree: 4, num_pq_chunks: 128,
695 build_quantization_type: QuantizationType::FP, l_build: 50,
697 data_path: TEST_DATA_FILE.to_string(),
698 index_path_prefix: INDEX_PATH_PREFIX.to_string(),
699 associated_data_path: None,
700 index_build_ram_gb: 1.0,
701 checkpoint_params: None,
702 num_threads: 1,
703 metric: L2,
704 block_size: DEFAULT_DISK_SECTOR_LEN,
705 }
706 }
707 }
708
709 impl TestParams {
710 fn truth_index_path_prefix(&self) -> &str {
712 match (self.max_degree, self.l_build, self.index_build_ram_gb) {
713 (4, 50, 1.0) => TRUTH_INDEX_PATH_PREFIX_R4_L50,
714 (max_degree, l_build, index_build_ram_gb) => panic!(
715 "Truth index path not found for max_degree={}, l_build={}, index_build_ram_gb={}",
716 max_degree, l_build, index_build_ram_gb
717 ),
718 }
719 }
720 pub fn truth_pq_compressed_path(&self) -> String {
721 let prefix = match self.num_pq_chunks {
722 128 => TRUTH_INDEX_PATH_PREFIX_R4_L50,
723 num_pq_chunks => panic!(
724 "Truth pq compressed path not found for num_pq_chunks={}",
725 num_pq_chunks,
726 ),
727 };
728 get_compressed_pq_file(prefix)
729 }
730
731 pub fn pq_compressed_path(&self) -> String {
732 get_compressed_pq_file(&self.index_path_prefix)
733 }
734 }
735
736 pub fn new_vfs() -> VirtualStorageProvider<OverlayFS> {
737 VirtualStorageProvider::new_overlay(test_data_root())
738 }
739
740 pub struct IndexBuildFixture<StorageProvider: StorageReadProvider + StorageWriteProvider> {
741 pub storage_provider: Arc<StorageProvider>,
742 pub params: TestParams,
743 }
744
745 impl<StorageProvider: StorageReadProvider + StorageWriteProvider + 'static>
746 IndexBuildFixture<StorageProvider>
747 {
748 pub fn new(storage_provider: StorageProvider, params: TestParams) -> ANNResult<Self> {
749 Ok(Self {
750 storage_provider: Arc::new(storage_provider),
751 params,
752 })
753 }
754
755 pub fn build<T>(&self) -> ANNResult<()>
756 where
757 T: GraphDataType<VectorIdType = u32>,
758 StorageProvider::Reader: std::marker::Send + Read,
759 {
760 let disk_index_build_parameters = DiskIndexBuildParameters::new(
762 MemoryBudget::try_from_gb(self.params.index_build_ram_gb)?,
763 self.params.build_quantization_type,
764 NumPQChunks::new_with(self.params.num_pq_chunks, self.params.full_dim)?,
765 );
766
767 let config = config::Builder::new_with(
768 self.params.max_degree.into_usize(),
769 config::MaxDegree::default_slack(),
770 self.params.l_build.into_usize(),
771 self.params.metric.into(),
772 |b| {
773 b.saturate_after_prune(true);
774 },
775 )
776 .build()?;
777
778 let metadata =
779 load_metadata_from_file(self.storage_provider.as_ref(), &self.params.data_path)
780 .unwrap();
781
782 assert_eq!(
783 self.params.dim,
784 metadata.ndims(),
785 "Parameters dimension {} and data dimension {} are not equal",
786 self.params.dim,
787 metadata.ndims(),
788 );
789
790 let config = IndexConfiguration::new(
791 self.params.metric,
792 self.params.dim,
793 metadata.npoints(),
794 ONE,
795 self.params.num_threads,
796 config,
797 )
798 .with_pseudo_rng_from_seed(100);
799
800 let disk_index_writer = DiskIndexWriter::new(
801 self.params.data_path.clone(),
802 self.params.index_path_prefix.clone(),
803 self.params.associated_data_path.clone(),
804 self.params.block_size,
805 )?;
806
807 let mut disk_index = match self.params.checkpoint_params {
808 Some(ref checkpoint_params) => {
809 let checkpoint_record_manager =
810 checkpoint_params.checkpoint_record_manager.clone_box();
811 let chunking_config = checkpoint_params.chunking_config.clone();
812 DiskIndexBuilder::<T, _>::new_with_chunking_config(
813 self.storage_provider.as_ref(),
814 disk_index_build_parameters,
815 config,
816 disk_index_writer,
817 chunking_config,
818 checkpoint_record_manager,
819 )
820 }
821 None => DiskIndexBuilder::<T, _>::new(
822 self.storage_provider.as_ref(),
823 disk_index_build_parameters,
824 config,
825 disk_index_writer,
826 ),
827 }?;
828
829 let timer = Instant::now();
830 disk_index.build()?;
831 println!("Indexing time: {} seconds", timer.elapsed().as_secs_f64());
832
833 Ok(())
834 }
835
836 pub fn compare_pq_compressed_files(&self) {
837 self.compare_files(
838 &self.params.pq_compressed_path(),
839 &self.params.truth_pq_compressed_path(),
840 );
841 }
842
843 pub fn assert_index_max_degree<T: GraphDataType>(&self) -> ANNResult<()> {
844 let index_file_path = get_disk_index_file(&self.params.index_path_prefix);
845 let file_data = load_file_to_vec(self.storage_provider.as_ref(), &index_file_path);
846 let graph_header = GraphHeader::try_from(&file_data[8..])?;
847 let max_degree = graph_header.max_degree::<T::VectorDataType>()?;
848 assert_eq!(
849 max_degree, self.params.max_degree as usize,
850 "Max degree mismatch: expected {}, got {}",
851 self.params.max_degree, max_degree
852 );
853
854 Ok(())
855 }
856
857 fn compare_disk_index_with_associated_data(
858 &self,
859 pivot_file_prefix_test: &str,
860 pivot_file_prefix_expected: &str,
861 index_file_suffix: &str,
862 ) {
863 let pq_pivot_path = pivot_file_prefix_test.to_string() + index_file_suffix;
864 let pq_pivot_path_truth = pivot_file_prefix_expected.to_string() + index_file_suffix;
865 let file1 = load_file_to_vec(self.storage_provider.as_ref(), &pq_pivot_path);
866 let file2 = load_file_to_vec(self.storage_provider.as_ref(), &pq_pivot_path_truth);
867 compare_disk_index_graphs(&file1, &file2)
868 }
869
870 pub fn compare_files(&self, file_path1: &str, file_path2: &str) {
871 let file1 = load_file_to_vec(self.storage_provider.as_ref(), file_path1);
872 let file2 = load_file_to_vec(self.storage_provider.as_ref(), file_path2);
873
874 assert_eq!(file1.len(), file2.len());
875 assert_eq!(file1, file2)
876 }
877 }
878
879 fn run_one_shot_test<F>(index_path_prefix: String, params_customizer: F)
881 where
882 F: FnOnce(TestParams) -> TestParams,
883 {
884 let l_build = 64;
885 let max_degree = 16;
886 let top_k = 10;
887 let search_l = 32;
888
889 let base_params = TestParams {
890 l_build,
891 max_degree,
892 index_path_prefix,
893 ..TestParams::default()
894 };
895
896 let params = params_customizer(base_params);
897
898 let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
899 fixture.build::<GraphDataF32VectorUnitData>().unwrap();
900
901 verify_search_result_with_ground_truth::<GraphDataF32VectorUnitData>(
903 &fixture.params,
904 top_k,
905 search_l,
906 &fixture.storage_provider,
907 )
908 .unwrap();
909
910 fixture
911 .assert_index_max_degree::<GraphDataF32VectorUnitData>()
912 .unwrap();
913
914 let mem_index_file_path = format!("{}_mem.index.data", fixture.params.index_path_prefix);
916 assert!(!fixture.storage_provider.exists(&mem_index_file_path));
917 }
918
919 #[rstest]
920 fn test_build_from_iter_one_shot_with_metric(
921 #[values(Metric::L2, Metric::InnerProduct, Metric::Cosine)] metric: Metric,
922 ) {
923 let index_path_prefix = format!("{}_metric_{:?}", INDEX_PATH_PREFIX, metric);
924
925 run_one_shot_test(index_path_prefix, |params| TestParams { metric, ..params });
926 }
927
928 #[test]
934 fn test_build_multi_sector_per_node() {
935 let params = TestParams {
936 block_size: 512,
937 max_degree: 16,
938 l_build: 64,
939 index_path_prefix: format!("{}_multi_sector", INDEX_PATH_PREFIX),
940 ..TestParams::default()
941 };
942 let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
943 fixture.build::<GraphDataF32VectorUnitData>().unwrap();
944 verify_search_result_with_ground_truth::<GraphDataF32VectorUnitData>(
945 &fixture.params,
946 10,
947 32,
948 &fixture.storage_provider,
949 )
950 .unwrap();
951 }
952
953 #[test]
954 fn test_build_from_iter_one_shot_with_associated_data() {
955 let params = TestParams {
957 associated_data_path: Some(
958 "/sift/siftsmall_learn_256pts_u32_associated_data.fbin".to_string(),
959 ),
960 ..TestParams::default()
961 };
962
963 let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
965
966 fixture.build::<GraphDataF32VectorU32Data>().unwrap();
968
969 let mem_index_file_path = format!("{}_mem.index.data", fixture.params.index_path_prefix);
971 let mem_index_associated_data_path = format!(
972 "{}_mem.index.associated_data",
973 fixture.params.index_path_prefix
974 );
975 assert!(!fixture.storage_provider.exists(&mem_index_file_path));
976 assert!(!fixture
977 .storage_provider
978 .exists(&mem_index_associated_data_path));
979
980 fixture.compare_disk_index_with_associated_data(
982 &fixture.params.index_path_prefix,
983 fixture.params.truth_index_path_prefix(),
984 "_disk.index",
985 );
986 }
987
988 #[test]
989 fn test_build_from_iter_merged_index() {
990 let l_build = 64;
992 let max_degree = 16;
993 let top_k = 10;
994 let search_l = 32;
995
996 let index_path_prefix =
997 "/disk_index_build/disk_index_sift_learn_test_disk_index_build_merged".to_string();
998 let params = TestParams {
999 l_build,
1000 max_degree,
1001 index_path_prefix,
1002 index_build_ram_gb: 0.0001, ..TestParams::default()
1004 };
1005
1006 let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
1007
1008 fixture.build::<GraphDataF32VectorUnitData>().unwrap();
1009
1010 verify_search_result_with_ground_truth::<GraphDataF32VectorUnitData>(
1011 &fixture.params,
1012 top_k,
1013 search_l,
1014 &fixture.storage_provider,
1015 )
1016 .unwrap();
1017
1018 fixture
1019 .assert_index_max_degree::<GraphDataF32VectorUnitData>()
1020 .unwrap();
1021 }
1022
1023 #[rstest]
1024 #[case(QuantizationType::SQ { nbits: 2, standard_deviation: None }, "SQ quantization is only supported for 1 bit")]
1025 fn test_build_quantization_type_failure_cases(
1026 #[case] build_quantization_type: QuantizationType,
1027 #[case] error_message: &str,
1028 ) {
1029 let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
1030 let disk_index_builder = create_disk_index_builder(
1031 1000, 128, 128, &storage_provider,
1035 build_quantization_type,
1036 );
1037
1038 let err = disk_index_builder.err().unwrap();
1039 assert!(err.to_string().contains(error_message));
1040 }
1041
1042 fn load_file_to_vec<StorageType: StorageReadProvider>(
1043 storage_provider: &StorageType,
1044 file_path: &str,
1045 ) -> Vec<u8> {
1046 let mut file = storage_provider.open_reader(file_path).unwrap();
1047 let mut buffer = vec![];
1048 file.read_to_end(&mut buffer).unwrap();
1049 buffer
1050 }
1051
1052 pub(crate) fn verify_search_result_with_ground_truth<
1059 G: GraphDataType<VectorIdType = u32, AssociatedDataType = ()>,
1060 >(
1061 params: &TestParams,
1062 top_k: usize,
1063 search_l: u32,
1064 storage_provider: &Arc<VirtualStorageProvider<OverlayFS>>,
1065 ) -> ANNResult<()> {
1066 let pq_pivot_path = get_pq_pivot_file(¶ms.index_path_prefix);
1067 let pq_compressed_path = get_compressed_pq_file(¶ms.index_path_prefix);
1068 let index_file_path = get_disk_index_file(¶ms.index_path_prefix);
1069
1070 let index_reader =
1071 DiskIndexReader::new(pq_pivot_path, pq_compressed_path, storage_provider.as_ref())?;
1072
1073 let vertex_provider_factory = DiskVertexProviderFactory::new(
1074 VirtualAlignedReaderFactory::new(index_file_path, Arc::clone(storage_provider)),
1075 CachingStrategy::None,
1076 )?;
1077
1078 let search_engine = DiskIndexSearcher::<G, DiskVertexProviderFactory<G, _>>::new(
1079 1,
1080 u32::MAX as usize,
1081 &index_reader,
1082 vertex_provider_factory,
1083 params.metric,
1084 None,
1085 )?;
1086
1087 let data =
1088 read_bin::<G::VectorDataType>(&mut storage_provider.open_reader(¶ms.data_path)?)?;
1089 let dim = data.ncols();
1090 let distance = <G::VectorDataType>::distance(params.metric, Some(dim));
1091
1092 for (q, query_data) in data.row_iter().enumerate() {
1099 let gt =
1100 diskann_providers::test_utils::groundtruth(data.as_view(), query_data, |a, b| {
1101 distance.evaluate_similarity(a, b)
1102 });
1103
1104 let mut query_stats = QueryStatistics::default();
1105
1106 let mut indices = vec![0u32; top_k];
1107 let mut distances = vec![0f32; top_k];
1108 let mut associated_data = vec![(); top_k];
1109
1110 _ = search_engine.search_internal(
1111 query_data,
1112 top_k,
1113 search_l,
1114 None, &mut query_stats,
1116 &mut indices,
1117 &mut distances,
1118 &mut associated_data,
1119 &crate::search::search_mode::SearchMode::graph(),
1120 );
1121
1122 diskann_providers::test_utils::assert_top_k_exactly_match(
1123 q, >, &indices, &distances, top_k,
1124 );
1125 }
1126
1127 Ok(())
1128 }
1129
1130 pub fn compare_disk_index_graphs(graph_data: &[u8], truth_graph_data: &[u8]) {
1132 let graph_header = GraphHeader::try_from(&graph_data[8..]).unwrap();
1133 let truth_graph_header = GraphHeader::try_from(&truth_graph_data[8..]).unwrap();
1134
1135 let test_node_per_block = graph_header.metadata().num_nodes_per_block;
1136 let test_max_node_length = graph_header.metadata().node_len;
1137
1138 let truth_node_per_block = truth_graph_header.metadata().num_nodes_per_block;
1139 let truth_max_node_length = truth_graph_header.metadata().node_len;
1140
1141 assert_eq!(
1142 graph_header.metadata().num_pts,
1143 truth_graph_header.metadata().num_pts
1144 );
1145
1146 assert_eq!(
1147 graph_header.metadata().dims,
1148 truth_graph_header.metadata().dims
1149 );
1150
1151 let num_pts = graph_header.metadata().num_pts as usize;
1152 let dim = graph_header.metadata().dims;
1153
1154 for idx in 0..num_pts {
1155 let test_node_id_offset = node_data_offset(
1156 idx,
1157 test_max_node_length as usize,
1158 test_node_per_block as usize,
1159 DEFAULT_DISK_SECTOR_LEN,
1160 );
1161
1162 let truth_node_id_offset = node_data_offset(
1163 idx,
1164 truth_max_node_length as usize,
1165 truth_node_per_block as usize,
1166 DEFAULT_DISK_SECTOR_LEN,
1167 );
1168
1169 assert_eq!(
1171 &graph_data
1172 [test_node_id_offset..test_node_id_offset + dim * std::mem::size_of::<f32>()],
1173 &truth_graph_data
1174 [truth_node_id_offset..truth_node_id_offset + dim * std::mem::size_of::<f32>()]
1175 );
1176
1177 let test_nbr_cnt_offset = test_node_id_offset + dim * std::mem::size_of::<f32>();
1179 let truth_nbr_cnt_offset = truth_node_id_offset + dim * std::mem::size_of::<f32>();
1180
1181 let test_nbr_count = u32::from_le_bytes([
1182 graph_data[test_nbr_cnt_offset],
1183 graph_data[test_nbr_cnt_offset + 1],
1184 graph_data[test_nbr_cnt_offset + 2],
1185 graph_data[test_nbr_cnt_offset + 3],
1186 ]);
1187
1188 let truth_nbr_count = u32::from_le_bytes([
1189 truth_graph_data[truth_nbr_cnt_offset],
1190 truth_graph_data[truth_nbr_cnt_offset + 1],
1191 truth_graph_data[truth_nbr_cnt_offset + 2],
1192 truth_graph_data[truth_nbr_cnt_offset + 3],
1193 ]);
1194
1195 assert_eq!(test_nbr_count, truth_nbr_count);
1196
1197 let test_nbr_offset = test_nbr_cnt_offset + 4;
1199 let truth_nbr_offset = truth_nbr_cnt_offset + 4;
1200 assert_eq!(
1201 graph_data[test_nbr_offset..test_nbr_offset + test_nbr_count as usize * 4],
1202 truth_graph_data[truth_nbr_offset..truth_nbr_offset + truth_nbr_count as usize * 4]
1203 );
1204 }
1205 }
1206
1207 pub fn node_data_offset(
1208 node_id: usize,
1209 node_length: usize,
1210 nodes_per_block: usize,
1211 block_size: usize,
1212 ) -> usize {
1213 let block_id = node_id / nodes_per_block;
1214 let node_id_in_block = node_id % nodes_per_block;
1215 let offset = block_id * block_size + node_id_in_block * node_length;
1216 offset + block_size
1217 }
1218
1219 fn create_disk_index_builder(
1220 num_points: usize,
1221 dim: usize,
1222 num_of_pq_chunks: usize,
1223 storage_provider: &VirtualStorageProvider<OverlayFS>,
1224 build_quantization_type: QuantizationType,
1225 ) -> ANNResult<
1226 DiskIndexBuilder<'_, GraphDataF32VectorUnitData, VirtualStorageProvider<OverlayFS>>,
1227 > {
1228 let memory_budget = MemoryBudget::try_from_gb(1.0)?;
1229 let num_pq_chunks = NumPQChunks::new_with(num_of_pq_chunks, dim)?;
1230
1231 let build_parameters =
1232 DiskIndexBuildParameters::new(memory_budget, build_quantization_type, num_pq_chunks);
1233
1234 let index_configuration = IndexConfiguration::new(
1235 L2,
1236 dim,
1237 num_points,
1238 ONE,
1239 1,
1240 config::Builder::new_with(4, config::MaxDegree::default_slack(), 50, L2.into(), |b| {
1241 b.saturate_after_prune(true);
1242 })
1243 .build()?,
1244 );
1245
1246 let disk_index_writer = DiskIndexWriter::new(
1247 "data_path".to_string(),
1248 "index_path_prefix".to_string(),
1249 None,
1250 DEFAULT_DISK_SECTOR_LEN,
1251 )?;
1252
1253 DiskIndexBuilder::<GraphDataF32VectorUnitData, VirtualStorageProvider<OverlayFS>>::new(
1254 storage_provider,
1255 build_parameters,
1256 index_configuration,
1257 disk_index_writer,
1258 )
1259 }
1260}
1261
1262#[cfg(test)]
1263mod ram_estimation_tests {
1264 use rstest::rstest;
1265
1266 use super::*;
1267 use crate::QuantizationType;
1268
1269 #[rstest]
1270 #[case(QuantizationType::FP)]
1271 #[case(QuantizationType::PQ { num_chunks: 15 })]
1272 #[case(QuantizationType::SQ { nbits: 1, standard_deviation: None })]
1273 fn test_estimate_build_index_ram_usage(#[case] build_quantization_type: QuantizationType) {
1274 let num_points = 1000;
1275 let dim = 128;
1276 let size_of_t = std::mem::size_of::<f32>() as u64;
1277 let graph_degree = 50;
1278
1279 let single_vec_size = match build_quantization_type {
1280 QuantizationType::FP => dim * size_of_t,
1281 QuantizationType::PQ { num_chunks } => num_chunks as u64,
1282 QuantizationType::SQ { nbits, .. } => {
1283 (nbits as u64 * dim).div_ceil(8) + std::mem::size_of::<f32>() as u64
1284 }
1285 };
1286 let mut expected_ram_usage = (num_points as f64)
1287 * (graph_degree as f64)
1288 * (std::mem::size_of::<u32>() as f64)
1289 * GRAPH_SLACK_FACTOR
1290 + (num_points * single_vec_size) as f64;
1291 expected_ram_usage *= OVERHEAD_FACTOR;
1292
1293 let actual_ram_usage = estimate_build_index_ram_usage(
1294 num_points,
1295 dim,
1296 size_of_t,
1297 graph_degree,
1298 &build_quantization_type,
1299 );
1300
1301 assert_eq!(actual_ram_usage, expected_ram_usage);
1302 }
1303}