#[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(any(target_arch = "x86", target_arch = "x86_64"))]
Avx2Fma,
#[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") && std::is_x86_feature_detected!("fma") {
return Self::Avx2Fma;
}
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 dot(self, left: &[f32], right: &[f32]) -> f32 {
debug_assert_eq!(left.len(), right.len());
match self {
Self::Scalar => dot_scalar(left, right),
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
Self::Sse2 => {
unsafe { dot_sse2(left, right) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
Self::Avx2 => {
unsafe { dot_avx2(left, right) }
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
Self::Avx2Fma => {
unsafe { dot_avx2_fma(left, right) }
}
#[cfg(target_arch = "aarch64")]
Self::Neon => {
unsafe { dot_neon(left, right) }
}
}
}
}
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
#[target_feature(enable = "avx2,fma")]
unsafe fn dot_avx2_fma(left: &[f32], right: &[f32]) -> f32 {
#[cfg(target_arch = "x86")]
use std::arch::x86::{
_mm256_add_ps, _mm256_fmadd_ps, _mm256_loadu_ps, _mm256_setzero_ps, _mm256_storeu_ps,
};
#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::{
_mm256_add_ps, _mm256_fmadd_ps, _mm256_loadu_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() {
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_fmadd_ps(left_a, right_a, sum_a);
sum_b = _mm256_fmadd_ps(left_b, right_b, sum_b);
index += 16;
}
let sum = _mm256_add_ps(sum_a, sum_b);
let mut lanes = [0.0_f32; 8];
unsafe { _mm256_storeu_ps(lanes.as_mut_ptr(), sum) };
let mut dot = lanes.into_iter().sum::<f32>();
while index < left.len() {
dot = left[index].mul_add(right[index], dot);
index += 1;
}
dot
}
#[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() {
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];
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() {
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];
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() {
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];
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 crate::simd::DistanceKernel;
#[test]
fn detected_kernel_matches_scalar_dot_product() {
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.dot(&left, &right);
let detected = DistanceKernel::detect().dot(&left, &right);
assert!(
(scalar - detected).abs() <= 1.0e-5,
"dimensions={dimensions} scalar={scalar} detected={detected}"
);
}
}
}