Skip to main content

diskann_disk/utils/
math_util.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5#![warn(missing_debug_implementations, missing_docs)]
6
7//! Mathematical utilities for distance computations and center finding.
8//!
9//! This module contains optimized functions for computing squared L2 norms,
10//! finding closest centers, and processing residuals. These are primarily
11//! used in k-means clustering and disk index partitioning.
12
13use 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
20// This is the chunk size applied when computing the closest centers in a block.
21// The chunk size is the number of points to process in a single iteration to reduce memory usage of
22// distance_matrix.
23// 1200 is a number we tested to be optimal for the number of points in a chunk that
24// * Large enough to take advantage of BLAS operations
25// * Small enough to avoid hefty memory allocations
26// the experiment performance of pq construction:
27// | Chunk Size   |1087932vector384dim  |8717820vector384dim  |
28// |--------------|---------------------|---------------------|
29// | 1            | 169.082s/3.181GB    | 202.175s/2.892GB    |
30// | 2            | 156.726s/1.704GB    | 189.860s/1.444GB    |
31// | 8            | 151.853s/0.996GB    | 185.035s/0.838GB    |
32// | 16           | 145.725s/0.995GB    | 185.756s/0.831GB    |
33// | 32           | 122.644s/0.996GB    | 141.831s/0.841GB    |
34// | 64           | 83.927s/0.994GB     | 97.761s/0.840GB     |
35// | 128          | 64.404s/0.994GB     | 79s/0.841GB         |
36// | 256          | 59.662s/0.995GB     | 73s/0.841GB         |
37// | 512          | 58.331s/0.996GB     | 70.552s/0.819GB     |
38// we are currently using the chunk size of 256 (about 1200 (256000 train data / 256))
39// test results are collected from i9-10900X 3.7GHz 10 cores 20 threads 32GB RAM
40// key parameters -M 1000 -R 59 -L 64 -T 8 -B 0.195 --dist_fn CosineNormalized
41const POINTS_PER_CHUNK: usize = 1200;
42
43struct PivotContainer {
44    piv_id: usize,
45    piv_dist: f32,
46}
47
48/// The PartialOrd trait is for types that can be partially ordered, i.e., where some pairs of values are incomparable (like with floating-point numbers when one of them is NaN).
49/// So the correct way to implement PartialOrd for a type that has Ord is to use self.cmp(other) directly.
50impl PartialOrd for PivotContainer {
51    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
52        Some(self.cmp(other))
53    }
54}
55
56/// The Ord trait is for types that have a total order, where every pair of values is comparable.
57impl Ord for PivotContainer {
58    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
59        // Treat NaN as less than all other values.
60        // piv_dist should never be NaN.
61        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
76/// The implementation of computing L2-squared norm of a vector
77fn 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
88/// Compute L2-squared norms of data stored in row-major num_points * dim,
89/// need to be pre-allocated
90pub 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/// Calculate k closest centers to data of num_points * dim (row-major)
129/// Centers is num_centers * dim (row-major)
130/// data_l2sq has pre-computed squared norms of data
131/// centers_l2sq has pre-computed squared norms of centers
132/// Pre-allocated center_index will contain id of nearest center
133/// Pre-allocated dist_matrix should be num_points * num_centers and contain squared distances
134/// Default value of k is 1
135/// Ideally used only by compute_closest_centers
136#[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, // Initialize the destination matrix
170        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), // Add to the destination matrix
184        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), // Add to the destination matrix.
198        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/// Given data in num_points * new_dim row major
246/// Pivots stored in full_pivot_data as num_centers * new_dim row major
247/// Calculate the k closest pivot for each point and store it in vector
248/// closest_centers_ivf (row major, num_points*k) (which needs to be allocated
249/// outside) Additionally, if inverted index is not null (and pre-allocated),
250/// it will return inverted index for each center, assuming each of the inverted
251/// indices is an empty vector. Additionally, if pts_norms_squared is not null,
252/// then it will assume that point norms are pre-computed and use those values
253#[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    // Validate data slice length
274    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    // Validate pivot_data slice length
291    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        // actual chunk size maybe less than the pt_num_per_chunk for the last chunk
352        let chunk_size = data_chunk.len() / dim;
353
354        // Potentially shrink scratch data structures.
355        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, &centers, 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            &centers,
489            num_centers,
490            &docs_l2sq,
491            &centers_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, &centers, 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            &centers,
536            num_centers,
537            &docs_l2sq,
538            &centers_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        // obviously, the order of distance [8000.0, 2000.0, 500.0, 125.0], is #3, #2, #1, #0
552        // so the top 2 closest centers for the first point are #3, #2
553        // obviously, the order of distance [10125.0, 3125.0, 1125.0, 0.0], is #3, #2, #1, #0
554        // so the top 2 closest centers for the second point are #3, #2
555        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        // Compute with None (baseline)
608        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        // Compute with pre-computed norms
624        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]; // Wrong length
657        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]; // Wrong length
683        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; // k > num_centers
699        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]; // Wrong length (should be num_points * k)
732        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; // k > num_centers
759        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            &centers,
772            num_centers,
773            &docs_l2sq,
774            &centers_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        // Create a scenario where vecs_l2sq_buffer.len() * dim overflows. usize::MAX * 2 overflows
791        let mut vecs_l2sq_buffer = [0.0f32; 2];
792        let data = &[];
793        let pool = create_thread_pool_for_test();
794
795        // 2 * usize::MAX overflows
796        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        // Test that num_points * k overflow is detected
807        // Note: This test checks data.len() validation overflow, not num_points * k,
808        // because data validation happens first
809        let num_points = usize::MAX;
810        let k = 2;
811        let dim = 2; // usize::MAX * 2 overflows
812        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        // Will hit num_points * dim overflow in data validation
832        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]; // Wrong length (too short)
845        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]; // Wrong length (too long)
876        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]; // num_centers * dim = 2 * 2
906        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}