1#![warn(missing_debug_implementations, missing_docs)]
6
7use std::{cmp::Ordering, collections::BinaryHeap};
14
15use diskann::{ANNError, ANNErrorKind, ANNResult};
16use diskann_linalg::{self, Transpose};
17use diskann_providers::utils::{ParallelIteratorInPool, RayonThreadPoolRef};
18use rayon::prelude::*;
19
20const POINTS_PER_CHUNK: usize = 1200;
42
43struct PivotContainer {
44 piv_id: usize,
45 piv_dist: f32,
46}
47
48impl PartialOrd for PivotContainer {
51 fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
52 Some(self.cmp(other))
53 }
54}
55
56impl Ord for PivotContainer {
58 fn cmp(&self, other: &Self) -> std::cmp::Ordering {
59 other
62 .piv_dist
63 .partial_cmp(&self.piv_dist)
64 .unwrap_or(Ordering::Less)
65 }
66}
67
68impl PartialEq for PivotContainer {
69 fn eq(&self, other: &Self) -> bool {
70 self.piv_dist == other.piv_dist
71 }
72}
73
74impl Eq for PivotContainer {}
75
76fn compute_vec_l2sq(data: &[f32], index: usize, dim: usize) -> f32 {
78 let start = index * dim;
79 let slice = unsafe { std::slice::from_raw_parts(data.as_ptr().add(start), dim) };
80 let mut sum_squared = 0.0;
81 for &value in slice {
82 sum_squared += value * value;
83 }
84
85 sum_squared
86}
87
88pub fn compute_vecs_l2sq(
91 vecs_l2sq: &mut [f32],
92 data: &[f32],
93 dim: usize,
94 pool: RayonThreadPoolRef<'_>,
95) -> ANNResult<()> {
96 let expected_data_len = vecs_l2sq.len().checked_mul(dim).ok_or_else(|| {
97 ANNError::log_index_error(format_args!(
98 "vecs_l2sq.len() * dim overflowed: vecs_l2sq.len() ({}) * dim ({})",
99 vecs_l2sq.len(),
100 dim
101 ))
102 })?;
103 if data.len() != expected_data_len {
104 return Err(ANNError::log_index_error(format_args!(
105 "data.len() ({}) should be vecs_l2sq.len() ({}) * dim ({})",
106 data.len(),
107 vecs_l2sq.len(),
108 dim
109 )));
110 }
111
112 if dim < 5 {
113 for (i, vec_l2sq) in vecs_l2sq.iter_mut().enumerate() {
114 *vec_l2sq = compute_vec_l2sq(data, i, dim);
115 }
116 } else {
117 vecs_l2sq
118 .par_iter_mut()
119 .enumerate()
120 .for_each_in_pool(pool, |(i, vec_l2sq)| {
121 *vec_l2sq = compute_vec_l2sq(data, i, dim);
122 });
123 }
124
125 Ok(())
126}
127
128#[allow(clippy::too_many_arguments)]
137pub fn compute_closest_centers_in_block(
138 data: &[f32],
139 num_points: usize,
140 dim: usize,
141 centers: &[f32],
142 num_centers: usize,
143 docs_l2sq: &[f32],
144 centers_l2sq: &[f32],
145 center_index: &mut [u32],
146 dist_matrix: &mut [f32],
147 k: usize,
148 pool: RayonThreadPoolRef<'_>,
149) -> ANNResult<()> {
150 if k > num_centers {
151 return Err(ANNError::log_index_error(format_args!(
152 "k ({}) should be equal or less than num_centers ({})",
153 k, num_centers
154 )));
155 }
156
157 let ones_a: Vec<f32> = vec![1.0; num_centers];
158 let ones_b: Vec<f32> = vec![1.0; num_points];
159
160 diskann_linalg::sgemm(
161 Transpose::None,
162 Transpose::Ordinary,
163 num_points,
164 num_centers,
165 1,
166 1.0,
167 docs_l2sq,
168 &ones_a,
169 None, dist_matrix,
171 )
172 .map_err(|e| ANNError::new(ANNErrorKind::IndexError, e))?;
173
174 diskann_linalg::sgemm(
175 Transpose::None,
176 Transpose::Ordinary,
177 num_points,
178 num_centers,
179 1,
180 1.0,
181 &ones_b,
182 centers_l2sq,
183 Some(1.0), dist_matrix,
185 )
186 .map_err(|e| ANNError::new(ANNErrorKind::IndexError, e))?;
187
188 diskann_linalg::sgemm(
189 Transpose::None,
190 Transpose::Ordinary,
191 num_points,
192 num_centers,
193 dim,
194 -2.0,
195 data,
196 centers,
197 Some(1.0), dist_matrix,
199 )
200 .map_err(|e| ANNError::new(ANNErrorKind::IndexError, e))?;
201
202 if k == 1 {
203 center_index
204 .par_iter_mut()
205 .enumerate()
206 .for_each_in_pool(pool, |(i, center_idx)| {
207 let mut min = f32::MAX;
208 let current = &dist_matrix[i * num_centers..(i + 1) * num_centers];
209 let mut min_idx = 0;
210 for (j, &distance) in current.iter().enumerate() {
211 if distance < min {
212 min = distance;
213 min_idx = j;
214 }
215 }
216 *center_idx = min_idx as u32;
217 });
218 } else {
219 center_index
220 .par_chunks_mut(k)
221 .enumerate()
222 .for_each_in_pool(pool, |(i, center_chunk)| {
223 let current = &dist_matrix[i * num_centers..(i + 1) * num_centers];
224 let mut top_k_queue = BinaryHeap::new();
225 for (j, &distance) in current.iter().enumerate() {
226 let this_piv = PivotContainer {
227 piv_id: j,
228 piv_dist: distance,
229 };
230 top_k_queue.push(this_piv);
231 }
232 for center_idx in center_chunk.iter_mut() {
233 if let Some(this_piv) = top_k_queue.pop() {
234 *center_idx = this_piv.piv_id as u32;
235 } else {
236 break;
237 }
238 }
239 });
240 }
241
242 Ok(())
243}
244
245#[allow(clippy::too_many_arguments)]
254pub fn compute_closest_centers(
255 data: &[f32],
256 num_points: usize,
257 dim: usize,
258 pivot_data: &[f32],
259 num_centers: usize,
260 k: usize,
261 closest_centers_ivf: &mut [u32],
262 mut inverted_index: Option<&mut Vec<Vec<usize>>>,
263 pts_norms_squared: Option<&[f32]>,
264 pool: RayonThreadPoolRef<'_>,
265) -> ANNResult<()> {
266 if k > num_centers {
267 return Err(ANNError::log_index_error(format_args!(
268 "k ({}) should be equal or less than num_centers ({})",
269 k, num_centers
270 )));
271 }
272
273 let expected_data_len = num_points.checked_mul(dim).ok_or_else(|| {
275 ANNError::log_index_error(format_args!(
276 "num_points * dim overflowed: num_points ({}) * dim ({})",
277 num_points, dim
278 ))
279 })?;
280
281 if data.len() != expected_data_len {
282 return Err(ANNError::log_index_error(format_args!(
283 "data.len() ({}) should equal num_points ({}) * dim ({})",
284 data.len(),
285 num_points,
286 dim
287 )));
288 }
289
290 let expected_pivot_len = num_centers.checked_mul(dim).ok_or_else(|| {
292 ANNError::log_index_error(format_args!(
293 "num_centers * dim overflowed: num_centers ({}) * dim ({})",
294 num_centers, dim
295 ))
296 })?;
297
298 if pivot_data.len() != expected_pivot_len {
299 return Err(ANNError::log_index_error(format_args!(
300 "pivot_data.len() ({}) should equal num_centers ({}) * dim ({})",
301 pivot_data.len(),
302 num_centers,
303 dim
304 )));
305 }
306
307 let expected_closest_centers_len = num_points.checked_mul(k).ok_or_else(|| {
308 ANNError::log_index_error(format_args!(
309 "num_points * k overflowed: num_points ({}) * k ({})",
310 num_points, k
311 ))
312 })?;
313
314 if closest_centers_ivf.len() != expected_closest_centers_len {
315 return Err(ANNError::log_index_error(format_args!(
316 "closest_centers_ivf.len() ({}) should equal num_points ({}) * k ({})",
317 closest_centers_ivf.len(),
318 num_points,
319 k
320 )));
321 }
322
323 let mut owned_pts_norms_squared;
324 let pts_norms_squared: &[f32] = if let Some(pts_norms) = pts_norms_squared {
325 if pts_norms.len() != num_points {
326 return Err(ANNError::log_index_error(format_args!(
327 "pts_norms_squared.len() ({}) should equal num_points ({})",
328 pts_norms.len(),
329 num_points
330 )));
331 }
332 pts_norms
333 } else {
334 owned_pts_norms_squared = vec![0.0; num_points];
335 compute_vecs_l2sq(&mut owned_pts_norms_squared, data, dim, pool)?;
336 &owned_pts_norms_squared
337 };
338
339 let mut pivs_norms_squared = vec![0.0; num_centers];
340 compute_vecs_l2sq(&mut pivs_norms_squared, pivot_data, dim, pool)?;
341
342 let mut distance_matrix = vec![0.0; POINTS_PER_CHUNK * num_centers];
343 let mut closest_center_indices = vec![0; POINTS_PER_CHUNK * k];
344 let pts_norms_squared_chunks = pts_norms_squared.chunks(POINTS_PER_CHUNK);
345
346 for (chunk_index, (data_chunk, pts_norms_squared_chunk)) in data
347 .chunks(dim * POINTS_PER_CHUNK)
348 .zip(pts_norms_squared_chunks)
349 .enumerate()
350 {
351 let chunk_size = data_chunk.len() / dim;
353
354 let this_distance_matrix = &mut distance_matrix[..num_centers * chunk_size];
356 let this_closest_center_indices = &mut closest_center_indices[..k * chunk_size];
357
358 compute_closest_centers_in_block(
359 data_chunk,
360 chunk_size,
361 dim,
362 pivot_data,
363 num_centers,
364 pts_norms_squared_chunk,
365 &pivs_norms_squared,
366 this_closest_center_indices,
367 this_distance_matrix,
368 k,
369 pool,
370 )?;
371
372 let point_start_index = chunk_index * POINTS_PER_CHUNK;
373
374 for point_index in point_start_index..point_start_index + chunk_size {
375 for l in 0..k {
376 let center_chunk_index = (point_index - point_start_index) * k + l;
377 let ivf_index = point_index * k + l;
378
379 let this_center_index = closest_center_indices[center_chunk_index];
380 closest_centers_ivf[ivf_index] = this_center_index;
381
382 if let Some(inverted_index) = &mut inverted_index {
383 inverted_index[this_center_index as usize].push(point_index);
384 }
385 }
386 }
387 }
388 Ok(())
389}
390
391#[cfg(test)]
392mod math_util_test {
393 use approx::assert_abs_diff_eq;
394
395 use super::*;
396 use diskann_providers::utils::create_thread_pool_for_test;
397
398 #[test]
399 fn partial_ord_test() {
400 let pviot1 = PivotContainer {
401 piv_id: 2,
402 piv_dist: f32::NAN,
403 };
404 let pivot2 = PivotContainer {
405 piv_id: 1,
406 piv_dist: 1.0,
407 };
408
409 assert_eq!(pviot1.partial_cmp(&pivot2), Some(Ordering::Less));
410 }
411
412 #[test]
413 fn ord_test() {
414 let pviot1 = PivotContainer {
415 piv_id: 1,
416 piv_dist: f32::NAN,
417 };
418 let pivot2 = PivotContainer {
419 piv_id: 2,
420 piv_dist: 1.0,
421 };
422
423 assert_eq!(pviot1.cmp(&pivot2), Ordering::Less);
424 }
425
426 #[test]
427 fn compute_vecs_l2sq_small_dim_test() {
428 let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
429 let num_points = 2;
430 let dim = 3;
431 let mut vecs_l2sq = vec![0.0; num_points];
432 let pool = create_thread_pool_for_test();
433
434 compute_vecs_l2sq(&mut vecs_l2sq, &data, dim, pool.as_ref()).unwrap();
435
436 let expected = [14.0, 77.0];
437
438 assert_eq!(vecs_l2sq.len(), num_points);
439 assert_abs_diff_eq!(vecs_l2sq[0], expected[0], epsilon = 1e-6);
440 assert_abs_diff_eq!(vecs_l2sq[1], expected[1], epsilon = 1e-6);
441 }
442
443 #[test]
444 fn compute_vecs_l2sq_large_dim_test() {
445 let data = vec![
446 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0,
447 ];
448 let num_points = 2;
449 let dim = 8;
450 let mut vecs_l2sq = vec![0.0; num_points];
451 let pool = create_thread_pool_for_test();
452 compute_vecs_l2sq(&mut vecs_l2sq, &data, dim, pool.as_ref()).unwrap();
453
454 let expected = [204.0, 1292.0];
455
456 assert_eq!(vecs_l2sq.len(), num_points);
457 assert_abs_diff_eq!(vecs_l2sq[0], expected[0], epsilon = 1e-6);
458 assert_abs_diff_eq!(vecs_l2sq[1], expected[1], epsilon = 1e-6);
459 }
460
461 #[test]
462 fn compute_closest_centers_in_block_test() {
463 let num_points = 10;
464 let dim = 5;
465 let num_centers = 3;
466 let data = vec![
467 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, 13.0, 14.0, 15.0, 16.0,
468 17.0, 18.0, 19.0, 20.0, 21.0, 22.0, 23.0, 24.0, 25.0, 26.0, 27.0, 28.0, 29.0, 30.0,
469 31.0, 32.0, 33.0, 34.0, 35.0, 36.0, 37.0, 38.0, 39.0, 40.0, 41.0, 42.0, 43.0, 44.0,
470 45.0, 46.0, 47.0, 48.0, 49.0, 50.0,
471 ];
472 let centers = vec![
473 1.0, 2.0, 3.0, 4.0, 5.0, 21.0, 22.0, 23.0, 24.0, 25.0, 31.0, 32.0, 33.0, 34.0, 35.0,
474 ];
475 let mut docs_l2sq = vec![0.0; num_points];
476 let pool = create_thread_pool_for_test();
477 compute_vecs_l2sq(&mut docs_l2sq, &data, dim, pool.as_ref()).unwrap();
478 let mut centers_l2sq = vec![0.0; num_centers];
479 compute_vecs_l2sq(&mut centers_l2sq, ¢ers, dim, pool.as_ref()).unwrap();
480 let mut center_index = vec![0; num_points];
481 let mut dist_matrix = vec![0.0; num_points * num_centers];
482 let k = 1;
483
484 compute_closest_centers_in_block(
485 &data,
486 num_points,
487 dim,
488 ¢ers,
489 num_centers,
490 &docs_l2sq,
491 ¢ers_l2sq,
492 &mut center_index,
493 &mut dist_matrix,
494 k,
495 pool.as_ref(),
496 )
497 .unwrap();
498
499 assert_eq!(center_index.len(), num_points);
500 let expected_center_index = vec![0, 0, 0, 1, 1, 1, 2, 2, 2, 2];
501 assert_abs_diff_eq!(*center_index, expected_center_index);
502
503 assert_eq!(dist_matrix.len(), num_points * num_centers);
504 let expected_dist_matrix = vec![
505 0.0, 2000.0, 4500.0, 125.0, 1125.0, 3125.0, 500.0, 500.0, 2000.0, 1125.0, 125.0,
506 1125.0, 2000.0, 0.0, 500.0, 3125.0, 125.0, 125.0, 4500.0, 500.0, 0.0, 6125.0, 1125.0,
507 125.0, 8000.0, 2000.0, 500.0, 10125.0, 3125.0, 1125.0,
508 ];
509 assert_abs_diff_eq!(*dist_matrix, expected_dist_matrix, epsilon = 1e-2);
510 }
511
512 #[test]
513 fn compute_closest_centers_in_block_test_k_equals_two() {
514 let num_points = 2;
515 let dim = 5;
516 let num_centers = 4;
517 let data = vec![41.0, 42.0, 43.0, 44.0, 45.0, 46.0, 47.0, 48.0, 49.0, 50.0];
518 let centers = vec![
519 1.0, 2.0, 3.0, 4.0, 5.0, 21.0, 22.0, 23.0, 24.0, 25.0, 31.0, 32.0, 33.0, 34.0, 35.0,
520 46.0, 47.0, 48.0, 49.0, 50.0,
521 ];
522 let mut docs_l2sq = vec![0.0; num_points];
523 let pool = create_thread_pool_for_test();
524 compute_vecs_l2sq(&mut docs_l2sq, &data, dim, pool.as_ref()).unwrap();
525 let mut centers_l2sq = vec![0.0; num_centers];
526 compute_vecs_l2sq(&mut centers_l2sq, ¢ers, dim, pool.as_ref()).unwrap();
527 let k = 2;
528 let mut center_index = vec![0; num_points * k];
529 let mut dist_matrix = vec![0.0; num_points * num_centers];
530
531 compute_closest_centers_in_block(
532 &data,
533 num_points,
534 dim,
535 ¢ers,
536 num_centers,
537 &docs_l2sq,
538 ¢ers_l2sq,
539 &mut center_index,
540 &mut dist_matrix,
541 k,
542 pool.as_ref(),
543 )
544 .unwrap();
545
546 assert_eq!(center_index.len(), num_points * k);
547 let expected_center_index = vec![3, 2, 3, 2];
548 assert_abs_diff_eq!(*center_index, expected_center_index);
549
550 assert_eq!(dist_matrix.len(), num_points * num_centers);
551 let expected_dist_matrix = vec![8000.0, 2000.0, 500.0, 125.0, 10125.0, 3125.0, 1125.0, 0.0];
556 assert_abs_diff_eq!(*dist_matrix, expected_dist_matrix, epsilon = 1e-2);
557 }
558
559 #[test]
560 fn test_compute_closest_centers() {
561 let num_points = 4;
562 let dim = 3;
563 let num_centers = 2;
564 let data = vec![
565 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
566 ];
567 let pivot_data = vec![1.0, 2.0, 3.0, 10.0, 11.0, 12.0];
568 let k = 1;
569
570 let mut closest_centers_ivf = vec![0u32; num_points * k];
571 let mut inverted_index: Vec<Vec<usize>> = vec![vec![], vec![]];
572 let pool = create_thread_pool_for_test();
573 compute_closest_centers(
574 &data,
575 num_points,
576 dim,
577 &pivot_data,
578 num_centers,
579 k,
580 &mut closest_centers_ivf,
581 Some(&mut inverted_index),
582 None,
583 pool.as_ref(),
584 )
585 .unwrap();
586
587 assert_eq!(closest_centers_ivf, vec![0, 0, 1, 1]);
588
589 for vec in inverted_index.iter_mut() {
590 vec.sort_unstable();
591 }
592 assert_eq!(inverted_index, vec![vec![0, 1], vec![2, 3]]);
593 }
594
595 #[test]
596 fn test_compute_closest_centers_with_precomputed_norms() {
597 let num_points = 4;
598 let dim = 3;
599 let num_centers = 2;
600 let data = vec![
601 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
602 ];
603 let pivot_data = vec![1.0, 2.0, 3.0, 10.0, 11.0, 12.0];
604 let k = 2;
605 let pool = create_thread_pool_for_test();
606
607 let mut closest_centers_none = vec![0u32; num_points * k];
609 compute_closest_centers(
610 &data,
611 num_points,
612 dim,
613 &pivot_data,
614 num_centers,
615 k,
616 &mut closest_centers_none,
617 None,
618 None,
619 pool.as_ref(),
620 )
621 .unwrap();
622
623 let mut pts_norms = vec![0.0; num_points];
625 compute_vecs_l2sq(&mut pts_norms, &data, dim, pool.as_ref()).unwrap();
626 let mut closest_centers_precomputed = vec![0u32; num_points * k];
627 compute_closest_centers(
628 &data,
629 num_points,
630 dim,
631 &pivot_data,
632 num_centers,
633 k,
634 &mut closest_centers_precomputed,
635 None,
636 Some(&pts_norms),
637 pool.as_ref(),
638 )
639 .unwrap();
640
641 assert_eq!(closest_centers_none, closest_centers_precomputed);
642 }
643
644 #[test]
645 fn test_compute_closest_centers_invalid_norms_length() {
646 let num_points = 4;
647 let dim = 3;
648 let num_centers = 2;
649 let data = vec![
650 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
651 ];
652 let pivot_data = vec![1.0, 2.0, 3.0, 10.0, 11.0, 12.0];
653 let k = 2;
654 let pool = create_thread_pool_for_test();
655
656 let invalid_norms = vec![0.0; num_points + 1]; let mut closest_centers = vec![0u32; num_points * k];
658 let result = compute_closest_centers(
659 &data,
660 num_points,
661 dim,
662 &pivot_data,
663 num_centers,
664 k,
665 &mut closest_centers,
666 None,
667 Some(&invalid_norms),
668 pool.as_ref(),
669 );
670
671 assert!(result
672 .unwrap_err()
673 .to_string()
674 .contains("pts_norms_squared.len() (5) should equal num_points (4)"));
675 }
676
677 #[test]
678 fn test_compute_vecs_l2sq_invalid_output_length() {
679 let num_points = 4;
680 let dim = 3;
681 let data = vec![1.0; num_points * dim];
682 let mut vecs_l2sq = vec![0.0; num_points + 1]; let pool = create_thread_pool_for_test();
684
685 let result = compute_vecs_l2sq(&mut vecs_l2sq, &data, dim, pool.as_ref());
686
687 assert!(result
688 .unwrap_err()
689 .to_string()
690 .contains("data.len() (12) should be vecs_l2sq.len() (5) * dim (3)"));
691 }
692
693 #[test]
694 fn test_compute_closest_centers_k_exceeds_num_centers() {
695 let num_points = 4;
696 let dim = 3;
697 let num_centers = 2;
698 let k = 3; let data = vec![1.0; num_points * dim];
700 let pivot_data = vec![1.0; num_centers * dim];
701 let mut closest_centers = vec![0u32; num_points * k];
702 let pool = create_thread_pool_for_test();
703
704 let result = compute_closest_centers(
705 &data,
706 num_points,
707 dim,
708 &pivot_data,
709 num_centers,
710 k,
711 &mut closest_centers,
712 None,
713 None,
714 pool.as_ref(),
715 );
716
717 assert!(result
718 .unwrap_err()
719 .to_string()
720 .contains("k (3) should be equal or less than num_centers (2)"));
721 }
722
723 #[test]
724 fn test_compute_closest_centers_invalid_output_length() {
725 let num_points = 4;
726 let dim = 3;
727 let num_centers = 2;
728 let k = 2;
729 let data = vec![1.0; num_points * dim];
730 let pivot_data = vec![1.0; num_centers * dim];
731 let mut closest_centers = vec![0u32; num_points]; let pool = create_thread_pool_for_test();
733
734 let result = compute_closest_centers(
735 &data,
736 num_points,
737 dim,
738 &pivot_data,
739 num_centers,
740 k,
741 &mut closest_centers,
742 None,
743 None,
744 pool.as_ref(),
745 );
746
747 assert!(result
748 .unwrap_err()
749 .to_string()
750 .contains("closest_centers_ivf.len() (4) should equal num_points (4) * k (2)"));
751 }
752
753 #[test]
754 fn test_compute_closest_centers_in_block_k_exceeds_num_centers() {
755 let num_points = 2;
756 let dim = 3;
757 let num_centers = 2;
758 let k = 3; let data = vec![1.0; num_points * dim];
760 let centers = vec![1.0; num_centers * dim];
761 let docs_l2sq = vec![1.0; num_points];
762 let centers_l2sq = vec![1.0; num_centers];
763 let mut center_index = vec![0u32; num_points * k];
764 let mut dist_matrix = vec![0.0; num_points * num_centers];
765 let pool = create_thread_pool_for_test();
766
767 let result = compute_closest_centers_in_block(
768 &data,
769 num_points,
770 dim,
771 ¢ers,
772 num_centers,
773 &docs_l2sq,
774 ¢ers_l2sq,
775 &mut center_index,
776 &mut dist_matrix,
777 k,
778 pool.as_ref(),
779 );
780
781 assert!(result
782 .unwrap_err()
783 .to_string()
784 .contains("k (3) should be equal or less than num_centers (2)"));
785 }
786
787 #[test]
788 fn test_compute_vecs_l2sq_overflow() {
789 let dim = usize::MAX;
790 let mut vecs_l2sq_buffer = [0.0f32; 2];
792 let data = &[];
793 let pool = create_thread_pool_for_test();
794
795 let result = compute_vecs_l2sq(&mut vecs_l2sq_buffer, data, dim, pool.as_ref());
797
798 assert!(result
799 .unwrap_err()
800 .to_string()
801 .contains("vecs_l2sq.len() * dim overflowed"));
802 }
803
804 #[test]
805 fn test_compute_closest_centers_output_buffer_overflow() {
806 let num_points = usize::MAX;
810 let k = 2;
811 let dim = 2; let num_centers = 2;
813 let data = &[];
814 let pivot_data = &[1.0f32; 4];
815 let mut closest_centers_buffer = [];
816 let pool = create_thread_pool_for_test();
817
818 let result = compute_closest_centers(
819 data,
820 num_points,
821 dim,
822 pivot_data,
823 num_centers,
824 k,
825 &mut closest_centers_buffer,
826 None,
827 None,
828 pool.as_ref(),
829 );
830
831 assert!(result
833 .unwrap_err()
834 .to_string()
835 .contains("num_points * dim overflowed"));
836 }
837
838 #[test]
839 fn test_compute_closest_centers_invalid_data_length() {
840 let num_points = 4;
841 let dim = 3;
842 let num_centers = 2;
843 let k = 1;
844 let data = vec![1.0; num_points * dim - 1]; let pivot_data = vec![1.0; num_centers * dim];
846 let mut closest_centers = vec![0u32; num_points * k];
847 let pool = create_thread_pool_for_test();
848
849 let result = compute_closest_centers(
850 &data,
851 num_points,
852 dim,
853 &pivot_data,
854 num_centers,
855 k,
856 &mut closest_centers,
857 None,
858 None,
859 pool.as_ref(),
860 );
861
862 assert!(result
863 .unwrap_err()
864 .to_string()
865 .contains("data.len() (11) should equal num_points (4) * dim (3)"));
866 }
867
868 #[test]
869 fn test_compute_closest_centers_invalid_pivot_data_length() {
870 let num_points = 4;
871 let dim = 3;
872 let num_centers = 2;
873 let k = 1;
874 let data = vec![1.0; num_points * dim];
875 let pivot_data = vec![1.0; num_centers * dim + 2]; let mut closest_centers = vec![0u32; num_points * k];
877 let pool = create_thread_pool_for_test();
878
879 let result = compute_closest_centers(
880 &data,
881 num_points,
882 dim,
883 &pivot_data,
884 num_centers,
885 k,
886 &mut closest_centers,
887 None,
888 None,
889 pool.as_ref(),
890 );
891
892 assert!(result
893 .unwrap_err()
894 .to_string()
895 .contains("pivot_data.len() (8) should equal num_centers (2) * dim (3)"));
896 }
897
898 #[test]
899 fn test_compute_closest_centers_data_overflow() {
900 let num_points = usize::MAX;
901 let dim = 2;
902 let num_centers = 2;
903 let k = 1;
904 let data = &[];
905 let pivot_data = &[1.0f32; 4]; let closest_centers = &mut [];
907 let pool = create_thread_pool_for_test();
908
909 let result = compute_closest_centers(
910 data,
911 num_points,
912 dim,
913 pivot_data,
914 num_centers,
915 k,
916 closest_centers,
917 None,
918 None,
919 pool.as_ref(),
920 );
921
922 assert!(result
923 .unwrap_err()
924 .to_string()
925 .contains("num_points * dim overflowed"));
926 }
927
928 #[test]
929 fn test_compute_closest_centers_pivot_overflow() {
930 let num_points = 4;
931 let dim = 3;
932 let num_centers = usize::MAX;
933 let k = 1;
934 let data = &[1.0f32; 12];
935 let pivot_data = &[];
936 let closest_centers = &mut [0u32; 4];
937 let pool = create_thread_pool_for_test();
938
939 let result = compute_closest_centers(
940 data,
941 num_points,
942 dim,
943 pivot_data,
944 num_centers,
945 k,
946 closest_centers,
947 None,
948 None,
949 pool.as_ref(),
950 );
951
952 assert!(result
953 .unwrap_err()
954 .to_string()
955 .contains("num_centers * dim overflowed"));
956 }
957}