Skip to main content

hermes_simd/tile_matmul/
int8.rs

1use super::{tile_loop_generic, validate_gemm_sizes, TiledGemm};
2#[cfg(target_arch = "x86_64")]
3use crate::cpu::{AmxSupport, Avx512Support};
4use eunomia::{I32, I8};
5use hermes_simd_core::view::{SimdError, TileMatrixMultiply};
6use hermes_simd_intrinsics::Scalar;
7#[cfg(target_arch = "x86_64")]
8use hermes_simd_intrinsics::{AmxInt8, Avx512, AvxVnni};
9
10/// Scalar cleanup for the rows/columns/K-depth the 16×16×64 tile loop leaves
11/// uncovered (`m % 16` rows, `n % 16` columns, `k % 64` tail on covered tiles).
12///
13/// One implementation shared by every dispatch arm and both element newtypes so
14/// the wrapping accumulation semantics stay identical across backends.
15fn gemm_i8_remainder(
16    m: usize,
17    n: usize,
18    k: usize,
19    a: &[i8],
20    a_stride: usize,
21    b: &[i8],
22    b_stride: usize,
23    c: &mut [i32],
24    c_stride: usize,
25) {
26    let tile_m_bound = (m / 16) * 16;
27    let tile_n_bound = (n / 16) * 16;
28    let tile_k_bound = (k / 64) * 64;
29    for r in 0..m {
30        for col in 0..n {
31            if r >= tile_m_bound || col >= tile_n_bound {
32                let mut sum = 0i32;
33                for kk in 0..k {
34                    sum = sum.wrapping_add(
35                        (a[r * a_stride + kk] as i32) * (b[kk * b_stride + col] as i32),
36                    );
37                }
38                c[r * c_stride + col] += sum;
39            } else if tile_k_bound < k {
40                let mut sum = 0i32;
41                for kk in tile_k_bound..k {
42                    sum = sum.wrapping_add(
43                        (a[r * a_stride + kk] as i32) * (b[kk * b_stride + col] as i32),
44                    );
45                }
46                c[r * c_stride + col] += sum;
47            }
48        }
49    }
50}
51
52/// Dispatched int8 GEMM body shared by the `i8` and `I8` trait impls (`I8`/`I32`
53/// are `#[repr(transparent)]` over `i8`/`i32`, so the newtype impl delegates via
54/// layout-preserving slice casts).
55///
56/// Backend ladder: AMX → AVX-512 VNNI → AVX-VNNI (256-bit) → scalar, each tier
57/// entered only after its runtime probe passes.
58///
59/// # Safety
60/// Caller must have validated operand extents via `validate_gemm_sizes`.
61unsafe fn gemm_i8_dispatched(
62    m: usize,
63    n: usize,
64    k: usize,
65    a: &[i8],
66    a_stride: usize,
67    b: &[i8],
68    b_stride: usize,
69    c: &mut [i32],
70    c_stride: usize,
71) -> Result<(), SimdError> {
72    #[cfg(target_arch = "x86_64")]
73    {
74        let decision = crate::dispatcher::AdaptiveDispatcher::select_backend(
75            m,
76            n,
77            k,
78            a.as_ptr(),
79            a.len(),
80            b.as_ptr(),
81            b.len(),
82        );
83
84        match decision {
85            crate::dispatcher::DispatchDecision::Amx => {
86                <AmxInt8 as hermes_simd_intrinsics::x86_64::amx::AmxGemm<i8, i8, i32>>::amx_gemm(
87                    m,
88                    n,
89                    k,
90                    a.as_ptr(),
91                    a_stride,
92                    b.as_ptr(),
93                    b_stride,
94                    c.as_mut_ptr(),
95                    c_stride,
96                );
97                return Ok(());
98            }
99            crate::dispatcher::DispatchDecision::Avx512 => {
100                tile_loop_generic::<i8, i8, i32, Avx512, 16, 16, 64>(
101                    m,
102                    n,
103                    k,
104                    a.as_ptr(),
105                    a_stride,
106                    b.as_ptr(),
107                    b_stride,
108                    c.as_mut_ptr(),
109                    c_stride,
110                );
111                gemm_i8_remainder(m, n, k, a, a_stride, b, b_stride, c, c_stride);
112                return Ok(());
113            }
114            crate::dispatcher::DispatchDecision::AvxVnni => {
115                tile_loop_generic::<i8, i8, i32, AvxVnni, 16, 16, 64>(
116                    m,
117                    n,
118                    k,
119                    a.as_ptr(),
120                    a_stride,
121                    b.as_ptr(),
122                    b_stride,
123                    c.as_mut_ptr(),
124                    c_stride,
125                );
126                gemm_i8_remainder(m, n, k, a, a_stride, b, b_stride, c, c_stride);
127                return Ok(());
128            }
129            crate::dispatcher::DispatchDecision::Scalar => {}
130        }
131    }
132
133    tile_loop_generic::<i8, i8, i32, Scalar, 16, 16, 64>(
134        m,
135        n,
136        k,
137        a.as_ptr(),
138        a_stride,
139        b.as_ptr(),
140        b_stride,
141        c.as_mut_ptr(),
142        c_stride,
143    );
144    gemm_i8_remainder(m, n, k, a, a_stride, b, b_stride, c, c_stride);
145    Ok(())
146}
147
148impl TiledGemm<i8, i8, i32> for (i8, i8, i32) {
149    #[inline]
150    unsafe fn dispatch_tile_matmul(
151        c: *mut i32,
152        c_stride: usize,
153        a: *const i8,
154        a_stride: usize,
155        b: *const i8,
156        b_stride: usize,
157    ) {
158        #[cfg(target_arch = "x86_64")]
159        {
160            if <i8 as AmxSupport>::has_amx() && hermes_simd_intrinsics::AmxSession::is_active() {
161                return <AmxInt8 as TileMatrixMultiply<
162                    i8,
163                    i8,
164                    i32,
165                    AmxInt8,
166                    AmxInt8,
167                    16,
168                    16,
169                    64,
170                >>::tile_matmul(c, c_stride, a, a_stride, b, b_stride);
171            }
172            if <i8 as Avx512Support>::has_avx512() {
173                return <Avx512 as TileMatrixMultiply<i8, i8, i32, Avx512, Avx512, 16, 16, 64>>::tile_matmul(c, c_stride, a, a_stride, b, b_stride);
174            }
175            if crate::cpu::has_avx_vnni() {
176                return <AvxVnni as TileMatrixMultiply<
177                    i8,
178                    i8,
179                    i32,
180                    AvxVnni,
181                    AvxVnni,
182                    16,
183                    16,
184                    64,
185                >>::tile_matmul(c, c_stride, a, a_stride, b, b_stride);
186            }
187        }
188        <Scalar as TileMatrixMultiply<i8, i8, i32, Scalar, Scalar, 16, 16, 64>>::tile_matmul(
189            c, c_stride, a, a_stride, b, b_stride,
190        );
191    }
192
193    #[inline]
194    unsafe fn gemm(
195        m: usize,
196        n: usize,
197        k: usize,
198        a: &[i8],
199        a_stride: usize,
200        b: &[i8],
201        b_stride: usize,
202        c: &mut [i32],
203        c_stride: usize,
204    ) -> Result<(), SimdError> {
205        validate_gemm_sizes(
206            a.len(),
207            b.len(),
208            c.len(),
209            m,
210            n,
211            k,
212            a_stride,
213            b_stride,
214            c_stride,
215        )?;
216        gemm_i8_dispatched(m, n, k, a, a_stride, b, b_stride, c, c_stride)
217    }
218}
219
220impl TiledGemm<I8, I8, I32> for (I8, I8, I32) {
221    #[inline]
222    unsafe fn dispatch_tile_matmul(
223        c: *mut I32,
224        c_stride: usize,
225        a: *const I8,
226        a_stride: usize,
227        b: *const I8,
228        b_stride: usize,
229    ) {
230        // SAFETY: `I8`/`I32` are `#[repr(transparent)]` over `i8`/`i32`, so the
231        // pointer casts preserve layout and the `i8` dispatcher's tile contract.
232        <(i8, i8, i32) as TiledGemm<i8, i8, i32>>::dispatch_tile_matmul(
233            c as *mut i32,
234            c_stride,
235            a as *const i8,
236            a_stride,
237            b as *const i8,
238            b_stride,
239        );
240    }
241
242    #[inline]
243    unsafe fn gemm(
244        m: usize,
245        n: usize,
246        k: usize,
247        a: &[I8],
248        a_stride: usize,
249        b: &[I8],
250        b_stride: usize,
251        c: &mut [I32],
252        c_stride: usize,
253    ) -> Result<(), SimdError> {
254        validate_gemm_sizes(
255            a.len(),
256            b.len(),
257            c.len(),
258            m,
259            n,
260            k,
261            a_stride,
262            b_stride,
263            c_stride,
264        )?;
265        // SAFETY: `I8`/`I32` are `#[repr(transparent)]` over `i8`/`i32`; the
266        // reborrowed slices alias the same memory with identical layout and
267        // length, and `c` is exclusively borrowed for the call's duration.
268        let a_raw = core::slice::from_raw_parts(a.as_ptr() as *const i8, a.len());
269        let b_raw = core::slice::from_raw_parts(b.as_ptr() as *const i8, b.len());
270        let c_raw = core::slice::from_raw_parts_mut(c.as_mut_ptr() as *mut i32, c.len());
271        gemm_i8_dispatched(m, n, k, a_raw, a_stride, b_raw, b_stride, c_raw, c_stride)
272    }
273}