Skip to main content

diskann_disk/utils/
partition.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5use diskann::{error::IntoANNResult, utils::VectorRepr, ANNError, ANNResult};
6use diskann_providers::storage::{StorageReadProvider, StorageWriteProvider};
7use diskann_providers::{
8    forward_threadpool,
9    utils::{gen_random_slice, AsThreadPool, RayonThreadPool, READ_WRITE_BLOCK_SIZE},
10};
11
12use crate::utils::{compute_closest_centers, k_meanspp_selecting_pivots, run_lloyds};
13use rand::Rng;
14use tracing::info;
15
16use crate::{
17    disk_index_build_parameter::BYTES_IN_GB,
18    storage::{CachedReader, CachedWriter, DiskIndexWriter},
19};
20
21/// Block size for reading/processing large files and matrices in blocks
22const BLOCK_SIZE_LARGE_FILE: u32 = 10_000;
23
24#[allow(clippy::too_many_arguments)]
25pub fn partition_with_ram_budget<T, StorageProvider, Pool, F>(
26    dataset_file: &str,
27    dim: usize,
28    sampling_rate: f64,
29    ram_budget_in_bytes: f64,
30    k_base: usize,
31    merged_index_prefix: &str,
32    storage_provider: &StorageProvider,
33    rng: &mut impl Rng,
34    pool: Pool,
35    ram_estimator: F,
36) -> ANNResult<usize>
37where
38    T: VectorRepr,
39    StorageProvider: StorageReadProvider + StorageWriteProvider,
40    Pool: AsThreadPool,
41    F: Fn(u64, u64) -> f64,
42{
43    forward_threadpool!(pool = pool);
44    // Find partition size and get pivot data
45    let (num_parts, pivot_data, train_dim) = find_partition_size::<T, StorageProvider, F>(
46        dataset_file,
47        sampling_rate,
48        ram_budget_in_bytes,
49        k_base,
50        storage_provider,
51        rng,
52        pool,
53        &ram_estimator,
54    )?;
55
56    info!("Saving shard data into clusters, with only ids");
57
58    shard_data_into_clusters_only_ids::<T, StorageProvider>(
59        dataset_file,
60        &pivot_data,
61        num_parts,
62        dim,
63        train_dim,
64        k_base,
65        merged_index_prefix,
66        storage_provider,
67        pool,
68    )?;
69
70    Ok(num_parts)
71}
72
73#[allow(clippy::too_many_arguments)]
74fn find_partition_size<T, StorageProvider, F>(
75    dataset_file: &str,
76    sampling_rate: f64,
77    ram_budget_in_bytes: f64,
78    k_base: usize,
79    storage_provider: &StorageProvider,
80    rng: &mut impl Rng,
81    pool: &RayonThreadPool,
82    ram_estimator: &F,
83) -> ANNResult<(usize, Vec<f32>, usize)>
84where
85    T: VectorRepr,
86    StorageProvider: StorageReadProvider + StorageWriteProvider,
87    F: Fn(u64, u64) -> f64,
88{
89    const MAX_K_MEANS_REPS: usize = 10;
90
91    let (train_data_float, num_train, train_dim) =
92        gen_random_slice::<T, StorageProvider>(dataset_file, sampling_rate, storage_provider, rng)?;
93    info!("Loaded {} points for train, dim: {}", num_train, train_dim);
94
95    let (test_data_float, num_test, test_dim) =
96        gen_random_slice::<T, StorageProvider>(dataset_file, sampling_rate, storage_provider, rng)?;
97    info!("Loaded {} points for test, dim: {}", num_test, test_dim);
98
99    // Calculate total points accounting for sampling rate
100    let total_points = (num_train as f64 / sampling_rate) as u64;
101    // Get initial partition count estimate
102    let initial_num_parts = estimate_initial_partition_count::<F>(
103        total_points,
104        train_dim as u64,
105        k_base,
106        ram_budget_in_bytes,
107        ram_estimator,
108    );
109
110    let mut num_parts = initial_num_parts;
111    let mut fit_in_ram = false;
112    let mut pivot_data = Vec::new();
113    // Iteratively find the right number of parts, kmeans_partitioning on training data
114    while !fit_in_ram {
115        fit_in_ram = true;
116
117        let mut max_ram_usage_in_bytes = 0.0;
118
119        pivot_data = vec![0.0; num_parts * train_dim];
120
121        // Process Global k-means for kmeans_partitioning Step
122        info!("Processing global k-means (kmeans_partitioning Step)");
123        k_meanspp_selecting_pivots(
124            &train_data_float,
125            num_train,
126            train_dim,
127            &mut pivot_data,
128            num_parts,
129            rng,
130            &mut (false),
131            pool,
132        )?;
133
134        run_lloyds(
135            &train_data_float,
136            num_train,
137            train_dim,
138            &mut pivot_data,
139            num_parts,
140            MAX_K_MEANS_REPS,
141            &mut (false),
142            pool,
143        )?;
144
145        // now pivots are ready. need to stream base points and assign them to closest clusters.
146
147        let mut cluster_sizes = Vec::new();
148        estimate_cluster_sizes(
149            &test_data_float,
150            num_test,
151            &pivot_data,
152            num_parts,
153            test_dim,
154            k_base,
155            &mut cluster_sizes,
156            pool,
157        )?;
158
159        let mut partition_stats = Vec::with_capacity(num_parts);
160        for p in &cluster_sizes {
161            // to account for the fact that p is the size of the shard over the testing sample.
162            let p = (*p as f64 / sampling_rate) as u64;
163            let cur_shard_ram_estimate_in_bytes = ram_estimator(p, train_dim as u64);
164            partition_stats.push((p, cur_shard_ram_estimate_in_bytes));
165
166            if cur_shard_ram_estimate_in_bytes > max_ram_usage_in_bytes {
167                max_ram_usage_in_bytes = cur_shard_ram_estimate_in_bytes;
168            }
169        }
170
171        info!(
172            "Partition RAM estimates (GB): {}",
173            partition_stats
174                .iter()
175                .map(|(size, ram)| format!("#{}: {:.2}", size, ram / BYTES_IN_GB))
176                .collect::<Vec<_>>()
177                .join(", ")
178        );
179
180        info!(
181            "With {} parts, max estimated RAM usage: {:.2} GB, budget given is {:.2} GB",
182            num_parts,
183            max_ram_usage_in_bytes / BYTES_IN_GB,
184            ram_budget_in_bytes / BYTES_IN_GB
185        );
186        if max_ram_usage_in_bytes > ram_budget_in_bytes {
187            fit_in_ram = false;
188            num_parts += 2;
189        } else {
190            info!(
191                "Found optimal partition count: [parts={}, initial={}, max_ram={:.2}GB, budget={:.2}GB]",
192                num_parts,
193                initial_num_parts,
194                max_ram_usage_in_bytes / BYTES_IN_GB,
195                ram_budget_in_bytes / BYTES_IN_GB
196            );
197        }
198    }
199
200    Ok((num_parts, pivot_data, train_dim))
201}
202
203/// Initial estimation of partition count based on dataset characteristics and RAM budget
204fn estimate_initial_partition_count<F>(
205    total_points: u64,
206    dimension: u64,
207    k_base: usize,
208    ram_budget_in_bytes: f64,
209    ram_estimator: &F,
210) -> usize
211where
212    F: Fn(u64, u64) -> f64,
213{
214    // Calculate total RAM needed without partitioning
215    let total_ram_estimate = ram_estimator(total_points * k_base as u64, dimension);
216
217    let mut partition_count = (total_ram_estimate / ram_budget_in_bytes).ceil() as usize;
218
219    // Ensure minimum of 3 partitions and odd number for balanced splitting
220    partition_count = std::cmp::max(3, partition_count);
221    if partition_count.is_multiple_of(2) {
222        partition_count += 1;
223    }
224
225    info!(
226        "Estimated initial partition count: {} (total points: {}, dimension: {}, k_base: {}, total_ram_estimate: {:.2} GB, ram_budget: {:.2} GB)",
227        partition_count,
228        total_points,
229        dimension,
230        k_base,
231        total_ram_estimate / BYTES_IN_GB,
232        ram_budget_in_bytes / BYTES_IN_GB
233    );
234
235    partition_count
236}
237
238#[allow(clippy::too_many_arguments)]
239fn shard_data_into_clusters_only_ids<T, StorageProvider>(
240    dataset_file: &str,
241    pivot_data: &[f32],
242    num_parts: usize,
243    dim: usize,
244    full_dim: usize,
245    k_base: usize,
246    merged_index_prefix: &str,
247    storage_provider: &StorageProvider,
248    pool: &RayonThreadPool,
249) -> ANNResult<()>
250where
251    T: VectorRepr,
252    StorageProvider: StorageReadProvider + StorageWriteProvider,
253{
254    let mut dataset_reader = CachedReader::<StorageProvider>::new(
255        dataset_file,
256        READ_WRITE_BLOCK_SIZE,
257        storage_provider,
258    )?;
259    let num_points = dataset_reader.read_u32()?;
260    let base_dim = dataset_reader.read_u32()?;
261    if base_dim != dim as u32 {
262        return Err(ANNError::log_index_error(
263            "dimensions dont match for train set and base set",
264        ));
265    }
266
267    let mut shard_counts = vec![0; num_parts];
268    let shard_idmaps_names = (0..num_parts)
269        .map(|shard| {
270            DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard)
271        })
272        .collect::<Vec<String>>();
273
274    // 8KB cache for small ID map files - matches default BufWriter size
275    const WRITE_ID_CACHE_SIZE: u64 = 8 * 1024;
276    let mut shard_idmap_cached_writers = Vec::new();
277    for name in &shard_idmaps_names {
278        let writer = storage_provider.create_for_write(name)?;
279        let cached_writer =
280            CachedWriter::<StorageProvider>::new(name, WRITE_ID_CACHE_SIZE, writer)?;
281        shard_idmap_cached_writers.push(cached_writer);
282    }
283
284    let dummy_size: u32 = 0;
285    let const_one: u32 = 1;
286    for writer in shard_idmap_cached_writers.iter_mut() {
287        writer.write(&dummy_size.to_le_bytes())?;
288        writer.write(&const_one.to_le_bytes())?;
289    }
290
291    let block_size = if num_points <= BLOCK_SIZE_LARGE_FILE {
292        num_points
293    } else {
294        BLOCK_SIZE_LARGE_FILE
295    };
296
297    let num_blocks = num_points.div_ceil(block_size);
298
299    let mut block_closest_centers = vec![0u32; block_size as usize * k_base];
300    let mut block_data_t: Vec<u8> = vec![0; block_size as usize * dim * std::mem::size_of::<T>()];
301    let mut block_data_float: Vec<f32> = vec![0.0; full_dim * block_size as usize];
302
303    for block in 0..num_blocks {
304        let start_id = (block * block_size) as usize;
305        let end_id = std::cmp::min((block + 1) * block_size, num_points) as usize;
306        let cur_blk_size = end_id - start_id;
307
308        dataset_reader.read(&mut block_data_t[..cur_blk_size * dim * std::mem::size_of::<T>()])?;
309
310        // convert data from type T to f32
311        let cur_vector_t: &[T] =
312            bytemuck::cast_slice(&block_data_t[..cur_blk_size * dim * std::mem::size_of::<T>()]);
313
314        for (v, dst) in cur_vector_t
315            .chunks_exact(dim)
316            .zip(block_data_float.chunks_exact_mut(full_dim))
317        {
318            T::as_f32_into(v, dst).into_ann_result()?;
319        }
320
321        compute_closest_centers(
322            &block_data_float[..full_dim * cur_blk_size],
323            cur_blk_size,
324            full_dim,
325            pivot_data,
326            num_parts,
327            k_base,
328            &mut block_closest_centers,
329            None,
330            None,
331            pool,
332        )?;
333
334        for p in 0..cur_blk_size {
335            for p1 in 0..k_base {
336                let shard_id = block_closest_centers[p * k_base + p1] as usize;
337                let original_point_map_id = (start_id + p) as u32;
338                shard_idmap_cached_writers[shard_id].write(&original_point_map_id.to_le_bytes())?;
339                shard_counts[shard_id] += 1;
340            }
341        }
342    }
343
344    let mut total_count = 0;
345
346    for i in 0..num_parts {
347        let cur_shard_count = shard_counts[i] as u32;
348        info!(" shard_{} with npts : {} ", i, cur_shard_count);
349        total_count += cur_shard_count;
350        shard_idmap_cached_writers[i].reset()?;
351        shard_idmap_cached_writers[i].write(&cur_shard_count.to_le_bytes())?;
352        shard_idmap_cached_writers[i].flush()?;
353    }
354
355    info!(
356        "Partitioned {} with replication factor {} to get {} points across {} shards",
357        num_points, k_base, total_count, num_parts
358    );
359
360    Ok(())
361}
362
363#[allow(clippy::too_many_arguments)]
364fn estimate_cluster_sizes(
365    data_float: &[f32],
366    num_pts: usize,
367    pivot_data: &[f32],
368    num_centers: usize,
369    dim: usize,
370    k_base: usize,
371    cluster_sizes: &mut Vec<u32>,
372    pool: &RayonThreadPool,
373) -> ANNResult<()> {
374    cluster_sizes.clear();
375    let mut shard_counts = vec![0; num_centers];
376
377    let block_size = if num_pts <= BLOCK_SIZE_LARGE_FILE as usize {
378        num_pts
379    } else {
380        BLOCK_SIZE_LARGE_FILE as usize
381    };
382
383    let mut block_closest_centers = vec![0; block_size * k_base];
384
385    let num_blocks = num_pts.div_ceil(block_size);
386
387    for block in 0..num_blocks {
388        let start_id = block * block_size;
389        let end_id = std::cmp::min((block + 1) * block_size, num_pts);
390        let cur_blk_size = end_id - start_id;
391
392        let block_data_float = &data_float[start_id * dim..(start_id + cur_blk_size) * dim];
393
394        compute_closest_centers(
395            block_data_float,
396            cur_blk_size,
397            dim,
398            pivot_data,
399            num_centers,
400            k_base,
401            &mut block_closest_centers,
402            None,
403            None,
404            pool,
405        )?;
406
407        for p in 0..cur_blk_size {
408            for p1 in 0..k_base {
409                let shard_id = block_closest_centers[p * k_base + p1] as usize;
410                shard_counts[shard_id] += 1;
411            }
412        }
413    }
414
415    (0..num_centers).for_each(|i| {
416        let cur_shard_count = shard_counts[i] as u32;
417        cluster_sizes.push(cur_shard_count);
418    });
419    info!("Estimated cluster sizes: {:?}", cluster_sizes);
420    Ok(())
421}
422
423#[cfg(test)]
424mod partition_test {
425    use std::io::Read;
426
427    use diskann_providers::storage::VirtualStorageProvider;
428    use diskann_providers::utils::create_thread_pool_for_test;
429    use diskann_utils::test_data_root;
430    use vfs::{MemoryFS, OverlayFS};
431
432    use super::*;
433
434    #[test]
435    fn test_estimate_cluster_sizes() {
436        let data_float = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
437        let num_pts = 3;
438        let pivot_data = &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
439        let num_centers = 3;
440        let dim = 2;
441        let k_base = 2;
442        let mut cluster_sizes = vec![];
443        let pool = create_thread_pool_for_test();
444
445        estimate_cluster_sizes(
446            &data_float,
447            num_pts,
448            pivot_data,
449            num_centers,
450            dim,
451            k_base,
452            &mut cluster_sizes,
453            &pool,
454        )
455        .unwrap();
456
457        assert_eq!(cluster_sizes.len(), num_centers);
458        assert_eq!(cluster_sizes, &[2, 3, 1]);
459    }
460
461    #[test]
462    fn test_shard_data_into_clusters_only_ids() {
463        // create a temporary file for the dataset
464        let dataset_path = "/dataset_file";
465        // write some dummy data to the dataset file
466        let mut data_float = Vec::new();
467        let num_points: u32 = 100;
468        let dim: usize = 10;
469
470        let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
471        {
472            let writer = storage_provider.create_for_write(dataset_path).unwrap();
473            let mut dataset_writer = CachedWriter::<VirtualStorageProvider<MemoryFS>>::new(
474                dataset_path,
475                READ_WRITE_BLOCK_SIZE,
476                writer,
477            )
478            .unwrap();
479            dataset_writer.write(&num_points.to_le_bytes()).unwrap();
480            dataset_writer.write(&dim.to_le_bytes()).unwrap();
481            for i in 0..num_points {
482                for j in 0..dim {
483                    let val = (i * dim as u32 + j as u32) as f32;
484                    data_float.push(val);
485                    dataset_writer.write(&val.to_le_bytes()).unwrap();
486                }
487            }
488        }
489
490        // create some dummy pivot data
491        let k_base: usize = 2;
492        let num_parts = 3;
493
494        // generate pivot data
495        let pivot_data: [f32; 30] = [
496            820.0, 821.0, 822.0, 823.0, 824.0, 825.0, 826.0, 827.0, 828.0, 829.0, 155.0, 156.0,
497            157.0, 158.0, 159.0, 160.0, 161.0, 162.0, 163.0, 164.0, 480.0, 481.0, 482.0, 483.0,
498            484.0, 485.0, 486.0, 487.0, 488.0, 489.0,
499        ];
500
501        // create a temporary prefix for the merged index prefix
502        let merged_index_prefix = "/merged_index";
503        let pool = create_thread_pool_for_test();
504        // call the function being tested
505        shard_data_into_clusters_only_ids::<f32, VirtualStorageProvider<OverlayFS>>(
506            dataset_path,
507            &pivot_data,
508            num_parts,
509            dim,
510            dim,
511            k_base,
512            merged_index_prefix,
513            &storage_provider,
514            &pool,
515        )
516        .unwrap();
517
518        // check that the output is as expected
519        let expected_prefix = "/partition/id_maps/merged_index_expected";
520        for shard in 0..num_parts {
521            let path1 =
522                DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, shard);
523            let path2 =
524                DiskIndexWriter::get_merged_index_subshard_id_map_file(expected_prefix, shard);
525            let file1 =
526                load_file_to_vec::<VirtualStorageProvider<OverlayFS>>(&path1, &storage_provider);
527            let file2 =
528                load_file_to_vec::<VirtualStorageProvider<OverlayFS>>(&path2, &storage_provider);
529
530            assert_eq!(file1.len(), file2.len());
531            assert_eq!(file1[..], file2[..]);
532
533            // clean up the temporary files and directory
534            storage_provider.delete(&path1).unwrap();
535        }
536
537        storage_provider.delete(dataset_path).unwrap();
538    }
539
540    fn load_file_to_vec<StorageProvider>(
541        file_path: &str,
542        storage_provider: &StorageProvider,
543    ) -> Vec<u8>
544    where
545        StorageProvider: StorageReadProvider,
546    {
547        let mut file = storage_provider.open_reader(file_path).unwrap();
548        let mut buffer = vec![];
549        file.read_to_end(&mut buffer).unwrap();
550        buffer
551    }
552
553    #[test]
554    fn test_partition_with_ram_budget() -> ANNResult<()> {
555        let storage_provider = VirtualStorageProvider::new_overlay(test_data_root());
556        let dataset_file = "/sift/siftsmall_learn.bin";
557        let mut file = storage_provider.open_reader(dataset_file).unwrap();
558        let mut data = vec![];
559        file.read_to_end(&mut data).unwrap();
560
561        let sampling_rate = 1.0;
562        let ram_budget_in_bytes = 15_000_000.0;
563        let max_degree = 64;
564        let k_base = 2;
565        let merged_index_prefix = "/test_merged_index_prefix";
566        let pool = create_thread_pool_for_test();
567
568        let num_parts = partition_with_ram_budget::<f32, _, _, _>(
569            dataset_file,
570            128, //sift is 128 dimensions
571            sampling_rate,
572            ram_budget_in_bytes,
573            k_base,
574            merged_index_prefix,
575            &storage_provider,
576            &mut diskann_providers::utils::create_rnd_in_tests(),
577            &pool,
578            |num_points, dim| {
579                // Simple RAM estimation for test - capture datasize and graph_degree from context
580                use diskann_providers::model::GRAPH_SLACK_FACTOR;
581
582                let datasize = std::mem::size_of::<f32>() as u64;
583                let graph_degree = max_degree as u64;
584                let dataset_size = (num_points * dim.next_multiple_of(8u64) * datasize) as f64;
585                let graph_size = (num_points * graph_degree * 4) as f64 * GRAPH_SLACK_FACTOR;
586                1.1 * (dataset_size + graph_size)
587            },
588        )?;
589
590        assert!(num_parts >= 3);
591
592        for i in 0..num_parts {
593            let idmap_filename =
594                DiskIndexWriter::get_merged_index_subshard_id_map_file(merged_index_prefix, i);
595            storage_provider.delete(&idmap_filename)?;
596        }
597
598        Ok(())
599    }
600}