#![warn(missing_debug_implementations, missing_docs)]
use std::{cmp::Ordering, collections::BinaryHeap};
use diskann::{ANNError, ANNErrorKind, ANNResult};
use diskann_linalg::{self, Transpose};
use diskann_providers::utils::{ParallelIteratorInPool, RayonThreadPoolRef};
use diskann_vector::{norm::FastL2NormSquared, Norm};
use rayon::prelude::*;
const POINTS_PER_CHUNK: usize = 1200;
struct PivotContainer {
piv_id: usize,
piv_dist: f32,
}
impl PartialOrd for PivotContainer {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
impl Ord for PivotContainer {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
other
.piv_dist
.partial_cmp(&self.piv_dist)
.unwrap_or(Ordering::Less)
}
}
impl PartialEq for PivotContainer {
fn eq(&self, other: &Self) -> bool {
self.piv_dist == other.piv_dist
}
}
impl Eq for PivotContainer {}
pub fn compute_vecs_l2sq(
vecs_l2sq: &mut [f32],
data: &[f32],
dim: usize,
pool: RayonThreadPoolRef<'_>,
) -> ANNResult<()> {
if dim == 0 {
return Err(ANNError::log_index_error(format_args!(
"dim must be non-zero"
)));
}
let expected_data_len = vecs_l2sq.len().checked_mul(dim).ok_or_else(|| {
ANNError::log_index_error(format_args!(
"vecs_l2sq.len() * dim overflowed: vecs_l2sq.len() ({}) * dim ({})",
vecs_l2sq.len(),
dim
))
})?;
if data.len() != expected_data_len {
return Err(ANNError::log_index_error(format_args!(
"data.len() ({}) should be vecs_l2sq.len() ({}) * dim ({})",
data.len(),
vecs_l2sq.len(),
dim
)));
}
if dim < 5 {
for (vec_l2sq, chunk) in vecs_l2sq.iter_mut().zip(data.chunks_exact(dim)) {
*vec_l2sq = FastL2NormSquared.evaluate(chunk);
}
} else {
vecs_l2sq
.par_iter_mut()
.zip(data.par_chunks_exact(dim))
.for_each_in_pool(pool, |(vec_l2sq, chunk)| {
*vec_l2sq = FastL2NormSquared.evaluate(chunk);
});
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn compute_closest_centers_in_block(
data: &[f32],
num_points: usize,
dim: usize,
centers: &[f32],
num_centers: usize,
docs_l2sq: &[f32],
centers_l2sq: &[f32],
center_index: &mut [u32],
dist_matrix: &mut [f32],
k: usize,
pool: RayonThreadPoolRef<'_>,
) -> ANNResult<()> {
if k > num_centers {
return Err(ANNError::log_index_error(format_args!(
"k ({}) should be equal or less than num_centers ({})",
k, num_centers
)));
}
let ones_a: Vec<f32> = vec![1.0; num_centers];
let ones_b: Vec<f32> = vec![1.0; num_points];
diskann_linalg::sgemm(
Transpose::None,
Transpose::Ordinary,
num_points,
num_centers,
1,
1.0,
docs_l2sq,
&ones_a,
None, dist_matrix,
)
.map_err(|e| ANNError::new(ANNErrorKind::IndexError, e))?;
diskann_linalg::sgemm(
Transpose::None,
Transpose::Ordinary,
num_points,
num_centers,
1,
1.0,
&ones_b,
centers_l2sq,
Some(1.0), dist_matrix,
)
.map_err(|e| ANNError::new(ANNErrorKind::IndexError, e))?;
diskann_linalg::sgemm(
Transpose::None,
Transpose::Ordinary,
num_points,
num_centers,
dim,
-2.0,
data,
centers,
Some(1.0), dist_matrix,
)
.map_err(|e| ANNError::new(ANNErrorKind::IndexError, e))?;
if k == 1 {
center_index
.par_iter_mut()
.enumerate()
.for_each_in_pool(pool, |(i, center_idx)| {
let mut min = f32::MAX;
let current = &dist_matrix[i * num_centers..(i + 1) * num_centers];
let mut min_idx = 0;
for (j, &distance) in current.iter().enumerate() {
if distance < min {
min = distance;
min_idx = j;
}
}
*center_idx = min_idx as u32;
});
} else {
center_index
.par_chunks_mut(k)
.enumerate()
.for_each_in_pool(pool, |(i, center_chunk)| {
let current = &dist_matrix[i * num_centers..(i + 1) * num_centers];
let mut top_k_queue = BinaryHeap::new();
for (j, &distance) in current.iter().enumerate() {
let this_piv = PivotContainer {
piv_id: j,
piv_dist: distance,
};
top_k_queue.push(this_piv);
}
for center_idx in center_chunk.iter_mut() {
if let Some(this_piv) = top_k_queue.pop() {
*center_idx = this_piv.piv_id as u32;
} else {
break;
}
}
});
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn compute_closest_centers(
data: &[f32],
num_points: usize,
dim: usize,
pivot_data: &[f32],
num_centers: usize,
k: usize,
closest_centers_ivf: &mut [u32],
mut inverted_index: Option<&mut Vec<Vec<usize>>>,
pts_norms_squared: Option<&[f32]>,
pool: RayonThreadPoolRef<'_>,
) -> ANNResult<()> {
if k > num_centers {
return Err(ANNError::log_index_error(format_args!(
"k ({}) should be equal or less than num_centers ({})",
k, num_centers
)));
}
let expected_data_len = num_points.checked_mul(dim).ok_or_else(|| {
ANNError::log_index_error(format_args!(
"num_points * dim overflowed: num_points ({}) * dim ({})",
num_points, dim
))
})?;
if data.len() != expected_data_len {
return Err(ANNError::log_index_error(format_args!(
"data.len() ({}) should equal num_points ({}) * dim ({})",
data.len(),
num_points,
dim
)));
}
let expected_pivot_len = num_centers.checked_mul(dim).ok_or_else(|| {
ANNError::log_index_error(format_args!(
"num_centers * dim overflowed: num_centers ({}) * dim ({})",
num_centers, dim
))
})?;
if pivot_data.len() != expected_pivot_len {
return Err(ANNError::log_index_error(format_args!(
"pivot_data.len() ({}) should equal num_centers ({}) * dim ({})",
pivot_data.len(),
num_centers,
dim
)));
}
let expected_closest_centers_len = num_points.checked_mul(k).ok_or_else(|| {
ANNError::log_index_error(format_args!(
"num_points * k overflowed: num_points ({}) * k ({})",
num_points, k
))
})?;
if closest_centers_ivf.len() != expected_closest_centers_len {
return Err(ANNError::log_index_error(format_args!(
"closest_centers_ivf.len() ({}) should equal num_points ({}) * k ({})",
closest_centers_ivf.len(),
num_points,
k
)));
}
let mut owned_pts_norms_squared;
let pts_norms_squared: &[f32] = if let Some(pts_norms) = pts_norms_squared {
if pts_norms.len() != num_points {
return Err(ANNError::log_index_error(format_args!(
"pts_norms_squared.len() ({}) should equal num_points ({})",
pts_norms.len(),
num_points
)));
}
pts_norms
} else {
owned_pts_norms_squared = vec![0.0; num_points];
compute_vecs_l2sq(&mut owned_pts_norms_squared, data, dim, pool)?;
&owned_pts_norms_squared
};
let mut pivs_norms_squared = vec![0.0; num_centers];
compute_vecs_l2sq(&mut pivs_norms_squared, pivot_data, dim, pool)?;
let mut distance_matrix = vec![0.0; POINTS_PER_CHUNK * num_centers];
let mut closest_center_indices = vec![0; POINTS_PER_CHUNK * k];
let pts_norms_squared_chunks = pts_norms_squared.chunks(POINTS_PER_CHUNK);
for (chunk_index, (data_chunk, pts_norms_squared_chunk)) in data
.chunks(dim * POINTS_PER_CHUNK)
.zip(pts_norms_squared_chunks)
.enumerate()
{
let chunk_size = data_chunk.len() / dim;
let this_distance_matrix = &mut distance_matrix[..num_centers * chunk_size];
let this_closest_center_indices = &mut closest_center_indices[..k * chunk_size];
compute_closest_centers_in_block(
data_chunk,
chunk_size,
dim,
pivot_data,
num_centers,
pts_norms_squared_chunk,
&pivs_norms_squared,
this_closest_center_indices,
this_distance_matrix,
k,
pool,
)?;
let point_start_index = chunk_index * POINTS_PER_CHUNK;
for point_index in point_start_index..point_start_index + chunk_size {
for l in 0..k {
let center_chunk_index = (point_index - point_start_index) * k + l;
let ivf_index = point_index * k + l;
let this_center_index = closest_center_indices[center_chunk_index];
closest_centers_ivf[ivf_index] = this_center_index;
if let Some(inverted_index) = &mut inverted_index {
inverted_index[this_center_index as usize].push(point_index);
}
}
}
}
Ok(())
}
#[cfg(test)]
mod math_util_test {
use approx::assert_abs_diff_eq;
use super::*;
use diskann_providers::utils::create_thread_pool_for_test;
#[test]
fn partial_ord_test() {
let pviot1 = PivotContainer {
piv_id: 2,
piv_dist: f32::NAN,
};
let pivot2 = PivotContainer {
piv_id: 1,
piv_dist: 1.0,
};
assert_eq!(pviot1.partial_cmp(&pivot2), Some(Ordering::Less));
}
#[test]
fn ord_test() {
let pviot1 = PivotContainer {
piv_id: 1,
piv_dist: f32::NAN,
};
let pivot2 = PivotContainer {
piv_id: 2,
piv_dist: 1.0,
};
assert_eq!(pviot1.cmp(&pivot2), Ordering::Less);
}
#[test]
fn compute_vecs_l2sq_small_dim_test() {
let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0];
let num_points = 2;
let dim = 3;
let mut vecs_l2sq = vec![0.0; num_points];
let pool = create_thread_pool_for_test();
compute_vecs_l2sq(&mut vecs_l2sq, &data, dim, pool.as_ref()).unwrap();
let expected = [14.0, 77.0];
assert_eq!(vecs_l2sq.len(), num_points);
assert_abs_diff_eq!(vecs_l2sq[0], expected[0], epsilon = 1e-6);
assert_abs_diff_eq!(vecs_l2sq[1], expected[1], epsilon = 1e-6);
}
#[test]
fn compute_vecs_l2sq_large_dim_test() {
let data = vec![
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,
];
let num_points = 2;
let dim = 8;
let mut vecs_l2sq = vec![0.0; num_points];
let pool = create_thread_pool_for_test();
compute_vecs_l2sq(&mut vecs_l2sq, &data, dim, pool.as_ref()).unwrap();
let expected = [204.0, 1292.0];
assert_eq!(vecs_l2sq.len(), num_points);
assert_abs_diff_eq!(vecs_l2sq[0], expected[0], epsilon = 1e-6);
assert_abs_diff_eq!(vecs_l2sq[1], expected[1], epsilon = 1e-6);
}
#[test]
fn compute_closest_centers_in_block_test() {
let num_points = 10;
let dim = 5;
let num_centers = 3;
let data = vec![
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,
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,
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,
45.0, 46.0, 47.0, 48.0, 49.0, 50.0,
];
let centers = vec![
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,
];
let mut docs_l2sq = vec![0.0; num_points];
let pool = create_thread_pool_for_test();
compute_vecs_l2sq(&mut docs_l2sq, &data, dim, pool.as_ref()).unwrap();
let mut centers_l2sq = vec![0.0; num_centers];
compute_vecs_l2sq(&mut centers_l2sq, ¢ers, dim, pool.as_ref()).unwrap();
let mut center_index = vec![0; num_points];
let mut dist_matrix = vec![0.0; num_points * num_centers];
let k = 1;
compute_closest_centers_in_block(
&data,
num_points,
dim,
¢ers,
num_centers,
&docs_l2sq,
¢ers_l2sq,
&mut center_index,
&mut dist_matrix,
k,
pool.as_ref(),
)
.unwrap();
assert_eq!(center_index.len(), num_points);
let expected_center_index = vec![0, 0, 0, 1, 1, 1, 2, 2, 2, 2];
assert_abs_diff_eq!(*center_index, expected_center_index);
assert_eq!(dist_matrix.len(), num_points * num_centers);
let expected_dist_matrix = vec![
0.0, 2000.0, 4500.0, 125.0, 1125.0, 3125.0, 500.0, 500.0, 2000.0, 1125.0, 125.0,
1125.0, 2000.0, 0.0, 500.0, 3125.0, 125.0, 125.0, 4500.0, 500.0, 0.0, 6125.0, 1125.0,
125.0, 8000.0, 2000.0, 500.0, 10125.0, 3125.0, 1125.0,
];
assert_abs_diff_eq!(*dist_matrix, expected_dist_matrix, epsilon = 1e-2);
}
#[test]
fn compute_closest_centers_in_block_test_k_equals_two() {
let num_points = 2;
let dim = 5;
let num_centers = 4;
let data = vec![41.0, 42.0, 43.0, 44.0, 45.0, 46.0, 47.0, 48.0, 49.0, 50.0];
let centers = vec![
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,
46.0, 47.0, 48.0, 49.0, 50.0,
];
let mut docs_l2sq = vec![0.0; num_points];
let pool = create_thread_pool_for_test();
compute_vecs_l2sq(&mut docs_l2sq, &data, dim, pool.as_ref()).unwrap();
let mut centers_l2sq = vec![0.0; num_centers];
compute_vecs_l2sq(&mut centers_l2sq, ¢ers, dim, pool.as_ref()).unwrap();
let k = 2;
let mut center_index = vec![0; num_points * k];
let mut dist_matrix = vec![0.0; num_points * num_centers];
compute_closest_centers_in_block(
&data,
num_points,
dim,
¢ers,
num_centers,
&docs_l2sq,
¢ers_l2sq,
&mut center_index,
&mut dist_matrix,
k,
pool.as_ref(),
)
.unwrap();
assert_eq!(center_index.len(), num_points * k);
let expected_center_index = vec![3, 2, 3, 2];
assert_abs_diff_eq!(*center_index, expected_center_index);
assert_eq!(dist_matrix.len(), num_points * num_centers);
let expected_dist_matrix = vec![8000.0, 2000.0, 500.0, 125.0, 10125.0, 3125.0, 1125.0, 0.0];
assert_abs_diff_eq!(*dist_matrix, expected_dist_matrix, epsilon = 1e-2);
}
#[test]
fn test_compute_closest_centers() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let data = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
];
let pivot_data = vec![1.0, 2.0, 3.0, 10.0, 11.0, 12.0];
let k = 1;
let mut closest_centers_ivf = vec![0u32; num_points * k];
let mut inverted_index: Vec<Vec<usize>> = vec![vec![], vec![]];
let pool = create_thread_pool_for_test();
compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers_ivf,
Some(&mut inverted_index),
None,
pool.as_ref(),
)
.unwrap();
assert_eq!(closest_centers_ivf, vec![0, 0, 1, 1]);
for vec in inverted_index.iter_mut() {
vec.sort_unstable();
}
assert_eq!(inverted_index, vec![vec![0, 1], vec![2, 3]]);
}
#[test]
fn test_compute_closest_centers_with_precomputed_norms() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let data = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
];
let pivot_data = vec![1.0, 2.0, 3.0, 10.0, 11.0, 12.0];
let k = 2;
let pool = create_thread_pool_for_test();
let mut closest_centers_none = vec![0u32; num_points * k];
compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers_none,
None,
None,
pool.as_ref(),
)
.unwrap();
let mut pts_norms = vec![0.0; num_points];
compute_vecs_l2sq(&mut pts_norms, &data, dim, pool.as_ref()).unwrap();
let mut closest_centers_precomputed = vec![0u32; num_points * k];
compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers_precomputed,
None,
Some(&pts_norms),
pool.as_ref(),
)
.unwrap();
assert_eq!(closest_centers_none, closest_centers_precomputed);
}
#[test]
fn test_compute_closest_centers_invalid_norms_length() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let data = vec![
1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0,
];
let pivot_data = vec![1.0, 2.0, 3.0, 10.0, 11.0, 12.0];
let k = 2;
let pool = create_thread_pool_for_test();
let invalid_norms = vec![0.0; num_points + 1]; let mut closest_centers = vec![0u32; num_points * k];
let result = compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers,
None,
Some(&invalid_norms),
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("pts_norms_squared.len() (5) should equal num_points (4)"));
}
#[test]
fn test_compute_vecs_l2sq_invalid_output_length() {
let num_points = 4;
let dim = 3;
let data = vec![1.0; num_points * dim];
let mut vecs_l2sq = vec![0.0; num_points + 1]; let pool = create_thread_pool_for_test();
let result = compute_vecs_l2sq(&mut vecs_l2sq, &data, dim, pool.as_ref());
assert!(result
.unwrap_err()
.to_string()
.contains("data.len() (12) should be vecs_l2sq.len() (5) * dim (3)"));
}
#[test]
fn test_compute_vecs_l2sq_zero_dim() {
let data: Vec<f32> = vec![];
let mut vecs_l2sq = vec![0.0; 4];
let pool = create_thread_pool_for_test();
let result = compute_vecs_l2sq(&mut vecs_l2sq, &data, 0, pool.as_ref());
assert!(result
.unwrap_err()
.to_string()
.contains("dim must be non-zero"));
}
#[test]
fn test_compute_closest_centers_k_exceeds_num_centers() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let k = 3; let data = vec![1.0; num_points * dim];
let pivot_data = vec![1.0; num_centers * dim];
let mut closest_centers = vec![0u32; num_points * k];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("k (3) should be equal or less than num_centers (2)"));
}
#[test]
fn test_compute_closest_centers_invalid_output_length() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let k = 2;
let data = vec![1.0; num_points * dim];
let pivot_data = vec![1.0; num_centers * dim];
let mut closest_centers = vec![0u32; num_points]; let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("closest_centers_ivf.len() (4) should equal num_points (4) * k (2)"));
}
#[test]
fn test_compute_closest_centers_in_block_k_exceeds_num_centers() {
let num_points = 2;
let dim = 3;
let num_centers = 2;
let k = 3; let data = vec![1.0; num_points * dim];
let centers = vec![1.0; num_centers * dim];
let docs_l2sq = vec![1.0; num_points];
let centers_l2sq = vec![1.0; num_centers];
let mut center_index = vec![0u32; num_points * k];
let mut dist_matrix = vec![0.0; num_points * num_centers];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers_in_block(
&data,
num_points,
dim,
¢ers,
num_centers,
&docs_l2sq,
¢ers_l2sq,
&mut center_index,
&mut dist_matrix,
k,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("k (3) should be equal or less than num_centers (2)"));
}
#[test]
fn test_compute_vecs_l2sq_overflow() {
let dim = usize::MAX;
let mut vecs_l2sq_buffer = [0.0f32; 2];
let data = &[];
let pool = create_thread_pool_for_test();
let result = compute_vecs_l2sq(&mut vecs_l2sq_buffer, data, dim, pool.as_ref());
assert!(result
.unwrap_err()
.to_string()
.contains("vecs_l2sq.len() * dim overflowed"));
}
#[test]
fn test_compute_closest_centers_output_buffer_overflow() {
let num_points = usize::MAX;
let k = 2;
let dim = 2; let num_centers = 2;
let data = &[];
let pivot_data = &[1.0f32; 4];
let mut closest_centers_buffer = [];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
data,
num_points,
dim,
pivot_data,
num_centers,
k,
&mut closest_centers_buffer,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("num_points * dim overflowed"));
}
#[test]
fn test_compute_closest_centers_invalid_data_length() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let k = 1;
let data = vec![1.0; num_points * dim - 1]; let pivot_data = vec![1.0; num_centers * dim];
let mut closest_centers = vec![0u32; num_points * k];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("data.len() (11) should equal num_points (4) * dim (3)"));
}
#[test]
fn test_compute_closest_centers_invalid_pivot_data_length() {
let num_points = 4;
let dim = 3;
let num_centers = 2;
let k = 1;
let data = vec![1.0; num_points * dim];
let pivot_data = vec![1.0; num_centers * dim + 2]; let mut closest_centers = vec![0u32; num_points * k];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
&data,
num_points,
dim,
&pivot_data,
num_centers,
k,
&mut closest_centers,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("pivot_data.len() (8) should equal num_centers (2) * dim (3)"));
}
#[test]
fn test_compute_closest_centers_data_overflow() {
let num_points = usize::MAX;
let dim = 2;
let num_centers = 2;
let k = 1;
let data = &[];
let pivot_data = &[1.0f32; 4]; let closest_centers = &mut [];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
data,
num_points,
dim,
pivot_data,
num_centers,
k,
closest_centers,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("num_points * dim overflowed"));
}
#[test]
fn test_compute_closest_centers_pivot_overflow() {
let num_points = 4;
let dim = 3;
let num_centers = usize::MAX;
let k = 1;
let data = &[1.0f32; 12];
let pivot_data = &[];
let closest_centers = &mut [0u32; 4];
let pool = create_thread_pool_for_test();
let result = compute_closest_centers(
data,
num_points,
dim,
pivot_data,
num_centers,
k,
closest_centers,
None,
None,
pool.as_ref(),
);
assert!(result
.unwrap_err()
.to_string()
.contains("num_centers * dim overflowed"));
}
}