Skip to main content

hermes_simd/dispatch/
gemv_transpose_strided.rs

1//! Generic runtime-dispatch transposed sub-matrix GEMV (`y += Aᵀ · x`, row
2//! stride `lda`).
3//!
4//! Generalizes [`super::gemv_transpose()`] to a row-major **sub-matrix**:
5//! `nrows × ncols` with leading dimension `lda ≥ ncols`. `lda = ncols` recovers
6//! the packed transpose. Computes `Σᵢ xᵢ·A[i,:]` (sum of the strided rows scaled
7//! by `x`), vectorizing across the `ncols` output lanes with no horizontal
8//! reduction. Admits the `Aᵀ·x` reduction over a trailing/leading block of a
9//! larger buffer — e.g. forming `Aw = Σⱼ wⱼ·colⱼ` in a reflector apply — without
10//! copying it out. Accumulates into `y`.
11
12use 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        // SAFETY: `Unaligned` skips the alignment check, so `SimdView::new` is
56        // `Some` for every slice — this arm is unreachable (mirrors `super::gemv`).
57        _ => unsafe { core::hint::unreachable_unchecked() },
58    }
59}
60
61#[cfg(test)]
62mod tests {
63    use crate::dispatch::{gemv_transpose, gemv_transpose_strided};
64
65    /// Naive `y = Aᵀ·x` over a sub-matrix with row stride `lda`.
66    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        // lda < ncols invalid.
115        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        // `(nrows-1)·lda + ncols` overflows `usize`; `x`/`y` are sized so only the
125        // A-span fails. Unchecked the product wraps and admits an OOB SIMD load —
126        // the checked span arithmetic rejects with the exact variant.
127        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}