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