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, RayonThreadPool, 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: &'a RayonThreadPool,
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: &'a RayonThreadPool,
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};
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 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, full_dim: 128,
692 max_degree: 4, num_pq_chunks: 128,
694 build_quantization_type: QuantizationType::FP, 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 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 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 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 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 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 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 let fixture = IndexBuildFixture::new(new_vfs(), params).unwrap();
938
939 fixture.build::<GraphDataF32VectorU32Data>().unwrap();
941
942 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 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 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, ..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, 128, 128, &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 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(¶ms.index_path_prefix);
1040 let pq_compressed_path = get_compressed_pq_file(¶ms.index_path_prefix);
1041 let index_file_path = get_disk_index_file(¶ms.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(¶ms.data_path)?)?;
1065 let dim = data.ncols();
1066 let distance = <G::VectorDataType>::distance(params.metric, Some(dim));
1067
1068 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, &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, >, &indices, &distances, top_k,
1101 );
1102 }
1103
1104 Ok(())
1105 }
1106
1107 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_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 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 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}