use super::VecPoint;
use instant_distance::Point;
use multiversion::multiversion;
const LANES: usize = 16;
#[inline(always)]
pub(crate) fn l2_sq_kernel(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len(), "L2 distance on mismatched dimensions");
let mut acc = [0.0f32; LANES];
let (a_chunks, a_rest) = a.as_chunks::<LANES>();
let (b_chunks, b_rest) = b.as_chunks::<LANES>();
for (ac, bc) in a_chunks.iter().zip(b_chunks) {
for l in 0..LANES {
let d = ac[l] - bc[l];
acc[l] += d * d;
}
}
let mut sum: f32 = acc.iter().sum();
for (x, y) in a_rest.iter().zip(b_rest) {
let d = x - y;
sum += d * d;
}
sum
}
#[multiversion(targets = "simd")]
pub(crate) fn sqdist_soa_range(
table: &[f32],
stride: usize,
q: &[f32],
lo: usize,
hi: usize,
dist: &mut [f32],
) {
let dist = &mut dist[..hi - lo];
dist.fill(0.0);
for (dim, &qd) in q.iter().enumerate() {
let src = &table[dim * stride + lo..dim * stride + hi];
for (a, &v) in dist.iter_mut().zip(src) {
let t = qd - v;
*a += t * t;
}
}
}
#[multiversion(targets = "simd")]
pub fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
l2_sq_kernel(a, b)
}
#[inline]
pub fn l2_simd(a: &[f32], b: &[f32]) -> f32 {
l2_sq(a, b).sqrt()
}
impl Point for VecPoint {
#[inline]
fn distance(&self, other: &Self) -> f32 {
l2_sq(&self.data, &other.data)
}
}