1#[cfg(target_arch = "x86_64")]
27const GEMV_TILE_THRESHOLD: usize = 8192;
28
29#[cfg(target_arch = "x86_64")]
43#[target_feature(enable = "avx2", enable = "fma")]
44pub unsafe fn gemv_avx2(k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
45 unsafe {
46 use std::arch::x86_64::*;
47
48 let n8 = n / 8 * 8;
49
50 let k4 = k / 4 * 4;
52 let mut ki = 0;
53 while ki < k4 {
54 let a0 = _mm256_set1_ps(*a.get_unchecked(ki));
55 let a1 = _mm256_set1_ps(*a.get_unchecked(ki + 1));
56 let a2 = _mm256_set1_ps(*a.get_unchecked(ki + 2));
57 let a3 = _mm256_set1_ps(*a.get_unchecked(ki + 3));
58 let b0_base = ki * n;
59 let b1_base = b0_base + n;
60 let b2_base = b1_base + n;
61 let b3_base = b2_base + n;
62
63 let mut j = 0;
64 let b_ptr = b.as_ptr();
65 let c_ptr = c.as_mut_ptr();
66 while j < n8 {
67 let cv = _mm256_loadu_ps(c_ptr.add(j));
68 let bv0 = _mm256_loadu_ps(b_ptr.add(b0_base + j));
69 let bv1 = _mm256_loadu_ps(b_ptr.add(b1_base + j));
70 let bv2 = _mm256_loadu_ps(b_ptr.add(b2_base + j));
71 let bv3 = _mm256_loadu_ps(b_ptr.add(b3_base + j));
72
73 let r = _mm256_fmadd_ps(a0, bv0, cv);
74 let r = _mm256_fmadd_ps(a1, bv1, r);
75 let r = _mm256_fmadd_ps(a2, bv2, r);
76 let r = _mm256_fmadd_ps(a3, bv3, r);
77
78 _mm256_storeu_ps(c_ptr.add(j), r);
79 j += 8;
80 }
81
82 while j < n {
84 *c.get_unchecked_mut(j) += *a.get_unchecked(ki) * *b.get_unchecked(b0_base + j)
85 + *a.get_unchecked(ki + 1) * *b.get_unchecked(b1_base + j)
86 + *a.get_unchecked(ki + 2) * *b.get_unchecked(b2_base + j)
87 + *a.get_unchecked(ki + 3) * *b.get_unchecked(b3_base + j);
88 j += 1;
89 }
90
91 ki += 4;
92 }
93
94 while ki < k {
96 let ak = *a.get_unchecked(ki);
97 let bk_base = ki * n;
98 let ak_v = _mm256_set1_ps(ak);
99
100 let mut j = 0;
101 let b_ptr = b.as_ptr();
102 let c_ptr = c.as_mut_ptr();
103 while j < n8 {
104 let cv = _mm256_loadu_ps(c_ptr.add(j));
105 let bv = _mm256_loadu_ps(b_ptr.add(bk_base + j));
106 let r = _mm256_fmadd_ps(ak_v, bv, cv);
107 _mm256_storeu_ps(c_ptr.add(j), r);
108 j += 8;
109 }
110 while j < n {
111 *c.get_unchecked_mut(j) += ak * *b.get_unchecked(bk_base + j);
112 j += 1;
113 }
114 ki += 1;
115 }
116 }
117}
118
119#[cfg(target_arch = "x86_64")]
133#[target_feature(enable = "avx2", enable = "fma")]
134unsafe fn gemv_tiled_avx2(k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
135 unsafe {
136 use std::arch::x86_64::*;
137
138 const NT: usize = 64;
141
142 let k4 = k / 4 * 4;
143 let nt_end = n / NT * NT;
144
145 for j0 in (0..nt_end).step_by(NT) {
146 let mut acc0 = _mm256_setzero_ps();
148 let mut acc1 = _mm256_setzero_ps();
149 let mut acc2 = _mm256_setzero_ps();
150 let mut acc3 = _mm256_setzero_ps();
151 let mut acc4 = _mm256_setzero_ps();
152 let mut acc5 = _mm256_setzero_ps();
153 let mut acc6 = _mm256_setzero_ps();
154 let mut acc7 = _mm256_setzero_ps();
155
156 let mut ki = 0;
158 while ki < k4 {
159 let a0 = _mm256_set1_ps(*a.get_unchecked(ki));
160 let a1 = _mm256_set1_ps(*a.get_unchecked(ki + 1));
161 let a2 = _mm256_set1_ps(*a.get_unchecked(ki + 2));
162 let a3 = _mm256_set1_ps(*a.get_unchecked(ki + 3));
163
164 let b0 = ki * n + j0;
165 let b1 = b0 + n;
166 let b2 = b1 + n;
167 let b3 = b2 + n;
168
169 if ki + 8 < k {
171 let pf = (ki + 8) * n + j0;
172 _mm_prefetch(b.as_ptr().add(pf) as *const i8, _MM_HINT_T0);
173 _mm_prefetch(b.as_ptr().add(pf + 32) as *const i8, _MM_HINT_T0);
174 }
175
176 let bv = _mm256_loadu_ps(b.get_unchecked(b0));
178 acc0 = _mm256_fmadd_ps(a0, bv, acc0);
179 let bv = _mm256_loadu_ps(b.get_unchecked(b1));
180 acc0 = _mm256_fmadd_ps(a1, bv, acc0);
181 let bv = _mm256_loadu_ps(b.get_unchecked(b2));
182 acc0 = _mm256_fmadd_ps(a2, bv, acc0);
183 let bv = _mm256_loadu_ps(b.get_unchecked(b3));
184 acc0 = _mm256_fmadd_ps(a3, bv, acc0);
185
186 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 8));
187 acc1 = _mm256_fmadd_ps(a0, bv, acc1);
188 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 8));
189 acc1 = _mm256_fmadd_ps(a1, bv, acc1);
190 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 8));
191 acc1 = _mm256_fmadd_ps(a2, bv, acc1);
192 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 8));
193 acc1 = _mm256_fmadd_ps(a3, bv, acc1);
194
195 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 16));
196 acc2 = _mm256_fmadd_ps(a0, bv, acc2);
197 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 16));
198 acc2 = _mm256_fmadd_ps(a1, bv, acc2);
199 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 16));
200 acc2 = _mm256_fmadd_ps(a2, bv, acc2);
201 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 16));
202 acc2 = _mm256_fmadd_ps(a3, bv, acc2);
203
204 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 24));
205 acc3 = _mm256_fmadd_ps(a0, bv, acc3);
206 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 24));
207 acc3 = _mm256_fmadd_ps(a1, bv, acc3);
208 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 24));
209 acc3 = _mm256_fmadd_ps(a2, bv, acc3);
210 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 24));
211 acc3 = _mm256_fmadd_ps(a3, bv, acc3);
212
213 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 32));
214 acc4 = _mm256_fmadd_ps(a0, bv, acc4);
215 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 32));
216 acc4 = _mm256_fmadd_ps(a1, bv, acc4);
217 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 32));
218 acc4 = _mm256_fmadd_ps(a2, bv, acc4);
219 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 32));
220 acc4 = _mm256_fmadd_ps(a3, bv, acc4);
221
222 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 40));
223 acc5 = _mm256_fmadd_ps(a0, bv, acc5);
224 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 40));
225 acc5 = _mm256_fmadd_ps(a1, bv, acc5);
226 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 40));
227 acc5 = _mm256_fmadd_ps(a2, bv, acc5);
228 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 40));
229 acc5 = _mm256_fmadd_ps(a3, bv, acc5);
230
231 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 48));
232 acc6 = _mm256_fmadd_ps(a0, bv, acc6);
233 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 48));
234 acc6 = _mm256_fmadd_ps(a1, bv, acc6);
235 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 48));
236 acc6 = _mm256_fmadd_ps(a2, bv, acc6);
237 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 48));
238 acc6 = _mm256_fmadd_ps(a3, bv, acc6);
239
240 let bv = _mm256_loadu_ps(b.get_unchecked(b0 + 56));
241 acc7 = _mm256_fmadd_ps(a0, bv, acc7);
242 let bv = _mm256_loadu_ps(b.get_unchecked(b1 + 56));
243 acc7 = _mm256_fmadd_ps(a1, bv, acc7);
244 let bv = _mm256_loadu_ps(b.get_unchecked(b2 + 56));
245 acc7 = _mm256_fmadd_ps(a2, bv, acc7);
246 let bv = _mm256_loadu_ps(b.get_unchecked(b3 + 56));
247 acc7 = _mm256_fmadd_ps(a3, bv, acc7);
248
249 ki += 4;
250 }
251
252 while ki < k {
254 let av = _mm256_set1_ps(*a.get_unchecked(ki));
255 let base = ki * n + j0;
256
257 acc0 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base)), acc0);
258 acc1 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 8)), acc1);
259 acc2 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 16)), acc2);
260 acc3 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 24)), acc3);
261 acc4 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 32)), acc4);
262 acc5 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 40)), acc5);
263 acc6 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 48)), acc6);
264 acc7 = _mm256_fmadd_ps(av, _mm256_loadu_ps(b.get_unchecked(base + 56)), acc7);
265 ki += 1;
266 }
267
268 _mm256_storeu_ps(c.get_unchecked_mut(j0), acc0);
270 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 8), acc1);
271 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 16), acc2);
272 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 24), acc3);
273 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 32), acc4);
274 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 40), acc5);
275 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 48), acc6);
276 _mm256_storeu_ps(c.get_unchecked_mut(j0 + 56), acc7);
277 }
278
279 if nt_end < n {
281 let rem_n = n - nt_end;
282 let rem8 = rem_n / 8 * 8;
283 let k4 = k / 4 * 4;
284
285 let mut ki = 0;
286 while ki < k4 {
287 let a0 = _mm256_set1_ps(*a.get_unchecked(ki));
288 let a1 = _mm256_set1_ps(*a.get_unchecked(ki + 1));
289 let a2 = _mm256_set1_ps(*a.get_unchecked(ki + 2));
290 let a3 = _mm256_set1_ps(*a.get_unchecked(ki + 3));
291 let b0 = ki * n + nt_end;
292 let b1 = b0 + n;
293 let b2 = b1 + n;
294 let b3 = b2 + n;
295
296 let mut j = 0;
297 while j < rem8 {
298 let cv = _mm256_loadu_ps(c.get_unchecked(nt_end + j));
299 let r = _mm256_fmadd_ps(a0, _mm256_loadu_ps(b.get_unchecked(b0 + j)), cv);
300 let r = _mm256_fmadd_ps(a1, _mm256_loadu_ps(b.get_unchecked(b1 + j)), r);
301 let r = _mm256_fmadd_ps(a2, _mm256_loadu_ps(b.get_unchecked(b2 + j)), r);
302 let r = _mm256_fmadd_ps(a3, _mm256_loadu_ps(b.get_unchecked(b3 + j)), r);
303 _mm256_storeu_ps(c.get_unchecked_mut(nt_end + j), r);
304 j += 8;
305 }
306 while j < rem_n {
307 let idx = nt_end + j;
308 *c.get_unchecked_mut(idx) += *a.get_unchecked(ki) * *b.get_unchecked(b0 + j)
309 + *a.get_unchecked(ki + 1) * *b.get_unchecked(b1 + j)
310 + *a.get_unchecked(ki + 2) * *b.get_unchecked(b2 + j)
311 + *a.get_unchecked(ki + 3) * *b.get_unchecked(b3 + j);
312 j += 1;
313 }
314 ki += 4;
315 }
316
317 while ki < k {
318 let ak = *a.get_unchecked(ki);
319 let bk = ki * n + nt_end;
320 let ak_v = _mm256_set1_ps(ak);
321
322 let mut j = 0;
323 while j < rem8 {
324 let cv = _mm256_loadu_ps(c.get_unchecked(nt_end + j));
325 let bv = _mm256_loadu_ps(b.get_unchecked(bk + j));
326 _mm256_storeu_ps(
327 c.get_unchecked_mut(nt_end + j),
328 _mm256_fmadd_ps(ak_v, bv, cv),
329 );
330 j += 8;
331 }
332 while j < rem_n {
333 *c.get_unchecked_mut(nt_end + j) += ak * *b.get_unchecked(bk + j);
334 j += 1;
335 }
336 ki += 1;
337 }
338 }
339 }
340}
341
342#[cfg(target_arch = "x86_64")]
351#[target_feature(enable = "avx512f", enable = "fma")]
352#[allow(dead_code)] unsafe fn gemv_tiled_avx512(k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
354 unsafe {
355 use std::arch::x86_64::*;
356
357 const NT: usize = 128;
360
361 let k4 = k / 4 * 4;
362 let nt_end = n / NT * NT;
363
364 for j0 in (0..nt_end).step_by(NT) {
365 let mut acc0 = _mm512_setzero_ps();
367 let mut acc1 = _mm512_setzero_ps();
368 let mut acc2 = _mm512_setzero_ps();
369 let mut acc3 = _mm512_setzero_ps();
370 let mut acc4 = _mm512_setzero_ps();
371 let mut acc5 = _mm512_setzero_ps();
372 let mut acc6 = _mm512_setzero_ps();
373 let mut acc7 = _mm512_setzero_ps();
374
375 let mut ki = 0;
377 while ki < k4 {
378 let a0 = _mm512_set1_ps(*a.get_unchecked(ki));
379 let a1 = _mm512_set1_ps(*a.get_unchecked(ki + 1));
380 let a2 = _mm512_set1_ps(*a.get_unchecked(ki + 2));
381 let a3 = _mm512_set1_ps(*a.get_unchecked(ki + 3));
382
383 let b0 = ki * n + j0;
384 let b1 = b0 + n;
385 let b2 = b1 + n;
386 let b3 = b2 + n;
387
388 if ki + 4 < k {
390 let pf = (ki + 4) * n + j0;
391 _mm_prefetch(b.as_ptr().add(pf) as *const i8, _MM_HINT_T0);
392 _mm_prefetch(b.as_ptr().add(pf + 64) as *const i8, _MM_HINT_T0);
393 }
394
395 let bv = _mm512_loadu_ps(b.get_unchecked(b0));
397 acc0 = _mm512_fmadd_ps(a0, bv, acc0);
398 let bv = _mm512_loadu_ps(b.get_unchecked(b1));
399 acc0 = _mm512_fmadd_ps(a1, bv, acc0);
400 let bv = _mm512_loadu_ps(b.get_unchecked(b2));
401 acc0 = _mm512_fmadd_ps(a2, bv, acc0);
402 let bv = _mm512_loadu_ps(b.get_unchecked(b3));
403 acc0 = _mm512_fmadd_ps(a3, bv, acc0);
404
405 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 16));
406 acc1 = _mm512_fmadd_ps(a0, bv, acc1);
407 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 16));
408 acc1 = _mm512_fmadd_ps(a1, bv, acc1);
409 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 16));
410 acc1 = _mm512_fmadd_ps(a2, bv, acc1);
411 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 16));
412 acc1 = _mm512_fmadd_ps(a3, bv, acc1);
413
414 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 32));
415 acc2 = _mm512_fmadd_ps(a0, bv, acc2);
416 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 32));
417 acc2 = _mm512_fmadd_ps(a1, bv, acc2);
418 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 32));
419 acc2 = _mm512_fmadd_ps(a2, bv, acc2);
420 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 32));
421 acc2 = _mm512_fmadd_ps(a3, bv, acc2);
422
423 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 48));
424 acc3 = _mm512_fmadd_ps(a0, bv, acc3);
425 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 48));
426 acc3 = _mm512_fmadd_ps(a1, bv, acc3);
427 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 48));
428 acc3 = _mm512_fmadd_ps(a2, bv, acc3);
429 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 48));
430 acc3 = _mm512_fmadd_ps(a3, bv, acc3);
431
432 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 64));
433 acc4 = _mm512_fmadd_ps(a0, bv, acc4);
434 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 64));
435 acc4 = _mm512_fmadd_ps(a1, bv, acc4);
436 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 64));
437 acc4 = _mm512_fmadd_ps(a2, bv, acc4);
438 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 64));
439 acc4 = _mm512_fmadd_ps(a3, bv, acc4);
440
441 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 80));
442 acc5 = _mm512_fmadd_ps(a0, bv, acc5);
443 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 80));
444 acc5 = _mm512_fmadd_ps(a1, bv, acc5);
445 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 80));
446 acc5 = _mm512_fmadd_ps(a2, bv, acc5);
447 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 80));
448 acc5 = _mm512_fmadd_ps(a3, bv, acc5);
449
450 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 96));
451 acc6 = _mm512_fmadd_ps(a0, bv, acc6);
452 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 96));
453 acc6 = _mm512_fmadd_ps(a1, bv, acc6);
454 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 96));
455 acc6 = _mm512_fmadd_ps(a2, bv, acc6);
456 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 96));
457 acc6 = _mm512_fmadd_ps(a3, bv, acc6);
458
459 let bv = _mm512_loadu_ps(b.get_unchecked(b0 + 112));
460 acc7 = _mm512_fmadd_ps(a0, bv, acc7);
461 let bv = _mm512_loadu_ps(b.get_unchecked(b1 + 112));
462 acc7 = _mm512_fmadd_ps(a1, bv, acc7);
463 let bv = _mm512_loadu_ps(b.get_unchecked(b2 + 112));
464 acc7 = _mm512_fmadd_ps(a2, bv, acc7);
465 let bv = _mm512_loadu_ps(b.get_unchecked(b3 + 112));
466 acc7 = _mm512_fmadd_ps(a3, bv, acc7);
467
468 ki += 4;
469 }
470
471 while ki < k {
473 let av = _mm512_set1_ps(*a.get_unchecked(ki));
474 let base = ki * n + j0;
475 acc0 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base)), acc0);
476 acc1 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 16)), acc1);
477 acc2 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 32)), acc2);
478 acc3 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 48)), acc3);
479 acc4 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 64)), acc4);
480 acc5 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 80)), acc5);
481 acc6 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 96)), acc6);
482 acc7 = _mm512_fmadd_ps(av, _mm512_loadu_ps(b.get_unchecked(base + 112)), acc7);
483 ki += 1;
484 }
485
486 let cp = c.as_mut_ptr().add(j0);
488 _mm512_storeu_ps(cp, acc0);
489 _mm512_storeu_ps(cp.add(16), acc1);
490 _mm512_storeu_ps(cp.add(32), acc2);
491 _mm512_storeu_ps(cp.add(48), acc3);
492 _mm512_storeu_ps(cp.add(64), acc4);
493 _mm512_storeu_ps(cp.add(80), acc5);
494 _mm512_storeu_ps(cp.add(96), acc6);
495 _mm512_storeu_ps(cp.add(112), acc7);
496 }
497
498 if nt_end < n {
501 let rem = n - nt_end;
502 let rem16 = rem / 16 * 16;
503
504 for j0 in (0..rem16).step_by(16) {
506 let j = nt_end + j0;
507 let mut acc = _mm512_setzero_ps();
508 for ki in 0..k {
509 let av = _mm512_set1_ps(*a.get_unchecked(ki));
510 let bv = _mm512_loadu_ps(b.get_unchecked(ki * n + j));
511 acc = _mm512_fmadd_ps(av, bv, acc);
512 }
513 _mm512_storeu_ps(c.as_mut_ptr().add(j), acc);
514 }
515
516 for j in (nt_end + rem16)..n {
518 let mut sum = 0.0f32;
519 for ki in 0..k {
520 sum += *a.get_unchecked(ki) * *b.get_unchecked(ki * n + j);
521 }
522 *c.get_unchecked_mut(j) = sum;
523 }
524 }
525 }
526}
527
528pub fn gemv_scalar(k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
530 let k4 = k / 4 * 4;
532 for ki in (0..k4).step_by(4) {
533 let a0 = a[ki];
534 let a1 = a[ki + 1];
535 let a2 = a[ki + 2];
536 let a3 = a[ki + 3];
537 let b0 = ki * n;
538 let b1 = b0 + n;
539 let b2 = b1 + n;
540 let b3 = b2 + n;
541 for j in 0..n {
542 c[j] += a0 * b[b0 + j] + a1 * b[b1 + j] + a2 * b[b2 + j] + a3 * b[b3 + j];
543 }
544 }
545
546 for ki in k4..k {
548 let a_k = a[ki];
549 let b_start = ki * n;
550 for j in 0..n {
551 c[j] += a_k * b[b_start + j];
552 }
553 }
554}
555
556pub fn gemv(k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
558 contract_pre_gemv!(a, b);
559 #[cfg(target_arch = "x86_64")]
560 {
561 if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
569 unsafe {
572 if n > GEMV_TILE_THRESHOLD {
573 gemv_tiled_avx2(k, n, a, b, c);
574 } else {
575 gemv_avx2(k, n, a, b, c);
576 }
577 }
578 return;
579 }
580 }
581 gemv_scalar(k, n, a, b, c);
582}
583
584#[cfg(test)]
585mod tests {
586 use super::*;
587
588 #[test]
589 fn test_gemv_basic() {
590 let a = [1.0, 2.0, 3.0];
592 let b = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0];
593 let mut c = [0.0f32; 4];
594
595 gemv(3, 4, &a, &b, &mut c);
596
597 assert!((c[0] - 38.0).abs() < 1e-5);
599 assert!((c[1] - 44.0).abs() < 1e-5);
600 assert!((c[2] - 50.0).abs() < 1e-5);
601 assert!((c[3] - 56.0).abs() < 1e-5);
602 }
603
604 #[test]
605 fn test_gemv_identity_row_select() {
606 let a = [0.0, 1.0, 0.0];
608 let b = [1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0];
609 let mut c = [0.0f32; 3];
610
611 gemv(3, 3, &a, &b, &mut c);
612
613 assert!((c[0] - 4.0).abs() < 1e-5);
614 assert!((c[1] - 5.0).abs() < 1e-5);
615 assert!((c[2] - 6.0).abs() < 1e-5);
616 }
617
618 #[test]
619 fn test_gemv_large_n() {
620 let k = 2;
622 let n = 17;
623 let a = [1.0f32, 2.0];
624 let b: Vec<f32> = (0..k * n).map(|i| i as f32).collect();
625 let mut c = vec![0.0f32; n];
626
627 gemv(k, n, &a, &b, &mut c);
628
629 for j in 0..n {
631 let expected = a[0] * b[j] + a[1] * b[n + j];
632 assert!((c[j] - expected).abs() < 1e-4, "c[{j}] = {} expected {expected}", c[j]);
633 }
634 }
635
636 #[test]
637 fn test_gemv_zeros() {
638 let a = [0.0f32; 4];
639 let b = vec![1.0f32; 4 * 8];
640 let mut c = vec![0.0f32; 8];
641
642 gemv(4, 8, &a, &b, &mut c);
643
644 for j in 0..8 {
645 assert!((c[j]).abs() < 1e-10);
646 }
647 }
648
649 #[test]
651 fn test_gemv_tiled_large_n() {
652 let k = 64;
653 let n = 8192; let a: Vec<f32> = (0..k).map(|i| ((i * 7 + 3) % 100) as f32 / 100.0 - 0.5).collect();
656 let b: Vec<f32> = (0..k * n).map(|i| ((i * 13 + 7) % 1000) as f32 / 1000.0 - 0.5).collect();
657 let mut c_tiled = vec![0.0f32; n];
658 let mut c_scalar = vec![0.0f32; n];
659
660 gemv(k, n, &a, &b, &mut c_tiled);
661 gemv_scalar(k, n, &a, &b, &mut c_scalar);
662
663 for j in 0..n {
664 let diff = (c_tiled[j] - c_scalar[j]).abs();
665 assert!(diff < 1e-2, "j={j}: tiled={} scalar={} diff={diff}", c_tiled[j], c_scalar[j]);
666 }
667 }
668
669 #[test]
671 fn test_gemv_tiled_llm_size() {
672 let k = 256; let n = 11008;
674
675 let a: Vec<f32> = (0..k).map(|i| ((i * 17 + 31) % 1000) as f32 / 1000.0 - 0.5).collect();
676 let b: Vec<f32> = (0..k * n).map(|i| ((i * 13 + 7) % 1000) as f32 / 1000.0 - 0.5).collect();
677 let mut c_tiled = vec![0.0f32; n];
678 let mut c_scalar = vec![0.0f32; n];
679
680 gemv(k, n, &a, &b, &mut c_tiled);
681 gemv_scalar(k, n, &a, &b, &mut c_scalar);
682
683 for j in 0..n {
684 let diff = (c_tiled[j] - c_scalar[j]).abs();
685 assert!(diff < 1e-1, "j={j}: tiled={} scalar={} diff={diff}", c_tiled[j], c_scalar[j]);
686 }
687 }
688
689 #[test]
691 fn test_gemv_tiled_remainder() {
692 let k = 32;
693 let n = 5000; let a: Vec<f32> = (0..k).map(|i| ((i * 7 + 3) % 100) as f32 / 100.0 - 0.5).collect();
696 let b: Vec<f32> = (0..k * n).map(|i| ((i * 13 + 7) % 1000) as f32 / 1000.0 - 0.5).collect();
697 let mut c_tiled = vec![0.0f32; n];
698 let mut c_scalar = vec![0.0f32; n];
699
700 gemv(k, n, &a, &b, &mut c_tiled);
701 gemv_scalar(k, n, &a, &b, &mut c_scalar);
702
703 for j in 0..n {
704 let diff = (c_tiled[j] - c_scalar[j]).abs();
705 assert!(diff < 1e-2, "j={j}: tiled={} scalar={} diff={diff}", c_tiled[j], c_scalar[j]);
706 }
707 }
708
709 #[test]
711 fn test_gemv_tiled_k_remainder() {
712 let k = 67; let n = 8192;
714
715 let a: Vec<f32> = (0..k).map(|i| ((i * 7 + 3) % 100) as f32 / 100.0 - 0.5).collect();
716 let b: Vec<f32> = (0..k * n).map(|i| ((i * 13 + 7) % 1000) as f32 / 1000.0 - 0.5).collect();
717 let mut c_tiled = vec![0.0f32; n];
718 let mut c_scalar = vec![0.0f32; n];
719
720 gemv(k, n, &a, &b, &mut c_tiled);
721 gemv_scalar(k, n, &a, &b, &mut c_scalar);
722
723 for j in 0..n {
724 let diff = (c_tiled[j] - c_scalar[j]).abs();
725 assert!(diff < 1e-2, "j={j}: tiled={} scalar={} diff={diff}", c_tiled[j], c_scalar[j]);
726 }
727 }
728
729 #[test]
732 fn test_gemv_avx512_attention_size() {
733 let k = 128;
734 let n = 512;
735
736 let a: Vec<f32> = (0..k).map(|i| ((i * 17 + 31) % 1000) as f32 / 1000.0 - 0.5).collect();
737 let b: Vec<f32> = (0..k * n).map(|i| ((i * 13 + 7) % 1000) as f32 / 1000.0 - 0.5).collect();
738 let mut c_gemv = vec![0.0f32; n];
739 let mut c_scalar = vec![0.0f32; n];
740
741 gemv(k, n, &a, &b, &mut c_gemv);
742 gemv_scalar(k, n, &a, &b, &mut c_scalar);
743
744 let max_diff =
745 c_gemv.iter().zip(c_scalar.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
746 assert!(max_diff < 1e-2, "FALSIFY-AVX512-GEMV-001: max diff {max_diff}");
747 }
748
749 #[test]
752 fn test_gemv_avx512_remainder() {
753 let k = 128;
754 let n = 300; let a: Vec<f32> = (0..k).map(|i| ((i * 7 + 3) % 100) as f32 / 100.0).collect();
757 let b: Vec<f32> = (0..k * n).map(|i| ((i * 13 + 7) % 1000) as f32 / 1000.0 - 0.5).collect();
758 let mut c_gemv = vec![0.0f32; n];
759 let mut c_scalar = vec![0.0f32; n];
760
761 gemv(k, n, &a, &b, &mut c_gemv);
762 gemv_scalar(k, n, &a, &b, &mut c_scalar);
763
764 let max_diff =
765 c_gemv.iter().zip(c_scalar.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
766 assert!(max_diff < 1e-2, "FALSIFY-AVX512-GEMV-002: max diff {max_diff}");
767 }
768}