Skip to main content

trueno/blis/
gemv.rs

1//! SIMD-accelerated GEMV (General Matrix-Vector Multiply)
2//!
3//! Specialized kernel for M=1 matrix-vector product: c = a × B
4//! where a is 1×K and B is K×N, both row-major.
5//!
6//! This bypasses the BLIS 5-loop packing overhead which dominates for M=1.
7//! Instead, uses direct AVX2 VFMADD on unpacked row-major data.
8//!
9//! # Algorithm
10//!
11//! Two strategies based on N:
12//!
13//! - **Small N (≤ 4096)**: Axpy pattern — outer K, inner N. c[] fits in L1.
14//! - **Large N (> 4096)**: N-tiled — outer N-tiles (64), inner K. c[] stays
15//!   in YMM registers for all K iterations, eliminating L1 thrashing.
16//!
17//! # References
18//!
19//! - GH-380: matvec (M=1) performance gap vs ndarray
20
21/// Threshold: when N > this, switch to tiled GEMV kernel.
22/// Raised from 4096 → 8192 (2026-04-05): tiled kernel has strided B access
23/// (stride=N*4 bytes between rows) which is TLB-unfriendly at large N.
24/// Measured: vecmat 4096×4096: tiled 9.3 GFLOPS vs axpy predicts better.
25/// 4096 path benchmarks to use axpy. c[] still fits L1 at N=8192 (32KB).
26#[cfg(target_arch = "x86_64")]
27const GEMV_TILE_THRESHOLD: usize = 8192;
28
29/// AVX2 GEMV using axpy pattern: c += a[k] * B[k,:] for each k
30///
31/// Outer loop over K (4-way unrolled), inner loop over N with AVX2 VFMADD.
32/// This matches row-major B access: B[k,:] is contiguous → sequential reads.
33///
34/// Best for small N where c[] fits in L1 cache.
35///
36/// # Safety
37///
38/// Requires AVX2+FMA CPU features. Caller must ensure:
39/// - `a` has length >= `k`
40/// - `b` has length >= `k * n`
41/// - `c` has length >= `n`
42#[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        // 4-way K-unrolled axpy with AVX2 VFMADD on inner N loop
51        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            // Scalar remainder for N % 8
83            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        // Remainder K (scalar axpy)
95        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/// AVX2 GEMV with N-dimension tiling for bandwidth-bound sizes.
120///
121/// Tiles the N dimension into strips of 64, keeping the c[] accumulator
122/// in 8 YMM registers for ALL K iterations. This eliminates the repeated
123/// L1 load/store of c[] that dominates the axpy pattern when N > L1.
124///
125/// For 4096×11008: original axpy does 1024 load-store sweeps of c[] (43KB).
126/// Tiled: each c[j0..j0+64] is loaded 0 times (initialized in registers)
127/// and stored once at the end. Saves ~88MB of c[] traffic.
128///
129/// # Safety
130///
131/// Requires AVX2+FMA CPU features.
132#[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        // NT=64: 8 YMM accumulators × 8 f32 = 64 elements.
139        // 4 registers for broadcast a, 1-2 for B loads = ~14 registers total.
140        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            // 8 YMM accumulators — stay in registers for ALL K iterations
147            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            // Process ALL K for this N-tile (4-way unrolled)
157            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                // Software prefetch: B rows 8 iterations ahead
170                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                // 8 chunks × 4 K iterations = 32 FMAs
177                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            // Remainder K (1 at a time)
253            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            // Store accumulators (one store per tile, not K/4 stores)
269            _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        // Remainder N (< 64 elements) — axpy is fine since c fits in L1
280        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/// AVX-512 GEMV with N-dimension tiling — 2× throughput vs AVX2.
343///
344/// NT=128: 8 ZMM accumulators × 16 f32 = 128 elements per tile.
345/// 4-way K-unrolled: 32 FMAs per tile per iteration.
346/// Software prefetch: B rows 4 iterations ahead.
347///
348/// For attention scoring (Q @ K_cache^T): head_dim=128, seq_len varies.
349/// This is the hottest path in LLM inference (44.3% of compute).
350#[cfg(target_arch = "x86_64")]
351#[target_feature(enable = "avx512f", enable = "fma")]
352#[allow(dead_code)] // Retained for Intel SPR (no AVX-512 throttle). See negative result above.
353unsafe 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        // NT=128: 8 ZMM accumulators × 16 f32 = 128 elements.
358        // Fits in 8 of 32 ZMM registers. 4 for A broadcasts + 1 for B load = 13 total.
359        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            // 8 ZMM accumulators — stay in registers for ALL K iterations
366            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            // Process ALL K for this N-tile (4-way unrolled)
376            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                // Prefetch B rows 4 iterations ahead
389                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                // 8 chunks × 4 K iterations = 32 FMAs
396                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            // Remainder K (no unroll)
472            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            // Store accumulators
487            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        // Remainder N (< 128 elements) — process with AVX-512 individual zmm loads
499        // and scalar fallback for < 16 elements. No allocation needed.
500        if nt_end < n {
501            let rem = n - nt_end;
502            let rem16 = rem / 16 * 16;
503
504            // Process 16-wide chunks
505            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            // Scalar remainder (< 16 elements)
517            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
528/// Scalar fallback GEMV for non-x86 or non-AVX2 platforms
529pub fn gemv_scalar(k: usize, n: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
530    // 4-way K-unrolled axpy (auto-vectorizable)
531    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    // Remainder K
547    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
556/// Dispatch GEMV to best available implementation
557pub 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        // NEGATIVE RESULT (2026-04-05): AVX-512 GEMV is slower than AVX2 at ALL sizes.
562        // GEMV is bandwidth-bound (not compute-bound like GEMM). Zen 4 reduces clock
563        // ~10-15% during AVX-512 ops, and the wider SIMD can't compensate since the
564        // bottleneck is DRAM bandwidth.
565        // Measured: 128×512 AVX2=74.7 vs 512=61.6, 4096×4096 AVX2=16.3 vs 512=10.2.
566        // AVX-512 GEMV disabled. AVX2 GEMV remains the optimal path.
567        // gemv_tiled_avx512() retained for future Intel SPR (no AVX-512 throttle).
568        if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
569            // SAFETY: AVX2+FMA verified by feature detection above.
570            // Slice bounds are checked by the caller (matmul_vector_matrix).
571            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        // 1×3 @ 3×4 → 1×4
591        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        // c[j] = 1*B[0,j] + 2*B[1,j] + 3*B[2,j]
598        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        // e_1 @ B should give B[1,:]
607        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        // K=2, N=17 (tests AVX2 8-element chunks + scalar remainder)
621        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        // Verify against scalar reference
630        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 tiled path: N > GEMV_TILE_THRESHOLD triggers tiled kernel
650    #[test]
651    fn test_gemv_tiled_large_n() {
652        let k = 64;
653        let n = 8192; // > 4096 → tiled path
654
655        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 tiled path with LLM-size dimensions
670    #[test]
671    fn test_gemv_tiled_llm_size() {
672        let k = 256; // reduced from 4096 for test speed
673        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 tiled path with N not a multiple of 64 (exercises remainder)
690    #[test]
691    fn test_gemv_tiled_remainder() {
692        let k = 32;
693        let n = 5000; // > 4096, not multiple of 64 → remainder = 5000 - 4992 = 8
694
695        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 tiled path with non-multiple-of-4 K (exercises K remainder)
710    #[test]
711    fn test_gemv_tiled_k_remainder() {
712        let k = 67; // not multiple of 4
713        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    /// FALSIFY-AVX512-GEMV-001: AVX-512 GEMV matches scalar for attention-size dims.
730    /// k=128 (head_dim), n=512 (seq_len) — exercises AVX-512 path (n >= 128).
731    #[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    /// FALSIFY-AVX512-GEMV-002: AVX-512 GEMV with N not a multiple of 128.
750    /// Exercises the 16-wide remainder and scalar tail paths.
751    #[test]
752    fn test_gemv_avx512_remainder() {
753        let k = 128;
754        let n = 300; // 300 = 256 (2 tiles) + 32 (2 zmm remainder) + 12 (scalar)
755
756        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}