#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
#[cfg(target_arch = "aarch64")]
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_neon_supported() -> bool {
cfg!(target_arch = "aarch64")
}
#[cfg(target_arch = "aarch64")]
pub unsafe fn dot_u8_to_i32_neon(a: &[u8], b: &[u8]) -> i32 {
unsafe {
let len = a.len();
let simd_len = len & !15; let mut acc = vdupq_n_u32(0);
let pa = a.as_ptr();
let pb = b.as_ptr();
let mut i = 0;
while i < simd_len {
let va = vld1q_u8(pa.add(i));
let vb = vld1q_u8(pb.add(i));
let a_lo = vmovl_u8(vget_low_u8(va)); let a_hi = vmovl_u8(vget_high_u8(va));
let b_lo = vmovl_u8(vget_low_u8(vb));
let b_hi = vmovl_u8(vget_high_u8(vb));
acc = vmlal_u16(acc, vget_low_u16(a_lo), vget_low_u16(b_lo));
acc = vmlal_u16(acc, vget_high_u16(a_lo), vget_high_u16(b_lo));
acc = vmlal_u16(acc, vget_low_u16(a_hi), vget_low_u16(b_hi));
acc = vmlal_u16(acc, vget_high_u16(a_hi), vget_high_u16(b_hi));
i += 16;
}
let mut total = vaddvq_u32(acc) as i32;
if simd_len < len {
total += dot_u8_to_i32_scalar(&a[simd_len..], &b[simd_len..]);
}
total
}
}
#[cfg(target_arch = "aarch64")]
pub unsafe fn sq_diff_u8_to_i32_neon(a: &[u8], b: &[u8]) -> i32 {
unsafe {
let len = a.len();
let simd_len = len & !15;
let mut acc = vdupq_n_u32(0);
let pa = a.as_ptr();
let pb = b.as_ptr();
let mut i = 0;
while i < simd_len {
let va = vld1q_u8(pa.add(i));
let vb = vld1q_u8(pb.add(i));
let ad = vabdq_u8(va, vb); let ad_lo = vmovl_u8(vget_low_u8(ad)); let ad_hi = vmovl_u8(vget_high_u8(ad));
acc = vmlal_u16(acc, vget_low_u16(ad_lo), vget_low_u16(ad_lo));
acc = vmlal_u16(acc, vget_high_u16(ad_lo), vget_high_u16(ad_lo));
acc = vmlal_u16(acc, vget_low_u16(ad_hi), vget_low_u16(ad_hi));
acc = vmlal_u16(acc, vget_high_u16(ad_hi), vget_high_u16(ad_hi));
i += 16;
}
let mut total = vaddvq_u32(acc) as i32;
if simd_len < len {
total += sq_diff_u8_to_i32_scalar(&a[simd_len..], &b[simd_len..]);
}
total
}
}
#[cfg(target_arch = "aarch64")]
pub unsafe fn abs_diff_u8_to_i32_neon(a: &[u8], b: &[u8]) -> i32 {
unsafe {
let len = a.len();
let simd_len = len & !15;
let mut acc = vdupq_n_u32(0);
let pa = a.as_ptr();
let pb = b.as_ptr();
let mut i = 0;
while i < simd_len {
let va = vld1q_u8(pa.add(i));
let vb = vld1q_u8(pb.add(i));
let ad = vabdq_u8(va, vb); let ad_lo = vmovl_u8(vget_low_u8(ad)); let ad_hi = vmovl_u8(vget_high_u8(ad));
acc = vpadalq_u16(acc, ad_lo);
acc = vpadalq_u16(acc, ad_hi);
i += 16;
}
let mut total = vaddvq_u32(acc) as i32;
if simd_len < len {
total += abs_diff_u8_to_i32_scalar(&a[simd_len..], &b[simd_len..]);
}
total
}
}