use super::traits::SimdFloat;
const MAX_WIDTH: usize = 16;
#[inline(always)]
pub fn dot<S: SimdFloat>(a: &[f32], b: &[f32]) -> f32 {
let len = a.len().min(b.len());
let chunks = len / S::WIDTH;
let mut acc = S::zero();
for i in 0..chunks {
unsafe {
let va = S::load(a.as_ptr().add(i * S::WIDTH));
let vb = S::load(b.as_ptr().add(i * S::WIDTH));
acc = va.fma(vb, acc);
}
}
let mut lanes = [0.0f32; MAX_WIDTH];
unsafe { acc.store(lanes.as_mut_ptr()) };
let mut sum: f32 = lanes[..S::WIDTH].iter().sum();
for i in (chunks * S::WIDTH)..len {
sum += a[i] * b[i];
}
sum
}
#[inline(always)]
pub fn mul_elementwise<S: SimdFloat>(dst: &mut [f32], a: &[f32], b: &[f32]) {
let len = dst.len().min(a.len()).min(b.len());
let chunks = len / S::WIDTH;
for i in 0..chunks {
unsafe {
let va = S::load(a.as_ptr().add(i * S::WIDTH));
let vb = S::load(b.as_ptr().add(i * S::WIDTH));
va.mul(vb).store(dst.as_mut_ptr().add(i * S::WIDTH));
}
}
for i in (chunks * S::WIDTH)..len {
dst[i] = a[i] * b[i];
}
}
#[inline(always)]
pub fn mix_scalar<S: SimdFloat>(dst: &mut [f32], dry: &[f32], wet: &[f32], mix: f32) {
let len = dst.len().min(dry.len()).min(wet.len());
let chunks = len / S::WIDTH;
let vmix = S::splat(mix);
let vinv = S::splat(1.0 - mix);
for i in 0..chunks {
unsafe {
let vdry = S::load(dry.as_ptr().add(i * S::WIDTH));
let vwet = S::load(wet.as_ptr().add(i * S::WIDTH));
let dry_term = vdry.mul(vinv);
vwet.fma(vmix, dry_term)
.store(dst.as_mut_ptr().add(i * S::WIDTH));
}
}
for i in (chunks * S::WIDTH)..len {
dst[i] = dry[i] * (1.0 - mix) + wet[i] * mix;
}
}