use crate::{
align::Alignment,
arch::SimdArch,
kernel::SimdKernel,
scalar::Scalar,
view::{SimdError, SimdView},
};
#[inline]
pub(super) fn dot_impl<T, Arch, Align, const TILE_M: usize>(
a: &SimdView<'_, T, Arch, Align>,
b: &SimdView<'_, T, Arch, Align>,
) -> Result<T, SimdError>
where
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
T: Scalar,
{
struct AssertM<const TILE_M: usize>;
impl<const TILE_M: usize> AssertM<TILE_M> {
const OK: () = assert!(TILE_M >= 1, "TILE_M must be at least 1");
}
let _ = AssertM::<TILE_M>::OK;
crate::view::check_lengths_equal(a.len(), b.len())?;
let len = a.len();
let lane_count = Arch::LANE_COUNT;
let tile_width = lane_count * TILE_M;
let tiled_len = (len / tile_width) * tile_width;
let load = |ptr: *const T| -> Arch::Vector {
if crate::align::is_aligned_for_arch::<Arch, Align>() {
unsafe { Arch::load_aligned(ptr) }
} else {
unsafe { Arch::load_unaligned(ptr) }
}
};
let base_a = a.as_slice().as_ptr();
let base_b = b.as_slice().as_ptr();
let mut total: T = unsafe {
let mut ptr_a = base_a;
let mut ptr_b = base_b;
let mut accumulators: [Arch::Vector; TILE_M] = [Arch::zero(); TILE_M];
if tiled_len > 0 {
for i in 0..TILE_M {
let va = load(ptr_a.add(i * lane_count));
let vb = load(ptr_b.add(i * lane_count));
accumulators[i] = Arch::mul(va, vb);
}
ptr_a = ptr_a.add(tile_width);
ptr_b = ptr_b.add(tile_width);
}
if tiled_len > tile_width {
let iterations = (tiled_len / tile_width) - 1;
for _ in 0..iterations {
for i in 0..TILE_M {
let va = load(ptr_a.add(i * lane_count));
let vb = load(ptr_b.add(i * lane_count));
accumulators[i] = Arch::fmadd(va, vb, accumulators[i]);
}
ptr_a = ptr_a.add(tile_width);
ptr_b = ptr_b.add(tile_width);
}
}
let mut total = T::ZERO;
if tiled_len > 0 {
let mut combined = accumulators[0];
for acc in accumulators.iter().take(TILE_M).skip(1) {
combined = Arch::add(combined, *acc);
}
total = Arch::sum_reduce(combined);
}
total
};
let a_slice = a.as_slice();
let b_slice = b.as_slice();
for i in tiled_len..len {
total += a_slice[i] * b_slice[i];
}
Ok(total)
}