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}