use hermes_simd_core::{
align::Unaligned,
arch::SimdArch,
execution::Unmasked,
kernel::SimdKernel,
scalar::Scalar,
view::{SimdError, SimdView},
};
use hermes_simd_macros::runtime_dispatch;
#[runtime_dispatch(avx512f, avx2, neon, scalar)]
pub(super) fn dispatch_gemv_kernel<T, A>(
a: &[T],
x: &[T],
y: &mut [T],
nrows: usize,
ncols: usize,
) -> Result<(), SimdError>
where
T: Scalar,
A: SimdArch + SimdKernel<T>,
{
match (
SimdView::<T, A, Unaligned, Unmasked, &[T]>::new(a),
SimdView::<T, A, Unaligned, Unmasked, &[T]>::new(x),
) {
(Some(va), Some(vx)) => {
use hermes_simd_core::tiling::{TilingPolicy, TilingStrategy};
if A::LANE_COUNT > 8 {
<TilingPolicy<8, 1> as TilingStrategy<T, A, Unaligned>>::gemv(
&va, &vx, y, nrows, ncols,
)
} else if A::LANE_COUNT > 1 {
<TilingPolicy<4, 1> as TilingStrategy<T, A, Unaligned>>::gemv(
&va, &vx, y, nrows, ncols,
)
} else {
<TilingPolicy<1, 1> as TilingStrategy<T, A, Unaligned>>::gemv(
&va, &vx, y, nrows, ncols,
)
}
}
_ => unsafe { core::hint::unreachable_unchecked() },
}
}
#[cfg(test)]
mod tests {
use crate::dispatch::gemv;
fn reference(a: &[f64], x: &[f64], nrows: usize, ncols: usize) -> Vec<f64> {
(0..nrows)
.map(|r| (0..ncols).map(|c| a[r * ncols + c] * x[c]).sum())
.collect()
}
fn run_case(nrows: usize, ncols: usize) {
let a: Vec<f64> = (0..nrows * ncols)
.map(|i| ((i % 7) as f64 - 3.0) * 0.25)
.collect();
let x: Vec<f64> = (0..ncols).map(|i| ((i % 5) as f64 - 2.0) * 0.5).collect();
let mut y = vec![0.0f64; nrows];
gemv::dispatch_gemv::<f64>(&a, &x, &mut y, nrows, ncols).unwrap();
let want = reference(&a, &x, nrows, ncols);
assert_eq!(y, want, "gemv {nrows}x{ncols} mismatch vs reference");
}
#[test]
fn gemv_matches_reference_across_shapes() {
for &(m, n) in &[
(1, 1),
(1, 17),
(3, 4),
(8, 8),
(9, 13),
(16, 1),
(17, 31),
(33, 64),
(64, 64),
] {
run_case(m, n);
}
}
#[test]
fn gemv_accumulates_into_y() {
let a = vec![1.0f64, 2.0, 3.0, 4.0]; let x = vec![1.0f64, 1.0];
let mut y = vec![10.0f64, 20.0];
gemv::dispatch_gemv::<f64>(&a, &x, &mut y, 2, 2).unwrap();
assert_eq!(y, vec![13.0, 27.0]);
}
#[test]
fn gemv_rejects_short_operands() {
let a = vec![1.0f64; 4];
let x = vec![1.0f64; 2];
let mut y = vec![0.0f64; 1]; assert!(gemv::dispatch_gemv::<f64>(&a, &x, &mut y, 2, 2).is_err());
}
}