1use 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
21const 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 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 let total_points = (num_train as f64 / sampling_rate) as u64;
101 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 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 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 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 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
203fn 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 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 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 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 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 let dataset_path = "/dataset_file";
465 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 let k_base: usize = 2;
492 let num_parts = 3;
493
494 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 let merged_index_prefix = "/merged_index";
503 let pool = create_thread_pool_for_test();
504 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 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 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, 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 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}