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 diskann_vector::{norm::FastL2NormSquared, Norm};
19use rayon::prelude::*;
20
21// This is the chunk size applied when computing the closest centers in a block.
22// The chunk size is the number of points to process in a single iteration to reduce memory usage of
23// distance_matrix.
24// 1200 is a number we tested to be optimal for the number of points in a chunk that
25// * Large enough to take advantage of BLAS operations
26// * Small enough to avoid hefty memory allocations
27// the experiment performance of pq construction:
28// | Chunk Size   |1087932vector384dim  |8717820vector384dim  |
29// |--------------|---------------------|---------------------|
30// | 1            | 169.082s/3.181GB    | 202.175s/2.892GB    |
31// | 2            | 156.726s/1.704GB    | 189.860s/1.444GB    |
32// | 8            | 151.853s/0.996GB    | 185.035s/0.838GB    |
33// | 16           | 145.725s/0.995GB    | 185.756s/0.831GB    |
34// | 32           | 122.644s/0.996GB    | 141.831s/0.841GB    |
35// | 64           | 83.927s/0.994GB     | 97.761s/0.840GB     |
36// | 128          | 64.404s/0.994GB     | 79s/0.841GB         |
37// | 256          | 59.662s/0.995GB     | 73s/0.841GB         |
38// | 512          | 58.331s/0.996GB     | 70.552s/0.819GB     |
39// we are currently using the chunk size of 256 (about 1200 (256000 train data / 256))
40// test results are collected from i9-10900X 3.7GHz 10 cores 20 threads 32GB RAM
41// key parameters -M 1000 -R 59 -L 64 -T 8 -B 0.195 --dist_fn CosineNormalized
42const POINTS_PER_CHUNK: usize = 1200;
43
44struct PivotContainer {
45    piv_id: usize,
46    piv_dist: f32,
47}
48
49/// 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).
50/// So the correct way to implement PartialOrd for a type that has Ord is to use self.cmp(other) directly.
51impl PartialOrd for PivotContainer {
52    fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
53        Some(self.cmp(other))
54    }
55}
56
57/// The Ord trait is for types that have a total order, where every pair of values is comparable.
58impl Ord for PivotContainer {
59    fn cmp(&self, other: &Self) -> std::cmp::Ordering {
60        // Treat NaN as less than all other values.
61        // piv_dist should never be NaN.
62        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
77/// Compute L2-squared norms of data stored in row-major num_points * dim,
78/// need to be pre-allocated
79pub 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/// Calculate k closest centers to data of num_points * dim (row-major)
125/// Centers is num_centers * dim (row-major)
126/// data_l2sq has pre-computed squared norms of data
127/// centers_l2sq has pre-computed squared norms of centers
128/// Pre-allocated center_index will contain id of nearest center
129/// Pre-allocated dist_matrix should be num_points * num_centers and contain squared distances
130/// Default value of k is 1
131/// Ideally used only by compute_closest_centers
132#[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, // Initialize the destination matrix
166        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), // Add to the destination matrix
180        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), // Add to the destination matrix.
194        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/// Given data in num_points * new_dim row major
242/// Pivots stored in full_pivot_data as num_centers * new_dim row major
243/// Calculate the k closest pivot for each point and store it in vector
244/// closest_centers_ivf (row major, num_points*k) (which needs to be allocated
245/// outside) Additionally, if inverted index is not null (and pre-allocated),
246/// it will return inverted index for each center, assuming each of the inverted
247/// indices is an empty vector. Additionally, if pts_norms_squared is not null,
248/// then it will assume that point norms are pre-computed and use those values
249#[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    // Validate data slice length
270    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    // Validate pivot_data slice length
287    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        // actual chunk size maybe less than the pt_num_per_chunk for the last chunk
348        let chunk_size = data_chunk.len() / dim;
349
350        // Potentially shrink scratch data structures.
351        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, &centers, 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            &centers,
485            num_centers,
486            &docs_l2sq,
487            &centers_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, &centers, 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            &centers,
532            num_centers,
533            &docs_l2sq,
534            &centers_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        // obviously, the order of distance [8000.0, 2000.0, 500.0, 125.0], is #3, #2, #1, #0
548        // so the top 2 closest centers for the first point are #3, #2
549        // obviously, the order of distance [10125.0, 3125.0, 1125.0, 0.0], is #3, #2, #1, #0
550        // so the top 2 closest centers for the second point are #3, #2
551        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        // Compute with None (baseline)
604        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        // Compute with pre-computed norms
620        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]; // Wrong length
653        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]; // Wrong length
679        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; // k > num_centers
709        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]; // Wrong length (should be num_points * k)
742        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; // k > num_centers
769        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            &centers,
782            num_centers,
783            &docs_l2sq,
784            &centers_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        // Create a scenario where vecs_l2sq_buffer.len() * dim overflows. usize::MAX * 2 overflows
801        let mut vecs_l2sq_buffer = [0.0f32; 2];
802        let data = &[];
803        let pool = create_thread_pool_for_test();
804
805        // 2 * usize::MAX overflows
806        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        // Test that num_points * k overflow is detected
817        // Note: This test checks data.len() validation overflow, not num_points * k,
818        // because data validation happens first
819        let num_points = usize::MAX;
820        let k = 2;
821        let dim = 2; // usize::MAX * 2 overflows
822        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        // Will hit num_points * dim overflow in data validation
842        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]; // Wrong length (too short)
855        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]; // Wrong length (too long)
886        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]; // num_centers * dim = 2 * 2
916        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}