Skip to main content

hermes_simd/dispatch/
gemm.rs

1//! Generic runtime-dispatch tiled GEMM kernel.
2//!
3//! # Tiling policy selection
4//!
5//! The tiling policy is selected solely from the **runtime-dispatched
6//! architecture** `A`, not from a secondary hardware-detection pass.  This
7//! avoids a historical double-dispatch bug where `AdaptiveDispatcher` would
8//! detect AVX-512 hardware and route the AVX2 kernel through `TilingPolicy<3,
9//! 4>` (17 live registers → spill on AVX2's 16-register file, measured
10//! 30–60 % slower at 256²).  The architecture parameter `A` already encodes
11//! the correct register file width via `A::LANE_COUNT`.
12//!
13//! | `A::LANE_COUNT` | Register file    | Tiling policy | Live registers |
14//! |-----------------|------------------|---------------|----------------|
15//! | > 8             | AVX-512 (32 regs)| `<6, 4>`      | 24+4+1 = 29    |
16//! | > 1             | AVX2 (16 regs)   | `<3, 3>`      | 9+3+1 = 13     |
17//! | > 1             | NEON (32 regs)   | `<3, 3>`      | 9+3+1 = 13     |
18//! | 1               | Scalar           | `<1, 1>`      | 1+1+1 = 3      |
19
20use hermes_simd_core::{
21    align::Unaligned,
22    arch::SimdArch,
23    execution::Unmasked,
24    kernel::SimdKernel,
25    scalar::Scalar,
26    view::{SimdError, SimdView},
27};
28use hermes_simd_macros::runtime_dispatch;
29
30#[runtime_dispatch(avx512f, avx2, neon, scalar)]
31pub(super) fn dispatch_tiled_gemm_kernel<T, A>(
32    a: &[T],
33    b: &[T],
34    c: &mut [T],
35    m: usize,
36    n: usize,
37    k: usize,
38) -> Result<(), SimdError>
39where
40    T: Scalar,
41    A: SimdArch + SimdKernel<T>,
42{
43    match (
44        SimdView::<T, A, Unaligned, Unmasked, &[T]>::new(a),
45        SimdView::<T, A, Unaligned, Unmasked, &[T]>::new(b),
46    ) {
47        (Some(v1), Some(v2)) => {
48            use hermes_simd_core::tiling::{TilingPolicy, TilingStrategy};
49
50            if A::LANE_COUNT > 8 && m >= 16 && n >= 16 && k >= 32 {
51                // AVX-512 class (32 vector registers). `<6, 4>` holds 24
52                // accumulators + 4 B-vectors + 1 broadcast A-scalar = 29
53                // registers, fitting within the 32-register file with
54                // headroom for loop temporaries.
55                <TilingPolicy<6, 4> as TilingStrategy<T, A, Unaligned>>::gemm(&v1, &v2, c, m, n, k)
56            } else if A::LANE_COUNT > 1 && m >= 16 && n >= 16 && k >= 32 {
57                // AVX2/NEON class (16 vector registers). `<3, 3>` holds 9
58                // accumulators + 3 B-vectors + 1 broadcast A-scalar = 13
59                // registers, leaving headroom for loop temporaries (no
60                // spill). `<3,4>` (12+4+1 = 17) and `<4,3>` (12+3+1 = 16,
61                // zero headroom) both spill on a 16-register file and were
62                // measured ~30-60% slower at 256².
63                <TilingPolicy<3, 3> as TilingStrategy<T, A, Unaligned>>::gemm(&v1, &v2, c, m, n, k)
64            } else {
65                <TilingPolicy<1, 1> as TilingStrategy<T, A, Unaligned>>::gemm(&v1, &v2, c, m, n, k)
66            }
67        }
68        _ => unsafe { core::hint::unreachable_unchecked() },
69    }
70}
71
72#[cfg(test)]
73mod tests {
74    use crate::dispatch::gemm;
75    use hermes_simd_core::view::SimdError;
76
77    #[test]
78    fn tiled_gemm_rejects_dimension_overflow() {
79        // `m·k` overflows `usize`. Unchecked (release `overflow-checks = false`)
80        // the product wraps, the `a_len < a_needed` guard passes, and the kernel
81        // issues an OOB SIMD load/store. The checked area arithmetic must reject
82        // with the exact variant. (Operand correctness is covered in
83        // `tests/tiling_tests.rs`.)
84        let a = vec![1.0f64; 16];
85        let b = vec![1.0f64; 16];
86        let mut c = vec![0.0f64; 16];
87        let r = gemm::dispatch_tiled_gemm::<f64>(&a, &b, &mut c, 2, 8, usize::MAX);
88        assert_eq!(r, Err(SimdError::LengthMismatch));
89    }
90}