weavatrix-search-vector 0.2.0

Persistent, mutable, bounded vector candidate search for Rust and Weavatrix
Documentation
#[derive(Debug, Clone, Copy)]
pub(crate) enum DistanceKernel {
    Scalar,
    #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
    Sse2,
    #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
    Avx2,
    #[cfg(target_arch = "aarch64")]
    Neon,
}

impl DistanceKernel {
    pub(crate) fn detect() -> Self {
        #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
        {
            if std::is_x86_feature_detected!("avx2") {
                return Self::Avx2;
            }
            if std::is_x86_feature_detected!("sse2") {
                return Self::Sse2;
            }
        }
        #[cfg(target_arch = "aarch64")]
        {
            if std::arch::is_aarch64_feature_detected!("neon") {
                return Self::Neon;
            }
        }
        Self::Scalar
    }

    #[inline]
    pub(crate) fn cosine_distance(
        self,
        normalized: &[f32],
        other: &[f32],
        other_inverse_norm: f32,
    ) -> f32 {
        debug_assert_eq!(normalized.len(), other.len());
        let dot = match self {
            Self::Scalar => dot_scalar(normalized, other),
            #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
            Self::Sse2 => {
                // SAFETY: this variant is selected only after runtime SSE2 detection.
                unsafe { dot_sse2(normalized, other) }
            }
            #[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
            Self::Avx2 => {
                // SAFETY: this variant is selected only after runtime AVX2 detection.
                unsafe { dot_avx2(normalized, other) }
            }
            #[cfg(target_arch = "aarch64")]
            Self::Neon => {
                // SAFETY: this variant is selected only after runtime NEON detection.
                unsafe { dot_neon(normalized, other) }
            }
        };
        (1.0 - dot * other_inverse_norm).clamp(0.0, 2.0)
    }
}

#[inline]
fn dot_scalar(left: &[f32], right: &[f32]) -> f32 {
    let mut sums = [0.0_f32; 8];
    let mut index = 0;
    while index + 8 <= left.len() {
        sums[0] += left[index] * right[index];
        sums[1] += left[index + 1] * right[index + 1];
        sums[2] += left[index + 2] * right[index + 2];
        sums[3] += left[index + 3] * right[index + 3];
        sums[4] += left[index + 4] * right[index + 4];
        sums[5] += left[index + 5] * right[index + 5];
        sums[6] += left[index + 6] * right[index + 6];
        sums[7] += left[index + 7] * right[index + 7];
        index += 8;
    }
    let mut dot = sums.into_iter().sum::<f32>();
    while index < left.len() {
        dot += left[index] * right[index];
        index += 1;
    }
    dot
}

#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2")]
unsafe fn dot_avx2(left: &[f32], right: &[f32]) -> f32 {
    #[cfg(target_arch = "x86")]
    use std::arch::x86::{
        _mm256_add_ps, _mm256_loadu_ps, _mm256_mul_ps, _mm256_setzero_ps, _mm256_storeu_ps,
    };
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::{
        _mm256_add_ps, _mm256_loadu_ps, _mm256_mul_ps, _mm256_setzero_ps, _mm256_storeu_ps,
    };

    let mut sum_a = _mm256_setzero_ps();
    let mut sum_b = _mm256_setzero_ps();
    let mut index = 0;
    while index + 16 <= left.len() {
        // SAFETY: both slices contain at least 16 elements from `index`.
        let (left_a, right_a, left_b, right_b) = unsafe {
            (
                _mm256_loadu_ps(left.as_ptr().add(index)),
                _mm256_loadu_ps(right.as_ptr().add(index)),
                _mm256_loadu_ps(left.as_ptr().add(index + 8)),
                _mm256_loadu_ps(right.as_ptr().add(index + 8)),
            )
        };
        sum_a = _mm256_add_ps(sum_a, _mm256_mul_ps(left_a, right_a));
        sum_b = _mm256_add_ps(sum_b, _mm256_mul_ps(left_b, right_b));
        index += 16;
    }
    let sum = _mm256_add_ps(sum_a, sum_b);
    let mut lanes = [0.0_f32; 8];
    // SAFETY: `lanes` has space for one 256-bit vector.
    unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sum) };
    let mut dot = lanes.into_iter().sum::<f32>();
    while index < left.len() {
        dot += left[index] * right[index];
        index += 1;
    }
    dot
}

