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
10fn 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
52unsafe 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 <(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 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}