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_transpose_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<1, 8> as TilingStrategy<T, A, Unaligned>>::gemv_transpose(
&va, &vx, y, nrows, ncols,
)
} else if A::LANE_COUNT > 1 {
<TilingPolicy<1, 4> as TilingStrategy<T, A, Unaligned>>::gemv_transpose(
&va, &vx, y, nrows, ncols,
)
} else {
<TilingPolicy<1, 1> as TilingStrategy<T, A, Unaligned>>::gemv_transpose(
&va, &vx, y, nrows, ncols,
)
}
}
_ => unsafe { core::hint::unreachable_unchecked() },
}
}
#[cfg(test)]
mod tests {
use crate::dispatch::gemv_transpose;
fn reference(a: &[f64], x: &[f64], nrows: usize, ncols: usize) -> Vec<f64> {
let mut y = vec![0.0f64; ncols];
for (i, &xi) in x.iter().enumerate().take(nrows) {
for (j, yj) in y.iter_mut().enumerate() {
*yj += a[i * ncols + j] * xi;
}
}
y
}
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..nrows).map(|i| ((i % 5) as f64 - 2.0) * 0.5).collect();
let mut y = vec![0.0f64; ncols];
gemv_transpose::dispatch_gemv_transpose::<f64>(&a, &x, &mut y, nrows, ncols).unwrap();
let want = reference(&a, &x, nrows, ncols);
assert_eq!(
y, want,
"gemv_transpose {nrows}x{ncols} mismatch vs reference"
);
}
#[test]
fn gemv_transpose_matches_reference_across_shapes() {
for &(m, n) in &[
(1, 1),
(1, 33),
(4, 3),
(8, 8),
(13, 9),
(1, 64),
(31, 17),
(64, 33),
(64, 64),
] {
run_case(m, n);
}
}
#[test]
fn gemv_transpose_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_transpose::dispatch_gemv_transpose::<f64>(&a, &x, &mut y, 2, 2).unwrap();
assert_eq!(y, vec![14.0, 26.0]);
}
#[test]
fn gemv_transpose_rejects_short_operands() {
let a = vec![1.0f64; 4];
let x = vec![1.0f64; 2];
let mut y = vec![0.0f64; 1]; assert!(gemv_transpose::dispatch_gemv_transpose::<f64>(&a, &x, &mut y, 2, 2).is_err());
}
}