1use hermes_simd_core::{
13 align::Unaligned,
14 arch::SimdArch,
15 execution::Unmasked,
16 kernel::SimdKernel,
17 scalar::Scalar,
18 view::{SimdError, SimdView},
19};
20use hermes_simd_macros::runtime_dispatch;
21
22#[runtime_dispatch(avx512f, avx2, neon, scalar)]
23pub(super) fn dispatch_gemv_transpose_strided_kernel<T, A>(
24 a: &[T],
25 x: &[T],
26 y: &mut [T],
27 nrows: usize,
28 ncols: usize,
29 lda: usize,
30) -> Result<(), SimdError>
31where
32 T: Scalar,
33 A: SimdArch + SimdKernel<T>,
34{
35 match (
36 SimdView::<T, A, Unaligned, Unmasked, &[T]>::new(a),
37 SimdView::<T, A, Unaligned, Unmasked, &[T]>::new(x),
38 ) {
39 (Some(va), Some(vx)) => {
40 use hermes_simd_core::tiling::{TilingPolicy, TilingStrategy};
41 if A::LANE_COUNT > 8 {
42 <TilingPolicy<1, 8> as TilingStrategy<T, A, Unaligned>>::gemv_transpose_strided(
43 &va, &vx, y, nrows, ncols, lda,
44 )
45 } else if A::LANE_COUNT > 1 {
46 <TilingPolicy<1, 4> as TilingStrategy<T, A, Unaligned>>::gemv_transpose_strided(
47 &va, &vx, y, nrows, ncols, lda,
48 )
49 } else {
50 <TilingPolicy<1, 1> as TilingStrategy<T, A, Unaligned>>::gemv_transpose_strided(
51 &va, &vx, y, nrows, ncols, lda,
52 )
53 }
54 }
55 _ => unsafe { core::hint::unreachable_unchecked() },
58 }
59}
60
61#[cfg(test)]
62mod tests {
63 use crate::dispatch::{gemv_transpose, gemv_transpose_strided};
64
65 fn reference(a: &[f64], x: &[f64], nrows: usize, ncols: usize, lda: usize) -> Vec<f64> {
67 let mut y = vec![0.0f64; ncols];
68 for (i, &xi) in x.iter().enumerate().take(nrows) {
69 for (j, yj) in y.iter_mut().enumerate() {
70 *yj += a[i * lda + j] * xi;
71 }
72 }
73 y
74 }
75
76 #[test]
77 fn gemv_transpose_strided_matches_reference_over_submatrix() {
78 let lda = 11usize;
79 let (nrows, ncols) = (5usize, 7usize);
80 let a: Vec<f64> = (0..nrows * lda)
81 .map(|i| ((i % 9) as f64 - 4.0) * 0.25)
82 .collect();
83 let x: Vec<f64> = (0..nrows).map(|i| ((i % 5) as f64 - 2.0) * 0.5).collect();
84 let mut y = vec![0.0f64; ncols];
85 gemv_transpose_strided::dispatch_gemv_transpose_strided::<f64>(
86 &a, &x, &mut y, nrows, ncols, lda,
87 )
88 .unwrap();
89 assert_eq!(y, reference(&a, &x, nrows, ncols, lda));
90 }
91
92 #[test]
93 fn gemv_transpose_strided_packed_equals_gemv_transpose() {
94 let (nrows, ncols) = (9usize, 13usize);
95 let a: Vec<f64> = (0..nrows * ncols)
96 .map(|i| ((i % 7) as f64 - 3.0) * 0.5)
97 .collect();
98 let x: Vec<f64> = (0..nrows).map(|i| ((i % 4) as f64 - 1.0) * 0.25).collect();
99 let mut y_s = vec![0.0f64; ncols];
100 let mut y_p = vec![0.0f64; ncols];
101 gemv_transpose_strided::dispatch_gemv_transpose_strided::<f64>(
102 &a, &x, &mut y_s, nrows, ncols, ncols,
103 )
104 .unwrap();
105 gemv_transpose::dispatch_gemv_transpose::<f64>(&a, &x, &mut y_p, nrows, ncols).unwrap();
106 assert_eq!(y_s, y_p);
107 }
108
109 #[test]
110 fn gemv_transpose_strided_rejects_invalid() {
111 let a = vec![1.0f64; 40];
112 let x = vec![1.0f64; 5];
113 let mut y = vec![0.0f64; 7];
114 assert!(
116 gemv_transpose_strided::dispatch_gemv_transpose_strided::<f64>(&a, &x, &mut y, 5, 7, 6)
117 .is_err()
118 );
119 }
120
121 #[test]
122 fn gemv_transpose_strided_rejects_dimension_overflow() {
123 use hermes_simd_core::view::SimdError;
124 let a = vec![1.0f64; 40];
128 let x = vec![1.0f64; 2];
129 let mut y = vec![0.0f64; 6];
130 let r = gemv_transpose_strided::dispatch_gemv_transpose_strided::<f64>(
131 &a,
132 &x,
133 &mut y,
134 2,
135 6,
136 usize::MAX,
137 );
138 assert_eq!(r, Err(SimdError::LengthMismatch));
139 }
140}