use crate::vec::VecData;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use half::bf16;
use half::f16;
use ndarray::ArrayView1;
use ndarray::ArrayView1 as AV;
#[cfg(target_arch = "aarch64")]
use std::arch::aarch64::*;
#[cfg(target_arch = "aarch64")]
use std::arch::is_aarch64_feature_detected;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use std::arch::x86_64::*;
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
use std::mem::transmute;
#[target_feature(enable = "avx512f,avx512bf16")]
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn dist_l2_bf16_avx512(a_slice: &[bf16], b_slice: &[bf16]) -> f64 {
let len = a_slice.len();
let mut acc_sum_ps = _mm512_setzero_ps();
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut i = 0;
let limit_avx512 = len - (len % 32);
while i < limit_avx512 {
let v_a_i = _mm512_loadu_si512(ptr_a.add(i) as *const _);
let v_b_i = _mm512_loadu_si512(ptr_b.add(i) as *const _);
let v_a_low_256bh = _mm512_castsi512_si256(v_a_i); let v_b_low_256bh = _mm512_castsi512_si256(v_b_i);
let v_a_low_ps = _mm512_cvtpbh_ps(transmute(v_a_low_256bh)); let v_b_low_ps = _mm512_cvtpbh_ps(transmute(v_b_low_256bh));
let v_a_high_256bh = _mm512_extracti64x4_epi64(v_a_i, 1); let v_b_high_256bh = _mm512_extracti64x4_epi64(v_b_i, 1);
let v_a_high_ps = _mm512_cvtpbh_ps(transmute(v_a_high_256bh)); let v_b_high_ps = _mm512_cvtpbh_ps(transmute(v_b_high_256bh));
let v_diff_low_ps = _mm512_sub_ps(v_a_low_ps, v_b_low_ps);
let v_diff_high_ps = _mm512_sub_ps(v_a_high_ps, v_b_high_ps);
let v_diff_bh = _mm512_cvtne2ps_pbh(v_diff_high_ps, v_diff_low_ps);
acc_sum_ps = _mm512_dpbf16_ps(acc_sum_ps, v_diff_bh, v_diff_bh);
i += 32;
}
let mut total_sum_f32 = _mm512_reduce_add_ps(acc_sum_ps);
if i < len {
let mut scalar_sum_f32: f32 = 0.0;
for k in i..len {
let val_a_f32 = (*ptr_a.add(k)).to_f32();
let val_b_f32 = (*ptr_b.add(k)).to_f32();
let diff_f32 = val_a_f32 - val_b_f32;
scalar_sum_f32 += diff_f32 * diff_f32;
}
total_sum_f32 += scalar_sum_f32;
}
let result_f32 = total_sum_f32.sqrt(); result_f32.into()
}
#[target_feature(enable = "avx512f,avx512fp16")]
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn dist_l2_f16_avx512(a_slice: &[f16], b_slice: &[f16]) -> f64 {
let len = a_slice.len();
let mut acc_sum_ps = _mm512_setzero_ps();
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut i = 0;
let limit_avx512 = len - (len % 16);
while i < limit_avx512 {
let v_a_ph_i = _mm256_loadu_si256(ptr_a.add(i) as *const _); let v_b_ph_i = _mm256_loadu_si256(ptr_b.add(i) as *const _);
let v_a_ps = _mm512_cvtph_ps(v_a_ph_i); let v_b_ps = _mm512_cvtph_ps(v_b_ph_i);
let v_diff_ps = _mm512_sub_ps(v_a_ps, v_b_ps);
acc_sum_ps = _mm512_fmadd_ps(v_diff_ps, v_diff_ps, acc_sum_ps);
i += 16;
}
let mut total_sum_f32 = _mm512_reduce_add_ps(acc_sum_ps);
if i < len {
let mut scalar_sum_f32: f32 = 0.0;
for k in i..len {
let val_a_f32 = (*ptr_a.add(k)).to_f32();
let val_b_f32 = (*ptr_b.add(k)).to_f32();
let diff_f32 = val_a_f32 - val_b_f32;
scalar_sum_f32 += diff_f32 * diff_f32;
}
total_sum_f32 += scalar_sum_f32;
}
let result_f32 = total_sum_f32.sqrt();
result_f32.into() }
#[target_feature(enable = "avx512f")]
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn dist_l2_f32_avx512(a_slice: &[f32], b_slice: &[f32]) -> f64 {
let len = a_slice.len();
let mut acc_sum_ps = _mm512_setzero_ps();
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut i = 0;
let limit_avx512 = len - (len % 16);
while i < limit_avx512 {
let v_a_ps = _mm512_loadu_ps(ptr_a.add(i));
let v_b_ps = _mm512_loadu_ps(ptr_b.add(i));
let v_diff_ps = _mm512_sub_ps(v_a_ps, v_b_ps);
acc_sum_ps = _mm512_fmadd_ps(v_diff_ps, v_diff_ps, acc_sum_ps);
i += 16;
}
let mut total_sum_f32 = _mm512_reduce_add_ps(acc_sum_ps);
if i < len {
let mut scalar_sum_f32: f32 = 0.0;
for k in i..len {
let val_a = *ptr_a.add(k);
let val_b = *ptr_b.add(k);
let diff = val_a - val_b;
scalar_sum_f32 += diff * diff;
}
total_sum_f32 += scalar_sum_f32;
}
let result_f32 = total_sum_f32.sqrt();
result_f32.into() }
#[target_feature(enable = "avx512f")]
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
unsafe fn dist_l2_f64_avx512(a_slice: &[f64], b_slice: &[f64]) -> f64 {
let len = a_slice.len();
let mut acc_sum_pd = _mm512_setzero_pd();
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut i = 0;
let limit_avx512 = len - (len % 8);
while i < limit_avx512 {
let v_a_pd = _mm512_loadu_pd(ptr_a.add(i));
let v_b_pd = _mm512_loadu_pd(ptr_b.add(i));
let v_diff_pd = _mm512_sub_pd(v_a_pd, v_b_pd);
acc_sum_pd = _mm512_fmadd_pd(v_diff_pd, v_diff_pd, acc_sum_pd);
i += 8;
}
let mut total_sum_f64 = _mm512_reduce_add_pd(acc_sum_pd);
if i < len {
let mut scalar_sum_f64: f64 = 0.0;
for k in i..len {
let val_a = *ptr_a.add(k);
let val_b = *ptr_b.add(k);
let diff = val_a - val_b;
scalar_sum_f64 += diff * diff;
}
total_sum_f64 += scalar_sum_f64;
}
total_sum_f64.sqrt() }
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon,fp16")]
unsafe fn dist_l2_f16_neon(a_slice: &[f16], b_slice: &[f16]) -> f64 {
let len = a_slice.len();
if len == 0 {
return 0.0;
}
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut acc_sum_f32 = vdupq_n_f32(0.0);
let mut i = 0;
let limit_neon = len - (len % 8);
while i < limit_neon {
let a_f16 = vld1q_f16(ptr_a.add(i) as *const _);
let b_f16 = vld1q_f16(ptr_b.add(i) as *const _);
let a_f32_low = vcvt_f32_f16(vget_low_f16(a_f16));
let a_f32_high = vcvt_f32_f16(vget_high_f16(a_f16));
let b_f32_low = vcvt_f32_f16(vget_low_f16(b_f16));
let b_f32_high = vcvt_f32_f16(vget_high_f16(b_f16));
let diff_low = vsubq_f32(a_f32_low, b_f32_low);
let diff_high = vsubq_f32(a_f32_high, b_f32_high);
acc_sum_f32 = vfmaq_f32(acc_sum_f32, diff_low, diff_low);
acc_sum_f32 = vfmaq_f32(acc_sum_f32, diff_high, diff_high);
i += 8;
}
let total_sum = vaddvq_f32(acc_sum_f32);
let mut scalar_sum = 0.0f32;
for k in i..len {
let a_f32 = (*ptr_a.add(k)).to_f32();
let b_f32 = (*ptr_b.add(k)).to_f32();
let diff = a_f32 - b_f32;
scalar_sum += diff * diff;
}
((total_sum + scalar_sum) as f64).sqrt()
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn dist_l2_f32_neon(a_slice: &[f32], b_slice: &[f32]) -> f64 {
let len = a_slice.len();
if len == 0 {
return 0.0;
}
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut acc_sum_f32 = vdupq_n_f32(0.0);
let mut i = 0;
let limit_neon = len - (len % 4);
while i < limit_neon {
let a_f32 = vld1q_f32(ptr_a.add(i));
let b_f32 = vld1q_f32(ptr_b.add(i));
let diff = vsubq_f32(a_f32, b_f32);
acc_sum_f32 = vfmaq_f32(acc_sum_f32, diff, diff);
i += 4;
}
let total_sum = vaddvq_f32(acc_sum_f32);
let mut scalar_sum = 0.0f32;
for k in i..len {
let a_val = *ptr_a.add(k);
let b_val = *ptr_b.add(k);
let diff = a_val - b_val;
scalar_sum += diff * diff;
}
((total_sum + scalar_sum) as f64).sqrt()
}
#[cfg(target_arch = "aarch64")]
#[target_feature(enable = "neon")]
unsafe fn dist_l2_f64_neon(a_slice: &[f64], b_slice: &[f64]) -> f64 {
let len = a_slice.len();
if len == 0 {
return 0.0;
}
let ptr_a = a_slice.as_ptr();
let ptr_b = b_slice.as_ptr();
let mut acc_sum_f64 = vdupq_n_f64(0.0);
let mut i = 0;
let limit_neon = len - (len % 2);
while i < limit_neon {
let a_f64 = vld1q_f64(ptr_a.add(i));
let b_f64 = vld1q_f64(ptr_b.add(i));
let diff = vsubq_f64(a_f64, b_f64);
acc_sum_f64 = vfmaq_f64(acc_sum_f64, diff, diff);
i += 2;
}
let total_sum = vaddvq_f64(acc_sum_f64);
let mut scalar_sum = 0.0f64;
for k in i..len {
let a_val = *ptr_a.add(k);
let b_val = *ptr_b.add(k);
let diff = a_val - b_val;
scalar_sum += diff * diff;
}
(total_sum + scalar_sum).sqrt()
}
fn dist_l2_scalar<T: num::Float>(a: &ArrayView1<T>, b: &ArrayView1<T>) -> f64 {
let diff = a - b;
let squared_diff = &diff * &diff;
let sum_squared_diff = squared_diff.sum();
sum_squared_diff.sqrt().to_f64().unwrap()
}
pub fn dist_l2(a: &VecData, b: &VecData) -> f64 {
assert_eq!(a.dim(), b.dim());
match (a, b) {
(VecData::BF16(a_arr), VecData::BF16(b_arr)) => {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512bf16") {
unsafe {
return dist_l2_bf16_avx512(a_arr, b_arr);
}
}
}
dist_l2_scalar(&AV::from(a_arr), &AV::from(b_arr))
}
(VecData::F16(a_arr), VecData::F16(b_arr)) => {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512fp16") {
unsafe {
return dist_l2_f16_avx512(a_arr, b_arr);
}
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") && is_aarch64_feature_detected!("fp16") {
unsafe {
return dist_l2_f16_neon(a_arr, b_arr);
}
}
}
dist_l2_scalar(&AV::from(a_arr), &AV::from(b_arr))
}
(VecData::F32(a_arr), VecData::F32(b_arr)) => {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if is_x86_feature_detected!("avx512f") {
unsafe {
return dist_l2_f32_avx512(a_arr, b_arr);
}
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
unsafe {
return dist_l2_f32_neon(a_arr, b_arr);
}
}
}
dist_l2_scalar(&AV::from(a_arr), &AV::from(b_arr))
}
(VecData::F64(a_arr), VecData::F64(b_arr)) => {
#[cfg(any(target_arch = "x86", target_arch = "x86_64"))]
{
if is_x86_feature_detected!("avx512f") {
unsafe {
return dist_l2_f64_avx512(a_arr, b_arr);
}
}
}
#[cfg(target_arch = "aarch64")]
{
if is_aarch64_feature_detected!("neon") {
unsafe {
return dist_l2_f64_neon(a_arr, b_arr);
}
}
}
dist_l2_scalar(&AV::from(a_arr), &AV::from(b_arr))
}
_ => panic!("differing dtypes"),
}
}