#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "sse2")]
unsafe fn dot_sse2(left: &[f32], right: &[f32]) -> f32 {
    #[cfg(target_arch = "x86")]
    use std::arch::x86::{_mm_add_ps, _mm_loadu_ps, _mm_mul_ps, _mm_setzero_ps, _mm_storeu_ps};
    #[cfg(target_arch = "x86_64")]
    use std::arch::x86_64::{_mm_add_ps, _mm_loadu_ps, _mm_mul_ps, _mm_setzero_ps, _mm_storeu_ps};

    let mut sum_a = _mm_setzero_ps();
    let mut sum_b = _mm_setzero_ps();
    let mut index = 0;
    while index + 8 <= left.len() {
        // SAFETY: both slices contain at least eight elements from `index`.
        let (left_a, right_a, left_b, right_b) = unsafe {
            (
                _mm_loadu_ps(left.as_ptr().add(index)),
                _mm_loadu_ps(right.as_ptr().add(index)),
                _mm_loadu_ps(left.as_ptr().add(index + 4)),
                _mm_loadu_ps(right.as_ptr().add(index + 4)),
            )
        };
        sum_a = _mm_add_ps(sum_a, _mm_mul_ps(left_a, right_a));
        sum_b = _mm_add_ps(sum_b, _mm_mul_ps(left_b, right_b));
        index += 8;
    }
    let sum = _mm_add_ps(sum_a, sum_b);
    let mut lanes = [0.0_f32; 4];
    // SAFETY: `lanes` has space for one 128-bit vector.
    unsafe { _mm_storeu_ps(lanes.as_mut_ptr(), sum) };
    let mut dot = lanes.into_iter().sum::<f32>();
    while index < left.len() {
        dot += left[index] * right[index];
        index += 1;
    }
    dot
}

#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn dot_neon(left: &[f32], right: &[f32]) -> f32 {
    use std::arch::aarch64::{vaddq_f32, vdupq_n_f32, vld1q_f32, vmulq_f32, vst1q_f32};

    let mut sum_a = vdupq_n_f32(0.0);
    let mut sum_b = vdupq_n_f32(0.0);
    let mut index = 0;
    while index + 8 <= left.len() {
        // SAFETY: both slices contain at least eight elements from `index`.
        let (left_a, right_a, left_b, right_b) = unsafe {
            (
                vld1q_f32(left.as_ptr().add(index)),
                vld1q_f32(right.as_ptr().add(index)),
                vld1q_f32(left.as_ptr().add(index + 4)),
                vld1q_f32(right.as_ptr().add(index + 4)),
            )
        };
        sum_a = vaddq_f32(sum_a, vmulq_f32(left_a, right_a));
        sum_b = vaddq_f32(sum_b, vmulq_f32(left_b, right_b));
        index += 8;
    }
    let sum = vaddq_f32(sum_a, sum_b);
    let mut lanes = [0.0_f32; 4];
    // SAFETY: `lanes` has space for one 128-bit vector.
    unsafe { vst1q_f32(lanes.as_mut_ptr(), sum) };
    let mut dot = lanes.into_iter().sum::<f32>();
    while index < left.len() {
        dot += left[index] * right[index];
        index += 1;
    }
    dot
}

#[cfg(test)]
mod tests {
    #![allow(clippy::cast_precision_loss)]

    use super::DistanceKernel;

    #[test]
    fn detected_kernel_matches_scalar_cosine() {
        for dimensions in [1, 3, 8, 17, 96, 384] {
            let left = (0..dimensions)
                .map(|index| ((index * 17 + 3) as f32).sin() * 0.1)
                .collect::<Vec<_>>();
            let right = (0..dimensions)
                .map(|index| ((index * 31 + 7) as f32).cos() * 0.2)
                .collect::<Vec<_>>();
            let scalar = DistanceKernel::Scalar.cosine_distance(&left, &right, 0.75);
            let detected = DistanceKernel::detect().cosine_distance(&left, &right, 0.75);
            assert!(
                (scalar - detected).abs() <= 1.0e-5,
                "dimensions={dimensions} scalar={scalar} detected={detected}"
            );
        }
    }
}