#![warn(missing_debug_implementations, missing_docs)]
use std::{cmp::Ordering, collections::BinaryHeap};
use diskann::{ANNError, ANNResult};
use diskann_linalg::{self, Transpose};
use diskann_providers::{
forward_threadpool,
utils::{AsThreadPool, ParallelIteratorInPool, RayonThreadPool},
};
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 {}
fn compute_vec_l2sq(data: &[f32], index: usize, dim: usize) -> f32 {
let start = index * dim;
let slice = unsafe { std::slice::from_raw_parts(data.as_ptr().add(start), dim) };
let mut sum_squared = 0.0;
for &value in slice {
sum_squared += value * value;
}
sum_squared
}
pub fn compute_vecs_l2sq<Pool: AsThreadPool>(
vecs_l2sq: &mut [f32],
data: &[f32],
num_points: usize,
dim: usize,
pool: Pool,
) -> ANNResult<()> {
if data.len() != num_points * dim {
return Err(ANNError::log_pq_error(format_args!(
"data.len() {} should be num_points {} * dim {}",
data.len(),
num_points,
dim
)));
}
if dim < 5 {
for (i, vec_l2sq) in vecs_l2sq.iter_mut().enumerate() {
*vec_l2sq = compute_vec_l2sq(data, i, dim);
}
} else {
forward_threadpool!(pool = pool);
vecs_l2sq
.par_iter_mut()
.enumerate()
.for_each_in_pool(pool, |(i, vec_l2sq)| {
*vec_l2sq = compute_vec_l2sq(data, i, dim);
});
}
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: &RayonThreadPool,
) -> ANNResult<()> {
if k > num_centers {
return Err(ANNError::log_index_error(format_args!(
"ERROR: k ({}) > 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,
);
diskann_linalg::sgemm(
Transpose::None,
Transpose::Ordinary,
num_points,
num_centers,
1,
1.0,
&ones_b,
centers_l2sq,
Some(1.0), dist_matrix,
);
diskann_linalg::sgemm(
Transpose::None,
Transpose::Ordinary,
num_points,
num_centers,
dim,
-2.0,
data,
centers,
Some(1.0), dist_matrix,
);
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<Pool: AsThreadPool>(
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: Pool,
) -> ANNResult<()> {
if k > num_centers {
return Err(ANNError::log_index_error(format_args!(
"ERROR: k ({}) > num_centers({})",
k, num_centers
)));
}
forward_threadpool!(pool = pool);
let pts_norms_squared = if let Some(pts_norms) = pts_norms_squared {
pts_norms.to_vec()
} else {
let mut norms_squared = vec![0.0; num_points];
compute_vecs_l2sq(&mut norms_squared, data, num_points, dim, pool)?;
norms_squared
};
let mut pivs_norms_squared = vec![0.0; num_centers];
compute_vecs_l2sq(&mut pivs_norms_squared, pivot_data, num_centers, 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, num_points, dim, &pool).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, num_points, dim, &pool).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, num_points, dim, &pool).unwrap();
let mut centers_l2sq = vec![0.0; num_centers];
compute_vecs_l2sq(&mut centers_l2sq, ¢ers, num_centers, dim, &pool).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,
)
.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, num_points, dim, &pool).unwrap();
let mut centers_l2sq = vec![0.0; num_centers];
compute_vecs_l2sq(&mut centers_l2sq, ¢ers, num_centers, dim, &pool).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,
)
.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,
)
.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]]);
}
}