use crate::{
align::Alignment,
arch::SimdArch,
kernel::SimdKernel,
scalar::Scalar,
view::{SimdError, SimdView},
};
#[inline(never)]
fn check_gemv_dimensions(
a_len: usize,
x_len: usize,
y_len: usize,
nrows: usize,
ncols: usize,
lda: usize,
) -> Result<(), SimdError> {
let a_needed =
super::dims::checked_strided_span(nrows, ncols, lda).ok_or(SimdError::LengthMismatch)?;
if lda < ncols || a_len < a_needed || x_len < ncols || y_len < nrows {
return Err(SimdError::LengthMismatch);
}
Ok(())
}
#[inline]
pub(super) fn gemv_impl<T, Arch, Align, const TILE_M: usize>(
a: &SimdView<'_, T, Arch, Align>,
x: &SimdView<'_, T, Arch, Align>,
y: &mut [T],
nrows: usize,
ncols: usize,
) -> Result<(), SimdError>
where
Arch: SimdArch + SimdKernel<T>,
Align: Alignment,
T: Scalar,
{
gemv_strided_impl::<T, Arch, Align, TILE_M>(a, x, y, nrows, ncols, ncols)
}
#[inline]
pub(super) fn gemv_strided_impl<T, Arch, Align, const TILE_M: usize>(
a: &SimdView<'_, T, Arch, Align>,
x: &SimdView<'_, T, Arch, Align>,
y: &mut [T],
nrows: usize,
ncols: usize,
lda: usize,
) -> Result<(), 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;
check_gemv_dimensions(a.len(), x.len(), y.len(), nrows, ncols, lda)?;
let a_slice = a.as_slice();
let x_slice = x.as_slice();
let lane_count = Arch::LANE_COUNT;
let simd_len = (ncols / lane_count) * lane_count;
let tail = ncols - simd_len;
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 (tail_mask, x_tail) = unsafe {
let tail_mask = Arch::leading_k_mask(tail);
let x_tail = if tail > 0 {
Arch::masked_load_unaligned(x_slice.as_ptr().add(simd_len), tail_mask, Arch::zero())
} else {
Arch::zero()
};
(tail_mask, x_tail)
};
let mut r = 0;
while r + TILE_M <= nrows {
unsafe {
let mut accumulators = [Arch::zero(); TILE_M];
let mut c = 0;
while c < simd_len {
let x_vec = load(x_slice.as_ptr().add(c));
for i in 0..TILE_M {
let a_vec = load(a_slice.as_ptr().add((r + i) * lda + c));
accumulators[i] = Arch::fmadd(a_vec, x_vec, accumulators[i]);
}
c += lane_count;
}
for i in 0..TILE_M {
let row_idx = r + i;
if tail > 0 {
let a_tail = Arch::masked_load_unaligned(
a_slice.as_ptr().add(row_idx * lda + simd_len),
tail_mask,
Arch::zero(),
);
accumulators[i] = Arch::fmadd(a_tail, x_tail, accumulators[i]);
}
y[row_idx] += Arch::sum_reduce(accumulators[i]);
}
}
r += TILE_M;
}
while r < nrows {
unsafe {
let mut acc = Arch::zero();
let mut c = 0;
while c < simd_len {
let x_vec = load(x_slice.as_ptr().add(c));
let a_vec = load(a_slice.as_ptr().add(r * lda + c));
acc = Arch::fmadd(a_vec, x_vec, acc);
c += lane_count;
}
if tail > 0 {
let a_tail = Arch::masked_load_unaligned(
a_slice.as_ptr().add(r * lda + simd_len),
tail_mask,
Arch::zero(),
);
acc = Arch::fmadd(a_tail, x_tail, acc);
}
y[r] += Arch::sum_reduce(acc);
}
r += 1;
}
Ok(())
}