use crate::{
align::Alignment,
arch::SimdArch,
kernel::SimdKernel,
scalar::Scalar,
view::{SimdError, SimdView},
};
#[inline(never)]
fn check_gemv_t_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 < nrows || y_len < ncols {
return Err(SimdError::LengthMismatch);
}
Ok(())
}
#[inline]
pub(super) fn gemv_transpose_impl<T, Arch, Align, const TILE_N: 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_transpose_strided_impl::<T, Arch, Align, TILE_N>(a, x, y, nrows, ncols, ncols)
}
#[inline]
pub(super) fn gemv_transpose_strided_impl<T, Arch, Align, const TILE_N: 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 AssertN<const TILE_N: usize>;
impl<const TILE_N: usize> AssertN<TILE_N> {
const OK: () = assert!(TILE_N >= 1, "TILE_N must be at least 1");
}
let _ = AssertN::<TILE_N>::OK;
check_gemv_t_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_cols = (ncols / lane_count) * lane_count;
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 mut c = 0;
while c + TILE_N * lane_count <= simd_cols {
unsafe {
let mut acc = [Arch::zero(); TILE_N];
for (t, slot) in acc.iter_mut().enumerate() {
*slot = Arch::load_unaligned(y.as_ptr().add(c + t * lane_count));
}
for i in 0..nrows {
let xi = Arch::splat(x_slice[i]);
let base = i * lda + c;
for (t, slot) in acc.iter_mut().enumerate() {
let a_vec = load(a_slice.as_ptr().add(base + t * lane_count));
*slot = Arch::fmadd(xi, a_vec, *slot);
}
}
for (t, &accv) in acc.iter().enumerate() {
Arch::store_unaligned(y.as_mut_ptr().add(c + t * lane_count), accv);
}
}
c += TILE_N * lane_count;
}
while c < simd_cols {
unsafe {
let mut acc = Arch::load_unaligned(y.as_ptr().add(c));
for i in 0..nrows {
let xi = Arch::splat(x_slice[i]);
let a_vec = load(a_slice.as_ptr().add(i * lda + c));
acc = Arch::fmadd(xi, a_vec, acc);
}
Arch::store_unaligned(y.as_mut_ptr().add(c), acc);
}
c += lane_count;
}
unsafe {
let a_ptr = a_slice.as_ptr();
let x_ptr = x_slice.as_ptr();
for c_tail in simd_cols..ncols {
let mut s = *y.as_mut_ptr().add(c_tail);
for i in 0..nrows {
s = s + *x_ptr.add(i) * *a_ptr.add(i * lda + c_tail);
}
*y.as_mut_ptr().add(c_tail) = s;
}
}
Ok(())
}