#[cfg_attr(all(target_arch = "wasm32", target_feature = "simd128"), allow(dead_code))]
#[inline]
pub(crate) fn dot_scalar(a: &[f32], b: &[f32]) -> f32 {
let mut sum = 0.0f32;
for (x, y) in a.iter().zip(b.iter()) {
sum += x * y;
}
sum
}
#[cfg_attr(all(target_arch = "wasm32", target_feature = "simd128"), allow(dead_code))]
#[inline]
pub(crate) fn l2_sq_scalar(a: &[f32], b: &[f32]) -> f32 {
let mut sum = 0.0f32;
for (x, y) in a.iter().zip(b.iter()) {
let diff = x - y;
sum += diff * diff;
}
sum
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[inline]
pub(crate) fn dot(a: &[f32], b: &[f32]) -> f32 {
simd128::dot(a, b)
}
#[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
#[inline]
pub(crate) fn dot(a: &[f32], b: &[f32]) -> f32 {
dot_scalar(a, b)
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[inline]
pub(crate) fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
simd128::l2_sq(a, b)
}
#[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
#[inline]
pub(crate) fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
l2_sq_scalar(a, b)
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
#[allow(unsafe_code)]
mod simd128 {
use core::arch::wasm32::*;
#[inline]
pub(crate) fn dot(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let chunks = n / 4;
let mut acc = f32x4_splat(0.0);
for i in 0..chunks {
let off = i * 4;
let va = unsafe { v128_load(a.as_ptr().add(off) as *const v128) };
let vb = unsafe { v128_load(b.as_ptr().add(off) as *const v128) };
acc = f32x4_add(acc, f32x4_mul(va, vb));
}
let mut sum = f32x4_extract_lane::<0>(acc)
+ f32x4_extract_lane::<1>(acc)
+ f32x4_extract_lane::<2>(acc)
+ f32x4_extract_lane::<3>(acc);
for i in (chunks * 4)..n {
sum += a[i] * b[i];
}
sum
}
#[inline]
pub(crate) fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let n = a.len();
let chunks = n / 4;
let mut acc = f32x4_splat(0.0);
for i in 0..chunks {
let off = i * 4;
let va = unsafe { v128_load(a.as_ptr().add(off) as *const v128) };
let vb = unsafe { v128_load(b.as_ptr().add(off) as *const v128) };
let diff = f32x4_sub(va, vb);
acc = f32x4_add(acc, f32x4_mul(diff, diff));
}
let mut sum = f32x4_extract_lane::<0>(acc)
+ f32x4_extract_lane::<1>(acc)
+ f32x4_extract_lane::<2>(acc)
+ f32x4_extract_lane::<3>(acc);
for i in (chunks * 4)..n {
let d = a[i] - b[i];
sum += d * d;
}
sum
}
}