#[cfg(target_arch = "x86_64")]
use std::arch::x86_64::*;
#[cfg(target_arch = "x86_64")]
use crate::vector::core::distance_quantized::{
abs_diff_u8_to_i32_scalar, dot_u8_to_i32_scalar, sq_diff_u8_to_i32_scalar,
};
#[inline]
pub fn is_avx2_supported() -> bool {
#[cfg(target_arch = "x86_64")]
{
is_x86_feature_detected!("avx2")
}
#[cfg(not(target_arch = "x86_64"))]
{
false
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub unsafe fn dot_u8_to_i32_avx2(a: &[u8], b: &[u8]) -> i32 {
unsafe {
let len = a.len();
let simd_len = len & !31; let zero = _mm256_setzero_si256();
let mut acc = _mm256_setzero_si256();
let pa = a.as_ptr();
let pb = b.as_ptr();
let mut i = 0;
while i < simd_len {
let va = _mm256_loadu_si256(pa.add(i) as *const __m256i);
let vb = _mm256_loadu_si256(pb.add(i) as *const __m256i);
let a_lo = _mm256_unpacklo_epi8(va, zero);
let a_hi = _mm256_unpackhi_epi8(va, zero);
let b_lo = _mm256_unpacklo_epi8(vb, zero);
let b_hi = _mm256_unpackhi_epi8(vb, zero);
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(a_lo, b_lo));
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(a_hi, b_hi));
i += 32;
}
let mut total = hsum_epi32(acc);
if simd_len < len {
total += dot_u8_to_i32_scalar(&a[simd_len..], &b[simd_len..]);
}
total
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub unsafe fn sq_diff_u8_to_i32_avx2(a: &[u8], b: &[u8]) -> i32 {
unsafe {
let len = a.len();
let simd_len = len & !31;
let zero = _mm256_setzero_si256();
let mut acc = _mm256_setzero_si256();
let pa = a.as_ptr();
let pb = b.as_ptr();
let mut i = 0;
while i < simd_len {
let va = _mm256_loadu_si256(pa.add(i) as *const __m256i);
let vb = _mm256_loadu_si256(pb.add(i) as *const __m256i);
let a_lo = _mm256_unpacklo_epi8(va, zero);
let a_hi = _mm256_unpackhi_epi8(va, zero);
let b_lo = _mm256_unpacklo_epi8(vb, zero);
let b_hi = _mm256_unpackhi_epi8(vb, zero);
let d_lo = _mm256_sub_epi16(a_lo, b_lo);
let d_hi = _mm256_sub_epi16(a_hi, b_hi);
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(d_lo, d_lo));
acc = _mm256_add_epi32(acc, _mm256_madd_epi16(d_hi, d_hi));
i += 32;
}
let mut total = hsum_epi32(acc);
if simd_len < len {
total += sq_diff_u8_to_i32_scalar(&a[simd_len..], &b[simd_len..]);
}
total
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
pub unsafe fn abs_diff_u8_to_i32_avx2(a: &[u8], b: &[u8]) -> i32 {
unsafe {
let len = a.len();
let simd_len = len & !31;
let mut acc = _mm256_setzero_si256();
let pa = a.as_ptr();
let pb = b.as_ptr();
let mut i = 0;
while i < simd_len {
let va = _mm256_loadu_si256(pa.add(i) as *const __m256i);
let vb = _mm256_loadu_si256(pb.add(i) as *const __m256i);
acc = _mm256_add_epi64(acc, _mm256_sad_epu8(va, vb));
i += 32;
}
let mut lanes = [0i64; 4];
_mm256_storeu_si256(lanes.as_mut_ptr() as *mut __m256i, acc);
let mut total = (lanes[0] + lanes[1] + lanes[2] + lanes[3]) as i32;
if simd_len < len {
total += abs_diff_u8_to_i32_scalar(&a[simd_len..], &b[simd_len..]);
}
total
}
}
#[cfg(target_arch = "x86_64")]
#[target_feature(enable = "avx2")]
unsafe fn hsum_epi32(v: __m256i) -> i32 {
unsafe {
let mut lanes = [0i32; 8];
_mm256_storeu_si256(lanes.as_mut_ptr() as *mut __m256i, v);
lanes.iter().sum()
}
}