use crate::data_source::DataSource;
use crate::dim::Dim;
use crate::scalar::{DistanceValue, Scalar};
pub trait Distance<T: Scalar>: Send + Sync {
type DistanceType: DistanceValue;
fn eval<DS: DataSource<T> + ?Sized, D: Dim>(
&self,
query: &[T],
ds: &DS,
idx: usize,
d: D,
) -> Self::DistanceType;
fn accum_dist(&self, a: T, b: T, axis: usize) -> Self::DistanceType;
}
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct L1;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct L2;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct L2Simple;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct L2Fma;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct SO2;
#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
pub struct SO3;
macro_rules! abs_ternary {
($e:expr) => {{
let v = $e;
if v < 0.0 {
-v
} else {
v
}
}};
}
#[inline]
fn debug_check_point_row_contract<T: Scalar, DS: DataSource<T> + ?Sized>(
ds: &DS,
idx: usize,
row: &[T],
dim: usize,
) {
if dim == 0 {
return;
}
#[inline]
fn same<T: Scalar>(a: T, b: T) -> bool {
a == b || (a.partial_cmp(&a).is_none() && b.partial_cmp(&b).is_none())
}
debug_assert!(
same(row[0], ds.point_component(idx, 0)),
"DataSource::point_row contract violated at idx={idx}: row[0] != point_component(idx, 0)"
);
debug_assert!(
same(row[dim - 1], ds.point_component(idx, dim - 1)),
"DataSource::point_row contract violated at idx={idx}: row[dim-1] != point_component(idx, dim-1)"
);
}
#[inline(always)]
fn l1_eval_row<T: Scalar>(query: &[T], row: &[T], dim: usize) -> T {
let zero = T::default();
#[inline]
fn abs_t<T: Scalar>(v: T, zero: T) -> T {
if v < zero {
zero - v
} else {
v
}
}
let mut result: T = zero;
let multof4 = (dim >> 2) << 2;
let (qc, _) = query[..multof4].as_chunks::<4>();
let (rc, _) = row[..multof4].as_chunks::<4>();
for (a, b) in qc.iter().zip(rc.iter()) {
let diff0 = abs_t(a[0] - b[0], zero);
let diff1 = abs_t(a[1] - b[1], zero);
let diff2 = abs_t(a[2] - b[2], zero);
let diff3 = abs_t(a[3] - b[3], zero);
result = result + ((diff0 + diff1) + (diff2 + diff3));
}
let d = multof4;
let rem = dim - multof4;
let qt = &query[d..d + rem];
let rt = &row[d..d + rem];
if rem >= 3 {
result = result + abs_t(qt[2] - rt[2], zero);
}
if rem >= 2 {
result = result + abs_t(qt[1] - rt[1], zero);
}
if rem >= 1 {
result = result + abs_t(qt[0] - rt[0], zero);
}
result
}
macro_rules! impl_l1 {
($t:ty) => {
impl Distance<$t> for L1 {
type DistanceType = $t;
fn eval<DS: DataSource<$t> + ?Sized, D: Dim>(
&self,
query: &[$t],
ds: &DS,
idx: usize,
dim: D,
) -> $t {
let dim = dim.dim();
if let Some(row) = ds.point_row(idx) {
if row.len() >= dim && query.len() >= dim {
debug_check_point_row_contract(ds, idx, row, dim);
return l1_eval_row(query, row, dim);
}
}
let mut result: $t = 0.0;
let multof4 = (dim >> 2) << 2; let mut d = 0usize;
while d < multof4 {
let diff0 = abs_ternary!(query[d] - ds.point_component(idx, d));
let diff1 = abs_ternary!(query[d + 1] - ds.point_component(idx, d + 1));
let diff2 = abs_ternary!(query[d + 2] - ds.point_component(idx, d + 2));
let diff3 = abs_ternary!(query[d + 3] - ds.point_component(idx, d + 3));
result += (diff0 + diff1) + (diff2 + diff3);
d += 4;
}
let rem = dim - multof4;
if rem >= 3 {
result += abs_ternary!(query[d + 2] - ds.point_component(idx, d + 2));
}
if rem >= 2 {
result += abs_ternary!(query[d + 1] - ds.point_component(idx, d + 1));
}
if rem >= 1 {
result += abs_ternary!(query[d] - ds.point_component(idx, d));
}
result
}
fn accum_dist(&self, a: $t, b: $t, _axis: usize) -> $t {
abs_ternary!(a - b)
}
}
};
}
#[inline(always)]
fn l2_eval_row<T: Scalar>(query: &[T], row: &[T], dim: usize) -> T {
let mut result: T = T::default(); let multof4 = (dim >> 2) << 2; let (qc, _) = query[..multof4].as_chunks::<4>();
let (rc, _) = row[..multof4].as_chunks::<4>();
for (a, b) in qc.iter().zip(rc.iter()) {
let diff0 = a[0] - b[0];
let diff1 = a[1] - b[1];
let diff2 = a[2] - b[2];
let diff3 = a[3] - b[3];
result = result + ((diff0 * diff0 + diff1 * diff1) + (diff2 * diff2 + diff3 * diff3));
}
let d = multof4;
let rem = dim - multof4;
let qt = &query[d..d + rem];
let rt = &row[d..d + rem];
if rem >= 3 {
let diff = qt[2] - rt[2];
result = result + diff * diff;
}
if rem >= 2 {
let diff = qt[1] - rt[1];
result = result + diff * diff;
}
if rem >= 1 {
let diff = qt[0] - rt[0];
result = result + diff * diff;
}
result
}
macro_rules! impl_l2 {
($t:ty) => {
impl Distance<$t> for L2 {
type DistanceType = $t;
#[inline(always)]
fn eval<DS: DataSource<$t> + ?Sized, D: Dim>(
&self,
query: &[$t],
ds: &DS,
idx: usize,
dim: D,
) -> $t {
let dim = dim.dim();
if let Some(row) = ds.point_row(idx) {
if row.len() >= dim && query.len() >= dim {
debug_check_point_row_contract(ds, idx, row, dim);
#[cfg(target_arch = "x86_64")]
{
if let Some(d) = <$t as crate::simd::L2Simd>::dispatch(query, row, dim)
{
return d;
}
}
return l2_eval_row(query, row, dim);
}
}
let mut result: $t = 0.0;
let multof4 = (dim >> 2) << 2; let mut d = 0usize;
while d < multof4 {
let diff0 = query[d] - ds.point_component(idx, d);
let diff1 = query[d + 1] - ds.point_component(idx, d + 1);
let diff2 = query[d + 2] - ds.point_component(idx, d + 2);
let diff3 = query[d + 3] - ds.point_component(idx, d + 3);
result += (diff0 * diff0 + diff1 * diff1) + (diff2 * diff2 + diff3 * diff3);
d += 4;
}
let rem = dim - multof4;
if rem >= 3 {
let diff = query[d + 2] - ds.point_component(idx, d + 2);
result += diff * diff;
}
if rem >= 2 {
let diff = query[d + 1] - ds.point_component(idx, d + 1);
result += diff * diff;
}
if rem >= 1 {
let diff = query[d] - ds.point_component(idx, d);
result += diff * diff;
}
result
}
fn accum_dist(&self, a: $t, b: $t, _axis: usize) -> $t {
let diff = a - b;
diff * diff
}
}
};
}
macro_rules! impl_l2_simple {
($t:ty) => {
impl Distance<$t> for L2Simple {
type DistanceType = $t;
fn eval<DS: DataSource<$t> + ?Sized, D: Dim>(
&self,
query: &[$t],
ds: &DS,
idx: usize,
dim: D,
) -> $t {
let dim = dim.dim();
let mut result: $t = 0.0;
for d in 0..dim {
let diff = query[d] - ds.point_component(idx, d);
result += diff * diff;
}
result
}
fn accum_dist(&self, a: $t, b: $t, _axis: usize) -> $t {
let diff = a - b;
diff * diff
}
}
};
}
macro_rules! impl_l2_fma {
($t:ty) => {
impl Distance<$t> for L2Fma {
type DistanceType = $t;
fn eval<DS: DataSource<$t> + ?Sized, D: Dim>(
&self,
query: &[$t],
ds: &DS,
idx: usize,
dim: D,
) -> $t {
let dim = dim.dim();
let mut result: $t = 0.0;
if let Some(row) = ds.point_row(idx) {
if row.len() >= dim && query.len() >= dim {
#[cfg(target_arch = "x86_64")]
{
if let Some(d) =
<$t as crate::simd::L2FmaSimd>::dispatch(query, row, dim)
{
return d;
}
}
for (q, p) in query[..dim].iter().zip(&row[..dim]) {
let diff = *q - *p;
result = diff.mul_add(diff, result);
}
return result;
}
}
for d in 0..dim {
let diff = query[d] - ds.point_component(idx, d);
result = diff.mul_add(diff, result);
}
result
}
fn accum_dist(&self, a: $t, b: $t, _axis: usize) -> $t {
let diff = a - b;
diff * diff
}
}
};
}
macro_rules! impl_so2 {
($t:ty, $pi:expr) => {
impl Distance<$t> for SO2 {
type DistanceType = $t;
fn eval<DS: DataSource<$t> + ?Sized, D: Dim>(
&self,
query: &[$t],
ds: &DS,
idx: usize,
dim: D,
) -> $t {
let dim = dim.dim();
self.accum_dist(query[dim - 1], ds.point_component(idx, dim - 1), dim - 1)
}
fn accum_dist(&self, a: $t, b: $t, _axis: usize) -> $t {
let mut diff = b - a;
let pi: $t = $pi;
if diff > pi {
diff -= 2.0 * pi;
} else if diff < -pi {
diff += 2.0 * pi;
}
abs_ternary!(diff)
}
}
};
}
macro_rules! impl_so3 {
($t:ty) => {
impl Distance<$t> for SO3 {
type DistanceType = $t;
fn eval<DS: DataSource<$t> + ?Sized, D: Dim>(
&self,
query: &[$t],
ds: &DS,
idx: usize,
dim: D,
) -> $t {
L2Simple.eval(query, ds, idx, dim)
}
fn accum_dist(&self, a: $t, b: $t, axis: usize) -> $t {
L2Simple.accum_dist(a, b, axis)
}
}
};
}
impl_l1!(f32);
impl_l1!(f64);
impl_l2!(f32);
impl_l2!(f64);
impl_l2_simple!(f32);
impl_l2_simple!(f64);
impl_l2_fma!(f32);
impl_l2_fma!(f64);
impl_so2!(f32, core::f32::consts::PI);
impl_so2!(f64, core::f64::consts::PI);
impl_so3!(f32);
impl_so3!(f64);
#[cfg(test)]
mod tests {
use super::*;
use crate::data_source::{FlatSlice, OwnedRows};
use crate::dim::{ConstDim, DynDim};
#[test]
fn l2_dim2_exact_f64() {
let q = [1.0f64, 2.0];
let p: &[[f64; 2]] = &[[4.0, 6.0]];
let got = L2.eval(&q, &p, 0, ConstDim::<2>);
assert_eq!(got, 25.0);
}
#[test]
fn l2_dim2_exact_f32() {
let q = [1.0f32, 2.0];
let p: &[[f32; 2]] = &[[4.0, 6.0]];
let got = L2.eval(&q, &p, 0, ConstDim::<2>);
assert_eq!(got, 25.0);
}
fn gen_int_points_f64(dim: usize) -> (Vec<f64>, Vec<f64>) {
let q: Vec<f64> = (0..dim).map(|i| (i as f64) * 7.0 + 3.0).collect();
let p: Vec<f64> = (0..dim).map(|i| (i as f64) * 3.0 + 11.0).collect();
(q, p)
}
fn gen_int_points_f32(dim: usize) -> (Vec<f32>, Vec<f32>) {
let q: Vec<f32> = (0..dim).map(|i| (i as f32) * 7.0 + 3.0).collect();
let p: Vec<f32> = (0..dim).map(|i| (i as f32) * 3.0 + 11.0).collect();
(q, p)
}
const TEST_DIMS: &[usize] = &[3, 5, 6, 7, 8, 16, 32];
#[test]
fn l2_matches_reference_various_dims_f64() {
for &dim in TEST_DIMS {
let (q, p) = gen_int_points_f64(dim);
let ds = FlatSlice::new(&p, dim);
let got = L2.eval(&q, &ds, 0, DynDim(dim));
let want: f64 = q.iter().zip(p.iter()).map(|(a, b)| (a - b) * (a - b)).sum();
assert_eq!(got, want, "dim={dim}");
}
}
#[test]
fn l2_matches_reference_various_dims_f32() {
for &dim in TEST_DIMS {
let (q, p) = gen_int_points_f32(dim);
let ds = FlatSlice::new(&p, dim);
let got = L2.eval(&q, &ds, 0, DynDim(dim));
let want: f32 = q.iter().zip(p.iter()).map(|(a, b)| (a - b) * (a - b)).sum();
assert_eq!(got, want, "dim={dim}");
}
}
#[test]
fn l1_matches_reference_various_dims_f64() {
for &dim in TEST_DIMS {
let (q, p) = gen_int_points_f64(dim);
let ds = FlatSlice::new(&p, dim);
let got = L1.eval(&q, &ds, 0, DynDim(dim));
let want: f64 = q.iter().zip(p.iter()).map(|(a, b)| (a - b).abs()).sum();
assert_eq!(got, want, "dim={dim}");
}
}
#[test]
fn l1_matches_reference_various_dims_f32() {
for &dim in TEST_DIMS {
let (q, p) = gen_int_points_f32(dim);
let ds = FlatSlice::new(&p, dim);
let got = L1.eval(&q, &ds, 0, DynDim(dim));
let want: f32 = q.iter().zip(p.iter()).map(|(a, b)| (a - b).abs()).sum();
assert_eq!(got, want, "dim={dim}");
}
}
#[test]
fn l2_simple_matches_l2_exact_integer_dim8() {
let (q, p) = gen_int_points_f64(8);
let ds = FlatSlice::new(&p, 8);
let l2 = L2.eval(&q, &ds, 0, ConstDim::<8>);
let l2s = L2Simple.eval(&q, &ds, 0, ConstDim::<8>);
assert_eq!(l2, l2s);
}
#[test]
fn l2_simple_matches_l2_within_4ulp_irrational_dim8() {
let dim = 8;
let q: Vec<f64> = (0..dim).map(|i| 0.1 * (i as f64)).collect();
let p: Vec<f64> = (0..dim).map(|i| 0.1 * ((i as f64) + 1.0)).collect();
let ds = FlatSlice::new(&p, dim);
let l2 = L2.eval(&q, &ds, 0, DynDim(dim));
let l2s = L2Simple.eval(&q, &ds, 0, DynDim(dim));
assert!((l2 - l2s).abs() < 4.0 * f64::EPSILON, "l2={l2} l2s={l2s}");
}
#[test]
fn l2_fma_matches_l2_exact_integer_dim8() {
let (q, p) = gen_int_points_f64(8);
let ds = FlatSlice::new(&p, 8);
let l2 = L2.eval(&q, &ds, 0, ConstDim::<8>);
let fma = L2Fma.eval(&q, &ds, 0, ConstDim::<8>);
assert_eq!(l2, fma);
}
#[test]
fn l2_fma_close_to_l2_irrational_dim8() {
let dim = 8;
let q: Vec<f64> = (0..dim).map(|i| 0.1 * (i as f64)).collect();
let p: Vec<f64> = (0..dim).map(|i| 0.1 * ((i as f64) + 1.0)).collect();
let ds = FlatSlice::new(&p, dim);
let l2 = L2.eval(&q, &ds, 0, DynDim(dim));
let fma = L2Fma.eval(&q, &ds, 0, DynDim(dim));
assert!((l2 - fma).abs() < 4.0 * f64::EPSILON, "l2={l2} fma={fma}");
}
#[test]
fn l2_fma_various_dims() {
for &dim in TEST_DIMS {
let (q, p) = gen_int_points_f64(dim);
let ds = FlatSlice::new(&p, dim);
let l2 = L2.eval(&q, &ds, 0, DynDim(dim));
let fma = L2Fma.eval(&q, &ds, 0, DynDim(dim));
assert!(
(l2 - fma).abs() < 4.0 * f64::EPSILON * l2.max(1.0),
"dim={dim} l2={l2} fma={fma}"
);
}
}
#[test]
fn so2_ignores_all_but_last_dim_and_is_unsquared() {
let q = [10.0f64, 0.1];
let p: &[[f64; 2]] = &[[-10.0, -0.1]];
let got = SO2.eval(&q, &p, 0, ConstDim::<2>);
assert!((got - 0.2).abs() < 1e-12, "got={got}");
}
#[test]
fn so2_wraps_from_positive_diff() {
let expected = (2.0 * core::f64::consts::PI - 6.0).abs();
let got = SO2.accum_dist(3.0f64, -3.0, 1);
assert!(
(got - expected).abs() < 1e-12,
"got={got} expected={expected}"
);
}
#[test]
fn so2_wraps_from_negative_diff() {
let expected = (6.0 - 2.0 * core::f64::consts::PI).abs();
let got = SO2.accum_dist(-3.0f64, 3.0, 1);
assert!(
(got - expected).abs() < 1e-12,
"got={got} expected={expected}"
);
}
#[test]
fn so2_no_wrap_when_within_pi() {
let got = SO2.accum_dist(0.0f64, 3.0, 1);
assert_eq!(got, 3.0);
}
#[test]
fn so2_result_always_in_0_pi_range() {
let pi = core::f64::consts::PI;
let n = 25;
for i in 0..n {
for j in 0..n {
let a = -pi + 2.0 * pi * (i as f64) / ((n - 1) as f64);
let b = -pi + 2.0 * pi * (j as f64) / ((n - 1) as f64);
let got = SO2.accum_dist(a, b, 1);
assert!(
got >= 0.0 && got <= pi + 1e-12,
"a={a} b={b} got={got} out of [0, pi]"
);
}
}
}
#[test]
fn so3_matches_l2_simple_bit_for_bit_dim4() {
let q = [0.3f64, -1.7, 2.25, 0.001];
let p: &[[f64; 4]] = &[[1.1, 0.05, -0.4, 3.333]];
let so3 = SO3.eval(&q, &p, 0, ConstDim::<4>);
let l2s = L2Simple.eval(&q, &p, 0, ConstDim::<4>);
assert_eq!(so3.to_bits(), l2s.to_bits());
let so3_acc = SO3.accum_dist(0.3f64, 1.1, 0);
let l2s_acc = L2Simple.accum_dist(0.3f64, 1.1, 0);
assert_eq!(so3_acc.to_bits(), l2s_acc.to_bits());
}
#[test]
fn accum_dist_l1_is_abs_diff() {
assert_eq!(L1.accum_dist(5.0f64, 2.0, 0), 3.0);
assert_eq!(L1.accum_dist(2.0f64, 5.0, 0), 3.0);
}
#[test]
fn accum_dist_l2_is_squared_diff() {
assert_eq!(L2.accum_dist(5.0f64, 2.0, 0), 9.0);
}
#[test]
fn accum_dist_l2_simple_is_squared_diff() {
assert_eq!(L2Simple.accum_dist(5.0f64, 2.0, 0), 9.0);
}
#[test]
fn accum_dist_so3_is_squared_diff() {
assert_eq!(SO3.accum_dist(5.0f64, 2.0, 0), 9.0);
}
#[test]
fn accum_dist_so2_wrapped_matches_wrap_test_values() {
let expected = (2.0 * core::f64::consts::PI - 6.0).abs();
assert!((SO2.accum_dist(3.0f64, -3.0, 1) - expected).abs() < 1e-12);
assert!((SO2.accum_dist(-3.0f64, 3.0, 1) - expected).abs() < 1e-12);
}
struct FallbackOnly<DS>(DS);
impl<T: Scalar, DS: DataSource<T>> DataSource<T> for FallbackOnly<DS> {
fn point_count(&self) -> usize {
self.0.point_count()
}
fn point_component(&self, idx: usize, dim: usize) -> T {
self.0.point_component(idx, dim)
}
}
const BIT_EQ_DIMS: &[usize] = &[1, 2, 3, 4, 5, 6, 7, 8, 15, 16, 17, 31, 32, 33, 64];
const SALTS_F64: &[f64] = &[0.4173, -1.9021, 3.5588, -4.2214, 5.7731, -6.6102];
const SALTS_F32: &[f32] = &[0.4173, -1.9021, 3.5588, -4.2214, 5.7731, -6.6102];
fn gen_random_ish_f64(dim: usize, salt: f64) -> (Vec<f64>, Vec<f64>) {
let q: Vec<f64> = (0..dim)
.map(|i| ((i as f64) * 0.837421 + salt).sin() * 137.035999)
.collect();
let p: Vec<f64> = (0..dim)
.map(|i| ((i as f64) * 1.928374 + salt * 1.5).cos() * 271.8281828 + 0.5)
.collect();
(q, p)
}
fn gen_random_ish_f32(dim: usize, salt: f32) -> (Vec<f32>, Vec<f32>) {
let q: Vec<f32> = (0..dim)
.map(|i| ((i as f32) * 0.837421 + salt).sin() * 137.036)
.collect();
let p: Vec<f32> = (0..dim)
.map(|i| ((i as f32) * 1.928374 + salt * 1.5).cos() * 271.828_2 + 0.5)
.collect();
(q, p)
}
#[test]
fn l2_eval_row_path_bit_equals_fallback_path_all_dims_f32() {
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F32 {
let (q, p) = gen_random_ish_f32(dim, salt);
let flat = FlatSlice::new(&p, dim);
let fallback = FallbackOnly(flat);
let row_result = L2.eval(&q, &flat, 0, DynDim(dim));
let fallback_result = L2.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l2_eval_row_path_bit_equals_fallback_path_all_dims_f64() {
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F64 {
let (q, p) = gen_random_ish_f64(dim, salt);
let flat = FlatSlice::new(&p, dim);
let fallback = FallbackOnly(flat);
let row_result = L2.eval(&q, &flat, 0, DynDim(dim));
let fallback_result = L2.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l1_eval_row_path_bit_equals_fallback_path_all_dims_f32() {
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F32 {
let (q, p) = gen_random_ish_f32(dim, salt);
let flat = FlatSlice::new(&p, dim);
let fallback = FallbackOnly(flat);
let row_result = L1.eval(&q, &flat, 0, DynDim(dim));
let fallback_result = L1.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l1_eval_row_path_bit_equals_fallback_path_all_dims_f64() {
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F64 {
let (q, p) = gen_random_ish_f64(dim, salt);
let flat = FlatSlice::new(&p, dim);
let fallback = FallbackOnly(flat);
let row_result = L1.eval(&q, &flat, 0, DynDim(dim));
let fallback_result = L1.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l2_eval_row_path_bit_equals_fallback_path_constdim3_array_f32() {
for &salt in SALTS_F32 {
let (q, p3) = gen_random_ish_f32(3, salt);
let points: [[f32; 3]; 1] = [[p3[0], p3[1], p3[2]]];
let arr: &[[f32; 3]] = &points;
let fallback = FallbackOnly(arr);
let row_result = L2.eval(&q, &arr, 0, ConstDim::<3>);
let fallback_result = L2.eval(&q, &fallback, 0, ConstDim::<3>);
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
#[test]
fn l2_eval_row_path_bit_equals_fallback_path_constdim3_array_f64() {
for &salt in SALTS_F64 {
let (q, p3) = gen_random_ish_f64(3, salt);
let points: [[f64; 3]; 1] = [[p3[0], p3[1], p3[2]]];
let arr: &[[f64; 3]] = &points;
let fallback = FallbackOnly(arr);
let row_result = L2.eval(&q, &arr, 0, ConstDim::<3>);
let fallback_result = L2.eval(&q, &fallback, 0, ConstDim::<3>);
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
const OWNED_ROWS_BIT_EQ_DIMS: &[usize] = &[3, 8, 32];
#[test]
fn l2_eval_row_path_bit_equals_fallback_path_owned_rows_f32() {
for &dim in OWNED_ROWS_BIT_EQ_DIMS {
for &salt in SALTS_F32 {
let (q, p) = gen_random_ish_f32(dim, salt);
let owned = OwnedRows::new(p, dim);
let fallback = FallbackOnly(&owned);
let row_result = L2.eval(&q, &owned, 0, DynDim(dim));
let fallback_result = L2.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l2_eval_row_path_bit_equals_fallback_path_owned_rows_f64() {
for &dim in OWNED_ROWS_BIT_EQ_DIMS {
for &salt in SALTS_F64 {
let (q, p) = gen_random_ish_f64(dim, salt);
let owned = OwnedRows::new(p, dim);
let fallback = FallbackOnly(&owned);
let row_result = L2.eval(&q, &owned, 0, DynDim(dim));
let fallback_result = L2.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l1_eval_row_path_bit_equals_fallback_path_owned_rows_f32() {
for &dim in OWNED_ROWS_BIT_EQ_DIMS {
for &salt in SALTS_F32 {
let (q, p) = gen_random_ish_f32(dim, salt);
let owned = OwnedRows::new(p, dim);
let fallback = FallbackOnly(&owned);
let row_result = L1.eval(&q, &owned, 0, DynDim(dim));
let fallback_result = L1.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l1_eval_row_path_bit_equals_fallback_path_owned_rows_f64() {
for &dim in OWNED_ROWS_BIT_EQ_DIMS {
for &salt in SALTS_F64 {
let (q, p) = gen_random_ish_f64(dim, salt);
let owned = OwnedRows::new(p, dim);
let fallback = FallbackOnly(&owned);
let row_result = L1.eval(&q, &owned, 0, DynDim(dim));
let fallback_result = L1.eval(&q, &fallback, 0, DynDim(dim));
assert_eq!(
row_result.to_bits(),
fallback_result.to_bits(),
"dim={dim} salt={salt} row={row_result} fallback={fallback_result}"
);
}
}
}
#[test]
fn l1_accum_dist_negative_zero_diff_stays_negative_zero() {
let got = L1.accum_dist(-0.0f64, 0.0, 0);
assert!(
got.is_sign_negative(),
"abs_ternary! must leave -0.0 unchanged, got {got} (bits {:#x})",
got.to_bits()
);
assert_eq!(got.to_bits(), (-0.0f64).to_bits());
}
#[test]
fn custom_stateful_metric_compiles_and_computes() {
struct WeightedL2 {
weights: Vec<f64>,
}
impl Distance<f64> for WeightedL2 {
type DistanceType = f64;
#[allow(clippy::needless_range_loop)]
fn eval<DS: DataSource<f64> + ?Sized, D: Dim>(
&self,
query: &[f64],
ds: &DS,
idx: usize,
dim: D,
) -> f64 {
let dim = dim.dim();
let mut result = 0.0;
for d in 0..dim {
let diff = query[d] - ds.point_component(idx, d);
result += self.weights[d] * diff * diff;
}
result
}
fn accum_dist(&self, a: f64, b: f64, axis: usize) -> f64 {
let diff = a - b;
self.weights[axis] * diff * diff
}
}
let metric = WeightedL2 {
weights: vec![1.0, 4.0, 0.5],
};
let q = [0.0f64, 0.0, 0.0];
let p: &[[f64; 3]] = &[[1.0, 1.0, 1.0]];
let got = metric.eval(&q, &p, 0, ConstDim::<3>);
assert_eq!(got, 5.5);
let acc = metric.accum_dist(0.0, 2.0, 1);
assert_eq!(acc, 16.0);
}
fn scalar_l2fma_f64(q: &[f64], r: &[f64], dim: usize) -> f64 {
let mut result = 0.0f64;
for i in 0..dim {
let d = q[i] - r[i];
result = d.mul_add(d, result);
}
result
}
fn scalar_l2fma_f32(q: &[f32], r: &[f32], dim: usize) -> f32 {
let mut result = 0.0f32;
for i in 0..dim {
let d = q[i] - r[i];
result = d.mul_add(d, result);
}
result
}
#[cfg(target_arch = "x86_64")]
#[test]
fn l2fma_simd_matches_scalar_all_dims_f64() {
use crate::simd::L2FmaSimd;
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F64 {
let (q, p) = gen_random_ish_f64(dim, salt);
let scalar = scalar_l2fma_f64(&q, &p, dim);
if let Some(simd) = f64::dispatch(&q, &p, dim) {
let tol = 4.0 * f64::EPSILON * scalar.abs().max(1.0);
assert!(
(simd - scalar).abs() <= tol,
"f64 dim={dim} salt={salt}: simd={simd} scalar={scalar} diff={} tol={tol}",
(simd - scalar).abs()
);
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn l2fma_simd_matches_scalar_all_dims_f32() {
use crate::simd::L2FmaSimd;
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F32 {
let (q, p) = gen_random_ish_f32(dim, salt);
let scalar = scalar_l2fma_f32(&q, &p, dim);
if let Some(simd) = f32::dispatch(&q, &p, dim) {
let tol = 4.0 * f32::EPSILON * scalar.abs().max(1.0);
assert!(
(simd - scalar).abs() <= tol,
"f32 dim={dim} salt={salt}: simd={simd} scalar={scalar} diff={} tol={tol}",
(simd - scalar).abs()
);
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn l2fma_simd_returns_none_below_dim8() {
use crate::simd::L2FmaSimd;
for dim in 0..8 {
let q = vec![1.0f64; dim];
let r = vec![2.0f64; dim];
assert!(f64::dispatch(&q, &r, dim).is_none(), "dim={dim}");
let q32 = vec![1.0f32; dim];
let r32 = vec![2.0f32; dim];
assert!(f32::dispatch(&q32, &r32, dim).is_none(), "dim={dim}");
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn l2_simd_bit_exact_vs_scalar_all_dims_f64() {
use crate::simd::L2Simd;
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F64 {
let (q, p) = gen_random_ish_f64(dim, salt);
let scalar = l2_eval_row(&q, &p[..dim], dim);
if let Some(simd) = <f64 as L2Simd>::dispatch(&q, &p[..dim], dim) {
assert_eq!(
simd.to_bits(),
scalar.to_bits(),
"f64 dim={dim} salt={salt}: simd={simd} scalar={scalar}"
);
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn l2_simd_bit_exact_vs_scalar_all_dims_f32() {
use crate::simd::L2Simd;
for &dim in BIT_EQ_DIMS {
for &salt in SALTS_F32 {
let (q, p) = gen_random_ish_f32(dim, salt);
let scalar = l2_eval_row(&q, &p[..dim], dim);
if let Some(simd) = <f32 as L2Simd>::dispatch(&q, &p[..dim], dim) {
assert_eq!(
simd.to_bits(),
scalar.to_bits(),
"f32 dim={dim} salt={salt}: simd={simd} scalar={scalar}"
);
}
}
}
}
#[cfg(target_arch = "x86_64")]
#[test]
fn l2_simd_returns_none_below_dim8() {
use crate::simd::L2Simd;
for dim in 0..8 {
let q = vec![1.0f64; dim];
let r = vec![2.0f64; dim];
assert!(
<f64 as L2Simd>::dispatch(&q, &r, dim).is_none(),
"dim={dim}"
);
let q32 = vec![1.0f32; dim];
let r32 = vec![2.0f32; dim];
assert!(
<f32 as L2Simd>::dispatch(&q32, &r32, dim).is_none(),
"dim={dim}"
);
}
}
}