#![expect(
dead_code,
reason = "graph construction integrates these kernels in the next PR"
)]
mod leaf_kernel;
mod leaf_metric;
mod partition_kernel;
mod partition_metric;
mod simd;
mod topk;
use crate::{ANNError, ANNResult};
use diskann_utils::views::rowmajor;
pub(super) struct L2;
pub(super) struct Cosine;
pub(super) struct CosineNormalized;
pub(super) struct InnerProduct;
#[inline(always)]
fn cosine_distance(dot: f32, source_norm: f32, target_norm: f32) -> f32 {
if source_norm < f32::MIN_POSITIVE.sqrt() || target_norm < f32::MIN_POSITIVE.sqrt() {
1.0
} else {
1.0 - (dot / (source_norm * target_norm)).clamp(-1.0, 1.0)
}
}
fn check_output_rows(points: usize, rows: usize) -> ANNResult<()> {
if rows == points {
Ok(())
} else {
Err(ANNError::message(format!(
"invalid kernel output row count {rows} for {points} points"
)))
}
}
fn distance_scratch(
storage: &mut Vec<f32>,
rows: usize,
columns: usize,
) -> ANNResult<rowmajor::Mut<'_, f32>> {
let len = rows.checked_mul(columns).ok_or_else(|| {
ANNError::message(format!(
"distance matrix size overflows for {rows} x {columns}"
))
})?;
if storage.len() < len {
storage.resize(len, 0.0);
}
Ok(rowmajor::Mut::try_from_data(
&mut storage[..len],
rows,
columns,
)?)
}
#[cfg(test)]
mod cosine_distance_tests {
use super::cosine_distance;
use rstest::rstest;
#[rstest]
#[case::same_direction(6.0, 2.0, 3.0, 0.0)]
#[case::opposite_directions(-6.0, 2.0, 3.0, 2.0)]
#[case::orthogonal(0.0, 2.0, 3.0, 1.0)]
#[case::positive_similarity(3.0, 2.0, 3.0, 0.5)]
#[case::negative_similarity(-3.0, 2.0, 3.0, 1.5)]
fn distance_is_one_minus_the_dot_divided_by_both_norms(
#[case] dot: f32,
#[case] source_norm: f32,
#[case] target_norm: f32,
#[case] expected: f32,
) {
assert_eq!(cosine_distance(dot, source_norm, target_norm), expected);
}
#[rstest]
#[case::above_one(1.0 + f32::EPSILON, 0.0)]
#[case::below_minus_one(-1.0 - f32::EPSILON, 2.0)]
fn rounding_outside_the_similarity_range_is_clamped(
#[case] similarity: f32,
#[case] expected: f32,
) {
assert_eq!(cosine_distance(similarity, 1.0, 1.0), expected);
}
#[rstest]
#[case::zero(0.0)]
#[case::below_cutoff(f32::from_bits(f32::MIN_POSITIVE.sqrt().to_bits() - 1))]
fn a_norm_below_the_cutoff_gives_distance_one_for_any_dot(#[case] small_norm: f32) {
for dot in [0.75, f32::NAN] {
for (source_norm, target_norm) in [(small_norm, 1.0), (1.0, small_norm)] {
assert_eq!(
cosine_distance(dot, source_norm, target_norm),
1.0,
"dot={dot}, norms=({source_norm:e}, {target_norm:e})"
);
}
}
}
#[rstest]
#[case::source(true)]
#[case::target(false)]
fn a_norm_at_the_cutoff_still_contributes_similarity(#[case] source_at_cutoff: bool) {
let norm = f32::MIN_POSITIVE.sqrt();
let (source_norm, target_norm) = if source_at_cutoff {
(norm, 1.0)
} else {
(1.0, norm)
};
assert_eq!(cosine_distance(norm / 4.0, source_norm, target_norm), 0.75);
}
#[rstest]
#[case::dot(f32::NAN, 2.0, 3.0)]
#[case::source_norm(1.0, f32::NAN, 3.0)]
#[case::target_norm(1.0, 2.0, f32::NAN)]
fn nan_propagates_when_neither_norm_is_below_the_cutoff(
#[case] dot: f32,
#[case] source_norm: f32,
#[case] target_norm: f32,
) {
assert!(cosine_distance(dot, source_norm, target_norm).is_nan());
}
}
#[cfg(test)]
mod test_support {
use super::simd::Simd;
use diskann_vector::distance::Metric;
pub(super) trait ArchCheck {
fn check<A: Simd>(&self, arch: A);
}
pub(super) fn for_each_arch(test: &impl ArchCheck) {
test.check(diskann_wide::arch::Scalar);
#[cfg(target_arch = "x86_64")]
{
use diskann_wide::arch::x86_64::{V3, V4};
if let Some(arch) = V3::new_checked() {
test.check(arch);
}
if let Some(arch) = V4::new_checked_miri() {
test.check(arch);
}
}
#[cfg(target_arch = "aarch64")]
if let Some(arch) = diskann_wide::arch::aarch64::Neon::new_checked() {
test.check(arch);
}
}
pub(super) fn dense_tolerance(metric: Metric, dimensions: usize) -> f64 {
match metric {
Metric::L2 | Metric::InnerProduct => 0.0,
Metric::Cosine => 16.0 * f64::from(f32::EPSILON),
Metric::CosineNormalized => {
let roundoff = dimensions as f64 * f64::from(f32::EPSILON);
roundoff / (1.0 - roundoff)
}
}
}
pub(super) fn dense_points(rows: usize, dimensions: usize, seed: u64) -> Vec<f32> {
use rand::{Rng, SeedableRng, rngs::StdRng};
let mut rng = StdRng::seed_from_u64(seed);
let mut values: Vec<_> = (0..rows * dimensions)
.map(|_| rng.random_range(-16..=16) as f32 / 8.0)
.collect();
for (point, row) in values.chunks_exact_mut(dimensions).enumerate() {
row[dimensions - 1] = 8.0 + (point % 7) as f32;
}
values
}
pub(super) fn distance(metric: Metric, point: &[f32], target: &[f32]) -> f64 {
let dot = |x: &[f32], y: &[f32]| {
x.iter()
.zip(y)
.map(|(&x, &y)| f64::from(x) * f64::from(y))
.sum::<f64>()
};
match metric {
Metric::L2 => point
.iter()
.zip(target)
.map(|(&x, &y)| (f64::from(x) - f64::from(y)).powi(2))
.sum(),
Metric::InnerProduct => -dot(point, target),
Metric::CosineNormalized => 1.0 - dot(point, target),
Metric::Cosine => {
let norm = |row: &[f32]| dot(row, row).sqrt();
1.0 - (dot(point, target) / (norm(point) * norm(target))).clamp(-1.0, 1.0)
}
}
}
pub(super) fn normalize(values: &mut [f32], dimensions: usize) {
for row in values.chunks_exact_mut(dimensions) {
let norm = row
.iter()
.map(|&x| f64::from(x).powi(2))
.sum::<f64>()
.sqrt();
for value in row {
*value = (f64::from(*value) / norm) as f32;
}
}
}
pub(super) fn packed_points(
coordinates: &[[f32; 2]],
dimensions: usize,
unit_norm: bool,
) -> Vec<f32> {
let mut values = vec![0.0; coordinates.len() * dimensions];
for (row, &[x, y]) in values.chunks_exact_mut(dimensions).zip(coordinates) {
row[0] = x;
row[dimensions - 1] = y;
}
if unit_norm {
normalize(&mut values, dimensions);
}
values
}
#[test]
fn dense_fixtures_are_deterministic_and_nonzero_in_the_last_dimension() {
let values = dense_points(3, 17, 1287);
assert_eq!(values, dense_points(3, 17, 1287));
assert_eq!(values.len(), 3 * 17);
assert_ne!(&values[..17], &values[17..34]);
assert_eq!([values[16], values[33], values[50]], [8.0, 9.0, 10.0]);
assert!(values[..16].iter().filter(|&&x| x != 0.0).count() > 8);
}
#[rstest::rstest]
#[case::squared_l2(Metric::L2, 18.0)]
#[case::negative_dot(Metric::InnerProduct, -2.0)]
#[case::one_minus_dot(Metric::CosineNormalized, -1.0)]
#[case::cosine(Metric::Cosine, 0.7830695421813438)]
fn scalar_reference_matches_hand_calculated_distances(
#[case] metric: Metric,
#[case] expected: f64,
) {
assert!((distance(metric, &[1.0, 2.0], &[4.0, -1.0]) - expected).abs() < 1.0e-14);
}
#[test]
fn normalized_packed_points_put_the_second_coordinate_last() {
let actual = packed_points(&[[3.0, 4.0], [0.0, -2.0]], 3, true);
assert_eq!(actual, [0.6, 0.0, 0.8, 0.0, 0.0, -1.0]);
assert_eq!(
distance(Metric::CosineNormalized, &actual[..3], &actual[3..]),
1.0 + f64::from(0.8_f32)
);
}
}