#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
pub(crate) trait L2Simd: Sized {
fn dispatch(q: &[Self], r: &[Self], dim: usize) -> Option<Self>;
}
#[cfg(target_arch = "x86_64")]
impl L2Simd for f64 {
#[inline]
fn dispatch(q: &[f64], r: &[f64], dim: usize) -> Option<f64> {
if dim < 8 {
return None;
}
if cfg!(target_feature = "avx512f") && dim >= 32 {
return Some(unsafe { l2_f64_avx512(q, r, dim) });
}
if cfg!(target_feature = "avx2") {
return Some(unsafe { l2_f64_avx2(q, r, dim) });
}
if is_x86_feature_detected!("avx512f") && dim >= 32 {
return Some(unsafe { l2_f64_avx512(q, r, dim) });
}
if is_x86_feature_detected!("avx2") {
return Some(unsafe { l2_f64_avx2(q, r, dim) });
}
None
}
}
#[cfg(target_arch = "x86_64")]
impl L2Simd for f32 {
#[inline]
fn dispatch(q: &[f32], r: &[f32], dim: usize) -> Option<f32> {
if dim < 8 {
return None;
}
if cfg!(target_feature = "avx512f") && dim >= 32 {
return Some(unsafe { l2_f32_avx512(q, r, dim) });
}
if cfg!(target_feature = "avx2") {
return Some(unsafe { l2_f32_avx2(q, r, dim) });
}
if is_x86_feature_detected!("avx512f") && dim >= 32 {
return Some(unsafe { l2_f32_avx512(q, r, dim) });
}
if is_x86_feature_detected!("avx2") {
return Some(unsafe { l2_f32_avx2(q, r, dim) });
}
None
}
}
pub(crate) trait L2FmaSimd: Sized {
fn dispatch(q: &[Self], r: &[Self], dim: usize) -> Option<Self>;
}
#[cfg(target_arch = "x86_64")]
impl L2FmaSimd for f64 {
#[inline]
fn dispatch(q: &[f64], r: &[f64], dim: usize) -> Option<f64> {
if dim < 8 {
return None;
}
if cfg!(target_feature = "avx512f") && dim >= 32 {
return Some(unsafe { l2fma_f64_avx512(q, r, dim) });
}
if cfg!(target_feature = "avx2") && cfg!(target_feature = "fma") {
return Some(unsafe { l2fma_f64_avx2(q, r, dim) });
}
if is_x86_feature_detected!("avx512f") && dim >= 32 {
return Some(unsafe { l2fma_f64_avx512(q, r, dim) });
}
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return Some(unsafe { l2fma_f64_avx2(q, r, dim) });
}
None
}
}
#[cfg(target_arch = "x86_64")]
impl L2FmaSimd for f32 {
#[inline]
fn dispatch(q: &[f32], r: &[f32], dim: usize) -> Option<f32> {
if dim < 8 {
return None;
}
if cfg!(target_feature = "avx512f") && dim >= 32 {
return Some(unsafe { l2fma_f32_avx512(q, r, dim) });
}
if cfg!(target_feature = "avx2") && cfg!(target_feature = "fma") {
return Some(unsafe { l2fma_f32_avx2(q, r, dim) });
}
if is_x86_feature_detected!("avx512f") && dim >= 32 {
return Some(unsafe { l2fma_f32_avx512(q, r, dim) });
}
if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
return Some(unsafe { l2fma_f32_avx2(q, r, dim) });
}
None
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn l2fma_f64_avx2(q: &[f64], r: &[f64], dim: usize) -> f64 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut acc0 = _mm256_setzero_pd();
let mut acc1 = _mm256_setzero_pd();
let mut i = 0usize;
while i + 8 <= dim {
let d0 = _mm256_sub_pd(_mm256_loadu_pd(qp.add(i)), _mm256_loadu_pd(rp.add(i)));
acc0 = _mm256_fmadd_pd(d0, d0, acc0);
let d1 = _mm256_sub_pd(
_mm256_loadu_pd(qp.add(i + 4)),
_mm256_loadu_pd(rp.add(i + 4)),
);
acc1 = _mm256_fmadd_pd(d1, d1, acc1);
i += 8;
}
if i + 4 <= dim {
let d0 = _mm256_sub_pd(_mm256_loadu_pd(qp.add(i)), _mm256_loadu_pd(rp.add(i)));
acc0 = _mm256_fmadd_pd(d0, d0, acc0);
i += 4;
}
acc0 = _mm256_add_pd(acc0, acc1);
let hi128 = _mm256_extractf128_pd(acc0, 1);
let lo128 = _mm256_castpd256_pd128(acc0);
let sum128 = _mm_add_pd(lo128, hi128);
let hi64 = _mm_unpackhi_pd(sum128, sum128);
let mut result = _mm_cvtsd_f64(_mm_add_sd(sum128, hi64));
while i < dim {
let d = *q.get_unchecked(i) - *r.get_unchecked(i);
result = d.mul_add(d, result);
i += 1;
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2,fma")]
unsafe fn l2fma_f32_avx2(q: &[f32], r: &[f32], dim: usize) -> f32 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut acc0 = _mm256_setzero_ps();
let mut acc1 = _mm256_setzero_ps();
let mut i = 0usize;
while i + 16 <= dim {
let d0 = _mm256_sub_ps(_mm256_loadu_ps(qp.add(i)), _mm256_loadu_ps(rp.add(i)));
acc0 = _mm256_fmadd_ps(d0, d0, acc0);
let d1 = _mm256_sub_ps(
_mm256_loadu_ps(qp.add(i + 8)),
_mm256_loadu_ps(rp.add(i + 8)),
);
acc1 = _mm256_fmadd_ps(d1, d1, acc1);
i += 16;
}
if i + 8 <= dim {
let d0 = _mm256_sub_ps(_mm256_loadu_ps(qp.add(i)), _mm256_loadu_ps(rp.add(i)));
acc0 = _mm256_fmadd_ps(d0, d0, acc0);
i += 8;
}
acc0 = _mm256_add_ps(acc0, acc1);
let hi128 = _mm256_extractf128_ps(acc0, 1);
let lo128 = _mm256_castps256_ps128(acc0);
let sum128 = _mm_add_ps(lo128, hi128);
let shuf1 = _mm_movehdup_ps(sum128);
let sum1 = _mm_add_ps(sum128, shuf1);
let shuf2 = _mm_movehl_ps(sum1, sum1);
let mut result = _mm_cvtss_f32(_mm_add_ss(sum1, shuf2));
while i < dim {
let d = *q.get_unchecked(i) - *r.get_unchecked(i);
result = d.mul_add(d, result);
i += 1;
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn l2fma_f64_avx512(q: &[f64], r: &[f64], dim: usize) -> f64 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut acc0 = _mm512_setzero_pd();
let mut acc1 = _mm512_setzero_pd();
let mut i = 0usize;
while i + 16 <= dim {
let d0 = _mm512_sub_pd(_mm512_loadu_pd(qp.add(i)), _mm512_loadu_pd(rp.add(i)));
acc0 = _mm512_fmadd_pd(d0, d0, acc0);
let d1 = _mm512_sub_pd(
_mm512_loadu_pd(qp.add(i + 8)),
_mm512_loadu_pd(rp.add(i + 8)),
);
acc1 = _mm512_fmadd_pd(d1, d1, acc1);
i += 16;
}
if i + 8 <= dim {
let d0 = _mm512_sub_pd(_mm512_loadu_pd(qp.add(i)), _mm512_loadu_pd(rp.add(i)));
acc0 = _mm512_fmadd_pd(d0, d0, acc0);
i += 8;
}
let tail = dim - i;
if tail > 0 {
let mask: __mmask8 = (1u8 << tail) - 1;
let qv = _mm512_maskz_loadu_pd(mask, qp.add(i));
let rv = _mm512_maskz_loadu_pd(mask, rp.add(i));
let d = _mm512_sub_pd(qv, rv);
acc1 = _mm512_fmadd_pd(d, d, acc1);
}
acc0 = _mm512_add_pd(acc0, acc1);
let lo256 = _mm512_castpd512_pd256(acc0);
let hi256 = _mm512_extractf64x4_pd(acc0, 1);
let sum256 = _mm256_add_pd(lo256, hi256);
let hi128 = _mm256_extractf128_pd(sum256, 1);
let lo128 = _mm256_castpd256_pd128(sum256);
let sum128 = _mm_add_pd(lo128, hi128);
let hi64 = _mm_unpackhi_pd(sum128, sum128);
_mm_cvtsd_f64(_mm_add_sd(sum128, hi64))
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn l2fma_f32_avx512(q: &[f32], r: &[f32], dim: usize) -> f32 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut acc0 = _mm512_setzero_ps();
let mut acc1 = _mm512_setzero_ps();
let mut i = 0usize;
while i + 32 <= dim {
let d0 = _mm512_sub_ps(_mm512_loadu_ps(qp.add(i)), _mm512_loadu_ps(rp.add(i)));
acc0 = _mm512_fmadd_ps(d0, d0, acc0);
let d1 = _mm512_sub_ps(
_mm512_loadu_ps(qp.add(i + 16)),
_mm512_loadu_ps(rp.add(i + 16)),
);
acc1 = _mm512_fmadd_ps(d1, d1, acc1);
i += 32;
}
if i + 16 <= dim {
let d0 = _mm512_sub_ps(_mm512_loadu_ps(qp.add(i)), _mm512_loadu_ps(rp.add(i)));
acc0 = _mm512_fmadd_ps(d0, d0, acc0);
i += 16;
}
let tail = dim - i;
if tail > 0 {
let mask: __mmask16 = (1u16 << tail) - 1;
let qv = _mm512_maskz_loadu_ps(mask, qp.add(i));
let rv = _mm512_maskz_loadu_ps(mask, rp.add(i));
let d = _mm512_sub_ps(qv, rv);
acc1 = _mm512_fmadd_ps(d, d, acc1);
}
acc0 = _mm512_add_ps(acc0, acc1);
let pd = _mm512_castps_pd(acc0);
let hi256 = _mm256_castpd_ps(_mm512_extractf64x4_pd(pd, 1));
let lo256 = _mm512_castps512_ps256(acc0);
let sum256 = _mm256_add_ps(lo256, hi256);
let hi128 = _mm256_extractf128_ps(sum256, 1);
let lo128 = _mm256_castps256_ps128(sum256);
let sum128 = _mm_add_ps(lo128, hi128);
let shuf1 = _mm_movehdup_ps(sum128);
let sum1 = _mm_add_ps(sum128, shuf1);
let shuf2 = _mm_movehl_ps(sum1, sum1);
_mm_cvtss_f32(_mm_add_ss(sum1, shuf2))
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn l2_f64_avx512(q: &[f64], r: &[f64], dim: usize) -> f64 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut result = 0.0f64;
let mut i = 0usize;
while i + 8 <= dim {
let d = _mm512_sub_pd(_mm512_loadu_pd(qp.add(i)), _mm512_loadu_pd(rp.add(i)));
let d2 = _mm512_mul_pd(d, d);
let lo256 = _mm512_castpd512_pd256(d2);
let hi256 = _mm512_extractf64x4_pd(d2, 1);
let lo_shuf = _mm256_permute_pd(lo256, 0b0101);
let lo_pairs = _mm256_add_pd(lo256, lo_shuf);
let lo_hi128 = _mm256_extractf128_pd(lo_pairs, 1);
let lo_lo128 = _mm256_castpd256_pd128(lo_pairs);
result += _mm_cvtsd_f64(_mm_add_sd(lo_lo128, lo_hi128));
let hi_shuf = _mm256_permute_pd(hi256, 0b0101);
let hi_pairs = _mm256_add_pd(hi256, hi_shuf);
let hi_hi128 = _mm256_extractf128_pd(hi_pairs, 1);
let hi_lo128 = _mm256_castpd256_pd128(hi_pairs);
result += _mm_cvtsd_f64(_mm_add_sd(hi_lo128, hi_hi128));
i += 8;
}
if i + 4 <= dim {
let d = _mm256_sub_pd(_mm256_loadu_pd(qp.add(i)), _mm256_loadu_pd(rp.add(i)));
let d2 = _mm256_mul_pd(d, d);
let shuf = _mm256_permute_pd(d2, 0b0101);
let pair_sums = _mm256_add_pd(d2, shuf);
let hi128 = _mm256_extractf128_pd(pair_sums, 1);
let lo128 = _mm256_castpd256_pd128(pair_sums);
result += _mm_cvtsd_f64(_mm_add_sd(lo128, hi128));
i += 4;
}
let rem = dim - i;
if rem >= 3 {
let d = *q.get_unchecked(i + 2) - *r.get_unchecked(i + 2);
result += d * d;
}
if rem >= 2 {
let d = *q.get_unchecked(i + 1) - *r.get_unchecked(i + 1);
result += d * d;
}
if rem >= 1 {
let d = *q.get_unchecked(i) - *r.get_unchecked(i);
result += d * d;
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx512f")]
unsafe fn l2_f32_avx512(q: &[f32], r: &[f32], dim: usize) -> f32 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut result = 0.0f32;
let mut i = 0usize;
while i + 16 <= dim {
let d = _mm512_sub_ps(_mm512_loadu_ps(qp.add(i)), _mm512_loadu_ps(rp.add(i)));
let d2 = _mm512_mul_ps(d, d);
let pd = _mm512_castps_pd(d2);
let hi256 = _mm256_castpd_ps(_mm512_extractf64x4_pd(pd, 1));
let lo256 = _mm512_castps512_ps256(d2);
for lane128 in [
_mm256_castps256_ps128(lo256),
_mm256_extractf128_ps(lo256, 1),
_mm256_castps256_ps128(hi256),
_mm256_extractf128_ps(hi256, 1),
] {
let s = _mm_movehdup_ps(lane128);
let p = _mm_add_ps(lane128, s);
let h = _mm_movehl_ps(p, p);
result += _mm_cvtss_f32(_mm_add_ss(p, h));
}
i += 16;
}
if i + 8 <= dim {
let d = _mm256_sub_ps(_mm256_loadu_ps(qp.add(i)), _mm256_loadu_ps(rp.add(i)));
let d2 = _mm256_mul_ps(d, d);
let lo = _mm256_castps256_ps128(d2);
let hi = _mm256_extractf128_ps(d2, 1);
let lo_shuf = _mm_movehdup_ps(lo);
let lo_pairs = _mm_add_ps(lo, lo_shuf);
let lo_high = _mm_movehl_ps(lo_pairs, lo_pairs);
result += _mm_cvtss_f32(_mm_add_ss(lo_pairs, lo_high));
let hi_shuf = _mm_movehdup_ps(hi);
let hi_pairs = _mm_add_ps(hi, hi_shuf);
let hi_high = _mm_movehl_ps(hi_pairs, hi_pairs);
result += _mm_cvtss_f32(_mm_add_ss(hi_pairs, hi_high));
i += 8;
}
if i + 4 <= dim {
let d = _mm_sub_ps(_mm_loadu_ps(qp.add(i)), _mm_loadu_ps(rp.add(i)));
let d2 = _mm_mul_ps(d, d);
let shuf = _mm_movehdup_ps(d2);
let pairs = _mm_add_ps(d2, shuf);
let high = _mm_movehl_ps(pairs, pairs);
result += _mm_cvtss_f32(_mm_add_ss(pairs, high));
i += 4;
}
let rem = dim - i;
if rem >= 3 {
let d = *q.get_unchecked(i + 2) - *r.get_unchecked(i + 2);
result += d * d;
}
if rem >= 2 {
let d = *q.get_unchecked(i + 1) - *r.get_unchecked(i + 1);
result += d * d;
}
if rem >= 1 {
let d = *q.get_unchecked(i) - *r.get_unchecked(i);
result += d * d;
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn l2_f64_avx2(q: &[f64], r: &[f64], dim: usize) -> f64 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut result = 0.0f64;
let mut i = 0usize;
while i + 4 <= dim {
let d = _mm256_sub_pd(_mm256_loadu_pd(qp.add(i)), _mm256_loadu_pd(rp.add(i)));
let d2 = _mm256_mul_pd(d, d);
let shuf = _mm256_permute_pd(d2, 0b0101);
let pair_sums = _mm256_add_pd(d2, shuf);
let hi128 = _mm256_extractf128_pd(pair_sums, 1);
let lo128 = _mm256_castpd256_pd128(pair_sums);
result += _mm_cvtsd_f64(_mm_add_sd(lo128, hi128));
i += 4;
}
let rem = dim - i;
if rem >= 3 {
let d = *q.get_unchecked(i + 2) - *r.get_unchecked(i + 2);
result += d * d;
}
if rem >= 2 {
let d = *q.get_unchecked(i + 1) - *r.get_unchecked(i + 1);
result += d * d;
}
if rem >= 1 {
let d = *q.get_unchecked(i) - *r.get_unchecked(i);
result += d * d;
}
result
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn l2_f32_avx2(q: &[f32], r: &[f32], dim: usize) -> f32 {
let qp = q.as_ptr();
let rp = r.as_ptr();
let mut result = 0.0f32;
let mut i = 0usize;
while i + 8 <= dim {
let d = _mm256_sub_ps(_mm256_loadu_ps(qp.add(i)), _mm256_loadu_ps(rp.add(i)));
let d2 = _mm256_mul_ps(d, d);
let lo = _mm256_castps256_ps128(d2);
let hi = _mm256_extractf128_ps(d2, 1);
let lo_shuf = _mm_movehdup_ps(lo);
let lo_pairs = _mm_add_ps(lo, lo_shuf);
let lo_high = _mm_movehl_ps(lo_pairs, lo_pairs);
result += _mm_cvtss_f32(_mm_add_ss(lo_pairs, lo_high));
let hi_shuf = _mm_movehdup_ps(hi);
let hi_pairs = _mm_add_ps(hi, hi_shuf);
let hi_high = _mm_movehl_ps(hi_pairs, hi_pairs);
result += _mm_cvtss_f32(_mm_add_ss(hi_pairs, hi_high));
i += 8;
}
if i + 4 <= dim {
let d = _mm_sub_ps(_mm_loadu_ps(qp.add(i)), _mm_loadu_ps(rp.add(i)));
let d2 = _mm_mul_ps(d, d);
let shuf = _mm_movehdup_ps(d2);
let pairs = _mm_add_ps(d2, shuf);
let high = _mm_movehl_ps(pairs, pairs);
result += _mm_cvtss_f32(_mm_add_ss(pairs, high));
i += 4;
}
let rem = dim - i;
if rem >= 3 {
let d = *q.get_unchecked(i + 2) - *r.get_unchecked(i + 2);
result += d * d;
}
if rem >= 2 {
let d = *q.get_unchecked(i + 1) - *r.get_unchecked(i + 1);
result += d * d;
}
if rem >= 1 {
let d = *q.get_unchecked(i) - *r.get_unchecked(i);
result += d * d;
}
result
}