Skip to main content

ferrox_quant/
repack.rs

1//! Interleaved Q4_K / Q5_K × Q8_K and Q8_0 × Q8 GEMV (llama.cpp repack layouts).
2//!
3//! - Q4_K: packs 8 rows into `block_q4_Kx8` (`make_block_q4_Kx8`).
4//! - Q5_K: packs 8 rows into `block_q5_Kx8` (`make_block_q5_Kx8`).
5//! - Q8_0: packs 4 rows into `block_q8_0x4` (`make_block_q8_0x4`) with
6//!   4-byte interleave for NEON SDOT `ggml_gemv_q8_0_4x4_q8_0`.
7//! - Q4_0: packs 4 rows into `block_q4_0x4` (`make_block_q4_0x4`) with
8//!   4-byte interleave + XOR `0x88888888` for `ggml_gemv_q4_0_4x4_q8_0`.
9//!
10//! Gated on `FERROX_CPU_INT_DOT`, which `ferrox` and `ferrox-server`
11//! turn on by default (`=0` opts out); off in the library so golden
12//! cross-validation stays reference-exact.
13
14use crate::{
15    Q8Activations, Q8KActivations, Q4_0_BLOCK_BYTES, Q4_0_BLOCK_ELEMS, Q4_K_BLOCK_BYTES,
16    Q4_K_BLOCK_ELEMS, Q5_K_BLOCK_BYTES, Q5_K_BLOCK_ELEMS, Q6_K_BLOCK_BYTES, Q6_K_BLOCK_ELEMS,
17    Q8_0_BLOCK_BYTES, Q8_0_BLOCK_ELEMS,
18};
19use half::f16;
20
21/// Bytes per interleaved `block_q4_Kx8` (8 × f16 d + 8 × f16 dmin + 96 scales + 1024 qs).
22pub const Q4_KX8_BLOCK_BYTES: usize = 1152;
23/// Number of Q4_K rows packed into one interleaved block.
24pub const Q4_KX8_NROWS: usize = 8;
25
26const KMASK1: u32 = 0x3f3f_3f3f;
27const KMASK2: u32 = 0x0f0f_0f0f;
28const KMASK3: u32 = 0x0303_0303;
29
30/// Preferred qs interleave width for this CPU: 8 on x86 AVX2 and ARM i8mm
31/// (`ggml_gemm_q4_K_8x8_q8_K`), 4 on DotProd-only NEON (`ggml_gemm_q4_K_8x4_q8_K`).
32#[inline]
33pub fn q4_kx8_interleave() -> usize {
34    if cfg!(target_arch = "x86_64") {
35        return 8;
36    }
37    #[cfg(target_arch = "aarch64")]
38    {
39        if std::arch::is_aarch64_feature_detected!("i8mm") {
40            return 8;
41        }
42    }
43    4
44}
45
46#[inline]
47fn f16_from_bytes(b: &[u8]) -> f32 {
48    f16::from_le_bytes([b[0], b[1]]).to_f32()
49}
50
51/// Which `×4` GEMM kernel this host runs, resolved once instead of once
52/// per call.
53///
54/// The `gemm_*_group_x4` entry points are called once per (row-group ×
55/// activation-quad) pair, which on a `pp512` projection is 10^4 to 10^5
56/// calls per GEMM. Each one used to re-run `is_aarch64_feature_detected!`,
57/// whose relaxed atomic load LLVM cannot hoist out of the caller's loop.
58/// Callers now probe once per matmul and pass the answer down through the
59/// `_on` variants; [`gemm_q4_kx8_group_x4`] and its siblings stay as
60/// probe-per-call wrappers so existing callers and tests are unchanged.
61///
62/// This is a dispatch decision only. Both arms compute the same values,
63/// bit-identically, which is what the `*_x4_portable_is_bit_exact_vs_scalar_gemv`
64/// tests assert. Forcing [`AccelX4::Portable`] on an i8mm host is therefore
65/// a valid (slow) way to run, and the tests use it that way.
66#[derive(Clone, Copy, PartialEq, Eq, Debug)]
67pub enum AccelX4 {
68    /// ARM i8mm `SMMLA` kernels in [`neon`].
69    NeonI8mm,
70    /// The portable scalar reference.
71    Portable,
72}
73
74impl AccelX4 {
75    /// The fastest kernel available on this host.
76    #[inline]
77    pub fn detect() -> Self {
78        #[cfg(target_arch = "aarch64")]
79        {
80            if std::arch::is_aarch64_feature_detected!("i8mm") {
81                return AccelX4::NeonI8mm;
82            }
83        }
84        AccelX4::Portable
85    }
86}
87
88/// Pack eight canonical Q4_K super-blocks (same column-block index) into
89/// one `block_q4_Kx8`. `interleave` is 4 (ARM DotProd) or 8 (x86 / ARM i8mm).
90pub fn make_block_q4_kx8(
91    rows: [&[u8]; Q4_KX8_NROWS],
92    interleave: usize,
93) -> [u8; Q4_KX8_BLOCK_BYTES] {
94    debug_assert!(interleave == 4 || interleave == 8);
95    for r in &rows {
96        debug_assert_eq!(r.len(), Q4_K_BLOCK_BYTES);
97    }
98    let mut out = [0u8; Q4_KX8_BLOCK_BYTES];
99    // d[8] at 0, dmin[8] at 16, scales[96] at 32, qs[1024] at 128.
100    for (i, row) in rows.iter().enumerate() {
101        out[i * 2] = row[0];
102        out[i * 2 + 1] = row[1];
103        out[16 + i * 2] = row[2];
104        out[16 + i * 2 + 1] = row[3];
105    }
106
107    let end = (Q4_K_BLOCK_ELEMS * 4) / interleave; // qs bytes * 8 rows / interleave
108    let qs_out = &mut out[128..];
109    for i in 0..end {
110        let src_id = i % Q4_KX8_NROWS;
111        let src_offset = (i / Q4_KX8_NROWS) * interleave;
112        let dst_offset = i * interleave;
113        let src_qs = &rows[src_id][16..144];
114        qs_out[dst_offset..dst_offset + interleave]
115            .copy_from_slice(&src_qs[src_offset..src_offset + interleave]);
116    }
117
118    // Rearrange 6-bit scales/mins across 8 rows into 96 packed bytes
119    // (llama.cpp `make_block_q4_Kx8`).
120    let mut s = [0u8; 8];
121    let mut m = [0u8; 8];
122    let scales_out = &mut out[32..128];
123
124    for i in 0..4 {
125        for j in 0..8 {
126            let sc = &rows[j][4..16];
127            s[j] = sc[i] & 63;
128            m[j] = sc[i + 4] & 63;
129        }
130        let base = i * 12;
131        scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
132        scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
133        scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
134        scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
135        scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
136        scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
137        scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
138        scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
139        scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
140        scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
141        scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
142        scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
143    }
144
145    for i in 0..4 {
146        for j in 0..8 {
147            let sc = &rows[j][4..16];
148            s[j] = ((sc[i] & 192) >> 2) | (sc[i + 8] & 15);
149            m[j] = ((sc[i + 4] & 192) >> 2) | ((sc[i + 8] & 240) >> 4);
150        }
151        let base = 48 + i * 12;
152        scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
153        scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
154        scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
155        scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
156        scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
157        scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
158        scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
159        scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
160        scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
161        scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
162        scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
163        scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
164    }
165
166    out
167}
168
169/// Repack a full Q4_K matrix (row-major canonical blocks) into interleaved
170/// `block_q4_Kx8` groups. Rows not divisible by 8 are left out (caller
171/// handles the tail with per-row dots). `interleave` defaults via
172/// [`q4_kx8_interleave`].
173pub fn pack_q4_k_matrix_x8(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
174    assert!(cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
175    let n_blocks = cols / Q4_K_BLOCK_ELEMS;
176    let row_bytes = n_blocks * Q4_K_BLOCK_BYTES;
177    assert_eq!(data.len(), rows * row_bytes);
178    let n_groups = rows / Q4_KX8_NROWS;
179    let mut out = Vec::with_capacity(n_groups * n_blocks * Q4_KX8_BLOCK_BYTES);
180    for g in 0..n_groups {
181        for b in 0..n_blocks {
182            let mut row_refs: [&[u8]; Q4_KX8_NROWS] = [&[]; Q4_KX8_NROWS];
183            for (r, slot) in row_refs.iter_mut().enumerate() {
184                let base = (g * Q4_KX8_NROWS + r) * row_bytes + b * Q4_K_BLOCK_BYTES;
185                *slot = &data[base..base + Q4_K_BLOCK_BYTES];
186            }
187            out.extend_from_slice(&make_block_q4_kx8(row_refs, interleave));
188        }
189    }
190    out
191}
192
193/// Decode one 12-byte packed scale/min group into 8 scales + 8 mins (u8).
194#[inline]
195fn decode_scales_mins(scales12: &[u8], scales_out: &mut [u8; 8], mins_out: &mut [u8; 8]) {
196    debug_assert!(scales12.len() >= 12);
197    let mut utmp = [0u32; 4];
198    utmp[0] = u32::from_le_bytes(scales12[0..4].try_into().unwrap());
199    utmp[1] = u32::from_le_bytes(scales12[4..8].try_into().unwrap());
200    utmp[2] = u32::from_le_bytes(scales12[8..12].try_into().unwrap());
201    utmp[3] = ((utmp[2] >> 4) & KMASK2) | (((utmp[1] >> 6) & KMASK3) << 4);
202    let uaux_0 = utmp[1] & KMASK1;
203    utmp[1] = (utmp[2] & KMASK2) | (((utmp[0] >> 6) & KMASK3) << 4);
204    utmp[2] = uaux_0;
205    utmp[0] &= KMASK1;
206    let bytes = unsafe { std::slice::from_raw_parts(utmp.as_ptr() as *const u8, 16) };
207    scales_out.copy_from_slice(&bytes[0..8]);
208    mins_out.copy_from_slice(&bytes[8..16]);
209}
210
211/// Scalar GEMV for interleave=4 (`ggml_gemv_q4_K_8x4_q8_K_generic`).
212fn gemv_q4_kx8_q8_k_scalar_4(
213    packed: &[u8],
214    act: &Q8KActivations,
215    n_cols: usize,
216    n_row_groups: usize,
217    out: &mut [f32],
218) {
219    let nb = n_cols / Q4_K_BLOCK_ELEMS;
220    let blocklen = 4;
221    let ncols_interleaved = Q4_KX8_NROWS;
222    debug_assert_eq!(act.n_blocks(), nb);
223    debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
224    debug_assert_eq!(packed.len(), n_row_groups * nb * Q4_KX8_BLOCK_BYTES);
225
226    for x in 0..n_row_groups {
227        let mut sumf = [0f32; 8];
228        let mut sum_minf = [0f32; 8];
229        let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
230        for l in 0..nb {
231            let blk = &packed[group_off + l * Q4_KX8_BLOCK_BYTES..][..Q4_KX8_BLOCK_BYTES];
232            let d = &blk[0..16];
233            let dmin = &blk[16..32];
234            let scales = &blk[32..128];
235            let qs = &blk[128..];
236            let da = act.d[l];
237            let q8 = &act.q[l * Q4_K_BLOCK_ELEMS..(l + 1) * Q4_K_BLOCK_ELEMS];
238            let bsums = &act.bsums[l * 16..(l + 1) * 16];
239
240            let mut all_scales = [[0u8; 8]; 8];
241            let mut all_mins = [[0u8; 8]; 8];
242            for sb in 0..8 {
243                decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
244            }
245
246            let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen); // 32
247            for k in 0..n_k {
248                let sb_pair = k / 8;
249                let sc0 = &all_scales[sb_pair * 2];
250                let sc1 = &all_scales[sb_pair * 2 + 1];
251                for j in 0..ncols_interleaved {
252                    let mut sumi = 0i32;
253                    for i in 0..blocklen {
254                        let qbyte = qs[k * ncols_interleaved * blocklen + j * blocklen + i];
255                        let v0 = (qbyte & 0x0F) as i32;
256                        let v1 = (qbyte >> 4) as i32;
257                        let a0 = q8[(k / 8) * 64 + (k % 8) * blocklen + i] as i32;
258                        let a1 = q8[(k / 8) * 64 + (k % 8) * blocklen + i + 32] as i32;
259                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
260                    }
261                    sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
262                }
263            }
264            for sb in 0..8 {
265                let mins = &all_mins[sb];
266                let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
267                for j in 0..ncols_interleaved {
268                    sum_minf[j] +=
269                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
270                }
271            }
272        }
273        let base = x * ncols_interleaved;
274        for j in 0..ncols_interleaved {
275            out[base + j] = sumf[j] - sum_minf[j];
276        }
277    }
278}
279
280/// Scalar GEMV for interleave=8 (`ggml_gemv_q4_K_8x8_q8_K_generic`).
281fn gemv_q4_kx8_q8_k_scalar_8(
282    packed: &[u8],
283    act: &Q8KActivations,
284    n_cols: usize,
285    n_row_groups: usize,
286    out: &mut [f32],
287) {
288    let nb = n_cols / Q4_K_BLOCK_ELEMS;
289    let blocklen = 8;
290    let ncols_interleaved = Q4_KX8_NROWS;
291    debug_assert_eq!(act.n_blocks(), nb);
292    debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
293
294    for x in 0..n_row_groups {
295        let mut sumf = [0f32; 8];
296        let mut sum_minf = [0f32; 8];
297        let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
298        for l in 0..nb {
299            let blk = &packed[group_off + l * Q4_KX8_BLOCK_BYTES..][..Q4_KX8_BLOCK_BYTES];
300            let d = &blk[0..16];
301            let dmin = &blk[16..32];
302            let scales = &blk[32..128];
303            let qs = &blk[128..];
304            let da = act.d[l];
305            let q8 = &act.q[l * Q4_K_BLOCK_ELEMS..(l + 1) * Q4_K_BLOCK_ELEMS];
306            let bsums = &act.bsums[l * 16..(l + 1) * 16];
307
308            let mut all_scales = [[0u8; 8]; 8];
309            let mut all_mins = [[0u8; 8]; 8];
310            for sb in 0..8 {
311                decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
312            }
313
314            let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen); // 16
315            for k in 0..n_k {
316                let sb_pair = k / 4;
317                let sc0 = &all_scales[sb_pair * 2];
318                let sc1 = &all_scales[sb_pair * 2 + 1];
319                for j in 0..ncols_interleaved {
320                    let mut sumi = 0i32;
321                    for i in 0..blocklen {
322                        let qbyte = qs[k * ncols_interleaved * blocklen + j * blocklen + i];
323                        let v0 = (qbyte & 0x0F) as i32;
324                        let v1 = (qbyte >> 4) as i32;
325                        let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
326                        let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
327                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
328                    }
329                    sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
330                }
331            }
332            for sb in 0..8 {
333                let mins = &all_mins[sb];
334                let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
335                for j in 0..ncols_interleaved {
336                    sum_minf[j] +=
337                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
338                }
339            }
340        }
341        let base = x * ncols_interleaved;
342        for j in 0..ncols_interleaved {
343            out[base + j] = sumf[j] - sum_minf[j];
344        }
345    }
346}
347
348/// GEMV: interleaved Q4_K weights × Q8_K activation → `n_row_groups * 8` f32s.
349/// Dispatches to NEON (both interleaves) / AVX2 (interleave 8) when available.
350pub fn gemv_q4_kx8_q8_k(
351    packed: &[u8],
352    act: &Q8KActivations,
353    n_cols: usize,
354    n_row_groups: usize,
355    interleave: usize,
356    out: &mut [f32],
357) {
358    assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
359    assert_eq!(out.len(), n_row_groups * Q4_KX8_NROWS);
360    match interleave {
361        4 => {
362            #[cfg(target_arch = "aarch64")]
363            {
364                if std::arch::is_aarch64_feature_detected!("dotprod") {
365                    unsafe {
366                        neon::gemv_q4_kx8_q8_k_neon_sdot(packed, act, n_cols, n_row_groups, out);
367                    }
368                    return;
369                }
370            }
371            gemv_q4_kx8_q8_k_scalar_4(packed, act, n_cols, n_row_groups, out);
372        }
373        8 => {
374            #[cfg(target_arch = "x86_64")]
375            {
376                if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
377                    unsafe {
378                        avx2::gemv_q4_kx8_q8_k_avx2(packed, act, n_cols, n_row_groups, out);
379                    }
380                    return;
381                }
382            }
383            #[cfg(target_arch = "aarch64")]
384            {
385                if std::arch::is_aarch64_feature_detected!("dotprod") {
386                    unsafe {
387                        neon::gemv_q4_kx8_q8_k_neon_8x8(packed, act, n_cols, n_row_groups, out);
388                    }
389                    return;
390                }
391            }
392            gemv_q4_kx8_q8_k_scalar_8(packed, act, n_cols, n_row_groups, out);
393        }
394        _ => panic!("q4_kx8 interleave must be 4 or 8, got {interleave}"),
395    }
396}
397
398/// One row-group (8 outputs) starting at `group` within a packed matrix.
399#[inline]
400pub fn gemv_q4_kx8_group(
401    packed: &[u8],
402    group: usize,
403    act: &Q8KActivations,
404    n_cols: usize,
405    interleave: usize,
406    out8: &mut [f32],
407) {
408    debug_assert_eq!(out8.len(), Q4_KX8_NROWS);
409    let nb = n_cols / Q4_K_BLOCK_ELEMS;
410    let off = group * nb * Q4_KX8_BLOCK_BYTES;
411    let slice = &packed[off..off + nb * Q4_KX8_BLOCK_BYTES];
412    gemv_q4_kx8_q8_k(slice, act, n_cols, 1, interleave, out8);
413}
414
415// ---------------------------------------------------------------------------
416// Q8_0 ×4 interleaved GEMV (llama.cpp `block_q8_0x4` / `ggml_gemv_q8_0_4x4`)
417// ---------------------------------------------------------------------------
418
419/// Bytes per interleaved `block_q8_0x4` (4 × f16 d + 128 qs).
420pub const Q8_0X4_BLOCK_BYTES: usize = 136;
421/// Number of Q8_0 rows packed into one interleaved block.
422pub const Q8_0X4_NROWS: usize = 4;
423/// qs interleave width for `ggml_gemv_q8_0_4x4_q8_0` (NEON SDOT). The
424/// DotProd-only default; [`q8_0x4_interleave`] picks 8 on i8mm hosts.
425pub const Q8_0X4_INTERLEAVE: usize = 4;
426
427/// Preferred qs interleave width: 8 on ARM i8mm (`ggml_gemm_q8_0_4x8_q8_0`
428/// via `ggml_repack_get_optimal_repack_type`), 4 on DotProd-only NEON and
429/// everywhere else (the scalar fallback handles either).
430#[inline]
431pub fn q8_0x4_interleave() -> usize {
432    #[cfg(target_arch = "aarch64")]
433    {
434        if std::arch::is_aarch64_feature_detected!("i8mm") {
435            return 8;
436        }
437    }
438    Q8_0X4_INTERLEAVE
439}
440
441/// Pack four canonical Q8_0 blocks (same column-block) into one
442/// `block_q8_0x4`. `interleave` is 4 (ARM 4x4) or 8 (4x8).
443pub fn make_block_q8_0x4(
444    rows: [&[u8]; Q8_0X4_NROWS],
445    interleave: usize,
446) -> [u8; Q8_0X4_BLOCK_BYTES] {
447    debug_assert!(interleave == 4 || interleave == 8);
448    for r in &rows {
449        debug_assert_eq!(r.len(), Q8_0_BLOCK_BYTES);
450    }
451    let mut out = [0u8; Q8_0X4_BLOCK_BYTES];
452    for (i, row) in rows.iter().enumerate() {
453        out[i * 2] = row[0];
454        out[i * 2 + 1] = row[1];
455    }
456    let end = (Q8_0_BLOCK_ELEMS * Q8_0X4_NROWS) / interleave;
457    let qs_out = &mut out[8..];
458    for i in 0..end {
459        let src_id = i % Q8_0X4_NROWS;
460        let src_offset = (i / Q8_0X4_NROWS) * interleave;
461        let dst_offset = i * interleave;
462        let src_qs = &rows[src_id][2..34];
463        qs_out[dst_offset..dst_offset + interleave]
464            .copy_from_slice(&src_qs[src_offset..src_offset + interleave]);
465    }
466    out
467}
468
469/// Repack a Q8_0 matrix into interleaved `block_q8_0x4` groups. Tail rows
470/// (not divisible by 4) are omitted; caller dots them with [`crate::dot_q8_0_q8`].
471pub fn pack_q8_0_matrix_x4(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
472    assert!(cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
473    let n_blocks = cols / Q8_0_BLOCK_ELEMS;
474    let row_bytes = n_blocks * Q8_0_BLOCK_BYTES;
475    assert_eq!(data.len(), rows * row_bytes);
476    let n_groups = rows / Q8_0X4_NROWS;
477    let mut out = Vec::with_capacity(n_groups * n_blocks * Q8_0X4_BLOCK_BYTES);
478    for g in 0..n_groups {
479        for b in 0..n_blocks {
480            let mut row_refs: [&[u8]; Q8_0X4_NROWS] = [&[]; Q8_0X4_NROWS];
481            for (r, slot) in row_refs.iter_mut().enumerate() {
482                let base = (g * Q8_0X4_NROWS + r) * row_bytes + b * Q8_0_BLOCK_BYTES;
483                *slot = &data[base..base + Q8_0_BLOCK_BYTES];
484            }
485            out.extend_from_slice(&make_block_q8_0x4(row_refs, interleave));
486        }
487    }
488    out
489}
490
491/// Scalar GEMV (`ggml_gemv_q8_0_4x{4,8}_q8_0_generic`); `blocklen` is the
492/// interleave the matrix was packed with.
493fn gemv_q8_0x4_q8_0_scalar(
494    packed: &[u8],
495    act: &Q8Activations,
496    n_cols: usize,
497    n_row_groups: usize,
498    blocklen: usize,
499    out: &mut [f32],
500) {
501    let nb = n_cols / Q8_0_BLOCK_ELEMS;
502    let ncols = Q8_0X4_NROWS;
503    debug_assert_eq!(act.n_blocks(), nb);
504    debug_assert_eq!(out.len(), n_row_groups * ncols);
505    debug_assert_eq!(packed.len(), n_row_groups * nb * Q8_0X4_BLOCK_BYTES);
506
507    for x in 0..n_row_groups {
508        let mut sumf = [0f32; 4];
509        let group_off = x * nb * Q8_0X4_BLOCK_BYTES;
510        for l in 0..nb {
511            let blk = &packed[group_off + l * Q8_0X4_BLOCK_BYTES..][..Q8_0X4_BLOCK_BYTES];
512            let qs = &blk[8..];
513            let da = act.d[l];
514            let q8 = &act.q[l * Q8_0_BLOCK_ELEMS..(l + 1) * Q8_0_BLOCK_ELEMS];
515            for k in 0..(Q8_0_BLOCK_ELEMS / blocklen) {
516                for j in 0..ncols {
517                    let mut sumi = 0i32;
518                    for i in 0..blocklen {
519                        let v0 = qs[k * ncols * blocklen + j * blocklen + i] as i8 as i32;
520                        sumi += v0 * q8[k * blocklen + i] as i32;
521                    }
522                    sumf[j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
523                }
524            }
525        }
526        let base = x * ncols;
527        out[base..base + ncols].copy_from_slice(&sumf);
528    }
529}
530
531/// GEMV: interleaved Q8_0 weights × Q8 activation → `n_row_groups * 4` f32s.
532/// `interleave` must match the packing (4: NEON SDOT `4x4`; 8: NEON `4x8`).
533pub fn gemv_q8_0x4_q8_0(
534    packed: &[u8],
535    act: &Q8Activations,
536    n_cols: usize,
537    n_row_groups: usize,
538    interleave: usize,
539    out: &mut [f32],
540) {
541    assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
542    assert_eq!(out.len(), n_row_groups * Q8_0X4_NROWS);
543    match interleave {
544        4 => {
545            #[cfg(target_arch = "aarch64")]
546            {
547                if std::arch::is_aarch64_feature_detected!("dotprod") {
548                    unsafe {
549                        neon::gemv_q8_0x4_q8_0_neon_sdot(packed, act, n_cols, n_row_groups, out);
550                    }
551                    return;
552                }
553            }
554            gemv_q8_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 4, out);
555        }
556        8 => {
557            #[cfg(target_arch = "aarch64")]
558            {
559                if std::arch::is_aarch64_feature_detected!("dotprod") {
560                    unsafe {
561                        neon::gemv_q8_0x4_q8_0_neon_4x8(packed, act, n_cols, n_row_groups, out);
562                    }
563                    return;
564                }
565            }
566            gemv_q8_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 8, out);
567        }
568        _ => panic!("q8_0x4 interleave must be 4 or 8, got {interleave}"),
569    }
570}
571
572/// How many activations one [`gemm_q8_0x4_group`] pass keeps in flight.
573/// Four f32x4 accumulators plus the eight loaded weight vectors fit
574/// comfortably in NEON's register file, so each weight load is amortized
575/// over four activations instead of being repeated per activation.
576pub const Q8_0X4_GEMM_NC: usize = 8;
577
578/// GEMM counterpart of [`gemv_q8_0x4_group`]: one row-group (4 rows)
579/// against `acts.len()` activations at once.
580///
581/// The difference that matters is register blocking over the *batch*
582/// dimension. Calling the GEMV once per activation reloads the group's
583/// eight `int8x16` weight vectors for every activation; this loads them
584/// once per `Q8_0X4_GEMM_NC` activations and issues the dot products
585/// back to back. That is the same reason llama.cpp ships
586/// `ggml_gemm_q8_0_4x4_q8_0` next to `ggml_gemv_q8_0_4x4_q8_0` rather
587/// than looping the GEMV.
588///
589/// `out` is `[row][act]`: `out[r * acts.len() + j]`, which is the layout
590/// `WeightMatrix::apply_batch` accumulates into.
591pub fn gemm_q8_0x4_group(
592    packed: &[u8],
593    group: usize,
594    acts: &[Q8Activations],
595    n_cols: usize,
596    interleave: usize,
597    out: &mut [f32],
598) {
599    assert_eq!(out.len(), Q8_0X4_NROWS * acts.len());
600    assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
601    if acts.is_empty() {
602        return;
603    }
604    let nb = n_cols / Q8_0_BLOCK_ELEMS;
605    let off = group * nb * Q8_0X4_BLOCK_BYTES;
606    let slice = &packed[off..off + nb * Q8_0X4_BLOCK_BYTES];
607
608    #[cfg(target_arch = "aarch64")]
609    {
610        if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
611            // Compatibility entry: interleaves the quads here, once per
612            // call. Batch callers should prepare them once per matmul and
613            // use [`gemm_q8_0x4_group_x4`] instead.
614            for (t, chunk) in acts.chunks(Q8K_ACTS_X4_NC).enumerate() {
615                let tile = prepare_q8_acts_x4(chunk, n_cols);
616                let mut tmp = [0f32; Q8_0X4_NROWS * Q8K_ACTS_X4_NC];
617                let n = chunk.len();
618                unsafe {
619                    neon::gemm_q8_0x4_q8_0_neon_i8mm(
620                        slice,
621                        &tile,
622                        n_cols,
623                        &mut tmp[..Q8_0X4_NROWS * n],
624                    );
625                }
626                for r in 0..Q8_0X4_NROWS {
627                    for j in 0..n {
628                        out[r * acts.len() + t * Q8K_ACTS_X4_NC + j] = tmp[r * n + j];
629                    }
630                }
631            }
632            return;
633        }
634        if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
635            unsafe {
636                neon::gemm_q8_0x4_q8_0_neon_sdot(slice, acts, n_cols, out);
637            }
638            return;
639        }
640    }
641    // Portable fallback: the GEMV, once per activation. Same results,
642    // none of the reuse.
643    let mut tmp = [0f32; Q8_0X4_NROWS];
644    for (j, act) in acts.iter().enumerate() {
645        gemv_q8_0x4_q8_0(slice, act, n_cols, 1, interleave, &mut tmp);
646        for (r, v) in tmp.iter().enumerate() {
647            out[r * acts.len() + j] = *v;
648        }
649    }
650}
651
652/// A quad of up to [`Q8K_ACTS_X4_NC`] Q8_0 activations, pre-interleaved
653/// into the layout llama.cpp's `ggml_quantize_mat_q8_0_4x8` writes into
654/// `block_q8_0x4` (`arch/arm/repack.cpp`): every 32-element block's qs in
655/// 8-byte runs, plus the per-block per-row scales. Consumed by the i8mm
656/// `4x8` GEMMs; prepared once per matmul, same hoist as [`Q8KActsX4`].
657pub struct Q8ActsX4 {
658    /// Real activations in the quad (≤ 4); rows `na..4` are zero padding.
659    pub na: usize,
660    /// Q8_0 blocks per activation (`n_cols / 32`).
661    pub n_blocks: usize,
662    /// Interleaved quants, `n_blocks * 128` long. Block `b`, 8-element run
663    /// `c`, quad row `a`, lane `k` ↦
664    /// `qs[b*128 + c*32 + a*8 + k] = acts[a].q[b*32 + c*8 + k]`.
665    pub qs: Vec<i8>,
666    /// Activation scales, `n_blocks * 4` long: `d[b*4 + a] = acts[a].d[b]`.
667    pub d: Vec<f32>,
668}
669
670/// Interleave a quad of Q8_0 activations for the `4x8` i8mm GEMMs
671/// (llama.cpp `ggml_quantize_mat_q8_0_4x8`, minus the quantization we
672/// already did). Zero-pads when `acts.len() < 4`. Available on every
673/// target so the portable GEMMs — and the tests pinning the NEON kernels
674/// to them — run anywhere.
675pub fn prepare_q8_acts_x4(acts: &[Q8Activations], n_cols: usize) -> Q8ActsX4 {
676    assert!(acts.len() <= Q8K_ACTS_X4_NC);
677    assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
678    let na = acts.len();
679    let nb = n_cols / Q8_0_BLOCK_ELEMS;
680    let mut qs = vec![0i8; nb * Q8_0_BLOCK_ELEMS * 4];
681    let mut d = vec![0f32; nb * 4];
682    for (a, act) in acts.iter().enumerate() {
683        debug_assert_eq!(act.d.len(), nb);
684        for b in 0..nb {
685            let src = &act.q[b * Q8_0_BLOCK_ELEMS..(b + 1) * Q8_0_BLOCK_ELEMS];
686            let dst = &mut qs[b * Q8_0_BLOCK_ELEMS * 4..(b + 1) * Q8_0_BLOCK_ELEMS * 4];
687            for (c, run) in src.as_chunks::<8>().0.iter().enumerate() {
688                dst[c * 32 + a * 8..c * 32 + a * 8 + 8].copy_from_slice(run);
689            }
690            d[b * 4 + a] = act.d[b];
691        }
692    }
693    Q8ActsX4 {
694        na,
695        n_blocks: nb,
696        qs,
697        d,
698    }
699}
700
701/// Whether [`gemm_q8_0x4_group_x4`] is the fast Q8_0 batch path on this
702/// CPU: ARM i8mm with the interleave-8 layout (`ggml_gemm_q8_0_4x8_q8_0`).
703#[inline]
704pub fn q8_0x4_gemm_uses_acts_x4(interleave: usize) -> bool {
705    #[cfg(target_arch = "aarch64")]
706    {
707        interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
708    }
709    #[cfg(not(target_arch = "aarch64"))]
710    {
711        let _ = interleave;
712        false
713    }
714}
715
716/// [`gemm_q8_0x4_group`] against a pre-interleaved activation quad;
717/// interleave-8 packing only, quad prepared once per matmul by
718/// [`prepare_q8_acts_x4`]. `out` is `[row][act]`: `out[r * tile.na + a]`.
719pub fn gemm_q8_0x4_group_x4(
720    packed: &[u8],
721    group: usize,
722    tile: &Q8ActsX4,
723    n_cols: usize,
724    interleave: usize,
725    out: &mut [f32],
726) {
727    gemm_q8_0x4_group_x4_on(
728        packed,
729        group,
730        tile,
731        n_cols,
732        interleave,
733        AccelX4::detect(),
734        out,
735    );
736}
737
738/// [`gemm_q8_0x4_group_x4`] with the kernel choice already made; see
739/// [`AccelX4`].
740#[inline]
741pub fn gemm_q8_0x4_group_x4_on(
742    packed: &[u8],
743    group: usize,
744    tile: &Q8ActsX4,
745    n_cols: usize,
746    interleave: usize,
747    accel: AccelX4,
748    out: &mut [f32],
749) {
750    assert_eq!(
751        interleave, 8,
752        "the x4 GEMM only exists for interleave-8 packing"
753    );
754    assert_eq!(out.len(), Q8_0X4_NROWS * tile.na);
755    assert!(n_cols.is_multiple_of(Q8_0_BLOCK_ELEMS));
756    debug_assert_eq!(tile.n_blocks, n_cols / Q8_0_BLOCK_ELEMS);
757    if tile.na == 0 {
758        return;
759    }
760    let nb = n_cols / Q8_0_BLOCK_ELEMS;
761    let off = group * nb * Q8_0X4_BLOCK_BYTES;
762    let slice = &packed[off..off + nb * Q8_0X4_BLOCK_BYTES];
763
764    #[cfg(target_arch = "aarch64")]
765    if accel == AccelX4::NeonI8mm {
766        unsafe {
767            neon::gemm_q8_0x4_q8_0_neon_i8mm(slice, tile, n_cols, out);
768        }
769        return;
770    }
771    let _ = accel;
772    gemm_q8_0x4_acts_x4_scalar_8(slice, tile, n_cols, out);
773}
774
775/// Portable reference for the Q8_0 ×4 GEMM: the same math as
776/// [`gemv_q8_0x4_q8_0_scalar`] at blocklen 8, per quad row, reading qs and
777/// d out of the pre-interleaved [`Q8ActsX4`]. Bit-identical to running
778/// that GEMV per activation, which is what the tests assert.
779fn gemm_q8_0x4_acts_x4_scalar_8(packed: &[u8], tile: &Q8ActsX4, n_cols: usize, out: &mut [f32]) {
780    let nb = n_cols / Q8_0_BLOCK_ELEMS;
781    let blocklen = 8;
782    let ncols = Q8_0X4_NROWS;
783    let na = tile.na;
784    let mut sumf = [[0f32; Q8_0X4_NROWS]; Q8K_ACTS_X4_NC];
785    for l in 0..nb {
786        let blk = &packed[l * Q8_0X4_BLOCK_BYTES..][..Q8_0X4_BLOCK_BYTES];
787        let qs = &blk[8..];
788        let q8 = &tile.qs[l * Q8_0_BLOCK_ELEMS * 4..][..Q8_0_BLOCK_ELEMS * 4];
789        for a in 0..na {
790            let da = tile.d[l * 4 + a];
791            for k in 0..(Q8_0_BLOCK_ELEMS / blocklen) {
792                for j in 0..ncols {
793                    let mut sumi = 0i32;
794                    for i in 0..blocklen {
795                        let v0 = qs[k * ncols * blocklen + j * blocklen + i] as i8 as i32;
796                        // Canonical q8 element `e` lives at run `e/8`, row
797                        // `a`, lane `e%8` of the interleaved block.
798                        let e = k * blocklen + i;
799                        sumi += v0 * q8[(e / 8) * 32 + a * 8 + (e % 8)] as i32;
800                    }
801                    sumf[a][j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
802                }
803            }
804        }
805    }
806    for j in 0..ncols {
807        for (a, row) in sumf.iter().take(na).enumerate() {
808            out[j * na + a] = row[j];
809        }
810    }
811}
812
813/// How many activations one [`gemm_q4_kx8_group`] pass keeps in flight.
814///
815/// Four is llama's shape for `ggml_gemm_q4_K_8x4_q8_K` (`q8_k_blocklen`),
816/// and it is what the register file allows here: eight `uint8x16` weight
817/// columns plus one activation's four `int8x16` and its accumulator pair
818/// stay resident while the batch loop turns.
819pub const Q4_KX8_GEMM_NC: usize = 4;
820
821/// GEMM counterpart of [`gemv_q4_kx8_group`]: one row-group (8 rows)
822/// against `acts.len()` activations at once.
823///
824/// Q4_K was the expensive omission. The GEMV repeats, *per activation*,
825/// work that depends only on the weights: 16 f16 scale conversions, 8
826/// `decode_scales_mins` calls and 16 `q4_cols` loads per 256-element
827/// super-block. At batch 512 that is the same 6-bit scale decode run 512
828/// times. Q8_0 already had `gemm_q8_0x4_group`; this is the same idea for
829/// the format that carries every `*_Q4_K_M` checkpoint's FFN.
830///
831/// `out` is `[row][act]`: `out[r * acts.len() + j]`, matching
832/// [`gemm_q8_0x4_group`] and what `WeightMatrix::apply_batch` writes.
833pub fn gemm_q4_kx8_group(
834    packed: &[u8],
835    group: usize,
836    acts: &[Q8KActivations],
837    n_cols: usize,
838    interleave: usize,
839    out: &mut [f32],
840) {
841    assert_eq!(out.len(), Q4_KX8_NROWS * acts.len());
842    assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
843    if acts.is_empty() {
844        return;
845    }
846    let nb = n_cols / Q4_K_BLOCK_ELEMS;
847    let off = group * nb * Q4_KX8_BLOCK_BYTES;
848    let slice = &packed[off..off + nb * Q4_KX8_BLOCK_BYTES];
849
850    #[cfg(target_arch = "aarch64")]
851    if acts.len() <= Q4_KX8_GEMM_NC {
852        if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
853            // Compatibility entry: interleaves the quad here, once per call.
854            // Batch callers should prepare the quad once per matmul and use
855            // [`gemm_q4_kx8_group_x4`] for every row-group instead.
856            let tile = prepare_q8_k_acts_x4(acts, n_cols);
857            unsafe {
858                neon::gemm_q4_kx8_q8_k_neon_i8mm(slice, &tile, n_cols, out);
859            }
860            return;
861        }
862        if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
863            unsafe {
864                neon::gemm_q4_kx8_q8_k_neon_sdot(slice, acts, n_cols, out);
865            }
866            return;
867        }
868    }
869    // Portable fallback: the GEMV, once per activation. Same results,
870    // none of the reuse.
871    let mut tmp = [0f32; Q4_KX8_NROWS];
872    for (j, act) in acts.iter().enumerate() {
873        gemv_q4_kx8_q8_k(slice, act, n_cols, 1, interleave, &mut tmp);
874        for (r, v) in tmp.iter().enumerate() {
875            out[r * acts.len() + j] = *v;
876        }
877    }
878}
879
880/// A quad of up to [`Q8K_ACTS_X4_NC`] Q8_K activations, pre-interleaved into
881/// the layout llama.cpp's `ggml_quantize_mat_q8_K_4x8` writes into
882/// `block_q8_Kx4` (`ggml-cpu/repack.cpp`): every super-block's qs, the folded
883/// `bsums` pairs, and the per-block per-row scales.
884///
885/// The i8mm GEMM consumes activations in this shape. Interleaving them once
886/// per matmul — instead of once per (row-group, block) inside the kernel —
887/// is the point: the old in-kernel repack was a scalar pass over
888/// `rows/8 · batch · cols` bytes with a div and a mod per element, roughly
889/// 4× the instruction count of the `vmmlaq_s32` math it fed.
890/// Activations per [`Q8KActsX4`] quad (llama.cpp's `q8_k_blocklen`).
891pub const Q8K_ACTS_X4_NC: usize = 4;
892
893pub struct Q8KActsX4 {
894    /// Real activations in the quad (≤ 4); rows `na..4` are zero padding.
895    pub na: usize,
896    /// Q8_K super-blocks per activation (`n_cols / 256`).
897    pub n_blocks: usize,
898    /// Interleaved quants, `n_blocks * 1024` long. Block `b`, 8-element run
899    /// `c`, quad row `a`, lane `k` ↦
900    /// `qs[b*1024 + c*32 + a*8 + k] = acts[a].q[b*256 + c*8 + k]`.
901    pub qs: Vec<i8>,
902    /// Folded `bsums` pairs, `n_blocks * 4 * 8` long:
903    /// `bsums[(b*4 + a)*8 + i] = acts[a].bsums[b*16 + 2i] + acts[a].bsums[b*16 + 2i + 1]`.
904    pub bsums: Vec<i16>,
905    /// Activation scales, `n_blocks * 4` long: `d[b*4 + a] = acts[a].d[b]`.
906    pub d: Vec<f32>,
907}
908
909/// Interleave a quad of activations for [`gemm_q4_kx8_group_x4`]
910/// (llama.cpp `ggml_quantize_mat_q8_K_4x8`, minus the quantization we
911/// already did). Zero-pads when `acts.len() < 4`, matching what the kernel's
912/// in-loop repack used to emit. Available on every target so the portable
913/// GEMM below — and the tests pinning the NEON kernel to it — run anywhere.
914pub fn prepare_q8_k_acts_x4(acts: &[Q8KActivations], n_cols: usize) -> Q8KActsX4 {
915    assert!(acts.len() <= Q8K_ACTS_X4_NC);
916    assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
917    let na = acts.len();
918    let nb = n_cols / Q4_K_BLOCK_ELEMS;
919    let mut qs = vec![0i8; nb * Q4_K_BLOCK_ELEMS * 4];
920    let mut bsums = vec![0i16; nb * 4 * 8];
921    let mut d = vec![0f32; nb * 4];
922    for (a, act) in acts.iter().enumerate() {
923        debug_assert_eq!(act.n_blocks(), nb);
924        for b in 0..nb {
925            let src = &act.q[b * Q4_K_BLOCK_ELEMS..(b + 1) * Q4_K_BLOCK_ELEMS];
926            let dst = &mut qs[b * Q4_K_BLOCK_ELEMS * 4..(b + 1) * Q4_K_BLOCK_ELEMS * 4];
927            for (c, run) in src.as_chunks::<8>().0.iter().enumerate() {
928                dst[c * 32 + a * 8..c * 32 + a * 8 + 8].copy_from_slice(run);
929            }
930            let src_bs = &act.bsums[b * 16..(b + 1) * 16];
931            let dst_bs = &mut bsums[(b * 4 + a) * 8..(b * 4 + a) * 8 + 8];
932            for (slot, pair) in dst_bs.iter_mut().zip(src_bs.as_chunks::<2>().0) {
933                *slot = pair[0] + pair[1];
934            }
935            d[b * 4 + a] = act.d[b];
936        }
937    }
938    Q8KActsX4 {
939        na,
940        n_blocks: nb,
941        qs,
942        bsums,
943        d,
944    }
945}
946
947/// Whether [`gemm_q4_kx8_group_x4`] is the fast Q4_K batch path on this CPU:
948/// ARM i8mm with the interleave-8 layout. Everywhere else preparing the quad
949/// buys nothing — x86 dispatches to AVX2 inside the per-activation GEMV — so
950/// callers should keep using [`gemm_q4_kx8_group`] and skip the tiles.
951#[inline]
952pub fn q4_kx8_gemm_uses_acts_x4(interleave: usize) -> bool {
953    #[cfg(target_arch = "aarch64")]
954    {
955        interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
956    }
957    #[cfg(not(target_arch = "aarch64"))]
958    {
959        let _ = interleave;
960        false
961    }
962}
963
964/// [`gemm_q4_kx8_group`] against a pre-interleaved activation quad.
965///
966/// Callers build the quad once per matmul with [`prepare_q8_k_acts_x4`] and
967/// pass it to every row-group, hoisting what the i8mm kernel used to redo
968/// `rows/8` times. Only the interleave-8 layout has this kernel; gate call
969/// sites with [`q4_kx8_gemm_uses_acts_x4`].
970///
971/// `out` is `[row][act]` over the `tile.na` real activations:
972/// `out[r * tile.na + a]`, matching [`gemm_q4_kx8_group`].
973pub fn gemm_q4_kx8_group_x4(
974    packed: &[u8],
975    group: usize,
976    tile: &Q8KActsX4,
977    n_cols: usize,
978    interleave: usize,
979    out: &mut [f32],
980) {
981    gemm_q4_kx8_group_x4_on(
982        packed,
983        group,
984        tile,
985        n_cols,
986        interleave,
987        AccelX4::detect(),
988        out,
989    );
990}
991
992/// [`gemm_q4_kx8_group_x4`] with the kernel choice already made. Hoist
993/// [`AccelX4::detect`] out of the (row-group × quad) loop and pass it here;
994/// see [`AccelX4`].
995#[inline]
996pub fn gemm_q4_kx8_group_x4_on(
997    packed: &[u8],
998    group: usize,
999    tile: &Q8KActsX4,
1000    n_cols: usize,
1001    interleave: usize,
1002    accel: AccelX4,
1003    out: &mut [f32],
1004) {
1005    assert_eq!(
1006        interleave, 8,
1007        "the x4 GEMM only exists for interleave-8 packing"
1008    );
1009    assert_eq!(out.len(), Q4_KX8_NROWS * tile.na);
1010    assert!(n_cols.is_multiple_of(Q4_K_BLOCK_ELEMS));
1011    debug_assert_eq!(tile.n_blocks, n_cols / Q4_K_BLOCK_ELEMS);
1012    if tile.na == 0 {
1013        return;
1014    }
1015    let nb = n_cols / Q4_K_BLOCK_ELEMS;
1016    let off = group * nb * Q4_KX8_BLOCK_BYTES;
1017    let slice = &packed[off..off + nb * Q4_KX8_BLOCK_BYTES];
1018
1019    #[cfg(target_arch = "aarch64")]
1020    if accel == AccelX4::NeonI8mm {
1021        unsafe {
1022            neon::gemm_q4_kx8_q8_k_neon_i8mm(slice, tile, n_cols, out);
1023        }
1024        return;
1025    }
1026    let _ = accel;
1027    gemm_q4_kx8_acts_x4_scalar_8(slice, tile, n_cols, out);
1028}
1029
1030/// Portable reference for the ×4 GEMM: the same math as
1031/// [`gemv_q4_kx8_q8_k_scalar_8`], per quad row, reading qs / folded bsums /
1032/// d straight out of the pre-interleaved [`Q8KActsX4`]. Bit-identical to
1033/// running that GEMV per activation, which is what the tests assert.
1034fn gemm_q4_kx8_acts_x4_scalar_8(packed: &[u8], tile: &Q8KActsX4, n_cols: usize, out: &mut [f32]) {
1035    let nb = n_cols / Q4_K_BLOCK_ELEMS;
1036    let blocklen = 8;
1037    let ncols_interleaved = Q4_KX8_NROWS;
1038    let na = tile.na;
1039    let mut sumf = [[0f32; Q4_KX8_NROWS]; Q4_KX8_GEMM_NC];
1040    let mut sum_minf = [[0f32; Q4_KX8_NROWS]; Q4_KX8_GEMM_NC];
1041    for l in 0..nb {
1042        let blk = &packed[l * Q4_KX8_BLOCK_BYTES..][..Q4_KX8_BLOCK_BYTES];
1043        let d = &blk[0..16];
1044        let dmin = &blk[16..32];
1045        let scales = &blk[32..128];
1046        let qs = &blk[128..];
1047        let q8 = &tile.qs[l * Q4_K_BLOCK_ELEMS * 4..][..Q4_K_BLOCK_ELEMS * 4];
1048
1049        let mut all_scales = [[0u8; 8]; 8];
1050        let mut all_mins = [[0u8; 8]; 8];
1051        for sb in 0..8 {
1052            decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
1053        }
1054
1055        for a in 0..na {
1056            let da = tile.d[l * 4 + a];
1057            let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen); // 16
1058            for k in 0..n_k {
1059                let sb_pair = k / 4;
1060                let sc0 = &all_scales[sb_pair * 2];
1061                let sc1 = &all_scales[sb_pair * 2 + 1];
1062                for j in 0..ncols_interleaved {
1063                    let mut sumi = 0i32;
1064                    for i in 0..blocklen {
1065                        let qbyte = qs[k * ncols_interleaved * blocklen + j * blocklen + i];
1066                        let v0 = (qbyte & 0x0F) as i32;
1067                        let v1 = (qbyte >> 4) as i32;
1068                        // Canonical q8 element `e` lives at run `e/8`, row
1069                        // `a`, lane `e%8` of the interleaved block.
1070                        let e0 = (k >> 2) * 64 + (k % 4) * blocklen + i;
1071                        let e1 = e0 + 32;
1072                        let a0 = q8[(e0 / 8) * 32 + a * 8 + (e0 % 8)] as i32;
1073                        let a1 = q8[(e1 / 8) * 32 + a * 8 + (e1 % 8)] as i32;
1074                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
1075                    }
1076                    sumf[a][j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
1077                }
1078            }
1079            for (sb, mins) in all_mins.iter().enumerate() {
1080                let bsum = tile.bsums[(l * 4 + a) * 8 + sb] as i32;
1081                for j in 0..ncols_interleaved {
1082                    sum_minf[a][j] +=
1083                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
1084                }
1085            }
1086        }
1087    }
1088    for j in 0..ncols_interleaved {
1089        for (a, row) in sumf.iter().take(na).enumerate() {
1090            out[j * na + a] = row[j] - sum_minf[a][j];
1091        }
1092    }
1093}
1094
1095// ---------------------------------------------------------------------------
1096// Q5_K ×8 interleaved GEMV/GEMM (llama.cpp `block_q5_Kx8` / `ggml_gemm_q5_K_8x4`)
1097// ---------------------------------------------------------------------------
1098
1099/// Bytes per interleaved `block_q5_Kx8` (8 × f16 d + 8 × f16 dmin + 96 scales +
1100/// 256 qh + 1024 qs).
1101pub const Q5_KX8_BLOCK_BYTES: usize = 1408;
1102/// Number of Q5_K rows packed into one interleaved block.
1103pub const Q5_KX8_NROWS: usize = 8;
1104
1105/// Preferred qs/qh interleave width: 8 on x86 AVX2 and ARM i8mm
1106/// (`ggml_gemm_q5_K_8x8_q8_K`), 4 on DotProd-only NEON (`8x4`).
1107#[inline]
1108pub fn q5_kx8_interleave() -> usize {
1109    if cfg!(target_arch = "x86_64") {
1110        return 8;
1111    }
1112    #[cfg(target_arch = "aarch64")]
1113    {
1114        if std::arch::is_aarch64_feature_detected!("i8mm") {
1115            return 8;
1116        }
1117    }
1118    4
1119}
1120
1121/// Pack eight canonical Q5_K super-blocks (same column-block index) into
1122/// one `block_q5_Kx8`. `interleave` is 4 (ARM DotProd) or 8 (x86 / ARM i8mm).
1123pub fn make_block_q5_kx8(
1124    rows: [&[u8]; Q5_KX8_NROWS],
1125    interleave: usize,
1126) -> [u8; Q5_KX8_BLOCK_BYTES] {
1127    debug_assert!(interleave == 4 || interleave == 8);
1128    for r in &rows {
1129        debug_assert_eq!(r.len(), Q5_K_BLOCK_BYTES);
1130    }
1131    let mut out = [0u8; Q5_KX8_BLOCK_BYTES];
1132    // d[8] at 0, dmin[8] at 16, scales[96] at 32, qh[256] at 128, qs[1024] at 384.
1133    for (i, row) in rows.iter().enumerate() {
1134        out[i * 2] = row[0];
1135        out[i * 2 + 1] = row[1];
1136        out[16 + i * 2] = row[2];
1137        out[16 + i * 2 + 1] = row[3];
1138    }
1139
1140    let end = (Q5_K_BLOCK_ELEMS * 4) / interleave;
1141    let qs_out = &mut out[384..];
1142    for i in 0..end {
1143        let src_id = i % Q5_KX8_NROWS;
1144        let src_offset = (i / Q5_KX8_NROWS) * interleave;
1145        let dst_offset = i * interleave;
1146        let src_qs = &rows[src_id][48..176];
1147        qs_out[dst_offset..dst_offset + interleave]
1148            .copy_from_slice(&src_qs[src_offset..src_offset + interleave]);
1149    }
1150
1151    let qh_end = end / 4;
1152    let qh_out = &mut out[128..384];
1153    for i in 0..qh_end {
1154        let src_id = i % Q5_KX8_NROWS;
1155        let src_offset = (i / Q5_KX8_NROWS) * interleave;
1156        let dst_offset = i * interleave;
1157        let src_qh = &rows[src_id][16..48];
1158        qh_out[dst_offset..dst_offset + interleave]
1159            .copy_from_slice(&src_qh[src_offset..src_offset + interleave]);
1160    }
1161
1162    // Scale/min rearrangement (same 6-bit packing as Q4_Kx8).
1163    let mut s = [0u8; 8];
1164    let mut m = [0u8; 8];
1165    let scales_out = &mut out[32..128];
1166
1167    for i in 0..4 {
1168        for j in 0..8 {
1169            let sc = &rows[j][4..16];
1170            s[j] = sc[i] & 63;
1171            m[j] = sc[i + 4] & 63;
1172        }
1173        let base = i * 12;
1174        scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
1175        scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
1176        scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
1177        scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
1178        scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
1179        scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
1180        scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
1181        scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
1182        scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
1183        scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
1184        scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
1185        scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
1186    }
1187
1188    for i in 0..4 {
1189        for j in 0..8 {
1190            let sc = &rows[j][4..16];
1191            s[j] = ((sc[i] & 192) >> 2) | (sc[i + 8] & 15);
1192            m[j] = ((sc[i + 4] & 192) >> 2) | ((sc[i + 8] & 240) >> 4);
1193        }
1194        let base = 48 + i * 12;
1195        scales_out[base] = (s[0] & 63) + ((s[4] & 48) << 2);
1196        scales_out[base + 1] = (s[1] & 63) + ((s[5] & 48) << 2);
1197        scales_out[base + 2] = (s[2] & 63) + ((s[6] & 48) << 2);
1198        scales_out[base + 3] = (s[3] & 63) + ((s[7] & 48) << 2);
1199        scales_out[base + 4] = (m[0] & 63) + ((m[4] & 48) << 2);
1200        scales_out[base + 5] = (m[1] & 63) + ((m[5] & 48) << 2);
1201        scales_out[base + 6] = (m[2] & 63) + ((m[6] & 48) << 2);
1202        scales_out[base + 7] = (m[3] & 63) + ((m[7] & 48) << 2);
1203        scales_out[base + 8] = (s[4] & 15) + ((m[4] & 15) << 4);
1204        scales_out[base + 9] = (s[5] & 15) + ((m[5] & 15) << 4);
1205        scales_out[base + 10] = (s[6] & 15) + ((m[6] & 15) << 4);
1206        scales_out[base + 11] = (s[7] & 15) + ((m[7] & 15) << 4);
1207    }
1208
1209    out
1210}
1211
1212/// Repack a full Q5_K matrix (row-major canonical blocks) into interleaved
1213/// `block_q5_Kx8` groups. Tail rows (not divisible by 8) are omitted.
1214pub fn pack_q5_k_matrix_x8(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
1215    assert!(cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
1216    let n_blocks = cols / Q5_K_BLOCK_ELEMS;
1217    let row_bytes = n_blocks * Q5_K_BLOCK_BYTES;
1218    assert_eq!(data.len(), rows * row_bytes);
1219    let n_groups = rows / Q5_KX8_NROWS;
1220    let mut out = Vec::with_capacity(n_groups * n_blocks * Q5_KX8_BLOCK_BYTES);
1221    for g in 0..n_groups {
1222        for b in 0..n_blocks {
1223            let mut row_refs: [&[u8]; Q5_KX8_NROWS] = [&[]; Q5_KX8_NROWS];
1224            for (r, slot) in row_refs.iter_mut().enumerate() {
1225                let base = (g * Q5_KX8_NROWS + r) * row_bytes + b * Q5_K_BLOCK_BYTES;
1226                *slot = &data[base..base + Q5_K_BLOCK_BYTES];
1227            }
1228            out.extend_from_slice(&make_block_q5_kx8(row_refs, interleave));
1229        }
1230    }
1231    out
1232}
1233
1234/// Scalar GEMV for interleave=4 (`ggml_gemv_q5_K_8x4_q8_K_generic`).
1235fn gemv_q5_kx8_q8_k_scalar_4(
1236    packed: &[u8],
1237    act: &Q8KActivations,
1238    n_cols: usize,
1239    n_row_groups: usize,
1240    out: &mut [f32],
1241) {
1242    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1243    let blocklen = 4;
1244    let ncols_interleaved = Q5_KX8_NROWS;
1245    debug_assert_eq!(act.n_blocks(), nb);
1246    debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
1247    debug_assert_eq!(packed.len(), n_row_groups * nb * Q5_KX8_BLOCK_BYTES);
1248
1249    for x in 0..n_row_groups {
1250        let mut sumf = [0f32; 8];
1251        let mut sum_minf = [0f32; 8];
1252        let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
1253        for l in 0..nb {
1254            let blk = &packed[group_off + l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
1255            let d = &blk[0..16];
1256            let dmin = &blk[16..32];
1257            let scales = &blk[32..128];
1258            let qh = &blk[128..384];
1259            let qs = &blk[384..];
1260            let da = act.d[l];
1261            let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
1262            let bsums = &act.bsums[l * 16..(l + 1) * 16];
1263
1264            let mut all_scales = [[0u8; 8]; 8];
1265            let mut all_mins = [[0u8; 8]; 8];
1266            for sb in 0..8 {
1267                decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
1268            }
1269
1270            let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen); // 32
1271            for k in 0..n_k {
1272                let sb_pair = k / 8;
1273                let sc0 = &all_scales[sb_pair * 2];
1274                let sc1 = &all_scales[sb_pair * 2 + 1];
1275                let qh_shift = sb_pair * 2;
1276                for j in 0..ncols_interleaved {
1277                    let mut sumi = 0i32;
1278                    for i in 0..blocklen {
1279                        let b_qs_offset = k * ncols_interleaved * blocklen + j * blocklen + i;
1280                        let qh_idx = (k * blocklen + i) % 32;
1281                        let qh_chunk = qh_idx / blocklen;
1282                        let qh_pos = qh_idx % blocklen;
1283                        let b_qh_offset =
1284                            qh_chunk * (blocklen * ncols_interleaved) + j * blocklen + qh_pos;
1285                        let qh_val = qh[b_qh_offset];
1286                        let h0 = (qh_val >> qh_shift) & 1;
1287                        let h1 = (qh_val >> (qh_shift + 1)) & 1;
1288                        let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
1289                        let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
1290                        let a0 = q8[(k / 8) * 64 + (k % 8) * blocklen + i] as i32;
1291                        let a1 = q8[(k / 8) * 64 + (k % 8) * blocklen + i + 32] as i32;
1292                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
1293                    }
1294                    sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
1295                }
1296            }
1297            for sb in 0..8 {
1298                let mins = &all_mins[sb];
1299                let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
1300                for j in 0..ncols_interleaved {
1301                    sum_minf[j] +=
1302                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
1303                }
1304            }
1305        }
1306        let base = x * ncols_interleaved;
1307        for j in 0..ncols_interleaved {
1308            out[base + j] = sumf[j] - sum_minf[j];
1309        }
1310    }
1311}
1312
1313/// Scalar GEMV for interleave=8 (`ggml_gemv_q5_K_8x8_q8_K_generic`).
1314fn gemv_q5_kx8_q8_k_scalar_8(
1315    packed: &[u8],
1316    act: &Q8KActivations,
1317    n_cols: usize,
1318    n_row_groups: usize,
1319    out: &mut [f32],
1320) {
1321    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1322    let blocklen = 8;
1323    let ncols_interleaved = Q5_KX8_NROWS;
1324    debug_assert_eq!(act.n_blocks(), nb);
1325    debug_assert_eq!(out.len(), n_row_groups * ncols_interleaved);
1326
1327    for x in 0..n_row_groups {
1328        let mut sumf = [0f32; 8];
1329        let mut sum_minf = [0f32; 8];
1330        let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
1331        for l in 0..nb {
1332            let blk = &packed[group_off + l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
1333            let d = &blk[0..16];
1334            let dmin = &blk[16..32];
1335            let scales = &blk[32..128];
1336            let qh = &blk[128..384];
1337            let qs = &blk[384..];
1338            let da = act.d[l];
1339            let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
1340            let bsums = &act.bsums[l * 16..(l + 1) * 16];
1341
1342            let mut all_scales = [[0u8; 8]; 8];
1343            let mut all_mins = [[0u8; 8]; 8];
1344            for sb in 0..8 {
1345                decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
1346            }
1347
1348            let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen); // 16
1349            for k in 0..n_k {
1350                let sb_pair = k / 4;
1351                let sc0 = &all_scales[sb_pair * 2];
1352                let sc1 = &all_scales[sb_pair * 2 + 1];
1353                let qh_shift = sb_pair * 2;
1354                for j in 0..ncols_interleaved {
1355                    let mut sumi = 0i32;
1356                    for i in 0..blocklen {
1357                        let b_qs_offset = k * ncols_interleaved * blocklen + j * blocklen + i;
1358                        let qh_idx = (k * blocklen + i) % 32;
1359                        let qh_chunk = qh_idx / blocklen;
1360                        let qh_pos = qh_idx % blocklen;
1361                        let b_qh_offset =
1362                            qh_chunk * (blocklen * ncols_interleaved) + j * blocklen + qh_pos;
1363                        let qh_val = qh[b_qh_offset];
1364                        let h0 = (qh_val >> qh_shift) & 1;
1365                        let h1 = (qh_val >> (qh_shift + 1)) & 1;
1366                        let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
1367                        let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
1368                        let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
1369                        let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
1370                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
1371                    }
1372                    sumf[j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
1373                }
1374            }
1375            for sb in 0..8 {
1376                let mins = &all_mins[sb];
1377                let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
1378                for j in 0..ncols_interleaved {
1379                    sum_minf[j] +=
1380                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
1381                }
1382            }
1383        }
1384        let base = x * ncols_interleaved;
1385        for j in 0..ncols_interleaved {
1386            out[base + j] = sumf[j] - sum_minf[j];
1387        }
1388    }
1389}
1390
1391/// GEMV: interleaved Q5_K weights × Q8_K activation → `n_row_groups * 8` f32s.
1392pub fn gemv_q5_kx8_q8_k(
1393    packed: &[u8],
1394    act: &Q8KActivations,
1395    n_cols: usize,
1396    n_row_groups: usize,
1397    interleave: usize,
1398    out: &mut [f32],
1399) {
1400    assert!(n_cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
1401    assert_eq!(out.len(), n_row_groups * Q5_KX8_NROWS);
1402    match interleave {
1403        4 => {
1404            #[cfg(target_arch = "aarch64")]
1405            {
1406                if std::arch::is_aarch64_feature_detected!("dotprod") {
1407                    unsafe {
1408                        neon::gemv_q5_kx8_q8_k_neon_sdot(packed, act, n_cols, n_row_groups, out);
1409                    }
1410                    return;
1411                }
1412            }
1413            gemv_q5_kx8_q8_k_scalar_4(packed, act, n_cols, n_row_groups, out);
1414        }
1415        8 => {
1416            #[cfg(target_arch = "aarch64")]
1417            {
1418                if std::arch::is_aarch64_feature_detected!("dotprod") {
1419                    unsafe {
1420                        neon::gemv_q5_kx8_q8_k_neon_8x8(packed, act, n_cols, n_row_groups, out);
1421                    }
1422                    return;
1423                }
1424            }
1425            gemv_q5_kx8_q8_k_scalar_8(packed, act, n_cols, n_row_groups, out);
1426        }
1427        _ => panic!("q5_kx8 interleave must be 4 or 8, got {interleave}"),
1428    }
1429}
1430
1431/// One row-group (8 outputs) starting at `group` within a packed Q5_K matrix.
1432#[inline]
1433pub fn gemv_q5_kx8_group(
1434    packed: &[u8],
1435    group: usize,
1436    act: &Q8KActivations,
1437    n_cols: usize,
1438    interleave: usize,
1439    out8: &mut [f32],
1440) {
1441    debug_assert_eq!(out8.len(), Q5_KX8_NROWS);
1442    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1443    let off = group * nb * Q5_KX8_BLOCK_BYTES;
1444    let slice = &packed[off..off + nb * Q5_KX8_BLOCK_BYTES];
1445    gemv_q5_kx8_q8_k(slice, act, n_cols, 1, interleave, out8);
1446}
1447
1448/// How many activations one [`gemm_q5_kx8_group`] pass keeps in flight.
1449pub const Q5_KX8_GEMM_NC: usize = 4;
1450
1451/// GEMM counterpart of [`gemv_q5_kx8_group`]: one row-group (8 rows)
1452/// against `acts.len()` activations at once. `out` is `[row][act]`:
1453/// `out[r * acts.len() + j]`.
1454///
1455/// Weight-side decode (scales/mins/qh/qs addressing) is amortized across
1456/// the activation tile — same motivation as llama `ggml_gemm_q5_K_*`.
1457pub fn gemm_q5_kx8_group(
1458    packed: &[u8],
1459    group: usize,
1460    acts: &[Q8KActivations],
1461    n_cols: usize,
1462    interleave: usize,
1463    out: &mut [f32],
1464) {
1465    assert_eq!(out.len(), Q5_KX8_NROWS * acts.len());
1466    assert!(n_cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
1467    if acts.is_empty() {
1468        return;
1469    }
1470    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1471    let off = group * nb * Q5_KX8_BLOCK_BYTES;
1472    let slice = &packed[off..off + nb * Q5_KX8_BLOCK_BYTES];
1473    #[cfg(target_arch = "aarch64")]
1474    {
1475        if acts.len() <= Q5_KX8_GEMM_NC {
1476            if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
1477                // Compatibility entry: interleaves the quad here, once per
1478                // call. Batch callers should prepare the quad once per
1479                // matmul and use [`gemm_q5_kx8_group_x4`] instead.
1480                let tile = prepare_q8_k_acts_x4(acts, n_cols);
1481                unsafe {
1482                    neon::gemm_q5_kx8_q8_k_neon_i8mm(slice, &tile, n_cols, out);
1483                }
1484                return;
1485            }
1486            if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
1487                unsafe {
1488                    neon::gemm_q5_kx8_q8_k_neon_sdot(slice, acts, n_cols, out);
1489                }
1490                return;
1491            }
1492        }
1493    }
1494    match interleave {
1495        4 => gemm_q5_kx8_q8_k_scalar_4(slice, acts, n_cols, out),
1496        8 => gemm_q5_kx8_q8_k_scalar_8(slice, acts, n_cols, out),
1497        _ => panic!("q5_kx8 interleave must be 4 or 8, got {interleave}"),
1498    }
1499}
1500
1501/// Whether [`gemm_q5_kx8_group_x4`] is the fast Q5_K batch path on this CPU:
1502/// ARM i8mm with the interleave-8 layout (`ggml_gemm_q5_K_8x8_q8_K`).
1503#[inline]
1504pub fn q5_kx8_gemm_uses_acts_x4(interleave: usize) -> bool {
1505    #[cfg(target_arch = "aarch64")]
1506    {
1507        interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
1508    }
1509    #[cfg(not(target_arch = "aarch64"))]
1510    {
1511        let _ = interleave;
1512        false
1513    }
1514}
1515
1516/// [`gemm_q5_kx8_group`] against a pre-interleaved activation quad; the
1517/// Q5_K counterpart of [`gemm_q4_kx8_group_x4`], with the same contract:
1518/// interleave-8 packing only, quad prepared once per matmul by
1519/// [`prepare_q8_k_acts_x4`], `out[r * tile.na + a]`.
1520pub fn gemm_q5_kx8_group_x4(
1521    packed: &[u8],
1522    group: usize,
1523    tile: &Q8KActsX4,
1524    n_cols: usize,
1525    interleave: usize,
1526    out: &mut [f32],
1527) {
1528    gemm_q5_kx8_group_x4_on(
1529        packed,
1530        group,
1531        tile,
1532        n_cols,
1533        interleave,
1534        AccelX4::detect(),
1535        out,
1536    );
1537}
1538
1539/// [`gemm_q5_kx8_group_x4`] with the kernel choice already made; see
1540/// [`AccelX4`].
1541#[inline]
1542pub fn gemm_q5_kx8_group_x4_on(
1543    packed: &[u8],
1544    group: usize,
1545    tile: &Q8KActsX4,
1546    n_cols: usize,
1547    interleave: usize,
1548    accel: AccelX4,
1549    out: &mut [f32],
1550) {
1551    assert_eq!(
1552        interleave, 8,
1553        "the x4 GEMM only exists for interleave-8 packing"
1554    );
1555    assert_eq!(out.len(), Q5_KX8_NROWS * tile.na);
1556    assert!(n_cols.is_multiple_of(Q5_K_BLOCK_ELEMS));
1557    debug_assert_eq!(tile.n_blocks, n_cols / Q5_K_BLOCK_ELEMS);
1558    if tile.na == 0 {
1559        return;
1560    }
1561    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1562    let off = group * nb * Q5_KX8_BLOCK_BYTES;
1563    let slice = &packed[off..off + nb * Q5_KX8_BLOCK_BYTES];
1564
1565    #[cfg(target_arch = "aarch64")]
1566    if accel == AccelX4::NeonI8mm {
1567        unsafe {
1568            neon::gemm_q5_kx8_q8_k_neon_i8mm(slice, tile, n_cols, out);
1569        }
1570        return;
1571    }
1572    let _ = accel;
1573    gemm_q5_kx8_acts_x4_scalar_8(slice, tile, n_cols, out);
1574}
1575
1576/// Portable reference for the Q5_K ×4 GEMM: the same math as
1577/// [`gemv_q5_kx8_q8_k_scalar_8`], per quad row, reading qs / qh / folded
1578/// bsums / d out of the pre-interleaved [`Q8KActsX4`]. Bit-identical to
1579/// running that GEMV per activation, which is what the tests assert.
1580fn gemm_q5_kx8_acts_x4_scalar_8(packed: &[u8], tile: &Q8KActsX4, n_cols: usize, out: &mut [f32]) {
1581    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1582    let blocklen = 8;
1583    let ncols_interleaved = Q5_KX8_NROWS;
1584    let na = tile.na;
1585    let mut sumf = [[0f32; Q5_KX8_NROWS]; Q5_KX8_GEMM_NC];
1586    let mut sum_minf = [[0f32; Q5_KX8_NROWS]; Q5_KX8_GEMM_NC];
1587    for l in 0..nb {
1588        let blk = &packed[l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
1589        let d = &blk[0..16];
1590        let dmin = &blk[16..32];
1591        let scales = &blk[32..128];
1592        let qh = &blk[128..384];
1593        let qs = &blk[384..];
1594        let q8 = &tile.qs[l * Q5_K_BLOCK_ELEMS * 4..][..Q5_K_BLOCK_ELEMS * 4];
1595
1596        let mut all_scales = [[0u8; 8]; 8];
1597        let mut all_mins = [[0u8; 8]; 8];
1598        for sb in 0..8 {
1599            decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
1600        }
1601
1602        for a in 0..na {
1603            let da = tile.d[l * 4 + a];
1604            let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen); // 16
1605            for k in 0..n_k {
1606                let sb_pair = k / 4;
1607                let sc0 = &all_scales[sb_pair * 2];
1608                let sc1 = &all_scales[sb_pair * 2 + 1];
1609                let qh_shift = sb_pair * 2;
1610                for j in 0..ncols_interleaved {
1611                    let mut sumi = 0i32;
1612                    for i in 0..blocklen {
1613                        let b_qs_offset = k * ncols_interleaved * blocklen + j * blocklen + i;
1614                        let qh_idx = (k * blocklen + i) % 32;
1615                        let qh_chunk = qh_idx / blocklen;
1616                        let qh_pos = qh_idx % blocklen;
1617                        let b_qh_offset =
1618                            qh_chunk * (blocklen * ncols_interleaved) + j * blocklen + qh_pos;
1619                        let qh_val = qh[b_qh_offset];
1620                        let h0 = (qh_val >> qh_shift) & 1;
1621                        let h1 = (qh_val >> (qh_shift + 1)) & 1;
1622                        let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
1623                        let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
1624                        // Canonical q8 element `e` lives at run `e/8`, row
1625                        // `a`, lane `e%8` of the interleaved block.
1626                        let e0 = (k >> 2) * 64 + (k % 4) * blocklen + i;
1627                        let e1 = e0 + 32;
1628                        let a0 = q8[(e0 / 8) * 32 + a * 8 + (e0 % 8)] as i32;
1629                        let a1 = q8[(e1 / 8) * 32 + a * 8 + (e1 % 8)] as i32;
1630                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
1631                    }
1632                    sumf[a][j] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
1633                }
1634            }
1635            for (sb, mins) in all_mins.iter().enumerate() {
1636                let bsum = tile.bsums[(l * 4 + a) * 8 + sb] as i32;
1637                for j in 0..ncols_interleaved {
1638                    sum_minf[a][j] +=
1639                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
1640                }
1641            }
1642        }
1643    }
1644    for j in 0..ncols_interleaved {
1645        for (a, row) in sumf.iter().take(na).enumerate() {
1646            out[j * na + a] = row[j] - sum_minf[a][j];
1647        }
1648    }
1649}
1650
1651fn gemm_q5_kx8_q8_k_scalar_4(
1652    packed: &[u8],
1653    acts: &[Q8KActivations],
1654    n_cols: usize,
1655    out: &mut [f32],
1656) {
1657    let na = acts.len();
1658    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1659    let blocklen = 4;
1660    let ncols = Q5_KX8_NROWS;
1661    out.fill(0.0);
1662    let mut sum_minf = vec![0f32; ncols * na];
1663    for l in 0..nb {
1664        let blk = &packed[l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
1665        let d = &blk[0..16];
1666        let dmin = &blk[16..32];
1667        let scales = &blk[32..128];
1668        let qh = &blk[128..384];
1669        let qs = &blk[384..];
1670        let mut all_scales = [[0u8; 8]; 8];
1671        let mut all_mins = [[0u8; 8]; 8];
1672        for sb in 0..8 {
1673            decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
1674        }
1675        let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen);
1676        for (a, act) in acts.iter().enumerate() {
1677            let da = act.d[l];
1678            let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
1679            let bsums = &act.bsums[l * 16..(l + 1) * 16];
1680            for k in 0..n_k {
1681                let sb_pair = k / 8;
1682                let sc0 = &all_scales[sb_pair * 2];
1683                let sc1 = &all_scales[sb_pair * 2 + 1];
1684                let qh_shift = sb_pair * 2;
1685                for j in 0..ncols {
1686                    let mut sumi = 0i32;
1687                    for i in 0..blocklen {
1688                        let b_qs_offset = k * ncols * blocklen + j * blocklen + i;
1689                        let qh_idx = (k * blocklen + i) % 32;
1690                        let qh_chunk = qh_idx / blocklen;
1691                        let qh_pos = qh_idx % blocklen;
1692                        let b_qh_offset = qh_chunk * (blocklen * ncols) + j * blocklen + qh_pos;
1693                        let qh_val = qh[b_qh_offset];
1694                        let h0 = (qh_val >> qh_shift) & 1;
1695                        let h1 = (qh_val >> (qh_shift + 1)) & 1;
1696                        let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
1697                        let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
1698                        let a0 = q8[(k / 8) * 64 + (k % 8) * blocklen + i] as i32;
1699                        let a1 = q8[(k / 8) * 64 + (k % 8) * blocklen + i + 32] as i32;
1700                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
1701                    }
1702                    out[j * na + a] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
1703                }
1704            }
1705            for sb in 0..8 {
1706                let mins = &all_mins[sb];
1707                let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
1708                for j in 0..ncols {
1709                    sum_minf[j * na + a] +=
1710                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
1711                }
1712            }
1713        }
1714    }
1715    for i in 0..ncols * na {
1716        out[i] -= sum_minf[i];
1717    }
1718}
1719
1720fn gemm_q5_kx8_q8_k_scalar_8(
1721    packed: &[u8],
1722    acts: &[Q8KActivations],
1723    n_cols: usize,
1724    out: &mut [f32],
1725) {
1726    let na = acts.len();
1727    let nb = n_cols / Q5_K_BLOCK_ELEMS;
1728    let blocklen = 8;
1729    let ncols = Q5_KX8_NROWS;
1730    out.fill(0.0);
1731    let mut sum_minf = vec![0f32; ncols * na];
1732    for l in 0..nb {
1733        let blk = &packed[l * Q5_KX8_BLOCK_BYTES..][..Q5_KX8_BLOCK_BYTES];
1734        let d = &blk[0..16];
1735        let dmin = &blk[16..32];
1736        let scales = &blk[32..128];
1737        let qh = &blk[128..384];
1738        let qs = &blk[384..];
1739        let mut all_scales = [[0u8; 8]; 8];
1740        let mut all_mins = [[0u8; 8]; 8];
1741        for sb in 0..8 {
1742            decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
1743        }
1744        let n_k = Q5_K_BLOCK_ELEMS / (2 * blocklen);
1745        for (a, act) in acts.iter().enumerate() {
1746            let da = act.d[l];
1747            let q8 = &act.q[l * Q5_K_BLOCK_ELEMS..(l + 1) * Q5_K_BLOCK_ELEMS];
1748            let bsums = &act.bsums[l * 16..(l + 1) * 16];
1749            for k in 0..n_k {
1750                let sb_pair = k / 4;
1751                let sc0 = &all_scales[sb_pair * 2];
1752                let sc1 = &all_scales[sb_pair * 2 + 1];
1753                let qh_shift = sb_pair * 2;
1754                for j in 0..ncols {
1755                    let mut sumi = 0i32;
1756                    for i in 0..blocklen {
1757                        let b_qs_offset = k * ncols * blocklen + j * blocklen + i;
1758                        let qh_idx = (k * blocklen + i) % 32;
1759                        let qh_chunk = qh_idx / blocklen;
1760                        let qh_pos = qh_idx % blocklen;
1761                        let b_qh_offset = qh_chunk * (blocklen * ncols) + j * blocklen + qh_pos;
1762                        let qh_val = qh[b_qh_offset];
1763                        let h0 = (qh_val >> qh_shift) & 1;
1764                        let h1 = (qh_val >> (qh_shift + 1)) & 1;
1765                        let v0 = ((qs[b_qs_offset] & 0x0F) | (h0 << 4)) as i32;
1766                        let v1 = ((qs[b_qs_offset] >> 4) | (h1 << 4)) as i32;
1767                        let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
1768                        let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
1769                        sumi += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
1770                    }
1771                    out[j * na + a] += sumi as f32 * f16_from_bytes(&d[j * 2..]) * da;
1772                }
1773            }
1774            for sb in 0..8 {
1775                let mins = &all_mins[sb];
1776                let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
1777                for j in 0..ncols {
1778                    sum_minf[j * na + a] +=
1779                        mins[j] as f32 * bsum as f32 * f16_from_bytes(&dmin[j * 2..]) * da;
1780                }
1781            }
1782        }
1783    }
1784    for i in 0..ncols * na {
1785        out[i] -= sum_minf[i];
1786    }
1787}
1788
1789// ---------------------------------------------------------------------------
1790// Q6_K ×8 interleaved GEMV/GEMM (llama.cpp `block_q6_Kx8`)
1791// ---------------------------------------------------------------------------
1792
1793/// Bytes per interleaved `block_q6_Kx8` (8×f16 d + 128 scales + 1024 ql + 512 qh).
1794pub const Q6_KX8_BLOCK_BYTES: usize = 1680;
1795/// Number of Q6_K rows packed into one interleaved block.
1796pub const Q6_KX8_NROWS: usize = 8;
1797
1798/// Preferred ql/qh interleave width: 8 on x86 AVX2 and ARM i8mm
1799/// (`ggml_gemm_q6_K_8x8_q8_K`), 4 on DotProd-only NEON.
1800#[inline]
1801pub fn q6_kx8_interleave() -> usize {
1802    if cfg!(target_arch = "x86_64") {
1803        return 8;
1804    }
1805    #[cfg(target_arch = "aarch64")]
1806    {
1807        if std::arch::is_aarch64_feature_detected!("i8mm") {
1808            return 8;
1809        }
1810    }
1811    4
1812}
1813
1814/// Pack eight canonical Q6_K super-blocks into one `block_q6_Kx8`.
1815pub fn make_block_q6_kx8(
1816    rows: [&[u8]; Q6_KX8_NROWS],
1817    interleave: usize,
1818) -> [u8; Q6_KX8_BLOCK_BYTES] {
1819    debug_assert!(interleave == 4 || interleave == 8);
1820    for r in &rows {
1821        debug_assert_eq!(r.len(), Q6_K_BLOCK_BYTES);
1822    }
1823    let mut out = [0u8; Q6_KX8_BLOCK_BYTES];
1824    // d[8] @0, scales[128] @16, ql[1024] @144, qh[512] @1168
1825    for (i, row) in rows.iter().enumerate() {
1826        out[i * 2] = row[208];
1827        out[i * 2 + 1] = row[209];
1828    }
1829    let end_ls = (Q6_K_BLOCK_ELEMS * 4) / interleave;
1830    let ql_out = &mut out[144..1168];
1831    for i in 0..end_ls {
1832        let src_id = i % Q6_KX8_NROWS;
1833        let src_offset = (i / Q6_KX8_NROWS) * interleave;
1834        let dst_offset = i * interleave;
1835        let src_ql = &rows[src_id][0..128];
1836        ql_out[dst_offset..dst_offset + interleave]
1837            .copy_from_slice(&src_ql[src_offset..src_offset + interleave]);
1838    }
1839    let end_hs = end_ls / 2;
1840    let qh_out = &mut out[1168..];
1841    for i in 0..end_hs {
1842        let src_id = i % Q6_KX8_NROWS;
1843        let src_offset = (i / Q6_KX8_NROWS) * interleave;
1844        let dst_offset = i * interleave;
1845        let src_qh = &rows[src_id][128..192];
1846        qh_out[dst_offset..dst_offset + interleave]
1847            .copy_from_slice(&src_qh[src_offset..src_offset + interleave]);
1848    }
1849    let n_scales = Q6_K_BLOCK_ELEMS / 16;
1850    let scales_out = &mut out[16..144];
1851    for i in 0..Q6_KX8_NROWS {
1852        let src_sc = &rows[i][192..208];
1853        for j in 0..n_scales {
1854            scales_out[j * Q6_KX8_NROWS + i] = src_sc[j];
1855        }
1856    }
1857    out
1858}
1859
1860pub fn pack_q6_k_matrix_x8(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
1861    // Rows past the last full group of 8 are left canonical, same as the
1862    // other pack_*_matrix helpers; callers dot them row-by-row.
1863    assert!(cols.is_multiple_of(Q6_K_BLOCK_ELEMS));
1864    let row_bytes = (cols / Q6_K_BLOCK_ELEMS) * Q6_K_BLOCK_BYTES;
1865    assert_eq!(data.len(), rows * row_bytes);
1866    let n_blocks = cols / Q6_K_BLOCK_ELEMS;
1867    let n_groups = rows / Q6_KX8_NROWS;
1868    let mut out = Vec::with_capacity(n_groups * n_blocks * Q6_KX8_BLOCK_BYTES);
1869    for g in 0..n_groups {
1870        for b in 0..n_blocks {
1871            let mut row_refs: [&[u8]; Q6_KX8_NROWS] = [&[]; Q6_KX8_NROWS];
1872            for (r, slot) in row_refs.iter_mut().enumerate() {
1873                let base = (g * Q6_KX8_NROWS + r) * row_bytes + b * Q6_K_BLOCK_BYTES;
1874                *slot = &data[base..base + Q6_K_BLOCK_BYTES];
1875            }
1876            out.extend_from_slice(&make_block_q6_kx8(row_refs, interleave));
1877        }
1878    }
1879    out
1880}
1881
1882fn gemv_q6_kx8_q8_k_scalar(
1883    packed: &[u8],
1884    act: &Q8KActivations,
1885    n_cols: usize,
1886    n_row_groups: usize,
1887    blocklen: usize,
1888    out: &mut [f32],
1889) {
1890    let nb = n_cols / Q6_K_BLOCK_ELEMS;
1891    let ncols = Q6_KX8_NROWS;
1892    let blocks_per_half = 64 / blocklen;
1893    debug_assert_eq!(act.n_blocks(), nb);
1894    debug_assert_eq!(out.len(), n_row_groups * ncols);
1895    for x in 0..n_row_groups {
1896        let mut sumf = [0f32; 8];
1897        let group_off = x * nb * Q6_KX8_BLOCK_BYTES;
1898        for l in 0..nb {
1899            let blk = &packed[group_off + l * Q6_KX8_BLOCK_BYTES..][..Q6_KX8_BLOCK_BYTES];
1900            let d = &blk[0..16];
1901            let scales = &blk[16..144];
1902            let ql = &blk[144..1168];
1903            let qh = &blk[1168..];
1904            let da = act.d[l];
1905            let q8 = &act.q[l * Q6_K_BLOCK_ELEMS..(l + 1) * Q6_K_BLOCK_ELEMS];
1906            for k in 0..(Q6_K_BLOCK_ELEMS / (2 * blocklen)) {
1907                let base_l = (k / blocks_per_half) * 128 + (k % blocks_per_half) * blocklen;
1908                let base_h = base_l + 64;
1909                let scale_idx_l = base_l / 16;
1910                let scale_idx_h = base_h / 16;
1911                let qh_shift_l = ((base_l % 128) / 32) * 2;
1912                let qh_shift_h = ((base_h % 128) / 32) * 2;
1913                let qh_half_l = (base_l / 128) * 32;
1914                let qh_half_h = (base_h / 128) * 32;
1915                for j in 0..ncols {
1916                    let scale_l = scales[scale_idx_l * ncols + j] as i8 as i32;
1917                    let scale_h = scales[scale_idx_h * ncols + j] as i8 as i32;
1918                    let mut sumi_l = 0i32;
1919                    let mut sumi_h = 0i32;
1920                    for i in 0..blocklen {
1921                        let ql_pos = k * ncols * blocklen + j * blocklen + i;
1922                        let l_4 = (ql[ql_pos] & 0x0F) as i32;
1923                        let hi_4 = ((ql[ql_pos] >> 4) & 0x0F) as i32;
1924                        let qh_idx_l = qh_half_l + ((base_l + i) % 32);
1925                        let qh_chunk_l = qh_idx_l / blocklen;
1926                        let qh_pos_l = qh_idx_l % blocklen;
1927                        let qh_offset_l = qh_chunk_l * (blocklen * ncols) + j * blocklen + qh_pos_l;
1928                        let hi_2_l = ((qh[qh_offset_l] >> qh_shift_l) & 0x3) as i32;
1929                        let qh_idx_h = qh_half_h + ((base_h + i) % 32);
1930                        let qh_chunk_h = qh_idx_h / blocklen;
1931                        let qh_pos_h = qh_idx_h % blocklen;
1932                        let qh_offset_h = qh_chunk_h * (blocklen * ncols) + j * blocklen + qh_pos_h;
1933                        let hi_2_h = ((qh[qh_offset_h] >> qh_shift_h) & 0x3) as i32;
1934                        let q_l = ((hi_2_l << 4) | l_4) - 32;
1935                        let q_h = ((hi_2_h << 4) | hi_4) - 32;
1936                        sumi_l += q_l * (q8[base_l + i] as i32);
1937                        sumi_h += q_h * (q8[base_h + i] as i32);
1938                    }
1939                    sumf[j] += (sumi_l * scale_l + sumi_h * scale_h) as f32
1940                        * f16_from_bytes(&d[j * 2..])
1941                        * da;
1942                }
1943            }
1944        }
1945        let base = x * ncols;
1946        out[base..base + ncols].copy_from_slice(&sumf);
1947    }
1948}
1949
1950pub fn gemv_q6_kx8_q8_k(
1951    packed: &[u8],
1952    act: &Q8KActivations,
1953    n_cols: usize,
1954    n_row_groups: usize,
1955    interleave: usize,
1956    out: &mut [f32],
1957) {
1958    assert_eq!(out.len(), n_row_groups * Q6_KX8_NROWS);
1959    match interleave {
1960        4 => gemv_q6_kx8_q8_k_scalar(packed, act, n_cols, n_row_groups, 4, out),
1961        8 => {
1962            #[cfg(target_arch = "aarch64")]
1963            {
1964                if std::arch::is_aarch64_feature_detected!("dotprod") {
1965                    unsafe {
1966                        neon::gemv_q6_kx8_q8_k_neon_8x8(packed, act, n_cols, n_row_groups, out);
1967                    }
1968                    return;
1969                }
1970            }
1971            gemv_q6_kx8_q8_k_scalar(packed, act, n_cols, n_row_groups, 8, out)
1972        }
1973        _ => panic!("q6_kx8 interleave must be 4 or 8, got {interleave}"),
1974    }
1975}
1976
1977pub fn gemv_q6_kx8_group(
1978    packed: &[u8],
1979    group: usize,
1980    act: &Q8KActivations,
1981    n_cols: usize,
1982    interleave: usize,
1983    out8: &mut [f32],
1984) {
1985    debug_assert_eq!(out8.len(), Q6_KX8_NROWS);
1986    let nb = n_cols / Q6_K_BLOCK_ELEMS;
1987    let off = group * nb * Q6_KX8_BLOCK_BYTES;
1988    let slice = &packed[off..off + nb * Q6_KX8_BLOCK_BYTES];
1989    gemv_q6_kx8_q8_k(slice, act, n_cols, 1, interleave, out8);
1990}
1991
1992pub const Q6_KX8_GEMM_NC: usize = 8;
1993
1994/// Multi-act GEMM for one Q6_Kx8 row-group; weight decode amortized across acts.
1995pub fn gemm_q6_kx8_group(
1996    packed: &[u8],
1997    group: usize,
1998    acts: &[Q8KActivations],
1999    n_cols: usize,
2000    interleave: usize,
2001    out: &mut [f32],
2002) {
2003    assert_eq!(out.len(), Q6_KX8_NROWS * acts.len());
2004    assert!(n_cols.is_multiple_of(Q6_K_BLOCK_ELEMS));
2005    if acts.is_empty() {
2006        return;
2007    }
2008    let nb = n_cols / Q6_K_BLOCK_ELEMS;
2009    let off = group * nb * Q6_KX8_BLOCK_BYTES;
2010    let slice = &packed[off..off + nb * Q6_KX8_BLOCK_BYTES];
2011    #[cfg(target_arch = "aarch64")]
2012    {
2013        if interleave == 8
2014            && acts.len() <= Q8K_ACTS_X4_NC
2015            && std::arch::is_aarch64_feature_detected!("i8mm")
2016        {
2017            // Compatibility entry: interleaves the quad here, once per
2018            // call. Batch callers should prepare the quad once per matmul
2019            // and use [`gemm_q6_kx8_group_x4`] instead.
2020            let tile = prepare_q8_k_acts_x4(acts, n_cols);
2021            unsafe {
2022                neon::gemm_q6_kx8_q8_k_neon_i8mm(slice, &tile, n_cols, out);
2023            }
2024            return;
2025        }
2026    }
2027    let blocklen = interleave;
2028    assert!(blocklen == 4 || blocklen == 8);
2029    let na = acts.len();
2030    let ncols = Q6_KX8_NROWS;
2031    let blocks_per_half = 64 / blocklen;
2032    out.fill(0.0);
2033    for l in 0..nb {
2034        let blk = &slice[l * Q6_KX8_BLOCK_BYTES..][..Q6_KX8_BLOCK_BYTES];
2035        let d = &blk[0..16];
2036        let scales = &blk[16..144];
2037        let ql = &blk[144..1168];
2038        let qh = &blk[1168..];
2039        let mut d_f = [0f32; 8];
2040        for j in 0..8 {
2041            d_f[j] = f16_from_bytes(&d[j * 2..]);
2042        }
2043        for k in 0..(Q6_K_BLOCK_ELEMS / (2 * blocklen)) {
2044            let base_l = (k / blocks_per_half) * 128 + (k % blocks_per_half) * blocklen;
2045            let base_h = base_l + 64;
2046            let scale_idx_l = base_l / 16;
2047            let scale_idx_h = base_h / 16;
2048            let qh_shift_l = ((base_l % 128) / 32) * 2;
2049            let qh_shift_h = ((base_h % 128) / 32) * 2;
2050            let qh_half_l = (base_l / 128) * 32;
2051            let qh_half_h = (base_h / 128) * 32;
2052            for j in 0..ncols {
2053                let scale_l = scales[scale_idx_l * ncols + j] as i8 as i32;
2054                let scale_h = scales[scale_idx_h * ncols + j] as i8 as i32;
2055                // Decode 8 weight quants for this (k,j) once.
2056                let mut q_l = [0i32; 8];
2057                let mut q_h = [0i32; 8];
2058                for i in 0..blocklen {
2059                    let ql_pos = k * ncols * blocklen + j * blocklen + i;
2060                    let l_4 = (ql[ql_pos] & 0x0F) as i32;
2061                    let hi_4 = ((ql[ql_pos] >> 4) & 0x0F) as i32;
2062                    let qh_idx_l = qh_half_l + ((base_l + i) % 32);
2063                    let qh_chunk_l = qh_idx_l / blocklen;
2064                    let qh_pos_l = qh_idx_l % blocklen;
2065                    let qh_offset_l = qh_chunk_l * (blocklen * ncols) + j * blocklen + qh_pos_l;
2066                    let hi_2_l = ((qh[qh_offset_l] >> qh_shift_l) & 0x3) as i32;
2067                    let qh_idx_h = qh_half_h + ((base_h + i) % 32);
2068                    let qh_chunk_h = qh_idx_h / blocklen;
2069                    let qh_pos_h = qh_idx_h % blocklen;
2070                    let qh_offset_h = qh_chunk_h * (blocklen * ncols) + j * blocklen + qh_pos_h;
2071                    let hi_2_h = ((qh[qh_offset_h] >> qh_shift_h) & 0x3) as i32;
2072                    q_l[i] = ((hi_2_l << 4) | l_4) - 32;
2073                    q_h[i] = ((hi_2_h << 4) | hi_4) - 32;
2074                }
2075                for (a, act) in acts.iter().enumerate() {
2076                    let da = act.d[l];
2077                    let q8 = &act.q[l * Q6_K_BLOCK_ELEMS..(l + 1) * Q6_K_BLOCK_ELEMS];
2078                    let mut sumi_l = 0i32;
2079                    let mut sumi_h = 0i32;
2080                    for i in 0..blocklen {
2081                        sumi_l += q_l[i] * (q8[base_l + i] as i32);
2082                        sumi_h += q_h[i] * (q8[base_h + i] as i32);
2083                    }
2084                    out[j * na + a] += (sumi_l * scale_l + sumi_h * scale_h) as f32 * d_f[j] * da;
2085                }
2086            }
2087        }
2088    }
2089}
2090
2091/// Whether [`gemm_q6_kx8_group_x4`] is the fast Q6_K batch path on this CPU:
2092/// ARM i8mm with the interleave-8 layout (`ggml_gemm_q6_K_8x8_q8_K`). The
2093/// scalar Kx8 GEMM measured slower than the per-row NEON dot on ARM, so
2094/// batch callers should use the Kx8 layout only when this returns true.
2095#[inline]
2096pub fn q6_kx8_gemm_uses_acts_x4(interleave: usize) -> bool {
2097    #[cfg(target_arch = "aarch64")]
2098    {
2099        interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
2100    }
2101    #[cfg(not(target_arch = "aarch64"))]
2102    {
2103        let _ = interleave;
2104        false
2105    }
2106}
2107
2108/// [`gemm_q6_kx8_group`] against a pre-interleaved activation quad; the
2109/// Q6_K counterpart of [`gemm_q4_kx8_group_x4`], with the same contract:
2110/// interleave-8 packing only, quad prepared once per matmul by
2111/// [`prepare_q8_k_acts_x4`] (up to 4 activations, not [`Q6_KX8_GEMM_NC`]),
2112/// `out[r * tile.na + a]`.
2113pub fn gemm_q6_kx8_group_x4(
2114    packed: &[u8],
2115    group: usize,
2116    tile: &Q8KActsX4,
2117    n_cols: usize,
2118    interleave: usize,
2119    out: &mut [f32],
2120) {
2121    gemm_q6_kx8_group_x4_on(
2122        packed,
2123        group,
2124        tile,
2125        n_cols,
2126        interleave,
2127        AccelX4::detect(),
2128        out,
2129    );
2130}
2131
2132/// [`gemm_q6_kx8_group_x4`] with the kernel choice already made; see
2133/// [`AccelX4`].
2134#[inline]
2135pub fn gemm_q6_kx8_group_x4_on(
2136    packed: &[u8],
2137    group: usize,
2138    tile: &Q8KActsX4,
2139    n_cols: usize,
2140    interleave: usize,
2141    accel: AccelX4,
2142    out: &mut [f32],
2143) {
2144    assert_eq!(
2145        interleave, 8,
2146        "the x4 GEMM only exists for interleave-8 packing"
2147    );
2148    assert_eq!(out.len(), Q6_KX8_NROWS * tile.na);
2149    assert!(n_cols.is_multiple_of(Q6_K_BLOCK_ELEMS));
2150    debug_assert_eq!(tile.n_blocks, n_cols / Q6_K_BLOCK_ELEMS);
2151    if tile.na == 0 {
2152        return;
2153    }
2154    let nb = n_cols / Q6_K_BLOCK_ELEMS;
2155    let off = group * nb * Q6_KX8_BLOCK_BYTES;
2156    let slice = &packed[off..off + nb * Q6_KX8_BLOCK_BYTES];
2157
2158    #[cfg(target_arch = "aarch64")]
2159    if accel == AccelX4::NeonI8mm {
2160        unsafe {
2161            neon::gemm_q6_kx8_q8_k_neon_i8mm(slice, tile, n_cols, out);
2162        }
2163        return;
2164    }
2165    let _ = accel;
2166    gemm_q6_kx8_acts_x4_scalar_8(slice, tile, n_cols, out);
2167}
2168
2169/// Portable reference for the Q6_K ×4 GEMM: the same math as
2170/// [`gemv_q6_kx8_q8_k_scalar`] at blocklen 8, per quad row, reading qs and
2171/// d out of the pre-interleaved [`Q8KActsX4`] (Q6_K has no mins, so the
2172/// folded bsums are unused). Bit-identical to running that GEMV per
2173/// activation, which is what the tests assert.
2174fn gemm_q6_kx8_acts_x4_scalar_8(packed: &[u8], tile: &Q8KActsX4, n_cols: usize, out: &mut [f32]) {
2175    let nb = n_cols / Q6_K_BLOCK_ELEMS;
2176    let blocklen = 8;
2177    let ncols = Q6_KX8_NROWS;
2178    let blocks_per_half = 64 / blocklen;
2179    let na = tile.na;
2180    let mut sumf = [[0f32; Q6_KX8_NROWS]; Q8K_ACTS_X4_NC];
2181    for l in 0..nb {
2182        let blk = &packed[l * Q6_KX8_BLOCK_BYTES..][..Q6_KX8_BLOCK_BYTES];
2183        let d = &blk[0..16];
2184        let scales = &blk[16..144];
2185        let ql = &blk[144..1168];
2186        let qh = &blk[1168..];
2187        let q8 = &tile.qs[l * Q6_K_BLOCK_ELEMS * 4..][..Q6_K_BLOCK_ELEMS * 4];
2188        for a in 0..na {
2189            let da = tile.d[l * 4 + a];
2190            for k in 0..(Q6_K_BLOCK_ELEMS / (2 * blocklen)) {
2191                let base_l = (k / blocks_per_half) * 128 + (k % blocks_per_half) * blocklen;
2192                let base_h = base_l + 64;
2193                let scale_idx_l = base_l / 16;
2194                let scale_idx_h = base_h / 16;
2195                let qh_shift_l = ((base_l % 128) / 32) * 2;
2196                let qh_shift_h = ((base_h % 128) / 32) * 2;
2197                let qh_half_l = (base_l / 128) * 32;
2198                let qh_half_h = (base_h / 128) * 32;
2199                for j in 0..ncols {
2200                    let scale_l = scales[scale_idx_l * ncols + j] as i8 as i32;
2201                    let scale_h = scales[scale_idx_h * ncols + j] as i8 as i32;
2202                    let mut sumi_l = 0i32;
2203                    let mut sumi_h = 0i32;
2204                    for i in 0..blocklen {
2205                        let ql_pos = k * ncols * blocklen + j * blocklen + i;
2206                        let l_4 = (ql[ql_pos] & 0x0F) as i32;
2207                        let hi_4 = ((ql[ql_pos] >> 4) & 0x0F) as i32;
2208                        let qh_idx_l = qh_half_l + ((base_l + i) % 32);
2209                        let qh_chunk_l = qh_idx_l / blocklen;
2210                        let qh_pos_l = qh_idx_l % blocklen;
2211                        let qh_offset_l = qh_chunk_l * (blocklen * ncols) + j * blocklen + qh_pos_l;
2212                        let hi_2_l = ((qh[qh_offset_l] >> qh_shift_l) & 0x3) as i32;
2213                        let qh_idx_h = qh_half_h + ((base_h + i) % 32);
2214                        let qh_chunk_h = qh_idx_h / blocklen;
2215                        let qh_pos_h = qh_idx_h % blocklen;
2216                        let qh_offset_h = qh_chunk_h * (blocklen * ncols) + j * blocklen + qh_pos_h;
2217                        let hi_2_h = ((qh[qh_offset_h] >> qh_shift_h) & 0x3) as i32;
2218                        let q_l = ((hi_2_l << 4) | l_4) - 32;
2219                        let q_h = ((hi_2_h << 4) | hi_4) - 32;
2220                        // Canonical q8 element `e` lives at run `e/8`, row
2221                        // `a`, lane `e%8` of the interleaved block.
2222                        let e_l = base_l + i;
2223                        let e_h = base_h + i;
2224                        sumi_l += q_l * (q8[(e_l / 8) * 32 + a * 8 + (e_l % 8)] as i32);
2225                        sumi_h += q_h * (q8[(e_h / 8) * 32 + a * 8 + (e_h % 8)] as i32);
2226                    }
2227                    sumf[a][j] += (sumi_l * scale_l + sumi_h * scale_h) as f32
2228                        * f16_from_bytes(&d[j * 2..])
2229                        * da;
2230                }
2231            }
2232        }
2233    }
2234    for j in 0..ncols {
2235        for (a, row) in sumf.iter().take(na).enumerate() {
2236            out[j * na + a] = row[j];
2237        }
2238    }
2239}
2240
2241// ---------------------------------------------------------------------------
2242// Q4_0 ×4 interleaved GEMV/GEMM (llama.cpp `block_q4_0x4` / `ggml_gemv_q4_0_4x4`)
2243// ---------------------------------------------------------------------------
2244
2245/// Bytes per interleaved `block_q4_0x4` (4 × f16 d + 64 qs).
2246pub const Q4_0X4_BLOCK_BYTES: usize = 72;
2247/// Number of Q4_0 rows packed into one interleaved block.
2248pub const Q4_0X4_NROWS: usize = 4;
2249/// qs interleave width for `ggml_gemv_q4_0_4x4_q8_0` (NEON SDOT). The
2250/// DotProd-only default; [`q4_0x4_interleave`] picks 8 on i8mm hosts.
2251pub const Q4_0X4_INTERLEAVE: usize = 4;
2252
2253/// Preferred qs interleave width: 8 on ARM i8mm (`ggml_gemm_q4_0_4x8_q8_0`
2254/// via `ggml_repack_get_optimal_repack_type`), 4 on DotProd-only NEON and
2255/// everywhere else (the scalar fallback handles either).
2256#[inline]
2257pub fn q4_0x4_interleave() -> usize {
2258    #[cfg(target_arch = "aarch64")]
2259    {
2260        if std::arch::is_aarch64_feature_detected!("i8mm") {
2261            return 8;
2262        }
2263    }
2264    Q4_0X4_INTERLEAVE
2265}
2266
2267const Q4_0X4_XOR_MASK_U32: u32 = 0x8888_8888;
2268const Q4_0X4_XOR_MASK_U64: u64 = 0x8888_8888_8888_8888;
2269
2270/// Pack four canonical Q4_0 blocks (same column-block) into one
2271/// `block_q4_0x4`. Nibble bytes are XOR-masked during interleave so
2272/// NEON can unpack without explicit `- 8` bias subtraction.
2273pub fn make_block_q4_0x4(
2274    rows: [&[u8]; Q4_0X4_NROWS],
2275    interleave: usize,
2276) -> [u8; Q4_0X4_BLOCK_BYTES] {
2277    debug_assert!(interleave == 4 || interleave == 8);
2278    for r in &rows {
2279        debug_assert_eq!(r.len(), Q4_0_BLOCK_BYTES);
2280    }
2281    let mut out = [0u8; Q4_0X4_BLOCK_BYTES];
2282    for (i, row) in rows.iter().enumerate() {
2283        out[i * 2] = row[0];
2284        out[i * 2 + 1] = row[1];
2285    }
2286    let end = (Q4_0_BLOCK_ELEMS * 2) / interleave;
2287    let qs_out = &mut out[8..];
2288    for i in 0..end {
2289        let src_id = i % Q4_0X4_NROWS;
2290        let src_offset = (i / Q4_0X4_NROWS) * interleave;
2291        let dst_offset = i * interleave;
2292        let src_qs = &rows[src_id][2..18];
2293        if interleave == 4 {
2294            let mut elems = u32::from_le_bytes(
2295                src_qs[src_offset..src_offset + 4]
2296                    .try_into()
2297                    .expect("4-byte interleave chunk"),
2298            );
2299            elems ^= Q4_0X4_XOR_MASK_U32;
2300            qs_out[dst_offset..dst_offset + 4].copy_from_slice(&elems.to_le_bytes());
2301        } else {
2302            let mut elems = u64::from_le_bytes(
2303                src_qs[src_offset..src_offset + 8]
2304                    .try_into()
2305                    .expect("8-byte interleave chunk"),
2306            );
2307            elems ^= Q4_0X4_XOR_MASK_U64;
2308            qs_out[dst_offset..dst_offset + 8].copy_from_slice(&elems.to_le_bytes());
2309        }
2310    }
2311    out
2312}
2313
2314/// Repack a Q4_0 matrix into interleaved `block_q4_0x4` groups. Tail rows
2315/// (not divisible by 4) are omitted; caller dots them with [`crate::dot_q4_0_q8`].
2316pub fn pack_q4_0_matrix_x4(data: &[u8], rows: usize, cols: usize, interleave: usize) -> Vec<u8> {
2317    assert!(cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
2318    let n_blocks = cols / Q4_0_BLOCK_ELEMS;
2319    let row_bytes = n_blocks * Q4_0_BLOCK_BYTES;
2320    assert_eq!(data.len(), rows * row_bytes);
2321    let n_groups = rows / Q4_0X4_NROWS;
2322    let mut out = Vec::with_capacity(n_groups * n_blocks * Q4_0X4_BLOCK_BYTES);
2323    for g in 0..n_groups {
2324        for b in 0..n_blocks {
2325            let mut row_refs: [&[u8]; Q4_0X4_NROWS] = [&[]; Q4_0X4_NROWS];
2326            for (r, slot) in row_refs.iter_mut().enumerate() {
2327                let base = (g * Q4_0X4_NROWS + r) * row_bytes + b * Q4_0_BLOCK_BYTES;
2328                *slot = &data[base..base + Q4_0_BLOCK_BYTES];
2329            }
2330            out.extend_from_slice(&make_block_q4_0x4(row_refs, interleave));
2331        }
2332    }
2333    out
2334}
2335
2336#[inline]
2337fn q4_0x4_nibble_dot(byte: u8, q8_lo: i32, q8_hi: i32) -> i32 {
2338    let v0 = ((byte << 4) as i8) as i32;
2339    let v1 = ((byte & 0xF0) as i8) as i32;
2340    ((v0 * q8_lo) + (v1 * q8_hi)) >> 4
2341}
2342
2343/// Scalar GEMV (`ggml_gemv_q4_0_4x{4,8}_q8_0_generic`); `blocklen` is the
2344/// interleave the matrix was packed with.
2345fn gemv_q4_0x4_q8_0_scalar(
2346    packed: &[u8],
2347    act: &Q8Activations,
2348    n_cols: usize,
2349    n_row_groups: usize,
2350    blocklen: usize,
2351    out: &mut [f32],
2352) {
2353    let nb = n_cols / Q4_0_BLOCK_ELEMS;
2354    let ncols = Q4_0X4_NROWS;
2355    debug_assert_eq!(act.n_blocks(), nb);
2356    debug_assert_eq!(out.len(), n_row_groups * ncols);
2357    debug_assert_eq!(packed.len(), n_row_groups * nb * Q4_0X4_BLOCK_BYTES);
2358
2359    for x in 0..n_row_groups {
2360        let mut sumf = [0f32; 4];
2361        let group_off = x * nb * Q4_0X4_BLOCK_BYTES;
2362        for l in 0..nb {
2363            let blk = &packed[group_off + l * Q4_0X4_BLOCK_BYTES..][..Q4_0X4_BLOCK_BYTES];
2364            let da = act.d[l];
2365            let q8 = &act.q[l * Q4_0_BLOCK_ELEMS..(l + 1) * Q4_0_BLOCK_ELEMS];
2366            for k in 0..(Q4_0_BLOCK_ELEMS / (2 * blocklen)) {
2367                for j in 0..ncols {
2368                    let mut sumi = 0i32;
2369                    for i in 0..blocklen {
2370                        let byte = blk[8 + k * ncols * blocklen + j * blocklen + i];
2371                        sumi += q4_0x4_nibble_dot(
2372                            byte,
2373                            q8[k * blocklen + i] as i32,
2374                            q8[k * blocklen + i + Q4_0_BLOCK_ELEMS / 2] as i32,
2375                        );
2376                    }
2377                    sumf[j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
2378                }
2379            }
2380        }
2381        let base = x * ncols;
2382        out[base..base + ncols].copy_from_slice(&sumf);
2383    }
2384}
2385
2386/// GEMV: interleaved Q4_0 weights × Q8 activation → `n_row_groups * 4` f32s.
2387/// `interleave` must match the packing (4: NEON SDOT `4x4`; 8: NEON `4x8`).
2388pub fn gemv_q4_0x4_q8_0(
2389    packed: &[u8],
2390    act: &Q8Activations,
2391    n_cols: usize,
2392    n_row_groups: usize,
2393    interleave: usize,
2394    out: &mut [f32],
2395) {
2396    assert!(n_cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
2397    assert_eq!(out.len(), n_row_groups * Q4_0X4_NROWS);
2398    match interleave {
2399        4 => {
2400            #[cfg(target_arch = "aarch64")]
2401            {
2402                if std::arch::is_aarch64_feature_detected!("dotprod") {
2403                    unsafe {
2404                        neon::gemv_q4_0x4_q8_0_neon_sdot(packed, act, n_cols, n_row_groups, out);
2405                    }
2406                    return;
2407                }
2408            }
2409            gemv_q4_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 4, out);
2410        }
2411        8 => {
2412            #[cfg(target_arch = "aarch64")]
2413            {
2414                if std::arch::is_aarch64_feature_detected!("dotprod") {
2415                    unsafe {
2416                        neon::gemv_q4_0x4_q8_0_neon_4x8(packed, act, n_cols, n_row_groups, out);
2417                    }
2418                    return;
2419                }
2420            }
2421            gemv_q4_0x4_q8_0_scalar(packed, act, n_cols, n_row_groups, 8, out);
2422        }
2423        _ => panic!("q4_0x4 interleave must be 4 or 8, got {interleave}"),
2424    }
2425}
2426
2427/// How many activations one [`gemm_q4_0x4_group`] pass keeps in flight.
2428pub const Q4_0X4_GEMM_NC: usize = 4;
2429
2430/// GEMM counterpart of [`gemv_q4_0x4_group`]: one row-group (4 rows)
2431/// against `acts.len()` activations at once. `out` is `[row][act]`:
2432/// `out[r * acts.len() + j]`.
2433pub fn gemm_q4_0x4_group(
2434    packed: &[u8],
2435    group: usize,
2436    acts: &[Q8Activations],
2437    n_cols: usize,
2438    interleave: usize,
2439    out: &mut [f32],
2440) {
2441    assert_eq!(out.len(), Q4_0X4_NROWS * acts.len());
2442    assert!(n_cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
2443    if acts.is_empty() {
2444        return;
2445    }
2446    let nb = n_cols / Q4_0_BLOCK_ELEMS;
2447    let off = group * nb * Q4_0X4_BLOCK_BYTES;
2448    let slice = &packed[off..off + nb * Q4_0X4_BLOCK_BYTES];
2449
2450    #[cfg(target_arch = "aarch64")]
2451    {
2452        if interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm") {
2453            // Compatibility entry: interleaves the quads here, once per
2454            // call. Batch callers should prepare them once per matmul and
2455            // use [`gemm_q4_0x4_group_x4`] instead.
2456            for (t, chunk) in acts.chunks(Q8K_ACTS_X4_NC).enumerate() {
2457                let tile = prepare_q8_acts_x4(chunk, n_cols);
2458                let mut tmp = [0f32; Q4_0X4_NROWS * Q8K_ACTS_X4_NC];
2459                let n = chunk.len();
2460                unsafe {
2461                    neon::gemm_q4_0x4_q8_0_neon_i8mm(
2462                        slice,
2463                        &tile,
2464                        n_cols,
2465                        &mut tmp[..Q4_0X4_NROWS * n],
2466                    );
2467                }
2468                for r in 0..Q4_0X4_NROWS {
2469                    for j in 0..n {
2470                        out[r * acts.len() + t * Q8K_ACTS_X4_NC + j] = tmp[r * n + j];
2471                    }
2472                }
2473            }
2474            return;
2475        }
2476        if interleave == 4 && std::arch::is_aarch64_feature_detected!("dotprod") {
2477            unsafe {
2478                neon::gemm_q4_0x4_q8_0_neon_sdot(slice, acts, n_cols, out);
2479            }
2480            return;
2481        }
2482    }
2483    let mut tmp = [0f32; Q4_0X4_NROWS];
2484    for (j, act) in acts.iter().enumerate() {
2485        gemv_q4_0x4_q8_0(slice, act, n_cols, 1, interleave, &mut tmp);
2486        for (r, v) in tmp.iter().enumerate() {
2487            out[r * acts.len() + j] = *v;
2488        }
2489    }
2490}
2491
2492/// Whether [`gemm_q4_0x4_group_x4`] is the fast Q4_0 batch path on this
2493/// CPU: ARM i8mm with the interleave-8 layout (`ggml_gemm_q4_0_4x8_q8_0`).
2494#[inline]
2495pub fn q4_0x4_gemm_uses_acts_x4(interleave: usize) -> bool {
2496    #[cfg(target_arch = "aarch64")]
2497    {
2498        interleave == 8 && std::arch::is_aarch64_feature_detected!("i8mm")
2499    }
2500    #[cfg(not(target_arch = "aarch64"))]
2501    {
2502        let _ = interleave;
2503        false
2504    }
2505}
2506
2507/// [`gemm_q4_0x4_group`] against a pre-interleaved activation quad;
2508/// interleave-8 packing only, quad prepared once per matmul by
2509/// [`prepare_q8_acts_x4`]. `out` is `[row][act]`: `out[r * tile.na + a]`.
2510pub fn gemm_q4_0x4_group_x4(
2511    packed: &[u8],
2512    group: usize,
2513    tile: &Q8ActsX4,
2514    n_cols: usize,
2515    interleave: usize,
2516    out: &mut [f32],
2517) {
2518    gemm_q4_0x4_group_x4_on(
2519        packed,
2520        group,
2521        tile,
2522        n_cols,
2523        interleave,
2524        AccelX4::detect(),
2525        out,
2526    );
2527}
2528
2529/// [`gemm_q4_0x4_group_x4`] with the kernel choice already made; see
2530/// [`AccelX4`].
2531#[inline]
2532pub fn gemm_q4_0x4_group_x4_on(
2533    packed: &[u8],
2534    group: usize,
2535    tile: &Q8ActsX4,
2536    n_cols: usize,
2537    interleave: usize,
2538    accel: AccelX4,
2539    out: &mut [f32],
2540) {
2541    assert_eq!(
2542        interleave, 8,
2543        "the x4 GEMM only exists for interleave-8 packing"
2544    );
2545    assert_eq!(out.len(), Q4_0X4_NROWS * tile.na);
2546    assert!(n_cols.is_multiple_of(Q4_0_BLOCK_ELEMS));
2547    debug_assert_eq!(tile.n_blocks, n_cols / Q4_0_BLOCK_ELEMS);
2548    if tile.na == 0 {
2549        return;
2550    }
2551    let nb = n_cols / Q4_0_BLOCK_ELEMS;
2552    let off = group * nb * Q4_0X4_BLOCK_BYTES;
2553    let slice = &packed[off..off + nb * Q4_0X4_BLOCK_BYTES];
2554
2555    #[cfg(target_arch = "aarch64")]
2556    if accel == AccelX4::NeonI8mm {
2557        unsafe {
2558            neon::gemm_q4_0x4_q8_0_neon_i8mm(slice, tile, n_cols, out);
2559        }
2560        return;
2561    }
2562    let _ = accel;
2563    gemm_q4_0x4_acts_x4_scalar_8(slice, tile, n_cols, out);
2564}
2565
2566/// Portable reference for the Q4_0 ×4 GEMM: the same math as
2567/// [`gemv_q4_0x4_q8_0_scalar`] at blocklen 8, per quad row, reading qs and
2568/// d out of the pre-interleaved [`Q8ActsX4`]. Bit-identical to running
2569/// that GEMV per activation, which is what the tests assert.
2570fn gemm_q4_0x4_acts_x4_scalar_8(packed: &[u8], tile: &Q8ActsX4, n_cols: usize, out: &mut [f32]) {
2571    let nb = n_cols / Q4_0_BLOCK_ELEMS;
2572    let blocklen = 8;
2573    let ncols = Q4_0X4_NROWS;
2574    let na = tile.na;
2575    let mut sumf = [[0f32; Q4_0X4_NROWS]; Q8K_ACTS_X4_NC];
2576    for l in 0..nb {
2577        let blk = &packed[l * Q4_0X4_BLOCK_BYTES..][..Q4_0X4_BLOCK_BYTES];
2578        let q8 = &tile.qs[l * Q4_0_BLOCK_ELEMS * 4..][..Q4_0_BLOCK_ELEMS * 4];
2579        for a in 0..na {
2580            let da = tile.d[l * 4 + a];
2581            for k in 0..(Q4_0_BLOCK_ELEMS / (2 * blocklen)) {
2582                for j in 0..ncols {
2583                    let mut sumi = 0i32;
2584                    for i in 0..blocklen {
2585                        let byte = blk[8 + k * ncols * blocklen + j * blocklen + i];
2586                        // Canonical q8 element `e` lives at run `e/8`, row
2587                        // `a`, lane `e%8` of the interleaved block.
2588                        let e0 = k * blocklen + i;
2589                        let e1 = e0 + Q4_0_BLOCK_ELEMS / 2;
2590                        sumi += q4_0x4_nibble_dot(
2591                            byte,
2592                            q8[(e0 / 8) * 32 + a * 8 + (e0 % 8)] as i32,
2593                            q8[(e1 / 8) * 32 + a * 8 + (e1 % 8)] as i32,
2594                        );
2595                    }
2596                    sumf[a][j] += sumi as f32 * f16_from_bytes(&blk[j * 2..]) * da;
2597                }
2598            }
2599        }
2600    }
2601    for j in 0..ncols {
2602        for (a, row) in sumf.iter().take(na).enumerate() {
2603            out[j * na + a] = row[j];
2604        }
2605    }
2606}
2607
2608/// One row-group (4 outputs) starting at `group` within a packed Q4_0x4 matrix.
2609#[inline]
2610pub fn gemv_q4_0x4_group(
2611    packed: &[u8],
2612    group: usize,
2613    act: &Q8Activations,
2614    n_cols: usize,
2615    interleave: usize,
2616    out4: &mut [f32],
2617) {
2618    debug_assert_eq!(out4.len(), Q4_0X4_NROWS);
2619    let nb = n_cols / Q4_0_BLOCK_ELEMS;
2620    let off = group * nb * Q4_0X4_BLOCK_BYTES;
2621    let slice = &packed[off..off + nb * Q4_0X4_BLOCK_BYTES];
2622    gemv_q4_0x4_q8_0(slice, act, n_cols, 1, interleave, out4);
2623}
2624
2625/// One row-group (4 outputs) starting at `group` within a packed Q8_0x4 matrix.
2626#[inline]
2627pub fn gemv_q8_0x4_group(
2628    packed: &[u8],
2629    group: usize,
2630    act: &Q8Activations,
2631    n_cols: usize,
2632    interleave: usize,
2633    out4: &mut [f32],
2634) {
2635    debug_assert_eq!(out4.len(), Q8_0X4_NROWS);
2636    let nb = n_cols / Q8_0_BLOCK_ELEMS;
2637    let off = group * nb * Q8_0X4_BLOCK_BYTES;
2638    let slice = &packed[off..off + nb * Q8_0X4_BLOCK_BYTES];
2639    gemv_q8_0x4_q8_0(slice, act, n_cols, 1, interleave, out4);
2640}
2641
2642#[cfg(target_arch = "aarch64")]
2643mod neon {
2644    use super::*;
2645    use std::arch::aarch64::*;
2646
2647    #[target_feature(enable = "neon,i8mm")]
2648    unsafe fn vmmla_s32(mut acc: int32x4_t, a: int8x16_t, b: int8x16_t) -> int32x4_t {
2649        std::arch::asm!(
2650            "smmla {acc:v}.4s, {a:v}.16b, {b:v}.16b",
2651            acc = inout(vreg) acc,
2652            a = in(vreg) a,
2653            b = in(vreg) b,
2654            options(pure, nomem, nostack),
2655        );
2656        acc
2657    }
2658
2659    #[target_feature(enable = "neon,dotprod")]
2660    unsafe fn sdot_lane(mut acc: int32x4_t, a: int8x16_t, b: int8x16_t, lane: u32) -> int32x4_t {
2661        // sdot Vd.4S, Vn.16B, Vm.4B[lane]
2662        match lane {
2663            0 => std::arch::asm!(
2664                "sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[0]",
2665                acc = inout(vreg) acc,
2666                a = in(vreg) a,
2667                b = in(vreg) b,
2668                options(pure, nomem, nostack),
2669            ),
2670            1 => std::arch::asm!(
2671                "sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[1]",
2672                acc = inout(vreg) acc,
2673                a = in(vreg) a,
2674                b = in(vreg) b,
2675                options(pure, nomem, nostack),
2676            ),
2677            2 => std::arch::asm!(
2678                "sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[2]",
2679                acc = inout(vreg) acc,
2680                a = in(vreg) a,
2681                b = in(vreg) b,
2682                options(pure, nomem, nostack),
2683            ),
2684            3 => std::arch::asm!(
2685                "sdot {acc:v}.4s, {a:v}.16b, {b:v}.4b[3]",
2686                acc = inout(vreg) acc,
2687                a = in(vreg) a,
2688                b = in(vreg) b,
2689                options(pure, nomem, nostack),
2690            ),
2691            _ => unreachable!(),
2692        }
2693        acc
2694    }
2695
2696    /// `sdot Vd.4S, Vn.16B, Vm.16B` -- the plain (non-lane) signed dot.
2697    #[target_feature(enable = "neon,dotprod")]
2698    unsafe fn sdot(mut acc: int32x4_t, a: int8x16_t, b: int8x16_t) -> int32x4_t {
2699        std::arch::asm!(
2700            "sdot {acc:v}.4s, {a:v}.16b, {b:v}.16b",
2701            acc = inout(vreg) acc,
2702            a = in(vreg) a,
2703            b = in(vreg) b,
2704            options(pure, nomem, nostack),
2705        );
2706        acc
2707    }
2708
2709    /// NEON DotProd GEMV for interleave-4 packed weights (Apple Silicon path).
2710    #[target_feature(enable = "neon,dotprod")]
2711    pub unsafe fn gemv_q4_kx8_q8_k_neon_sdot(
2712        packed: &[u8],
2713        act: &Q8KActivations,
2714        n_cols: usize,
2715        n_row_groups: usize,
2716        out: &mut [f32],
2717    ) {
2718        let nb = n_cols / Q4_K_BLOCK_ELEMS;
2719        let m4b = vdupq_n_u8(0x0f);
2720
2721        for x in 0..n_row_groups {
2722            let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
2723            let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
2724
2725            for b in 0..nb {
2726                let blk = packed.as_ptr().add(group_off + b * Q4_KX8_BLOCK_BYTES);
2727                let mut d_arr = [0f32; 8];
2728                let mut dmin_arr = [0f32; 8];
2729                for j in 0..8 {
2730                    d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
2731                    dmin_arr[j] =
2732                        f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
2733                }
2734                let q8_d = act.d[b];
2735                let sb_scale_0123 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d);
2736                let sb_scale_4567 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d);
2737                let sb_min_0123 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d);
2738                let sb_min_4567 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d);
2739
2740                let mut bias_acc = [vdupq_n_s32(0), vdupq_n_s32(0)];
2741                let q8_base = act.q.as_ptr().add(b * Q4_K_BLOCK_ELEMS);
2742                let bsums_ptr = act.bsums.as_ptr().add(b * 16);
2743                // Pairwise-add 16 bsums → 8 (matching llama vpaddq_s16).
2744                let mut bsums_arr = [0i16; 8];
2745                for (i, slot) in bsums_arr.iter_mut().enumerate() {
2746                    *slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
2747                }
2748
2749                let scales_base = blk.add(32);
2750                let qs_base = blk.add(128);
2751
2752                for sb in 0..4 {
2753                    let mut acc_lo = [vdupq_n_s32(0), vdupq_n_s32(0)];
2754                    let mut acc_hi = [vdupq_n_s32(0), vdupq_n_s32(0)];
2755
2756                    let mut q4sb_mins = [vdupq_n_s16(0); 2];
2757                    let mut q4sb_scales = [vdupq_n_s16(0); 2];
2758                    for i in 0..2 {
2759                        let mut sc = [0u8; 8];
2760                        let mut mn = [0u8; 8];
2761                        let offset = sb * 24 + i * 12;
2762                        decode_scales_mins(
2763                            std::slice::from_raw_parts(scales_base.add(offset), 12),
2764                            &mut sc,
2765                            &mut mn,
2766                        );
2767                        let mut sc_i8 = [0i8; 8];
2768                        let mut mn_i8 = [0i8; 8];
2769                        for t in 0..8 {
2770                            sc_i8[t] = sc[t] as i8;
2771                            mn_i8[t] = mn[t] as i8;
2772                        }
2773                        q4sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
2774                        q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
2775                    }
2776
2777                    let mut q8_qs = [vdupq_n_s8(0); 4];
2778                    for (i, slot) in q8_qs.iter_mut().enumerate() {
2779                        *slot = vld1q_s8(q8_base.add(sb * 64 + i * 16));
2780                    }
2781
2782                    for c in 0..2 {
2783                        let mut q4_cols = [vdupq_n_u8(0); 8];
2784                        for (i, slot) in q4_cols.iter_mut().enumerate() {
2785                            *slot = vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + i * 32 + 16 * c));
2786                        }
2787
2788                        acc_lo[c] = sdot_lane(
2789                            acc_lo[c],
2790                            vreinterpretq_s8_u8(vandq_u8(q4_cols[0], m4b)),
2791                            q8_qs[0],
2792                            0,
2793                        );
2794                        acc_lo[c] = sdot_lane(
2795                            acc_lo[c],
2796                            vreinterpretq_s8_u8(vandq_u8(q4_cols[1], m4b)),
2797                            q8_qs[0],
2798                            1,
2799                        );
2800                        acc_lo[c] = sdot_lane(
2801                            acc_lo[c],
2802                            vreinterpretq_s8_u8(vandq_u8(q4_cols[2], m4b)),
2803                            q8_qs[0],
2804                            2,
2805                        );
2806                        acc_lo[c] = sdot_lane(
2807                            acc_lo[c],
2808                            vreinterpretq_s8_u8(vandq_u8(q4_cols[3], m4b)),
2809                            q8_qs[0],
2810                            3,
2811                        );
2812                        acc_lo[c] = sdot_lane(
2813                            acc_lo[c],
2814                            vreinterpretq_s8_u8(vandq_u8(q4_cols[4], m4b)),
2815                            q8_qs[1],
2816                            0,
2817                        );
2818                        acc_lo[c] = sdot_lane(
2819                            acc_lo[c],
2820                            vreinterpretq_s8_u8(vandq_u8(q4_cols[5], m4b)),
2821                            q8_qs[1],
2822                            1,
2823                        );
2824                        acc_lo[c] = sdot_lane(
2825                            acc_lo[c],
2826                            vreinterpretq_s8_u8(vandq_u8(q4_cols[6], m4b)),
2827                            q8_qs[1],
2828                            2,
2829                        );
2830                        acc_lo[c] = sdot_lane(
2831                            acc_lo[c],
2832                            vreinterpretq_s8_u8(vandq_u8(q4_cols[7], m4b)),
2833                            q8_qs[1],
2834                            3,
2835                        );
2836
2837                        acc_hi[c] = sdot_lane(
2838                            acc_hi[c],
2839                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[0], 4)),
2840                            q8_qs[2],
2841                            0,
2842                        );
2843                        acc_hi[c] = sdot_lane(
2844                            acc_hi[c],
2845                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[1], 4)),
2846                            q8_qs[2],
2847                            1,
2848                        );
2849                        acc_hi[c] = sdot_lane(
2850                            acc_hi[c],
2851                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[2], 4)),
2852                            q8_qs[2],
2853                            2,
2854                        );
2855                        acc_hi[c] = sdot_lane(
2856                            acc_hi[c],
2857                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[3], 4)),
2858                            q8_qs[2],
2859                            3,
2860                        );
2861                        acc_hi[c] = sdot_lane(
2862                            acc_hi[c],
2863                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[4], 4)),
2864                            q8_qs[3],
2865                            0,
2866                        );
2867                        acc_hi[c] = sdot_lane(
2868                            acc_hi[c],
2869                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[5], 4)),
2870                            q8_qs[3],
2871                            1,
2872                        );
2873                        acc_hi[c] = sdot_lane(
2874                            acc_hi[c],
2875                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[6], 4)),
2876                            q8_qs[3],
2877                            2,
2878                        );
2879                        acc_hi[c] = sdot_lane(
2880                            acc_hi[c],
2881                            vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[7], 4)),
2882                            q8_qs[3],
2883                            3,
2884                        );
2885                    }
2886
2887                    let sc_0123_lo = vget_low_s16(q4sb_scales[0]);
2888                    let sc_0123_hi = vget_low_s16(q4sb_scales[1]);
2889                    let sumf_0123 = vcvtq_f32_s32(vaddq_s32(
2890                        vmulq_s32(vmovl_s16(sc_0123_lo), acc_lo[0]),
2891                        vmulq_s32(vmovl_s16(sc_0123_hi), acc_hi[0]),
2892                    ));
2893                    acc_f32[0] = vfmaq_f32(acc_f32[0], sb_scale_0123, sumf_0123);
2894
2895                    let sc_4567_lo = vget_high_s16(q4sb_scales[0]);
2896                    let sc_4567_hi = vget_high_s16(q4sb_scales[1]);
2897                    let sumf_4567 = vcvtq_f32_s32(vaddq_s32(
2898                        vmulq_s32(vmovl_s16(sc_4567_lo), acc_lo[1]),
2899                        vmulq_s32(vmovl_s16(sc_4567_hi), acc_hi[1]),
2900                    ));
2901                    acc_f32[1] = vfmaq_f32(acc_f32[1], sb_scale_4567, sumf_4567);
2902
2903                    let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
2904                    let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
2905                    bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_lo, vget_low_s16(q4sb_mins[0]));
2906                    bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_hi, vget_low_s16(q4sb_mins[1]));
2907                    bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_lo, vget_high_s16(q4sb_mins[0]));
2908                    bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_hi, vget_high_s16(q4sb_mins[1]));
2909                }
2910
2911                acc_f32[0] = vmlsq_f32(acc_f32[0], vcvtq_f32_s32(bias_acc[0]), sb_min_0123);
2912                acc_f32[1] = vmlsq_f32(acc_f32[1], vcvtq_f32_s32(bias_acc[1]), sb_min_4567);
2913            }
2914
2915            let base = x * Q4_KX8_NROWS;
2916            vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
2917            vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
2918        }
2919    }
2920
2921    /// NEON DotProd GEMV for interleave-4 `block_q5_Kx8` (llama
2922    /// `ggml_gemv_q5_K_8x4_q8_K`).
2923    #[target_feature(enable = "neon,dotprod")]
2924    pub unsafe fn gemv_q5_kx8_q8_k_neon_sdot(
2925        packed: &[u8],
2926        act: &Q8KActivations,
2927        n_cols: usize,
2928        n_row_groups: usize,
2929        out: &mut [f32],
2930    ) {
2931        let nb = n_cols / Q5_K_BLOCK_ELEMS;
2932        let m4b = vdupq_n_u8(0x0f);
2933        let mone = vdupq_n_u8(1);
2934        let mtwo = vdupq_n_u8(2);
2935
2936        for x in 0..n_row_groups {
2937            let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
2938            let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
2939
2940            for b in 0..nb {
2941                let blk = packed.as_ptr().add(group_off + b * Q5_KX8_BLOCK_BYTES);
2942                let mut d_arr = [0f32; 8];
2943                let mut dmin_arr = [0f32; 8];
2944                for j in 0..8 {
2945                    d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
2946                    dmin_arr[j] =
2947                        f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
2948                }
2949                let q8_d = act.d[b];
2950                let sb_scale_0123 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d);
2951                let sb_scale_4567 = vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d);
2952                let sb_min_0123 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d);
2953                let sb_min_4567 = vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d);
2954
2955                let mut bias_acc = [vdupq_n_s32(0), vdupq_n_s32(0)];
2956                let q8_base = act.q.as_ptr().add(b * Q5_K_BLOCK_ELEMS);
2957                let bsums_ptr = act.bsums.as_ptr().add(b * 16);
2958                let mut bsums_arr = [0i16; 8];
2959                for (i, slot) in bsums_arr.iter_mut().enumerate() {
2960                    *slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
2961                }
2962
2963                let scales_base = blk.add(32);
2964                let qh_base = blk.add(128);
2965                let qs_base = blk.add(384);
2966
2967                // qh[c][i]: 2 col-groups × 8 vectors; shift down 2 bits each sb.
2968                let mut qh = [[vdupq_n_u8(0); 8]; 2];
2969                for (c, qh_c) in qh.iter_mut().enumerate() {
2970                    for (i, slot) in qh_c.iter_mut().enumerate() {
2971                        *slot = vld1q_u8(qh_base.add(i * 32 + 16 * c));
2972                    }
2973                }
2974
2975                for sb in 0..4 {
2976                    let mut acc_lo = [vdupq_n_s32(0), vdupq_n_s32(0)];
2977                    let mut acc_hi = [vdupq_n_s32(0), vdupq_n_s32(0)];
2978
2979                    let mut q5sb_mins = [vdupq_n_s16(0); 2];
2980                    let mut q5sb_scales = [vdupq_n_s16(0); 2];
2981                    for i in 0..2 {
2982                        let mut sc = [0u8; 8];
2983                        let mut mn = [0u8; 8];
2984                        let offset = sb * 24 + i * 12;
2985                        decode_scales_mins(
2986                            std::slice::from_raw_parts(scales_base.add(offset), 12),
2987                            &mut sc,
2988                            &mut mn,
2989                        );
2990                        let mut sc_i8 = [0i8; 8];
2991                        let mut mn_i8 = [0i8; 8];
2992                        for t in 0..8 {
2993                            sc_i8[t] = sc[t] as i8;
2994                            mn_i8[t] = mn[t] as i8;
2995                        }
2996                        q5sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
2997                        q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
2998                    }
2999
3000                    let mut q8_qs = [vdupq_n_s8(0); 4];
3001                    for (i, slot) in q8_qs.iter_mut().enumerate() {
3002                        *slot = vld1q_s8(q8_base.add(sb * 64 + i * 16));
3003                    }
3004
3005                    for c in 0..2 {
3006                        let mut q5_lo = [vdupq_n_s8(0); 8];
3007                        let mut q5_hi = [vdupq_n_s8(0); 8];
3008                        for i in 0..8 {
3009                            let q5_cols =
3010                                vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + i * 32 + 16 * c));
3011                            let hbit_lo = vandq_u8(qh[c][i], mone);
3012                            let hbit_hi = vshlq_n_u8(vandq_u8(qh[c][i], mtwo), 3);
3013                            qh[c][i] = vshrq_n_u8(qh[c][i], 2);
3014                            q5_lo[i] =
3015                                vreinterpretq_s8_u8(vsliq_n_u8(vandq_u8(q5_cols, m4b), hbit_lo, 4));
3016                            q5_hi[i] =
3017                                vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_cols, 4), hbit_hi));
3018                        }
3019                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[0], q8_qs[0], 0);
3020                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[1], q8_qs[0], 1);
3021                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[2], q8_qs[0], 2);
3022                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[3], q8_qs[0], 3);
3023                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[4], q8_qs[1], 0);
3024                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[5], q8_qs[1], 1);
3025                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[6], q8_qs[1], 2);
3026                        acc_lo[c] = sdot_lane(acc_lo[c], q5_lo[7], q8_qs[1], 3);
3027                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[0], q8_qs[2], 0);
3028                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[1], q8_qs[2], 1);
3029                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[2], q8_qs[2], 2);
3030                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[3], q8_qs[2], 3);
3031                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[4], q8_qs[3], 0);
3032                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[5], q8_qs[3], 1);
3033                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[6], q8_qs[3], 2);
3034                        acc_hi[c] = sdot_lane(acc_hi[c], q5_hi[7], q8_qs[3], 3);
3035                    }
3036
3037                    let sc_0123_lo = vget_low_s16(q5sb_scales[0]);
3038                    let sc_0123_hi = vget_low_s16(q5sb_scales[1]);
3039                    let sumf_0123 = vcvtq_f32_s32(vaddq_s32(
3040                        vmulq_s32(vmovl_s16(sc_0123_lo), acc_lo[0]),
3041                        vmulq_s32(vmovl_s16(sc_0123_hi), acc_hi[0]),
3042                    ));
3043                    acc_f32[0] = vfmaq_f32(acc_f32[0], sb_scale_0123, sumf_0123);
3044
3045                    let sc_4567_lo = vget_high_s16(q5sb_scales[0]);
3046                    let sc_4567_hi = vget_high_s16(q5sb_scales[1]);
3047                    let sumf_4567 = vcvtq_f32_s32(vaddq_s32(
3048                        vmulq_s32(vmovl_s16(sc_4567_lo), acc_lo[1]),
3049                        vmulq_s32(vmovl_s16(sc_4567_hi), acc_hi[1]),
3050                    ));
3051                    acc_f32[1] = vfmaq_f32(acc_f32[1], sb_scale_4567, sumf_4567);
3052
3053                    let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
3054                    let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
3055                    bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_lo, vget_low_s16(q5sb_mins[0]));
3056                    bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_hi, vget_low_s16(q5sb_mins[1]));
3057                    bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_lo, vget_high_s16(q5sb_mins[0]));
3058                    bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_hi, vget_high_s16(q5sb_mins[1]));
3059                }
3060
3061                acc_f32[0] = vmlsq_f32(acc_f32[0], vcvtq_f32_s32(bias_acc[0]), sb_min_0123);
3062                acc_f32[1] = vmlsq_f32(acc_f32[1], vcvtq_f32_s32(bias_acc[1]), sb_min_4567);
3063            }
3064
3065            let base = x * Q5_KX8_NROWS;
3066            vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
3067            vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
3068        }
3069    }
3070
3071    /// NEON DotProd GEMV for interleave-8 packed Q4_K weights (llama.cpp
3072    /// `ggml_gemv_q4_K_8x8_q8_K` in `arch/arm/repack.cpp`). Each 8-byte q8
3073    /// run is broadcast to both vector halves so one `sdot` covers two
3074    /// interleaved columns at once.
3075    #[target_feature(enable = "neon,dotprod")]
3076    pub unsafe fn gemv_q4_kx8_q8_k_neon_8x8(
3077        packed: &[u8],
3078        act: &Q8KActivations,
3079        n_cols: usize,
3080        n_row_groups: usize,
3081        out: &mut [f32],
3082    ) {
3083        let nb = n_cols / Q4_K_BLOCK_ELEMS;
3084        let m4b = vdupq_n_u8(0x0f);
3085
3086        for x in 0..n_row_groups {
3087            let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
3088            let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
3089
3090            for b in 0..nb {
3091                let blk = packed.as_ptr().add(group_off + b * Q4_KX8_BLOCK_BYTES);
3092                let mut d_arr = [0f32; 8];
3093                let mut dmin_arr = [0f32; 8];
3094                for j in 0..8 {
3095                    d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
3096                    dmin_arr[j] =
3097                        f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
3098                }
3099                let q8_d = act.d[b];
3100                let sb_scale = [
3101                    vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d),
3102                    vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d),
3103                ];
3104                let sb_min = [
3105                    vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d),
3106                    vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d),
3107                ];
3108
3109                let q8_base = act.q.as_ptr().add(b * Q4_K_BLOCK_ELEMS);
3110                let bsums_ptr = act.bsums.as_ptr().add(b * 16);
3111                let mut bsums_arr = [0i16; 8];
3112                for (i, slot) in bsums_arr.iter_mut().enumerate() {
3113                    *slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
3114                }
3115
3116                let scales_base = blk.add(32);
3117                let qs_base = blk.add(128);
3118
3119                let mut bias_acc = [vdupq_n_s32(0), vdupq_n_s32(0)];
3120
3121                for sb in 0..4 {
3122                    let mut acc_lo = [vdupq_n_s32(0); 4];
3123                    let mut acc_hi = [vdupq_n_s32(0); 4];
3124
3125                    let mut q4sb_scales = [vdupq_n_s16(0); 2];
3126                    let mut q4sb_mins = [vdupq_n_s16(0); 2];
3127                    for i in 0..2 {
3128                        let mut sc = [0u8; 8];
3129                        let mut mn = [0u8; 8];
3130                        let offset = sb * 24 + i * 12;
3131                        decode_scales_mins(
3132                            std::slice::from_raw_parts(scales_base.add(offset), 12),
3133                            &mut sc,
3134                            &mut mn,
3135                        );
3136                        let mut sc_i8 = [0i8; 8];
3137                        let mut mn_i8 = [0i8; 8];
3138                        for t in 0..8 {
3139                            sc_i8[t] = sc[t] as i8;
3140                            mn_i8[t] = mn[t] as i8;
3141                        }
3142                        q4sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
3143                        q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
3144                    }
3145
3146                    let q8_sb = q8_base.add(sb * 64);
3147                    let mut q8_qs = [vdupq_n_s8(0); 8];
3148                    for (i, slot) in q8_qs.iter_mut().enumerate() {
3149                        *slot = vreinterpretq_s8_s64(vld1q_dup_s64(q8_sb.add(i * 8) as *const i64));
3150                    }
3151
3152                    for cp in 0..4 {
3153                        let q4_qs = [
3154                            vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp)),
3155                            vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 64)),
3156                            vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 128)),
3157                            vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 192)),
3158                        ];
3159                        for m in 0..4 {
3160                            let q4_lo = vreinterpretq_s8_u8(vandq_u8(q4_qs[m], m4b));
3161                            let q4_hi = vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[m], 4));
3162                            acc_lo[cp] = sdot(acc_lo[cp], q4_lo, q8_qs[m]);
3163                            acc_hi[cp] = sdot(acc_hi[cp], q4_hi, q8_qs[m + 4]);
3164                        }
3165                    }
3166
3167                    for i in 0..2 {
3168                        let p = i * 2;
3169                        let (scales_lo, scales_hi) = if i == 0 {
3170                            (vget_low_s16(q4sb_scales[0]), vget_low_s16(q4sb_scales[1]))
3171                        } else {
3172                            (vget_high_s16(q4sb_scales[0]), vget_high_s16(q4sb_scales[1]))
3173                        };
3174                        let sumf_0 = vcvtq_f32_s32(vmulq_s32(
3175                            vmovl_s16(scales_lo),
3176                            vpaddq_s32(acc_lo[p], acc_lo[p + 1]),
3177                        ));
3178                        acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_0);
3179                        let sumf_1 = vcvtq_f32_s32(vmulq_s32(
3180                            vmovl_s16(scales_hi),
3181                            vpaddq_s32(acc_hi[p], acc_hi[p + 1]),
3182                        ));
3183                        acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_1);
3184                    }
3185
3186                    let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
3187                    let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
3188                    bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_lo, vget_low_s16(q4sb_mins[0]));
3189                    bias_acc[0] = vmlal_s16(bias_acc[0], bsums_vec_hi, vget_low_s16(q4sb_mins[1]));
3190                    bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_lo, vget_high_s16(q4sb_mins[0]));
3191                    bias_acc[1] = vmlal_s16(bias_acc[1], bsums_vec_hi, vget_high_s16(q4sb_mins[1]));
3192                }
3193
3194                acc_f32[0] = vmlsq_f32(acc_f32[0], vcvtq_f32_s32(bias_acc[0]), sb_min[0]);
3195                acc_f32[1] = vmlsq_f32(acc_f32[1], vcvtq_f32_s32(bias_acc[1]), sb_min[1]);
3196            }
3197
3198            let base = x * Q4_KX8_NROWS;
3199            vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
3200            vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
3201        }
3202    }
3203
3204    /// NEON DotProd GEMV for interleave-8 packed Q5_K weights (llama.cpp
3205    /// `ggml_gemv_q5_K_8x8_q8_K` in `arch/arm/repack.cpp`). Each 8-byte q8
3206    /// run is broadcast to both vector halves so one `sdot` covers two
3207    /// interleaved columns at once.
3208    #[target_feature(enable = "neon,dotprod")]
3209    pub unsafe fn gemv_q5_kx8_q8_k_neon_8x8(
3210        packed: &[u8],
3211        act: &Q8KActivations,
3212        n_cols: usize,
3213        n_row_groups: usize,
3214        out: &mut [f32],
3215    ) {
3216        let nb = n_cols / Q5_K_BLOCK_ELEMS;
3217        let m4b = vdupq_n_u8(0x0f);
3218        let mone = vdupq_n_u8(1);
3219        let mtwo = vdupq_n_u8(2);
3220
3221        for x in 0..n_row_groups {
3222            let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
3223            let group_off = x * nb * Q5_KX8_BLOCK_BYTES;
3224
3225            for b in 0..nb {
3226                let blk = packed.as_ptr().add(group_off + b * Q5_KX8_BLOCK_BYTES);
3227                let mut d_arr = [0f32; 8];
3228                let mut dmin_arr = [0f32; 8];
3229                for j in 0..8 {
3230                    d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
3231                    dmin_arr[j] =
3232                        f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
3233                }
3234                let q8_d = act.d[b];
3235                let sb_scale = [
3236                    vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d),
3237                    vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d),
3238                ];
3239                let sb_min = [
3240                    vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr()), q8_d),
3241                    vmulq_n_f32(vld1q_f32(dmin_arr.as_ptr().add(4)), q8_d),
3242                ];
3243
3244                let q8_base = act.q.as_ptr().add(b * Q5_K_BLOCK_ELEMS);
3245                let bsums_ptr = act.bsums.as_ptr().add(b * 16);
3246                let mut bsums_arr = [0i16; 8];
3247                for (i, slot) in bsums_arr.iter_mut().enumerate() {
3248                    *slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
3249                }
3250
3251                let scales_base = blk.add(32);
3252                let qh_base = blk.add(128);
3253                let qs_base = blk.add(384);
3254
3255                // qh state per column pair; two bits consumed per sub-block.
3256                let mut qh = [[vdupq_n_u8(0); 4]; 4];
3257                for (cp, qh_cp) in qh.iter_mut().enumerate() {
3258                    for (m, slot) in qh_cp.iter_mut().enumerate() {
3259                        *slot = vld1q_u8(qh_base.add(16 * cp + 64 * m));
3260                    }
3261                }
3262
3263                for sb in 0..4 {
3264                    let mut acc_lo = [vdupq_n_s32(0); 4];
3265                    let mut acc_hi = [vdupq_n_s32(0); 4];
3266
3267                    let mut q5sb_scales = [vdupq_n_s16(0); 2];
3268                    let mut q5sb_mins = [vdupq_n_s16(0); 2];
3269                    for i in 0..2 {
3270                        let mut sc = [0u8; 8];
3271                        let mut mn = [0u8; 8];
3272                        let offset = sb * 24 + i * 12;
3273                        decode_scales_mins(
3274                            std::slice::from_raw_parts(scales_base.add(offset), 12),
3275                            &mut sc,
3276                            &mut mn,
3277                        );
3278                        let mut sc_i8 = [0i8; 8];
3279                        let mut mn_i8 = [0i8; 8];
3280                        for t in 0..8 {
3281                            sc_i8[t] = sc[t] as i8;
3282                            mn_i8[t] = mn[t] as i8;
3283                        }
3284                        q5sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
3285                        q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
3286                    }
3287
3288                    let q8_sb = q8_base.add(sb * 64);
3289                    let mut q8_qs = [vdupq_n_s8(0); 8];
3290                    for (i, slot) in q8_qs.iter_mut().enumerate() {
3291                        *slot = vreinterpretq_s8_s64(vld1q_dup_s64(q8_sb.add(i * 8) as *const i64));
3292                    }
3293
3294                    for cp in 0..4 {
3295                        let q5_qs = [
3296                            vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp)),
3297                            vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 64)),
3298                            vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 128)),
3299                            vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 192)),
3300                        ];
3301                        for m in 0..4 {
3302                            let hbit_lo = vandq_u8(qh[cp][m], mone);
3303                            let hbit_hi = vshlq_n_u8(vandq_u8(qh[cp][m], mtwo), 3);
3304                            qh[cp][m] = vshrq_n_u8(qh[cp][m], 2);
3305                            let q5_lo = vreinterpretq_s8_u8(vsliq_n_u8(
3306                                vandq_u8(q5_qs[m], m4b),
3307                                hbit_lo,
3308                                4,
3309                            ));
3310                            let q5_hi =
3311                                vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_qs[m], 4), hbit_hi));
3312                            acc_lo[cp] = sdot(acc_lo[cp], q5_lo, q8_qs[m]);
3313                            acc_hi[cp] = sdot(acc_hi[cp], q5_hi, q8_qs[m + 4]);
3314                        }
3315                    }
3316
3317                    let bsums_vec_lo = vdup_n_s16(bsums_arr[2 * sb]);
3318                    let bsums_vec_hi = vdup_n_s16(bsums_arr[2 * sb + 1]);
3319                    for i in 0..2 {
3320                        let p = i * 2;
3321                        let (scales_lo, scales_hi, mins_lo, mins_hi) = if i == 0 {
3322                            (
3323                                vget_low_s16(q5sb_scales[0]),
3324                                vget_low_s16(q5sb_scales[1]),
3325                                vget_low_s16(q5sb_mins[0]),
3326                                vget_low_s16(q5sb_mins[1]),
3327                            )
3328                        } else {
3329                            (
3330                                vget_high_s16(q5sb_scales[0]),
3331                                vget_high_s16(q5sb_scales[1]),
3332                                vget_high_s16(q5sb_mins[0]),
3333                                vget_high_s16(q5sb_mins[1]),
3334                            )
3335                        };
3336                        let sumf_0 = vcvtq_f32_s32(vmulq_s32(
3337                            vmovl_s16(scales_lo),
3338                            vpaddq_s32(acc_lo[p], acc_lo[p + 1]),
3339                        ));
3340                        acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_0);
3341                        let sumf_1 = vcvtq_f32_s32(vmulq_s32(
3342                            vmovl_s16(scales_hi),
3343                            vpaddq_s32(acc_hi[p], acc_hi[p + 1]),
3344                        ));
3345                        acc_f32[i] = vfmaq_f32(acc_f32[i], sb_scale[i], sumf_1);
3346
3347                        let mut bias = vmull_s16(bsums_vec_lo, mins_lo);
3348                        bias = vmlal_s16(bias, bsums_vec_hi, mins_hi);
3349                        acc_f32[i] = vmlsq_f32(acc_f32[i], sb_min[i], vcvtq_f32_s32(bias));
3350                    }
3351                }
3352            }
3353
3354            let base = x * Q5_KX8_NROWS;
3355            vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
3356            vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
3357        }
3358    }
3359
3360    /// NEON DotProd GEMV for interleave-8 packed Q6_K weights (llama.cpp
3361    /// `ggml_gemv_q6_K_8x8_q8_K` in `arch/arm/repack.cpp`). The -32 offset
3362    /// is folded into a bsums × scales bias (shifted left 5) instead of
3363    /// being subtracted from every value.
3364    #[target_feature(enable = "neon,dotprod")]
3365    pub unsafe fn gemv_q6_kx8_q8_k_neon_8x8(
3366        packed: &[u8],
3367        act: &Q8KActivations,
3368        n_cols: usize,
3369        n_row_groups: usize,
3370        out: &mut [f32],
3371    ) {
3372        let nb = n_cols / Q6_K_BLOCK_ELEMS;
3373        let m4b = vdupq_n_u8(0x0f);
3374        let mask_lo = vdupq_n_u8(0x03);
3375        let mask_hi = vdupq_n_u8(0x30);
3376
3377        for x in 0..n_row_groups {
3378            let mut acc_f32 = [vdupq_n_f32(0.0), vdupq_n_f32(0.0)];
3379            let group_off = x * nb * Q6_KX8_BLOCK_BYTES;
3380
3381            for b in 0..nb {
3382                let blk = packed.as_ptr().add(group_off + b * Q6_KX8_BLOCK_BYTES);
3383                let scales_base = blk.add(16) as *const i8;
3384                let ql_blk = blk.add(144);
3385                let qh_blk = blk.add(1168);
3386
3387                let mut d_arr = [0f32; 8];
3388                for (j, slot) in d_arr.iter_mut().enumerate() {
3389                    *slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
3390                }
3391                let q8_d = act.d[b];
3392                let sb_scale = [
3393                    vmulq_n_f32(vld1q_f32(d_arr.as_ptr()), q8_d),
3394                    vmulq_n_f32(vld1q_f32(d_arr.as_ptr().add(4)), q8_d),
3395                ];
3396
3397                let mut acc = [vdup_n_s32(0); 4];
3398
3399                // 16 groups of 8 i8 scales, widened once per block.
3400                let mut q6_scales = [0i16; 16 * 8];
3401                for i in 0..16 {
3402                    let s16 = vmovl_s8(vld1_s8(scales_base.add(i * 8)));
3403                    vst1q_s16(q6_scales.as_mut_ptr().add(i * 8), s16);
3404                }
3405
3406                // Bias: bsums × scales × 32 replaces subtracting 32 from
3407                // every 6-bit value.
3408                let mut bias_lo = vdupq_n_s32(0);
3409                let mut bias_hi = vdupq_n_s32(0);
3410                for i in (0..16).step_by(4) {
3411                    let bsums_vec = vld1_s16(act.bsums.as_ptr().add(b * 16 + i));
3412                    let sc = q6_scales.as_ptr();
3413                    bias_lo = vmlal_lane_s16::<0>(bias_lo, vld1_s16(sc.add(i * 8)), bsums_vec);
3414                    bias_hi = vmlal_lane_s16::<0>(bias_hi, vld1_s16(sc.add(i * 8 + 4)), bsums_vec);
3415                    bias_lo =
3416                        vmlal_lane_s16::<1>(bias_lo, vld1_s16(sc.add((i + 1) * 8)), bsums_vec);
3417                    bias_hi =
3418                        vmlal_lane_s16::<1>(bias_hi, vld1_s16(sc.add((i + 1) * 8 + 4)), bsums_vec);
3419                    bias_lo =
3420                        vmlal_lane_s16::<2>(bias_lo, vld1_s16(sc.add((i + 2) * 8)), bsums_vec);
3421                    bias_hi =
3422                        vmlal_lane_s16::<2>(bias_hi, vld1_s16(sc.add((i + 2) * 8 + 4)), bsums_vec);
3423                    bias_lo =
3424                        vmlal_lane_s16::<3>(bias_lo, vld1_s16(sc.add((i + 3) * 8)), bsums_vec);
3425                    bias_hi =
3426                        vmlal_lane_s16::<3>(bias_hi, vld1_s16(sc.add((i + 3) * 8 + 4)), bsums_vec);
3427                }
3428                bias_lo = vshlq_n_s32(bias_lo, 5);
3429                bias_hi = vshlq_n_s32(bias_hi, 5);
3430
3431                for half in 0..2 {
3432                    let ql_base = ql_blk.add(half * 512);
3433                    let qh_base = qh_blk.add(half * 256);
3434                    let q8_half = act.q.as_ptr().add(b * Q6_K_BLOCK_ELEMS + half * 128);
3435
3436                    for sb in 0..4 {
3437                        let q8_base_l = q8_half.add(sb * 16);
3438                        let q8_base_h = q8_base_l.add(64);
3439                        let mut q8_l = [vdupq_n_s8(0); 2];
3440                        let mut q8_h = [vdupq_n_s8(0); 2];
3441                        for i in 0..2 {
3442                            q8_l[i] = vreinterpretq_s8_s64(vld1q_dup_s64(
3443                                q8_base_l.add(i * 8) as *const i64
3444                            ));
3445                            q8_h[i] = vreinterpretq_s8_s64(vld1q_dup_s64(
3446                                q8_base_h.add(i * 8) as *const i64
3447                            ));
3448                        }
3449
3450                        let ql_off = sb * (Q6_K_BLOCK_ELEMS / 2);
3451                        let qh_off = ql_off & 255; // wraps after 256 bytes
3452                        let mut q6_ql_0 = [vdupq_n_u8(0); 4];
3453                        let mut q6_ql_1 = [vdupq_n_u8(0); 4];
3454                        let mut q6_qh_0 = [vdupq_n_u8(0); 4];
3455                        let mut q6_qh_1 = [vdupq_n_u8(0); 4];
3456                        for k in 0..4 {
3457                            q6_ql_0[k] = vld1q_u8(ql_base.add(ql_off + 16 * k));
3458                            q6_ql_1[k] = vld1q_u8(ql_base.add(ql_off + 64 + 16 * k));
3459                            q6_qh_0[k] = vld1q_u8(qh_base.add(qh_off + 16 * k));
3460                            q6_qh_1[k] = vld1q_u8(qh_base.add(qh_off + 64 + 16 * k));
3461                        }
3462                        // High bits for sub-blocks 2 and 3 sit two bits up.
3463                        if sb > 1 {
3464                            for k in 0..4 {
3465                                q6_qh_0[k] = vshrq_n_u8(q6_qh_0[k], 2);
3466                                q6_qh_1[k] = vshrq_n_u8(q6_qh_1[k], 2);
3467                            }
3468                        }
3469
3470                        for cp in 0..4 {
3471                            let hh_0 = vandq_u8(q6_qh_0[cp], mask_hi);
3472                            let hh_1 = vandq_u8(q6_qh_1[cp], mask_hi);
3473
3474                            // q6 = low4 | high2<<4; no -32 here, the bias
3475                            // pass above already carries it.
3476                            let q6_l0 = vreinterpretq_s8_u8(vsliq_n_u8(
3477                                vandq_u8(q6_ql_0[cp], m4b),
3478                                vandq_u8(q6_qh_0[cp], mask_lo),
3479                                4,
3480                            ));
3481                            let q6_l1 = vreinterpretq_s8_u8(vsliq_n_u8(
3482                                vandq_u8(q6_ql_1[cp], m4b),
3483                                vandq_u8(q6_qh_1[cp], mask_lo),
3484                                4,
3485                            ));
3486                            let q6_h0 =
3487                                vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_0[cp], 4), hh_0));
3488                            let q6_h1 =
3489                                vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_1[cp], 4), hh_1));
3490
3491                            let mut sb_acc_l = vdupq_n_s32(0);
3492                            sb_acc_l = sdot(sb_acc_l, q6_l0, q8_l[0]);
3493                            sb_acc_l = sdot(sb_acc_l, q6_l1, q8_l[1]);
3494                            let mut sb_acc_h = vdupq_n_s32(0);
3495                            sb_acc_h = sdot(sb_acc_h, q6_h0, q8_h[0]);
3496                            sb_acc_h = sdot(sb_acc_h, q6_h1, q8_h[1]);
3497
3498                            let sum_l = vpadd_s32(vget_low_s32(sb_acc_l), vget_high_s32(sb_acc_l));
3499                            let sum_h = vpadd_s32(vget_low_s32(sb_acc_h), vget_high_s32(sb_acc_h));
3500
3501                            let scale_idx_l = half * 8 + sb;
3502                            let scale_idx_h = half * 8 + sb + 4;
3503                            let scale_vec_l = vset_lane_s32::<1>(
3504                                i32::from(q6_scales[scale_idx_l * 8 + cp * 2 + 1]),
3505                                vdup_n_s32(i32::from(q6_scales[scale_idx_l * 8 + cp * 2])),
3506                            );
3507                            let scale_vec_h = vset_lane_s32::<1>(
3508                                i32::from(q6_scales[scale_idx_h * 8 + cp * 2 + 1]),
3509                                vdup_n_s32(i32::from(q6_scales[scale_idx_h * 8 + cp * 2])),
3510                            );
3511
3512                            acc[cp] = vmla_s32(acc[cp], sum_l, scale_vec_l);
3513                            acc[cp] = vmla_s32(acc[cp], sum_h, scale_vec_h);
3514                        }
3515                    }
3516                }
3517
3518                acc[0] = vsub_s32(acc[0], vget_low_s32(bias_lo));
3519                acc[1] = vsub_s32(acc[1], vget_high_s32(bias_lo));
3520                acc[2] = vsub_s32(acc[2], vget_low_s32(bias_hi));
3521                acc[3] = vsub_s32(acc[3], vget_high_s32(bias_hi));
3522
3523                let w_01 = vmul_f32(vcvt_f32_s32(acc[0]), vget_low_f32(sb_scale[0]));
3524                let w_23 = vmul_f32(vcvt_f32_s32(acc[1]), vget_high_f32(sb_scale[0]));
3525                let w_45 = vmul_f32(vcvt_f32_s32(acc[2]), vget_low_f32(sb_scale[1]));
3526                let w_67 = vmul_f32(vcvt_f32_s32(acc[3]), vget_high_f32(sb_scale[1]));
3527
3528                acc_f32[0] = vaddq_f32(acc_f32[0], vcombine_f32(w_01, w_23));
3529                acc_f32[1] = vaddq_f32(acc_f32[1], vcombine_f32(w_45, w_67));
3530            }
3531
3532            let base = x * Q6_KX8_NROWS;
3533            vst1q_f32(out.as_mut_ptr().add(base), acc_f32[0]);
3534            vst1q_f32(out.as_mut_ptr().add(base + 4), acc_f32[1]);
3535        }
3536    }
3537
3538    /// NEON DotProd **GEMM** for interleave-4 packed Q4_K weights: one
3539    /// row-group (8 rows) against up to [`Q4_KX8_GEMM_NC`] activations.
3540    ///
3541    /// Same arithmetic as [`gemv_q4_kx8_q8_k_neon_sdot`], reordered so
3542    /// the weight-side unpack happens once per activation *tile* rather
3543    /// than once per activation. Per 256-element super-block that hoists
3544    /// 16 f16 scale conversions, 8 `decode_scales_mins` calls and 16
3545    /// `q4_cols` loads out of the batch loop -- which is the whole point,
3546    /// and the same reason llama.cpp ships `ggml_gemm_q4_K_8x4_q8_K`
3547    /// beside its GEMV rather than looping the GEMV.
3548    ///
3549    /// `out` is `[row][act]`: `out[r * na + j]`.
3550    #[target_feature(enable = "neon,dotprod")]
3551    pub unsafe fn gemm_q4_kx8_q8_k_neon_sdot(
3552        packed: &[u8],
3553        acts: &[Q8KActivations],
3554        n_cols: usize,
3555        out: &mut [f32],
3556    ) {
3557        let na = acts.len();
3558        debug_assert!(na <= Q4_KX8_GEMM_NC);
3559        let nb = n_cols / Q4_K_BLOCK_ELEMS;
3560        let m4b = vdupq_n_u8(0x0f);
3561
3562        // [act][row-half]; row-half 0 is rows 0..3, 1 is rows 4..7.
3563        let mut acc_f32 = [[vdupq_n_f32(0.0); 2]; Q4_KX8_GEMM_NC];
3564        let mut bias_acc = [[vdupq_n_s32(0); 2]; Q4_KX8_GEMM_NC];
3565
3566        for b in 0..nb {
3567            let blk = packed.as_ptr().add(b * Q4_KX8_BLOCK_BYTES);
3568
3569            // --- weight-side, once per block (was once per activation) ---
3570            let mut d_arr = [0f32; 8];
3571            let mut dmin_arr = [0f32; 8];
3572            for j in 0..8 {
3573                d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
3574                dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
3575            }
3576            let d_lo = vld1q_f32(d_arr.as_ptr());
3577            let d_hi = vld1q_f32(d_arr.as_ptr().add(4));
3578            let dmin_lo = vld1q_f32(dmin_arr.as_ptr());
3579            let dmin_hi = vld1q_f32(dmin_arr.as_ptr().add(4));
3580
3581            // Per-activation scaling of those, plus the pairwise-added
3582            // bsums this block needs (llama's vpaddq_s16).
3583            let mut sb_scale = [[vdupq_n_f32(0.0); 2]; Q4_KX8_GEMM_NC];
3584            let mut sb_min = [[vdupq_n_f32(0.0); 2]; Q4_KX8_GEMM_NC];
3585            let mut bsums_arr = [[0i16; 8]; Q4_KX8_GEMM_NC];
3586            for (a, act) in acts.iter().enumerate() {
3587                let q8_d = act.d[b];
3588                sb_scale[a] = [vmulq_n_f32(d_lo, q8_d), vmulq_n_f32(d_hi, q8_d)];
3589                sb_min[a] = [vmulq_n_f32(dmin_lo, q8_d), vmulq_n_f32(dmin_hi, q8_d)];
3590                let bsums_ptr = act.bsums.as_ptr().add(b * 16);
3591                for (i, slot) in bsums_arr[a].iter_mut().enumerate() {
3592                    *slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
3593                }
3594            }
3595
3596            let scales_base = blk.add(32);
3597            let qs_base = blk.add(128);
3598
3599            for sb in 0..4 {
3600                // 6-bit scale/min decode: once per block-quarter, not
3601                // once per (block-quarter, activation).
3602                let mut q4sb_mins = [vdupq_n_s16(0); 2];
3603                let mut q4sb_scales = [vdupq_n_s16(0); 2];
3604                for i in 0..2 {
3605                    let mut sc = [0u8; 8];
3606                    let mut mn = [0u8; 8];
3607                    let offset = sb * 24 + i * 12;
3608                    decode_scales_mins(
3609                        std::slice::from_raw_parts(scales_base.add(offset), 12),
3610                        &mut sc,
3611                        &mut mn,
3612                    );
3613                    let mut sc_i8 = [0i8; 8];
3614                    let mut mn_i8 = [0i8; 8];
3615                    for t in 0..8 {
3616                        sc_i8[t] = sc[t] as i8;
3617                        mn_i8[t] = mn[t] as i8;
3618                    }
3619                    q4sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
3620                    q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
3621                }
3622
3623                // `c` selects the row half, so each pass owns one output
3624                // quad and the accumulators can be consumed immediately
3625                // instead of all eight staying live.
3626                for c in 0..2 {
3627                    let mut q4_cols = [vdupq_n_u8(0); 8];
3628                    for (i, slot) in q4_cols.iter_mut().enumerate() {
3629                        *slot = vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + i * 32 + 16 * c));
3630                    }
3631                    let (sc_lo, sc_hi) = if c == 0 {
3632                        (vget_low_s16(q4sb_scales[0]), vget_low_s16(q4sb_scales[1]))
3633                    } else {
3634                        (vget_high_s16(q4sb_scales[0]), vget_high_s16(q4sb_scales[1]))
3635                    };
3636
3637                    // Mask once per weight tile, not once per
3638                    // activation, and keep the lane indices literal --
3639                    // a runtime lane forces a real call per `sdot`
3640                    // instead of the single instruction it should be.
3641                    let lo0 = vreinterpretq_s8_u8(vandq_u8(q4_cols[0], m4b));
3642                    let lo1 = vreinterpretq_s8_u8(vandq_u8(q4_cols[1], m4b));
3643                    let lo2 = vreinterpretq_s8_u8(vandq_u8(q4_cols[2], m4b));
3644                    let lo3 = vreinterpretq_s8_u8(vandq_u8(q4_cols[3], m4b));
3645                    let lo4 = vreinterpretq_s8_u8(vandq_u8(q4_cols[4], m4b));
3646                    let lo5 = vreinterpretq_s8_u8(vandq_u8(q4_cols[5], m4b));
3647                    let lo6 = vreinterpretq_s8_u8(vandq_u8(q4_cols[6], m4b));
3648                    let lo7 = vreinterpretq_s8_u8(vandq_u8(q4_cols[7], m4b));
3649                    let hi0 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[0], 4));
3650                    let hi1 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[1], 4));
3651                    let hi2 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[2], 4));
3652                    let hi3 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[3], 4));
3653                    let hi4 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[4], 4));
3654                    let hi5 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[5], 4));
3655                    let hi6 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[6], 4));
3656                    let hi7 = vreinterpretq_s8_u8(vshrq_n_u8(q4_cols[7], 4));
3657                    let sc_lo_w = vmovl_s16(sc_lo);
3658                    let sc_hi_w = vmovl_s16(sc_hi);
3659
3660                    for a in 0..na {
3661                        let q8_base = acts[a].q.as_ptr().add(b * Q4_K_BLOCK_ELEMS);
3662                        let y0 = vld1q_s8(q8_base.add(sb * 64));
3663                        let y1 = vld1q_s8(q8_base.add(sb * 64 + 16));
3664                        let y2 = vld1q_s8(q8_base.add(sb * 64 + 32));
3665                        let y3 = vld1q_s8(q8_base.add(sb * 64 + 48));
3666                        let mut acc_lo = vdupq_n_s32(0);
3667                        let mut acc_hi = vdupq_n_s32(0);
3668                        acc_lo = sdot_lane(acc_lo, lo0, y0, 0);
3669                        acc_lo = sdot_lane(acc_lo, lo1, y0, 1);
3670                        acc_lo = sdot_lane(acc_lo, lo2, y0, 2);
3671                        acc_lo = sdot_lane(acc_lo, lo3, y0, 3);
3672                        acc_lo = sdot_lane(acc_lo, lo4, y1, 0);
3673                        acc_lo = sdot_lane(acc_lo, lo5, y1, 1);
3674                        acc_lo = sdot_lane(acc_lo, lo6, y1, 2);
3675                        acc_lo = sdot_lane(acc_lo, lo7, y1, 3);
3676                        acc_hi = sdot_lane(acc_hi, hi0, y2, 0);
3677                        acc_hi = sdot_lane(acc_hi, hi1, y2, 1);
3678                        acc_hi = sdot_lane(acc_hi, hi2, y2, 2);
3679                        acc_hi = sdot_lane(acc_hi, hi3, y2, 3);
3680                        acc_hi = sdot_lane(acc_hi, hi4, y3, 0);
3681                        acc_hi = sdot_lane(acc_hi, hi5, y3, 1);
3682                        acc_hi = sdot_lane(acc_hi, hi6, y3, 2);
3683                        acc_hi = sdot_lane(acc_hi, hi7, y3, 3);
3684                        let sumf = vcvtq_f32_s32(vaddq_s32(
3685                            vmulq_s32(sc_lo_w, acc_lo),
3686                            vmulq_s32(sc_hi_w, acc_hi),
3687                        ));
3688                        acc_f32[a][c] = vfmaq_f32(acc_f32[a][c], sb_scale[a][c], sumf);
3689                    }
3690                }
3691
3692                for a in 0..na {
3693                    let bs_lo = vdup_n_s16(bsums_arr[a][2 * sb]);
3694                    let bs_hi = vdup_n_s16(bsums_arr[a][2 * sb + 1]);
3695                    bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_lo, vget_low_s16(q4sb_mins[0]));
3696                    bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_hi, vget_low_s16(q4sb_mins[1]));
3697                    bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_lo, vget_high_s16(q4sb_mins[0]));
3698                    bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_hi, vget_high_s16(q4sb_mins[1]));
3699                }
3700            }
3701
3702            for a in 0..na {
3703                for c in 0..2 {
3704                    acc_f32[a][c] =
3705                        vmlsq_f32(acc_f32[a][c], vcvtq_f32_s32(bias_acc[a][c]), sb_min[a][c]);
3706                    bias_acc[a][c] = vdupq_n_s32(0);
3707                }
3708            }
3709        }
3710
3711        for a in 0..na {
3712            let mut row = [0f32; Q4_KX8_NROWS];
3713            vst1q_f32(row.as_mut_ptr(), acc_f32[a][0]);
3714            vst1q_f32(row.as_mut_ptr().add(4), acc_f32[a][1]);
3715            for (r, v) in row.iter().enumerate() {
3716                out[r * na + a] = *v;
3717            }
3718        }
3719    }
3720
3721    /// NEON DotProd **GEMM** for interleave-4 `block_q5_Kx8` — weight unpack
3722    /// once per act tile (llama `ggml_gemm_q5_K_8x4_q8_K` motivation).
3723    #[target_feature(enable = "neon,dotprod")]
3724    pub unsafe fn gemm_q5_kx8_q8_k_neon_sdot(
3725        packed: &[u8],
3726        acts: &[Q8KActivations],
3727        n_cols: usize,
3728        out: &mut [f32],
3729    ) {
3730        let na = acts.len();
3731        debug_assert!(na <= Q5_KX8_GEMM_NC);
3732        let nb = n_cols / Q5_K_BLOCK_ELEMS;
3733        let m4b = vdupq_n_u8(0x0f);
3734        let mone = vdupq_n_u8(1);
3735        let mtwo = vdupq_n_u8(2);
3736
3737        let mut acc_f32 = [[vdupq_n_f32(0.0); 2]; Q5_KX8_GEMM_NC];
3738        let mut bias_acc = [[vdupq_n_s32(0); 2]; Q5_KX8_GEMM_NC];
3739
3740        for b in 0..nb {
3741            let blk = packed.as_ptr().add(b * Q5_KX8_BLOCK_BYTES);
3742            let mut d_arr = [0f32; 8];
3743            let mut dmin_arr = [0f32; 8];
3744            for j in 0..8 {
3745                d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
3746                dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
3747            }
3748            let d_lo = vld1q_f32(d_arr.as_ptr());
3749            let d_hi = vld1q_f32(d_arr.as_ptr().add(4));
3750            let dmin_lo = vld1q_f32(dmin_arr.as_ptr());
3751            let dmin_hi = vld1q_f32(dmin_arr.as_ptr().add(4));
3752
3753            let mut sb_scale = [[vdupq_n_f32(0.0); 2]; Q5_KX8_GEMM_NC];
3754            let mut sb_min = [[vdupq_n_f32(0.0); 2]; Q5_KX8_GEMM_NC];
3755            let mut bsums_arr = [[0i16; 8]; Q5_KX8_GEMM_NC];
3756            for (a, act) in acts.iter().enumerate() {
3757                let q8_d = act.d[b];
3758                sb_scale[a] = [vmulq_n_f32(d_lo, q8_d), vmulq_n_f32(d_hi, q8_d)];
3759                sb_min[a] = [vmulq_n_f32(dmin_lo, q8_d), vmulq_n_f32(dmin_hi, q8_d)];
3760                let bsums_ptr = act.bsums.as_ptr().add(b * 16);
3761                for (i, slot) in bsums_arr[a].iter_mut().enumerate() {
3762                    *slot = *bsums_ptr.add(2 * i) + *bsums_ptr.add(2 * i + 1);
3763                }
3764            }
3765
3766            let scales_base = blk.add(32);
3767            let qh_base = blk.add(128);
3768            let qs_base = blk.add(384);
3769
3770            let mut qh = [[vdupq_n_u8(0); 8]; 2];
3771            for (c, qh_c) in qh.iter_mut().enumerate() {
3772                for (i, slot) in qh_c.iter_mut().enumerate() {
3773                    *slot = vld1q_u8(qh_base.add(i * 32 + 16 * c));
3774                }
3775            }
3776
3777            for sb in 0..4 {
3778                let mut q5sb_mins = [vdupq_n_s16(0); 2];
3779                let mut q5sb_scales = [vdupq_n_s16(0); 2];
3780                for i in 0..2 {
3781                    let mut sc = [0u8; 8];
3782                    let mut mn = [0u8; 8];
3783                    let offset = sb * 24 + i * 12;
3784                    decode_scales_mins(
3785                        std::slice::from_raw_parts(scales_base.add(offset), 12),
3786                        &mut sc,
3787                        &mut mn,
3788                    );
3789                    let mut sc_i8 = [0i8; 8];
3790                    let mut mn_i8 = [0i8; 8];
3791                    for t in 0..8 {
3792                        sc_i8[t] = sc[t] as i8;
3793                        mn_i8[t] = mn[t] as i8;
3794                    }
3795                    q5sb_scales[i] = vmovl_s8(vld1_s8(sc_i8.as_ptr()));
3796                    q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
3797                }
3798
3799                for c in 0..2 {
3800                    let mut q5_lo = [vdupq_n_s8(0); 8];
3801                    let mut q5_hi = [vdupq_n_s8(0); 8];
3802                    for i in 0..8 {
3803                        let q5_cols =
3804                            vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + i * 32 + 16 * c));
3805                        let hbit_lo = vandq_u8(qh[c][i], mone);
3806                        let hbit_hi = vshlq_n_u8(vandq_u8(qh[c][i], mtwo), 3);
3807                        qh[c][i] = vshrq_n_u8(qh[c][i], 2);
3808                        q5_lo[i] =
3809                            vreinterpretq_s8_u8(vsliq_n_u8(vandq_u8(q5_cols, m4b), hbit_lo, 4));
3810                        q5_hi[i] = vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_cols, 4), hbit_hi));
3811                    }
3812                    let (sc_lo, sc_hi) = if c == 0 {
3813                        (vget_low_s16(q5sb_scales[0]), vget_low_s16(q5sb_scales[1]))
3814                    } else {
3815                        (vget_high_s16(q5sb_scales[0]), vget_high_s16(q5sb_scales[1]))
3816                    };
3817                    let sc_lo_w = vmovl_s16(sc_lo);
3818                    let sc_hi_w = vmovl_s16(sc_hi);
3819
3820                    for a in 0..na {
3821                        let q8_base = acts[a].q.as_ptr().add(b * Q5_K_BLOCK_ELEMS);
3822                        let y0 = vld1q_s8(q8_base.add(sb * 64));
3823                        let y1 = vld1q_s8(q8_base.add(sb * 64 + 16));
3824                        let y2 = vld1q_s8(q8_base.add(sb * 64 + 32));
3825                        let y3 = vld1q_s8(q8_base.add(sb * 64 + 48));
3826                        let mut acc_lo = vdupq_n_s32(0);
3827                        let mut acc_hi = vdupq_n_s32(0);
3828                        acc_lo = sdot_lane(acc_lo, q5_lo[0], y0, 0);
3829                        acc_lo = sdot_lane(acc_lo, q5_lo[1], y0, 1);
3830                        acc_lo = sdot_lane(acc_lo, q5_lo[2], y0, 2);
3831                        acc_lo = sdot_lane(acc_lo, q5_lo[3], y0, 3);
3832                        acc_lo = sdot_lane(acc_lo, q5_lo[4], y1, 0);
3833                        acc_lo = sdot_lane(acc_lo, q5_lo[5], y1, 1);
3834                        acc_lo = sdot_lane(acc_lo, q5_lo[6], y1, 2);
3835                        acc_lo = sdot_lane(acc_lo, q5_lo[7], y1, 3);
3836                        acc_hi = sdot_lane(acc_hi, q5_hi[0], y2, 0);
3837                        acc_hi = sdot_lane(acc_hi, q5_hi[1], y2, 1);
3838                        acc_hi = sdot_lane(acc_hi, q5_hi[2], y2, 2);
3839                        acc_hi = sdot_lane(acc_hi, q5_hi[3], y2, 3);
3840                        acc_hi = sdot_lane(acc_hi, q5_hi[4], y3, 0);
3841                        acc_hi = sdot_lane(acc_hi, q5_hi[5], y3, 1);
3842                        acc_hi = sdot_lane(acc_hi, q5_hi[6], y3, 2);
3843                        acc_hi = sdot_lane(acc_hi, q5_hi[7], y3, 3);
3844                        let sumf = vcvtq_f32_s32(vaddq_s32(
3845                            vmulq_s32(sc_lo_w, acc_lo),
3846                            vmulq_s32(sc_hi_w, acc_hi),
3847                        ));
3848                        acc_f32[a][c] = vfmaq_f32(acc_f32[a][c], sb_scale[a][c], sumf);
3849                    }
3850                }
3851
3852                for a in 0..na {
3853                    let bs_lo = vdup_n_s16(bsums_arr[a][2 * sb]);
3854                    let bs_hi = vdup_n_s16(bsums_arr[a][2 * sb + 1]);
3855                    bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_lo, vget_low_s16(q5sb_mins[0]));
3856                    bias_acc[a][0] = vmlal_s16(bias_acc[a][0], bs_hi, vget_low_s16(q5sb_mins[1]));
3857                    bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_lo, vget_high_s16(q5sb_mins[0]));
3858                    bias_acc[a][1] = vmlal_s16(bias_acc[a][1], bs_hi, vget_high_s16(q5sb_mins[1]));
3859                }
3860            }
3861
3862            for a in 0..na {
3863                for c in 0..2 {
3864                    acc_f32[a][c] =
3865                        vmlsq_f32(acc_f32[a][c], vcvtq_f32_s32(bias_acc[a][c]), sb_min[a][c]);
3866                    bias_acc[a][c] = vdupq_n_s32(0);
3867                }
3868            }
3869        }
3870
3871        for a in 0..na {
3872            let mut row = [0f32; Q5_KX8_NROWS];
3873            vst1q_f32(row.as_mut_ptr(), acc_f32[a][0]);
3874            vst1q_f32(row.as_mut_ptr().add(4), acc_f32[a][1]);
3875            for (r, v) in row.iter().enumerate() {
3876                out[r * na + a] = *v;
3877            }
3878        }
3879    }
3880
3881    /// NEON i8mm **GEMM** for interleave-8 packed Q4_K weights (llama.cpp
3882    /// `ggml_gemm_q4_K_8x8_q8_K` in `arch/arm/repack.cpp`). Uses `vmmlaq_s32`
3883    /// on 2×8×8 tiles; the activation quad arrives pre-interleaved as
3884    /// [`Q8KActsX4`] (llama's `block_q8_Kx4` in `wdata`), so nothing here is
3885    /// repacked per row-group.
3886    #[target_feature(enable = "neon,i8mm")]
3887    pub unsafe fn gemm_q4_kx8_q8_k_neon_i8mm(
3888        packed: &[u8],
3889        tile: &Q8KActsX4,
3890        n_cols: usize,
3891        out: &mut [f32],
3892    ) {
3893        let na = tile.na;
3894        debug_assert!(na <= Q4_KX8_GEMM_NC);
3895        let nb = n_cols / Q4_K_BLOCK_ELEMS;
3896        debug_assert_eq!(tile.n_blocks, nb);
3897        let m4b = vdupq_n_u8(0x0f);
3898        const Q8_K_BLOCKLEN: usize = 4;
3899
3900        let mut acc_f32 = [vdupq_n_f32(0.0); Q4_KX8_GEMM_NC * 2];
3901
3902        for b in 0..nb {
3903            let blk = packed.as_ptr().add(b * Q4_KX8_BLOCK_BYTES);
3904            let bsums_base = tile.bsums.as_ptr().add(b * Q8_K_BLOCKLEN * 8);
3905
3906            let mut acc = [vdupq_n_s32(0); 8];
3907            let mut bias_acc = [vdupq_n_s32(0); 8];
3908            for i in 0..8 {
3909                acc[i] = vdupq_n_s32(0);
3910                bias_acc[i] = vdupq_n_s32(0);
3911            }
3912
3913            let scales_base = blk.add(32);
3914            let qs_base = blk.add(128);
3915            let q8_base = tile.qs.as_ptr().add(b * Q4_K_BLOCK_ELEMS * 4);
3916
3917            for sb in 0..4 {
3918                let mut q4sb_scales = [[0i8; 8]; 2];
3919                let mut q4sb_mins = [vdupq_n_s16(0); 2];
3920                for i in 0..2 {
3921                    let mut sc = [0u8; 8];
3922                    let mut mn = [0u8; 8];
3923                    let offset = sb * 24 + i * 12;
3924                    decode_scales_mins(
3925                        std::slice::from_raw_parts(scales_base.add(offset), 12),
3926                        &mut sc,
3927                        &mut mn,
3928                    );
3929                    let mut mn_i8 = [0i8; 8];
3930                    for t in 0..8 {
3931                        q4sb_scales[i][t] = sc[t] as i8;
3932                        mn_i8[t] = mn[t] as i8;
3933                    }
3934                    q4sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
3935                }
3936
3937                let q8_sb = q8_base.add(sb * 256);
3938                let mut q8_qs_01 = [vdupq_n_s8(0); 8];
3939                let mut q8_qs_23 = [vdupq_n_s8(0); 8];
3940                for i in 0..8 {
3941                    q8_qs_01[i] = vld1q_s8(q8_sb.add(i * 32));
3942                    q8_qs_23[i] = vld1q_s8(q8_sb.add(i * 32 + 16));
3943                }
3944                let q8s = [q8_qs_01, q8_qs_23];
3945
3946                for cp in 0..4 {
3947                    let mut sb_acc = [vdupq_n_s32(0); 4];
3948
3949                    let q4_qs = [
3950                        vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp)),
3951                        vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 64)),
3952                        vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 128)),
3953                        vld1q_u8(qs_base.add(sb * Q4_K_BLOCK_ELEMS + 16 * cp + 192)),
3954                    ];
3955                    let q4_nibbles = [
3956                        [
3957                            vreinterpretq_s8_u8(vandq_u8(q4_qs[0], m4b)),
3958                            vreinterpretq_s8_u8(vandq_u8(q4_qs[1], m4b)),
3959                            vreinterpretq_s8_u8(vandq_u8(q4_qs[2], m4b)),
3960                            vreinterpretq_s8_u8(vandq_u8(q4_qs[3], m4b)),
3961                        ],
3962                        [
3963                            vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[0], 4)),
3964                            vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[1], 4)),
3965                            vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[2], 4)),
3966                            vreinterpretq_s8_u8(vshrq_n_u8(q4_qs[3], 4)),
3967                        ],
3968                    ];
3969
3970                    for rp in 0..2 {
3971                        for blk in 0..2 {
3972                            let q8 = &q8s[rp][4 * blk..4 * blk + 4];
3973                            let q4 = &q4_nibbles[blk];
3974                            let mut tile_acc = sb_acc[2 * rp + blk];
3975                            for qs_offset in 0..4 {
3976                                tile_acc = vmmla_s32(tile_acc, q4[qs_offset], q8[qs_offset]);
3977                            }
3978                            sb_acc[2 * rp + blk] = tile_acc;
3979                        }
3980                    }
3981
3982                    let scale_offset = cp * 2;
3983                    let block_scale_0 = vcombine_s32(
3984                        vdup_n_s32(i32::from(q4sb_scales[0][scale_offset])),
3985                        vdup_n_s32(i32::from(q4sb_scales[0][scale_offset + 1])),
3986                    );
3987                    let block_scale_1 = vcombine_s32(
3988                        vdup_n_s32(i32::from(q4sb_scales[1][scale_offset])),
3989                        vdup_n_s32(i32::from(q4sb_scales[1][scale_offset + 1])),
3990                    );
3991
3992                    acc[cp] = vmlaq_s32(acc[cp], sb_acc[0], block_scale_0);
3993                    acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[2], block_scale_0);
3994                    acc[cp] = vmlaq_s32(acc[cp], sb_acc[1], block_scale_1);
3995                    acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[3], block_scale_1);
3996                }
3997
3998                for q8_row in 0..Q8_K_BLOCKLEN {
3999                    let bs_lo = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb));
4000                    let bs_hi = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb + 1));
4001                    bias_acc[2 * q8_row] =
4002                        vmlal_s16(bias_acc[2 * q8_row], bs_lo, vget_low_s16(q4sb_mins[0]));
4003                    bias_acc[2 * q8_row] =
4004                        vmlal_s16(bias_acc[2 * q8_row], bs_hi, vget_low_s16(q4sb_mins[1]));
4005                    bias_acc[2 * q8_row + 1] =
4006                        vmlal_s16(bias_acc[2 * q8_row + 1], bs_lo, vget_high_s16(q4sb_mins[0]));
4007                    bias_acc[2 * q8_row + 1] =
4008                        vmlal_s16(bias_acc[2 * q8_row + 1], bs_hi, vget_high_s16(q4sb_mins[1]));
4009                }
4010            }
4011
4012            for lane in acc.iter_mut() {
4013                let aux = vzip_s32(vget_low_s32(*lane), vget_high_s32(*lane));
4014                *lane = vcombine_s32(aux.0, aux.1);
4015            }
4016            let reorder_acc = [
4017                vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
4018                vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
4019                vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
4020                vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
4021                vcombine_s32(vget_low_s32(acc[4]), vget_low_s32(acc[5])),
4022                vcombine_s32(vget_low_s32(acc[6]), vget_low_s32(acc[7])),
4023                vcombine_s32(vget_high_s32(acc[4]), vget_high_s32(acc[5])),
4024                vcombine_s32(vget_high_s32(acc[6]), vget_high_s32(acc[7])),
4025            ];
4026
4027            let mut d_arr = [0f32; 8];
4028            let mut dmin_arr = [0f32; 8];
4029            for j in 0..8 {
4030                d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4031                dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
4032            }
4033
4034            for i in 0..na {
4035                for j in 0..2 {
4036                    let q8_d = vdupq_n_f32(*tile.d.as_ptr().add(b * Q8_K_BLOCKLEN + i));
4037                    let dmins = vmulq_f32(vld1q_f32(dmin_arr.as_ptr().add(j * 4)), q8_d);
4038                    let scale = vmulq_f32(vld1q_f32(d_arr.as_ptr().add(j * 4)), q8_d);
4039                    let idx = 2 * i + j;
4040                    acc_f32[idx] = vmlsq_f32(acc_f32[idx], vcvtq_f32_s32(bias_acc[idx]), dmins);
4041                    acc_f32[idx] = vmlaq_f32(acc_f32[idx], vcvtq_f32_s32(reorder_acc[idx]), scale);
4042                }
4043            }
4044        }
4045
4046        for a in 0..na {
4047            let mut row = [0f32; Q4_KX8_NROWS];
4048            vst1q_f32(row.as_mut_ptr(), acc_f32[2 * a]);
4049            vst1q_f32(row.as_mut_ptr().add(4), acc_f32[2 * a + 1]);
4050            for (r, v) in row.iter().enumerate() {
4051                out[r * na + a] = *v;
4052            }
4053        }
4054    }
4055
4056    /// NEON i8mm **GEMM** for interleave-8 packed Q5_K weights (llama.cpp
4057    /// `ggml_gemm_q5_K_8x8_q8_K` in `arch/arm/repack.cpp`). Same 2×8×8
4058    /// `vmmlaq_s32` tiling as [`gemm_q4_kx8_q8_k_neon_i8mm`]; the only
4059    /// difference is splicing the fifth bit in from `qh` before the MMLAs,
4060    /// two bits consumed per sub-block.
4061    #[target_feature(enable = "neon,i8mm")]
4062    pub unsafe fn gemm_q5_kx8_q8_k_neon_i8mm(
4063        packed: &[u8],
4064        tile: &Q8KActsX4,
4065        n_cols: usize,
4066        out: &mut [f32],
4067    ) {
4068        let na = tile.na;
4069        debug_assert!(na <= Q5_KX8_GEMM_NC);
4070        let nb = n_cols / Q5_K_BLOCK_ELEMS;
4071        debug_assert_eq!(tile.n_blocks, nb);
4072        let m4b = vdupq_n_u8(0x0f);
4073        let mone = vdupq_n_u8(1);
4074        let mtwo = vdupq_n_u8(2);
4075        const Q8_K_BLOCKLEN: usize = 4;
4076
4077        let mut acc_f32 = [vdupq_n_f32(0.0); Q5_KX8_GEMM_NC * 2];
4078
4079        for b in 0..nb {
4080            let blk = packed.as_ptr().add(b * Q5_KX8_BLOCK_BYTES);
4081            let bsums_base = tile.bsums.as_ptr().add(b * Q8_K_BLOCKLEN * 8);
4082
4083            let mut acc = [vdupq_n_s32(0); 8];
4084            let mut bias_acc = [vdupq_n_s32(0); 8];
4085
4086            let scales_base = blk.add(32);
4087            let qh_base = blk.add(128);
4088            let qs_base = blk.add(384);
4089            let q8_base = tile.qs.as_ptr().add(b * Q5_K_BLOCK_ELEMS * 4);
4090
4091            // qh state per column pair; two bits consumed per sub-block.
4092            let mut qh = [[vdupq_n_u8(0); 4]; 4];
4093            for (cp, qh_cp) in qh.iter_mut().enumerate() {
4094                for (m, slot) in qh_cp.iter_mut().enumerate() {
4095                    *slot = vld1q_u8(qh_base.add(16 * cp + 64 * m));
4096                }
4097            }
4098
4099            for sb in 0..4 {
4100                let mut q5sb_scales = [[0i8; 8]; 2];
4101                let mut q5sb_mins = [vdupq_n_s16(0); 2];
4102                for i in 0..2 {
4103                    let mut sc = [0u8; 8];
4104                    let mut mn = [0u8; 8];
4105                    let offset = sb * 24 + i * 12;
4106                    decode_scales_mins(
4107                        std::slice::from_raw_parts(scales_base.add(offset), 12),
4108                        &mut sc,
4109                        &mut mn,
4110                    );
4111                    let mut mn_i8 = [0i8; 8];
4112                    for t in 0..8 {
4113                        q5sb_scales[i][t] = sc[t] as i8;
4114                        mn_i8[t] = mn[t] as i8;
4115                    }
4116                    q5sb_mins[i] = vmovl_s8(vld1_s8(mn_i8.as_ptr()));
4117                }
4118
4119                let q8_sb = q8_base.add(sb * 256);
4120                let mut q8_qs_01 = [vdupq_n_s8(0); 8];
4121                let mut q8_qs_23 = [vdupq_n_s8(0); 8];
4122                for i in 0..8 {
4123                    q8_qs_01[i] = vld1q_s8(q8_sb.add(i * 32));
4124                    q8_qs_23[i] = vld1q_s8(q8_sb.add(i * 32 + 16));
4125                }
4126                let q8s = [q8_qs_01, q8_qs_23];
4127
4128                for cp in 0..4 {
4129                    let mut sb_acc = [vdupq_n_s32(0); 4];
4130
4131                    let q5_qs = [
4132                        vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp)),
4133                        vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 64)),
4134                        vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 128)),
4135                        vld1q_u8(qs_base.add(sb * Q5_K_BLOCK_ELEMS + 16 * cp + 192)),
4136                    ];
4137                    let mut q5_lo = [vdupq_n_s8(0); 4];
4138                    let mut q5_hi = [vdupq_n_s8(0); 4];
4139                    for m in 0..4 {
4140                        let hbit_lo = vandq_u8(qh[cp][m], mone);
4141                        let hbit_hi = vshlq_n_u8(vandq_u8(qh[cp][m], mtwo), 3);
4142                        qh[cp][m] = vshrq_n_u8(qh[cp][m], 2);
4143                        q5_lo[m] =
4144                            vreinterpretq_s8_u8(vsliq_n_u8(vandq_u8(q5_qs[m], m4b), hbit_lo, 4));
4145                        q5_hi[m] = vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q5_qs[m], 4), hbit_hi));
4146                    }
4147                    let q5_vals = [q5_lo, q5_hi];
4148
4149                    for rp in 0..2 {
4150                        for half in 0..2 {
4151                            let q8 = &q8s[rp][4 * half..4 * half + 4];
4152                            let q5 = &q5_vals[half];
4153                            let mut tile_acc = sb_acc[2 * rp + half];
4154                            for m in 0..4 {
4155                                tile_acc = vmmla_s32(tile_acc, q5[m], q8[m]);
4156                            }
4157                            sb_acc[2 * rp + half] = tile_acc;
4158                        }
4159                    }
4160
4161                    let scale_offset = cp * 2;
4162                    let block_scale_0 = vcombine_s32(
4163                        vdup_n_s32(i32::from(q5sb_scales[0][scale_offset])),
4164                        vdup_n_s32(i32::from(q5sb_scales[0][scale_offset + 1])),
4165                    );
4166                    let block_scale_1 = vcombine_s32(
4167                        vdup_n_s32(i32::from(q5sb_scales[1][scale_offset])),
4168                        vdup_n_s32(i32::from(q5sb_scales[1][scale_offset + 1])),
4169                    );
4170
4171                    acc[cp] = vmlaq_s32(acc[cp], sb_acc[0], block_scale_0);
4172                    acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[2], block_scale_0);
4173                    acc[cp] = vmlaq_s32(acc[cp], sb_acc[1], block_scale_1);
4174                    acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc[3], block_scale_1);
4175                }
4176
4177                for q8_row in 0..Q8_K_BLOCKLEN {
4178                    let bs_lo = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb));
4179                    let bs_hi = vdup_n_s16(*bsums_base.add(q8_row * 8 + 2 * sb + 1));
4180                    bias_acc[2 * q8_row] =
4181                        vmlal_s16(bias_acc[2 * q8_row], bs_lo, vget_low_s16(q5sb_mins[0]));
4182                    bias_acc[2 * q8_row] =
4183                        vmlal_s16(bias_acc[2 * q8_row], bs_hi, vget_low_s16(q5sb_mins[1]));
4184                    bias_acc[2 * q8_row + 1] =
4185                        vmlal_s16(bias_acc[2 * q8_row + 1], bs_lo, vget_high_s16(q5sb_mins[0]));
4186                    bias_acc[2 * q8_row + 1] =
4187                        vmlal_s16(bias_acc[2 * q8_row + 1], bs_hi, vget_high_s16(q5sb_mins[1]));
4188                }
4189            }
4190
4191            for lane in acc.iter_mut() {
4192                let aux = vzip_s32(vget_low_s32(*lane), vget_high_s32(*lane));
4193                *lane = vcombine_s32(aux.0, aux.1);
4194            }
4195            let reorder_acc = [
4196                vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
4197                vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
4198                vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
4199                vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
4200                vcombine_s32(vget_low_s32(acc[4]), vget_low_s32(acc[5])),
4201                vcombine_s32(vget_low_s32(acc[6]), vget_low_s32(acc[7])),
4202                vcombine_s32(vget_high_s32(acc[4]), vget_high_s32(acc[5])),
4203                vcombine_s32(vget_high_s32(acc[6]), vget_high_s32(acc[7])),
4204            ];
4205
4206            let mut d_arr = [0f32; 8];
4207            let mut dmin_arr = [0f32; 8];
4208            for j in 0..8 {
4209                d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4210                dmin_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
4211            }
4212
4213            for i in 0..na {
4214                for j in 0..2 {
4215                    let q8_d = vdupq_n_f32(*tile.d.as_ptr().add(b * Q8_K_BLOCKLEN + i));
4216                    let dmins = vmulq_f32(vld1q_f32(dmin_arr.as_ptr().add(j * 4)), q8_d);
4217                    let scale = vmulq_f32(vld1q_f32(d_arr.as_ptr().add(j * 4)), q8_d);
4218                    let idx = 2 * i + j;
4219                    acc_f32[idx] = vmlsq_f32(acc_f32[idx], vcvtq_f32_s32(bias_acc[idx]), dmins);
4220                    acc_f32[idx] = vmlaq_f32(acc_f32[idx], vcvtq_f32_s32(reorder_acc[idx]), scale);
4221                }
4222            }
4223        }
4224
4225        for a in 0..na {
4226            let mut row = [0f32; Q5_KX8_NROWS];
4227            vst1q_f32(row.as_mut_ptr(), acc_f32[2 * a]);
4228            vst1q_f32(row.as_mut_ptr().add(4), acc_f32[2 * a + 1]);
4229            for (r, v) in row.iter().enumerate() {
4230                out[r * na + a] = *v;
4231            }
4232        }
4233    }
4234
4235    /// NEON i8mm **GEMM** for interleave-8 packed Q6_K weights (llama.cpp
4236    /// `ggml_gemm_q6_K_8x8_q8_K` in `arch/arm/repack.cpp`). Q6_K has no
4237    /// mins: the -32 offset is folded into the i8 values before the MMLAs
4238    /// (63 - 32 fits i8), so there is no bias pass at all.
4239    #[target_feature(enable = "neon,i8mm")]
4240    pub unsafe fn gemm_q6_kx8_q8_k_neon_i8mm(
4241        packed: &[u8],
4242        tile: &Q8KActsX4,
4243        n_cols: usize,
4244        out: &mut [f32],
4245    ) {
4246        let na = tile.na;
4247        debug_assert!(na <= Q8K_ACTS_X4_NC);
4248        let nb = n_cols / Q6_K_BLOCK_ELEMS;
4249        debug_assert_eq!(tile.n_blocks, nb);
4250        let m4b = vdupq_n_u8(0x0f);
4251        let mask_lo = vdupq_n_u8(0x03);
4252        let mask_hi = vdupq_n_u8(0x30);
4253        let m32s = vdupq_n_s8(32);
4254        const Q8_K_BLOCKLEN: usize = 4;
4255
4256        let mut acc_f32 = [vdupq_n_f32(0.0); Q8K_ACTS_X4_NC * 2];
4257
4258        for b in 0..nb {
4259            let blk = packed.as_ptr().add(b * Q6_KX8_BLOCK_BYTES);
4260            let scales_base = blk.add(16) as *const i8;
4261            let ql_blk = blk.add(144);
4262            let qh_blk = blk.add(1168);
4263            let q8_blk = tile.qs.as_ptr().add(b * Q6_K_BLOCK_ELEMS * 4);
4264
4265            let mut acc = [vdupq_n_s32(0); 8];
4266
4267            // 16 groups of 8 i8 scales, widened once per block.
4268            let mut q6_scales = [0i16; 16 * 8];
4269            for i in 0..16 {
4270                let s16 = vmovl_s8(vld1_s8(scales_base.add(i * 8)));
4271                vst1q_s16(q6_scales.as_mut_ptr().add(i * 8), s16);
4272            }
4273
4274            for half in 0..2 {
4275                let ql_base = ql_blk.add(half * 512);
4276                let qh_base = qh_blk.add(half * 256);
4277
4278                for sb in 0..4 {
4279                    let q8_base_l = q8_blk.add(half * 512 + sb * 64);
4280                    let q8_base_h = q8_blk.add(half * 512 + 256 + sb * 64);
4281
4282                    let mut q8_l_01 = [vdupq_n_s8(0); 2];
4283                    let mut q8_l_23 = [vdupq_n_s8(0); 2];
4284                    let mut q8_h_01 = [vdupq_n_s8(0); 2];
4285                    let mut q8_h_23 = [vdupq_n_s8(0); 2];
4286                    for i in 0..2 {
4287                        q8_l_01[i] = vld1q_s8(q8_base_l.add(i * 32));
4288                        q8_l_23[i] = vld1q_s8(q8_base_l.add(i * 32 + 16));
4289                        q8_h_01[i] = vld1q_s8(q8_base_h.add(i * 32));
4290                        q8_h_23[i] = vld1q_s8(q8_base_h.add(i * 32 + 16));
4291                    }
4292
4293                    let ql_off = sb * (Q6_K_BLOCK_ELEMS / 2);
4294                    let qh_off = ql_off & 255; // wraps after 256 bytes
4295                    let mut q6_ql_0 = [vdupq_n_u8(0); 4];
4296                    let mut q6_ql_1 = [vdupq_n_u8(0); 4];
4297                    let mut q6_qh_0 = [vdupq_n_u8(0); 4];
4298                    let mut q6_qh_1 = [vdupq_n_u8(0); 4];
4299                    for k in 0..4 {
4300                        q6_ql_0[k] = vld1q_u8(ql_base.add(ql_off + 16 * k));
4301                        q6_ql_1[k] = vld1q_u8(ql_base.add(ql_off + 64 + 16 * k));
4302                        q6_qh_0[k] = vld1q_u8(qh_base.add(qh_off + 16 * k));
4303                        q6_qh_1[k] = vld1q_u8(qh_base.add(qh_off + 64 + 16 * k));
4304                    }
4305                    // High bits for sub-blocks 2 and 3 sit two bits up.
4306                    if sb > 1 {
4307                        for k in 0..4 {
4308                            q6_qh_0[k] = vshrq_n_u8(q6_qh_0[k], 2);
4309                            q6_qh_1[k] = vshrq_n_u8(q6_qh_1[k], 2);
4310                        }
4311                    }
4312
4313                    for cp in 0..4 {
4314                        let hh_0 = vandq_u8(q6_qh_0[cp], mask_hi);
4315                        let hh_1 = vandq_u8(q6_qh_1[cp], mask_hi);
4316
4317                        // q6 = (low4 | high2<<4) - 32
4318                        let q6_l0 = vsubq_s8(
4319                            vreinterpretq_s8_u8(vsliq_n_u8(
4320                                vandq_u8(q6_ql_0[cp], m4b),
4321                                vandq_u8(q6_qh_0[cp], mask_lo),
4322                                4,
4323                            )),
4324                            m32s,
4325                        );
4326                        let q6_l1 = vsubq_s8(
4327                            vreinterpretq_s8_u8(vsliq_n_u8(
4328                                vandq_u8(q6_ql_1[cp], m4b),
4329                                vandq_u8(q6_qh_1[cp], mask_lo),
4330                                4,
4331                            )),
4332                            m32s,
4333                        );
4334                        let q6_h0 = vsubq_s8(
4335                            vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_0[cp], 4), hh_0)),
4336                            m32s,
4337                        );
4338                        let q6_h1 = vsubq_s8(
4339                            vreinterpretq_s8_u8(vorrq_u8(vshrq_n_u8(q6_ql_1[cp], 4), hh_1)),
4340                            m32s,
4341                        );
4342
4343                        let mut sb_acc_0l = vmmla_s32(vdupq_n_s32(0), q6_l0, q8_l_01[0]);
4344                        sb_acc_0l = vmmla_s32(sb_acc_0l, q6_l1, q8_l_01[1]);
4345                        let mut sb_acc_0h = vmmla_s32(vdupq_n_s32(0), q6_h0, q8_h_01[0]);
4346                        sb_acc_0h = vmmla_s32(sb_acc_0h, q6_h1, q8_h_01[1]);
4347                        let mut sb_acc_1l = vmmla_s32(vdupq_n_s32(0), q6_l0, q8_l_23[0]);
4348                        sb_acc_1l = vmmla_s32(sb_acc_1l, q6_l1, q8_l_23[1]);
4349                        let mut sb_acc_1h = vmmla_s32(vdupq_n_s32(0), q6_h0, q8_h_23[0]);
4350                        sb_acc_1h = vmmla_s32(sb_acc_1h, q6_h1, q8_h_23[1]);
4351
4352                        let scale_idx_l = half * 8 + sb;
4353                        let scale_idx_h = half * 8 + sb + 4;
4354                        let scale_l = vcombine_s32(
4355                            vdup_n_s32(i32::from(q6_scales[scale_idx_l * 8 + cp * 2])),
4356                            vdup_n_s32(i32::from(q6_scales[scale_idx_l * 8 + cp * 2 + 1])),
4357                        );
4358                        let scale_h = vcombine_s32(
4359                            vdup_n_s32(i32::from(q6_scales[scale_idx_h * 8 + cp * 2])),
4360                            vdup_n_s32(i32::from(q6_scales[scale_idx_h * 8 + cp * 2 + 1])),
4361                        );
4362
4363                        acc[cp] = vmlaq_s32(acc[cp], sb_acc_0l, scale_l);
4364                        acc[cp] = vmlaq_s32(acc[cp], sb_acc_0h, scale_h);
4365                        acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc_1l, scale_l);
4366                        acc[cp + 4] = vmlaq_s32(acc[cp + 4], sb_acc_1h, scale_h);
4367                    }
4368                }
4369            }
4370
4371            for lane in acc.iter_mut() {
4372                let aux = vzip_s32(vget_low_s32(*lane), vget_high_s32(*lane));
4373                *lane = vcombine_s32(aux.0, aux.1);
4374            }
4375            let reorder_acc = [
4376                vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
4377                vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
4378                vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
4379                vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
4380                vcombine_s32(vget_low_s32(acc[4]), vget_low_s32(acc[5])),
4381                vcombine_s32(vget_low_s32(acc[6]), vget_low_s32(acc[7])),
4382                vcombine_s32(vget_high_s32(acc[4]), vget_high_s32(acc[5])),
4383                vcombine_s32(vget_high_s32(acc[6]), vget_high_s32(acc[7])),
4384            ];
4385
4386            let mut d_arr = [0f32; 8];
4387            for (j, slot) in d_arr.iter_mut().enumerate() {
4388                *slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4389            }
4390
4391            for i in 0..na {
4392                for j in 0..2 {
4393                    let q8_d = vdupq_n_f32(*tile.d.as_ptr().add(b * Q8_K_BLOCKLEN + i));
4394                    let scale = vmulq_f32(vld1q_f32(d_arr.as_ptr().add(j * 4)), q8_d);
4395                    let idx = 2 * i + j;
4396                    acc_f32[idx] = vmlaq_f32(acc_f32[idx], vcvtq_f32_s32(reorder_acc[idx]), scale);
4397                }
4398            }
4399        }
4400
4401        for a in 0..na {
4402            let mut row = [0f32; Q6_KX8_NROWS];
4403            vst1q_f32(row.as_mut_ptr(), acc_f32[2 * a]);
4404            vst1q_f32(row.as_mut_ptr().add(4), acc_f32[2 * a + 1]);
4405            for (r, v) in row.iter().enumerate() {
4406                out[r * na + a] = *v;
4407            }
4408        }
4409    }
4410
4411    /// NEON DotProd GEMV for `block_q8_0x4` (llama `ggml_gemv_q8_0_4x4_q8_0`).
4412    #[target_feature(enable = "neon,dotprod")]
4413    pub unsafe fn gemv_q8_0x4_q8_0_neon_sdot(
4414        packed: &[u8],
4415        act: &Q8Activations,
4416        n_cols: usize,
4417        n_row_groups: usize,
4418        out: &mut [f32],
4419    ) {
4420        let nb = n_cols / Q8_0_BLOCK_ELEMS;
4421        for x in 0..n_row_groups {
4422            let mut acc = vdupq_n_f32(0.0);
4423            let group_off = x * nb * Q8_0X4_BLOCK_BYTES;
4424            for b in 0..nb {
4425                let blk = packed.as_ptr().add(group_off + b * Q8_0X4_BLOCK_BYTES);
4426                let qs = blk.add(8);
4427                // Four int8x16: first 64 qs bytes (k=0..3 × 4 rows × 4).
4428                let b0 = vld1q_s8(qs as *const i8);
4429                let b1 = vld1q_s8(qs.add(16) as *const i8);
4430                let b2 = vld1q_s8(qs.add(32) as *const i8);
4431                let b3 = vld1q_s8(qs.add(48) as *const i8);
4432                let b4 = vld1q_s8(qs.add(64) as *const i8);
4433                let b5 = vld1q_s8(qs.add(80) as *const i8);
4434                let b6 = vld1q_s8(qs.add(96) as *const i8);
4435                let b7 = vld1q_s8(qs.add(112) as *const i8);
4436
4437                let a_ptr = act.q.as_ptr().add(b * Q8_0_BLOCK_ELEMS);
4438                let a0 = vld1q_s8(a_ptr);
4439                let a1 = vld1q_s8(a_ptr.add(16));
4440
4441                let mut ret = vdupq_n_s32(0);
4442                ret = sdot_lane(ret, b0, a0, 0);
4443                ret = sdot_lane(ret, b1, a0, 1);
4444                ret = sdot_lane(ret, b2, a0, 2);
4445                ret = sdot_lane(ret, b3, a0, 3);
4446                ret = sdot_lane(ret, b4, a1, 0);
4447                ret = sdot_lane(ret, b5, a1, 1);
4448                ret = sdot_lane(ret, b6, a1, 2);
4449                ret = sdot_lane(ret, b7, a1, 3);
4450
4451                // Four f16 weight scales at blk[0..8] — load as u16 then
4452                // convert (avoids 4× scalar half::f16 path per block).
4453                let d_bits = vld1_u16(blk as *const u16);
4454                let mut dw = [0f32; 4];
4455                dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
4456                dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
4457                dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
4458                dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
4459                let scale = vmulq_n_f32(vld1q_f32(dw.as_ptr()), act.d[b]);
4460                acc = vfmaq_f32(acc, vcvtq_f32_s32(ret), scale);
4461            }
4462            vst1q_f32(out.as_mut_ptr().add(x * Q8_0X4_NROWS), acc);
4463        }
4464    }
4465
4466    /// NEON DotProd GEMM for one `block_q8_0x4` row-group against
4467    /// several activations (llama `ggml_gemm_q8_0_4x4_q8_0` in shape).
4468    ///
4469    /// The eight weight vectors of a block are loaded once and reused
4470    /// across a tile of [`Q8_0X4_GEMM_NC`] activations, which is the
4471    /// whole point of having a GEMM rather than a loop over the GEMV.
4472    #[target_feature(enable = "neon,dotprod")]
4473    pub unsafe fn gemm_q8_0x4_q8_0_neon_sdot(
4474        group: &[u8],
4475        acts: &[Q8Activations],
4476        n_cols: usize,
4477        out: &mut [f32],
4478    ) {
4479        let nb = n_cols / Q8_0_BLOCK_ELEMS;
4480        let n_acts = acts.len();
4481        let mut j0 = 0;
4482        while j0 < n_acts {
4483            let tile = Q8_0X4_GEMM_NC.min(n_acts - j0);
4484            let mut acc = [vdupq_n_f32(0.0); Q8_0X4_GEMM_NC];
4485            for b in 0..nb {
4486                let blk = group.as_ptr().add(b * Q8_0X4_BLOCK_BYTES);
4487                let qs = blk.add(8);
4488                let w = [
4489                    vld1q_s8(qs as *const i8),
4490                    vld1q_s8(qs.add(16) as *const i8),
4491                    vld1q_s8(qs.add(32) as *const i8),
4492                    vld1q_s8(qs.add(48) as *const i8),
4493                    vld1q_s8(qs.add(64) as *const i8),
4494                    vld1q_s8(qs.add(80) as *const i8),
4495                    vld1q_s8(qs.add(96) as *const i8),
4496                    vld1q_s8(qs.add(112) as *const i8),
4497                ];
4498                let d_bits = vld1_u16(blk as *const u16);
4499                let mut dw = [0f32; 4];
4500                dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
4501                dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
4502                dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
4503                dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
4504                let dw_v = vld1q_f32(dw.as_ptr());
4505
4506                for t in 0..tile {
4507                    let act = &acts[j0 + t];
4508                    let a_ptr = act.q.as_ptr().add(b * Q8_0_BLOCK_ELEMS);
4509                    let a0 = vld1q_s8(a_ptr);
4510                    let a1 = vld1q_s8(a_ptr.add(16));
4511                    let mut ret = vdupq_n_s32(0);
4512                    ret = sdot_lane(ret, w[0], a0, 0);
4513                    ret = sdot_lane(ret, w[1], a0, 1);
4514                    ret = sdot_lane(ret, w[2], a0, 2);
4515                    ret = sdot_lane(ret, w[3], a0, 3);
4516                    ret = sdot_lane(ret, w[4], a1, 0);
4517                    ret = sdot_lane(ret, w[5], a1, 1);
4518                    ret = sdot_lane(ret, w[6], a1, 2);
4519                    ret = sdot_lane(ret, w[7], a1, 3);
4520                    let scale = vmulq_n_f32(dw_v, act.d[b]);
4521                    acc[t] = vfmaq_f32(acc[t], vcvtq_f32_s32(ret), scale);
4522                }
4523            }
4524            for t in 0..tile {
4525                let mut lanes = [0f32; Q8_0X4_NROWS];
4526                vst1q_f32(lanes.as_mut_ptr(), acc[t]);
4527                for (r, v) in lanes.iter().enumerate() {
4528                    out[r * n_acts + j0 + t] = *v;
4529                }
4530            }
4531            j0 += tile;
4532        }
4533    }
4534
4535    /// NEON DotProd GEMV for `block_q4_0x4` (llama `ggml_gemv_q4_0_4x4_q8_0`).
4536    #[target_feature(enable = "neon,dotprod")]
4537    pub unsafe fn gemv_q4_0x4_q8_0_neon_sdot(
4538        packed: &[u8],
4539        act: &Q8Activations,
4540        n_cols: usize,
4541        n_row_groups: usize,
4542        out: &mut [f32],
4543    ) {
4544        let nb = n_cols / Q4_0_BLOCK_ELEMS;
4545        let maskf0 = vdupq_n_u8(0xF0);
4546        for x in 0..n_row_groups {
4547            let mut acc = vdupq_n_f32(0.0);
4548            let group_off = x * nb * Q4_0X4_BLOCK_BYTES;
4549            for b in 0..nb {
4550                let blk = packed.as_ptr().add(group_off + b * Q4_0X4_BLOCK_BYTES);
4551                let qs = blk.add(8);
4552                let a_ptr = act.q.as_ptr().add(b * Q4_0_BLOCK_ELEMS);
4553                let a0 = vld1q_s8(a_ptr);
4554                let a1 = vld1q_s8(a_ptr.add(16));
4555
4556                let mut ret = vdupq_n_s32(0);
4557                for wi in 0..4u32 {
4558                    let w = vld1q_u8(qs.add(wi as usize * 16));
4559                    let hi = vreinterpretq_s8_u8(vshlq_n_u8(w, 4));
4560                    let lo = vreinterpretq_s8_u8(vandq_u8(w, maskf0));
4561                    ret = sdot_lane(ret, hi, a0, wi);
4562                    ret = sdot_lane(ret, lo, a1, wi);
4563                }
4564
4565                let d_bits = vld1_u16(blk as *const u16);
4566                let mut dw = [0f32; 4];
4567                dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
4568                dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
4569                dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
4570                dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
4571                let scale = vmulq_n_f32(vld1q_f32(dw.as_ptr()), act.d[b]);
4572                acc = vfmaq_f32(acc, vcvtq_f32_s32(vshrq_n_s32(ret, 4)), scale);
4573            }
4574            vst1q_f32(out.as_mut_ptr().add(x * Q4_0X4_NROWS), acc);
4575        }
4576    }
4577
4578    /// NEON DotProd GEMM for one `block_q4_0x4` row-group against several
4579    /// activations (llama `ggml_gemm_q4_0_4x4_q8_0` in shape).
4580    #[target_feature(enable = "neon,dotprod")]
4581    pub unsafe fn gemm_q4_0x4_q8_0_neon_sdot(
4582        group: &[u8],
4583        acts: &[Q8Activations],
4584        n_cols: usize,
4585        out: &mut [f32],
4586    ) {
4587        let nb = n_cols / Q4_0_BLOCK_ELEMS;
4588        let n_acts = acts.len();
4589        let maskf0 = vdupq_n_u8(0xF0);
4590        let mut j0 = 0;
4591        while j0 < n_acts {
4592            let tile = Q4_0X4_GEMM_NC.min(n_acts - j0);
4593            let mut acc = [vdupq_n_f32(0.0); Q4_0X4_GEMM_NC];
4594            for b in 0..nb {
4595                let blk = group.as_ptr().add(b * Q4_0X4_BLOCK_BYTES);
4596                let qs = blk.add(8);
4597                let w = [
4598                    vld1q_u8(qs),
4599                    vld1q_u8(qs.add(16)),
4600                    vld1q_u8(qs.add(32)),
4601                    vld1q_u8(qs.add(48)),
4602                ];
4603                let d_bits = vld1_u16(blk as *const u16);
4604                let mut dw = [0f32; 4];
4605                dw[0] = f16::from_bits(vget_lane_u16(d_bits, 0)).to_f32();
4606                dw[1] = f16::from_bits(vget_lane_u16(d_bits, 1)).to_f32();
4607                dw[2] = f16::from_bits(vget_lane_u16(d_bits, 2)).to_f32();
4608                dw[3] = f16::from_bits(vget_lane_u16(d_bits, 3)).to_f32();
4609                let dw_v = vld1q_f32(dw.as_ptr());
4610
4611                for t in 0..tile {
4612                    let act = &acts[j0 + t];
4613                    let a_ptr = act.q.as_ptr().add(b * Q4_0_BLOCK_ELEMS);
4614                    let a0 = vld1q_s8(a_ptr);
4615                    let a1 = vld1q_s8(a_ptr.add(16));
4616                    let mut ret = vdupq_n_s32(0);
4617                    for (wi, wchunk) in w.iter().enumerate() {
4618                        let hi = vreinterpretq_s8_u8(vshlq_n_u8(*wchunk, 4));
4619                        let lo = vreinterpretq_s8_u8(vandq_u8(*wchunk, maskf0));
4620                        ret = sdot_lane(ret, hi, a0, wi as u32);
4621                        ret = sdot_lane(ret, lo, a1, wi as u32);
4622                    }
4623                    let scale = vmulq_n_f32(dw_v, act.d[b]);
4624                    acc[t] = vfmaq_f32(acc[t], vcvtq_f32_s32(vshrq_n_s32(ret, 4)), scale);
4625                }
4626            }
4627            for t in 0..tile {
4628                let mut lanes = [0f32; Q4_0X4_NROWS];
4629                vst1q_f32(lanes.as_mut_ptr(), acc[t]);
4630                for (r, v) in lanes.iter().enumerate() {
4631                    out[r * n_acts + j0 + t] = *v;
4632                }
4633            }
4634            j0 += tile;
4635        }
4636    }
4637
4638    /// NEON DotProd GEMV for interleave-8 packed Q8_0 weights (llama.cpp
4639    /// `ggml_gemv_q8_0_4x8_q8_0` in `arch/arm/repack.cpp`). Each 8-byte
4640    /// activation run is broadcast to both vector halves so one `sdot`
4641    /// covers two interleaved rows.
4642    #[target_feature(enable = "neon,dotprod")]
4643    pub unsafe fn gemv_q8_0x4_q8_0_neon_4x8(
4644        packed: &[u8],
4645        act: &Q8Activations,
4646        n_cols: usize,
4647        n_row_groups: usize,
4648        out: &mut [f32],
4649    ) {
4650        let nb = n_cols / Q8_0_BLOCK_ELEMS;
4651
4652        for x in 0..n_row_groups {
4653            let mut acc = vdupq_n_f32(0.0);
4654            let group_off = x * nb * Q8_0X4_BLOCK_BYTES;
4655
4656            for b in 0..nb {
4657                let blk = packed.as_ptr().add(group_off + b * Q8_0X4_BLOCK_BYTES);
4658                let qs = blk.add(8) as *const i8;
4659                let mut d_arr = [0f32; 4];
4660                for (j, slot) in d_arr.iter_mut().enumerate() {
4661                    *slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4662                }
4663                let b_d = vld1q_f32(d_arr.as_ptr());
4664                let a_base = act.q.as_ptr().add(b * Q8_0_BLOCK_ELEMS);
4665
4666                let mut ret0 = vdupq_n_s32(0);
4667                let mut ret1 = vdupq_n_s32(0);
4668                for c in 0..4 {
4669                    let a = vreinterpretq_s8_s64(vld1q_dup_s64(a_base.add(c * 8) as *const i64));
4670                    ret0 = sdot(ret0, vld1q_s8(qs.add(c * 32)), a);
4671                    ret1 = sdot(ret1, vld1q_s8(qs.add(c * 32 + 16)), a);
4672                }
4673                let ret = vpaddq_s32(ret0, ret1);
4674
4675                acc = vfmaq_f32(acc, vcvtq_f32_s32(ret), vmulq_n_f32(b_d, act.d[b]));
4676            }
4677
4678            vst1q_f32(out.as_mut_ptr().add(x * Q8_0X4_NROWS), acc);
4679        }
4680    }
4681
4682    /// NEON DotProd GEMV for interleave-8 packed Q4_0 weights (llama.cpp
4683    /// `ggml_gemv_q4_0_4x8_q8_0` in `arch/arm/repack.cpp`). Nibbles are
4684    /// consumed at 16× their value (`<< 4` for the low half, `& 0xf0` for
4685    /// the high half — the pack's 0x88 XOR already folded in the -8), and
4686    /// the fixed-point convert divides the 16 back out.
4687    #[target_feature(enable = "neon,dotprod")]
4688    pub unsafe fn gemv_q4_0x4_q8_0_neon_4x8(
4689        packed: &[u8],
4690        act: &Q8Activations,
4691        n_cols: usize,
4692        n_row_groups: usize,
4693        out: &mut [f32],
4694    ) {
4695        let nb = n_cols / Q4_0_BLOCK_ELEMS;
4696        let m4b = vdupq_n_u8(0xf0);
4697
4698        for x in 0..n_row_groups {
4699            let mut acc = vdupq_n_f32(0.0);
4700            let group_off = x * nb * Q4_0X4_BLOCK_BYTES;
4701
4702            for b in 0..nb {
4703                let blk = packed.as_ptr().add(group_off + b * Q4_0X4_BLOCK_BYTES);
4704                let qs = blk.add(8);
4705                let mut d_arr = [0f32; 4];
4706                for (j, slot) in d_arr.iter_mut().enumerate() {
4707                    *slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4708                }
4709                let b_d = vld1q_f32(d_arr.as_ptr());
4710                let a_base = act.q.as_ptr().add(b * Q4_0_BLOCK_ELEMS);
4711
4712                let b0 = vld1q_u8(qs);
4713                let b1 = vld1q_u8(qs.add(16));
4714                let b2 = vld1q_u8(qs.add(32));
4715                let b3 = vld1q_u8(qs.add(48));
4716
4717                let mut a = [vdupq_n_s8(0); 4];
4718                for (c, slot) in a.iter_mut().enumerate() {
4719                    *slot = vreinterpretq_s8_s64(vld1q_dup_s64(a_base.add(c * 8) as *const i64));
4720                }
4721
4722                let mut ret0 = vdupq_n_s32(0);
4723                let mut ret1 = vdupq_n_s32(0);
4724                ret0 = sdot(ret0, vreinterpretq_s8_u8(vshlq_n_u8(b0, 4)), a[0]);
4725                ret1 = sdot(ret1, vreinterpretq_s8_u8(vshlq_n_u8(b1, 4)), a[0]);
4726                ret0 = sdot(ret0, vreinterpretq_s8_u8(vshlq_n_u8(b2, 4)), a[1]);
4727                ret1 = sdot(ret1, vreinterpretq_s8_u8(vshlq_n_u8(b3, 4)), a[1]);
4728                ret0 = sdot(ret0, vreinterpretq_s8_u8(vandq_u8(b0, m4b)), a[2]);
4729                ret1 = sdot(ret1, vreinterpretq_s8_u8(vandq_u8(b1, m4b)), a[2]);
4730                ret0 = sdot(ret0, vreinterpretq_s8_u8(vandq_u8(b2, m4b)), a[3]);
4731                ret1 = sdot(ret1, vreinterpretq_s8_u8(vandq_u8(b3, m4b)), a[3]);
4732                let ret = vpaddq_s32(ret0, ret1);
4733
4734                acc = vfmaq_f32(acc, vcvtq_n_f32_s32::<4>(ret), vmulq_n_f32(b_d, act.d[b]));
4735            }
4736
4737            vst1q_f32(out.as_mut_ptr().add(x * Q4_0X4_NROWS), acc);
4738        }
4739    }
4740
4741    /// NEON i8mm **GEMM** for interleave-8 packed Q8_0 weights (llama.cpp
4742    /// `ggml_gemm_q8_0_4x8_q8_0` in `arch/arm/repack.cpp`, NEON branch).
4743    /// The activation quad arrives pre-interleaved as [`Q8ActsX4`], so
4744    /// every `vmmlaq_s32` covers a 2×2 (activation × weight-row) tile.
4745    #[target_feature(enable = "neon,i8mm")]
4746    pub unsafe fn gemm_q8_0x4_q8_0_neon_i8mm(
4747        packed: &[u8],
4748        tile: &Q8ActsX4,
4749        n_cols: usize,
4750        out: &mut [f32],
4751    ) {
4752        let na = tile.na;
4753        debug_assert!(na <= Q8K_ACTS_X4_NC);
4754        let nb = n_cols / Q8_0_BLOCK_ELEMS;
4755        debug_assert_eq!(tile.n_blocks, nb);
4756
4757        // acc_f32[a] holds activation row a's 4 weight-row outputs.
4758        let mut acc_f32 = [vdupq_n_f32(0.0); Q8K_ACTS_X4_NC];
4759
4760        for b in 0..nb {
4761            let blk = packed.as_ptr().add(b * Q8_0X4_BLOCK_BYTES);
4762            let qs = blk.add(8) as *const i8;
4763            let a_base = tile.qs.as_ptr().add(b * Q8_0_BLOCK_ELEMS * 4);
4764
4765            let mut acc = [vdupq_n_s32(0); 4];
4766            for chunk in 0..4 {
4767                let a01 = vld1q_s8(a_base.add(chunk * 32));
4768                let a23 = vld1q_s8(a_base.add(chunk * 32 + 16));
4769                let b01 = vld1q_s8(qs.add(chunk * 32));
4770                let b23 = vld1q_s8(qs.add(chunk * 32 + 16));
4771
4772                acc[0] = vmmla_s32(acc[0], a01, b01);
4773                acc[1] = vmmla_s32(acc[1], a01, b23);
4774                acc[2] = vmmla_s32(acc[2], a23, b01);
4775                acc[3] = vmmla_s32(acc[3], a23, b23);
4776            }
4777
4778            // 2×2 tiles → per-activation-row vectors of 4 weight rows.
4779            let rows = [
4780                vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
4781                vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
4782                vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
4783                vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
4784            ];
4785
4786            let mut d_arr = [0f32; 4];
4787            for (j, slot) in d_arr.iter_mut().enumerate() {
4788                *slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4789            }
4790            let b_d = vld1q_f32(d_arr.as_ptr());
4791
4792            for a in 0..na {
4793                acc_f32[a] = vfmaq_f32(
4794                    acc_f32[a],
4795                    vcvtq_f32_s32(rows[a]),
4796                    vmulq_n_f32(b_d, *tile.d.as_ptr().add(b * 4 + a)),
4797                );
4798            }
4799        }
4800
4801        for a in 0..na {
4802            let mut lanes = [0f32; Q8_0X4_NROWS];
4803            vst1q_f32(lanes.as_mut_ptr(), acc_f32[a]);
4804            for (r, v) in lanes.iter().enumerate() {
4805                out[r * na + a] = *v;
4806            }
4807        }
4808    }
4809
4810    /// NEON i8mm **GEMM** for interleave-8 packed Q4_0 weights. llama.cpp
4811    /// ships `ggml_gemm_q4_0_4x8_q8_0` (`arch/arm/repack.cpp`) as raw
4812    /// inline asm; this is the same computation with intrinsics, following
4813    /// the `4x8` GEMV's nibble handling and the Q8_0 GEMM's MMLA tiling.
4814    #[target_feature(enable = "neon,i8mm")]
4815    pub unsafe fn gemm_q4_0x4_q8_0_neon_i8mm(
4816        packed: &[u8],
4817        tile: &Q8ActsX4,
4818        n_cols: usize,
4819        out: &mut [f32],
4820    ) {
4821        let na = tile.na;
4822        debug_assert!(na <= Q8K_ACTS_X4_NC);
4823        let nb = n_cols / Q4_0_BLOCK_ELEMS;
4824        debug_assert_eq!(tile.n_blocks, nb);
4825        let m4b = vdupq_n_u8(0xf0);
4826
4827        let mut acc_f32 = [vdupq_n_f32(0.0); Q8K_ACTS_X4_NC];
4828
4829        for b in 0..nb {
4830            let blk = packed.as_ptr().add(b * Q4_0X4_BLOCK_BYTES);
4831            let qs = blk.add(8);
4832            let a_base = tile.qs.as_ptr().add(b * Q4_0_BLOCK_ELEMS * 4);
4833
4834            let bv = [
4835                vld1q_u8(qs),
4836                vld1q_u8(qs.add(16)),
4837                vld1q_u8(qs.add(32)),
4838                vld1q_u8(qs.add(48)),
4839            ];
4840            // Weight vectors per activation chunk: lo nibbles cover elems
4841            // 0..16 (chunks 0,1), hi nibbles elems 16..32 (chunks 2,3),
4842            // all at 16× their value until the fixed-point convert.
4843            let w = [
4844                [
4845                    vreinterpretq_s8_u8(vshlq_n_u8(bv[0], 4)),
4846                    vreinterpretq_s8_u8(vshlq_n_u8(bv[1], 4)),
4847                ],
4848                [
4849                    vreinterpretq_s8_u8(vshlq_n_u8(bv[2], 4)),
4850                    vreinterpretq_s8_u8(vshlq_n_u8(bv[3], 4)),
4851                ],
4852                [
4853                    vreinterpretq_s8_u8(vandq_u8(bv[0], m4b)),
4854                    vreinterpretq_s8_u8(vandq_u8(bv[1], m4b)),
4855                ],
4856                [
4857                    vreinterpretq_s8_u8(vandq_u8(bv[2], m4b)),
4858                    vreinterpretq_s8_u8(vandq_u8(bv[3], m4b)),
4859                ],
4860            ];
4861
4862            let mut acc = [vdupq_n_s32(0); 4];
4863            for (chunk, w_pair) in w.iter().enumerate() {
4864                let a01 = vld1q_s8(a_base.add(chunk * 32));
4865                let a23 = vld1q_s8(a_base.add(chunk * 32 + 16));
4866
4867                acc[0] = vmmla_s32(acc[0], a01, w_pair[0]);
4868                acc[1] = vmmla_s32(acc[1], a01, w_pair[1]);
4869                acc[2] = vmmla_s32(acc[2], a23, w_pair[0]);
4870                acc[3] = vmmla_s32(acc[3], a23, w_pair[1]);
4871            }
4872
4873            let rows = [
4874                vcombine_s32(vget_low_s32(acc[0]), vget_low_s32(acc[1])),
4875                vcombine_s32(vget_high_s32(acc[0]), vget_high_s32(acc[1])),
4876                vcombine_s32(vget_low_s32(acc[2]), vget_low_s32(acc[3])),
4877                vcombine_s32(vget_high_s32(acc[2]), vget_high_s32(acc[3])),
4878            ];
4879
4880            let mut d_arr = [0f32; 4];
4881            for (j, slot) in d_arr.iter_mut().enumerate() {
4882                *slot = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4883            }
4884            let b_d = vld1q_f32(d_arr.as_ptr());
4885
4886            for a in 0..na {
4887                acc_f32[a] = vfmaq_f32(
4888                    acc_f32[a],
4889                    vcvtq_n_f32_s32::<4>(rows[a]),
4890                    vmulq_n_f32(b_d, *tile.d.as_ptr().add(b * 4 + a)),
4891                );
4892            }
4893        }
4894
4895        for a in 0..na {
4896            let mut lanes = [0f32; Q4_0X4_NROWS];
4897            vst1q_f32(lanes.as_mut_ptr(), acc_f32[a]);
4898            for (r, v) in lanes.iter().enumerate() {
4899                out[r * na + a] = *v;
4900            }
4901        }
4902    }
4903}
4904
4905#[cfg(target_arch = "x86_64")]
4906mod avx2 {
4907    use super::*;
4908    use std::arch::x86_64::*;
4909
4910    /// AVX2 GEMV for interleave-8 packed weights. Accumulates 8 f32 outputs
4911    /// in `__m256` lanes; inner int dots use maddubs over nibble×act pairs.
4912    #[target_feature(enable = "avx2,fma")]
4913    pub unsafe fn gemv_q4_kx8_q8_k_avx2(
4914        packed: &[u8],
4915        act: &Q8KActivations,
4916        n_cols: usize,
4917        n_row_groups: usize,
4918        out: &mut [f32],
4919    ) {
4920        let nb = n_cols / Q4_K_BLOCK_ELEMS;
4921        let blocklen = 8;
4922        let ncols = Q4_KX8_NROWS;
4923
4924        for x in 0..n_row_groups {
4925            let mut acc = _mm256_setzero_ps();
4926            let mut acc_min = _mm256_setzero_ps();
4927            let group_off = x * nb * Q4_KX8_BLOCK_BYTES;
4928
4929            for l in 0..nb {
4930                let blk = packed.as_ptr().add(group_off + l * Q4_KX8_BLOCK_BYTES);
4931                let mut d_arr = [0f32; 8];
4932                let mut dmin_arr = [0f32; 8];
4933                for j in 0..8 {
4934                    d_arr[j] = f16_from_bytes(std::slice::from_raw_parts(blk.add(j * 2), 2));
4935                    dmin_arr[j] =
4936                        f16_from_bytes(std::slice::from_raw_parts(blk.add(16 + j * 2), 2));
4937                }
4938                let da = act.d[l];
4939                let d_vec = _mm256_mul_ps(_mm256_loadu_ps(d_arr.as_ptr()), _mm256_set1_ps(da));
4940                let dmin_vec =
4941                    _mm256_mul_ps(_mm256_loadu_ps(dmin_arr.as_ptr()), _mm256_set1_ps(da));
4942
4943                let scales = std::slice::from_raw_parts(blk.add(32), 96);
4944                let qs = std::slice::from_raw_parts(blk.add(128), 1024);
4945                let q8 = &act.q[l * Q4_K_BLOCK_ELEMS..(l + 1) * Q4_K_BLOCK_ELEMS];
4946                let bsums = &act.bsums[l * 16..(l + 1) * 16];
4947
4948                let mut all_scales = [[0u8; 8]; 8];
4949                let mut all_mins = [[0u8; 8]; 8];
4950                for sb in 0..8 {
4951                    decode_scales_mins(&scales[sb * 12..], &mut all_scales[sb], &mut all_mins[sb]);
4952                }
4953
4954                let mut isum = [0i32; 8];
4955                let n_k = Q4_K_BLOCK_ELEMS / (2 * blocklen);
4956                for k in 0..n_k {
4957                    let sb_pair = k / 4;
4958                    let sc0 = &all_scales[sb_pair * 2];
4959                    let sc1 = &all_scales[sb_pair * 2 + 1];
4960                    for j in 0..ncols {
4961                        let mut s = 0i32;
4962                        for i in 0..blocklen {
4963                            let qbyte = qs[k * ncols * blocklen + j * blocklen + i];
4964                            let v0 = (qbyte & 0x0F) as i32;
4965                            let v1 = (qbyte >> 4) as i32;
4966                            let a0 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i] as i32;
4967                            let a1 = q8[(k >> 2) * 64 + (k % 4) * blocklen + i + 32] as i32;
4968                            s += v0 * a0 * sc0[j] as i32 + v1 * a1 * sc1[j] as i32;
4969                        }
4970                        isum[j] += s;
4971                    }
4972                }
4973
4974                let isum_ps =
4975                    _mm256_cvtepi32_ps(_mm256_loadu_si256(isum.as_ptr() as *const __m256i));
4976                acc = _mm256_fmadd_ps(isum_ps, d_vec, acc);
4977
4978                let mut minsum = [0i32; 8];
4979                for sb in 0..8 {
4980                    let bsum = bsums[sb * 2] as i32 + bsums[sb * 2 + 1] as i32;
4981                    for j in 0..ncols {
4982                        minsum[j] += all_mins[sb][j] as i32 * bsum;
4983                    }
4984                }
4985                let minsum_ps =
4986                    _mm256_cvtepi32_ps(_mm256_loadu_si256(minsum.as_ptr() as *const __m256i));
4987                acc_min = _mm256_fmadd_ps(minsum_ps, dmin_vec, acc_min);
4988            }
4989
4990            _mm256_storeu_ps(out.as_mut_ptr().add(x * ncols), _mm256_sub_ps(acc, acc_min));
4991        }
4992    }
4993}
4994
4995#[cfg(test)]
4996mod tests {
4997    use super::*;
4998    use crate::{
4999        dot_q4_0_q8_scalar, dot_q4_k_q8_scalar, dot_q5_k_q8_scalar, dot_q6_k_q8_scalar,
5000        dot_q8_0_q8_scalar, quantize_activations_q8, quantize_activations_q8_k, Q4_0_BLOCK_BYTES,
5001        Q5_K_BLOCK_BYTES, Q6_K_BLOCK_BYTES,
5002    };
5003
5004    fn synth_q5_k_row(n_blocks: usize, seed: u8) -> Vec<u8> {
5005        let mut weights = Vec::with_capacity(n_blocks * Q5_K_BLOCK_BYTES);
5006        for b in 0..n_blocks {
5007            weights.extend_from_slice(
5008                &f16::from_f32(0.05 + (b as f32 + seed as f32) * 0.01).to_le_bytes(),
5009            );
5010            weights.extend_from_slice(
5011                &f16::from_f32(0.01 + (b as f32 + seed as f32) * 0.002).to_le_bytes(),
5012            );
5013            for i in 0..12u8 {
5014                weights.push(20 + i.wrapping_mul(3).wrapping_add(seed));
5015            }
5016            for i in 0..32u8 {
5017                weights.push(i.wrapping_mul(11).wrapping_add(b as u8).wrapping_add(seed));
5018            }
5019            for i in 0..128u8 {
5020                weights.push(i.wrapping_mul(19).wrapping_add(b as u8).wrapping_add(seed));
5021            }
5022        }
5023        weights
5024    }
5025
5026    fn synth_q6_k_row(n_blocks: usize, seed: u8) -> Vec<u8> {
5027        let mut weights = Vec::with_capacity(n_blocks * Q6_K_BLOCK_BYTES);
5028        for b in 0..n_blocks {
5029            for i in 0..128u8 {
5030                weights.push(i.wrapping_mul(17).wrapping_add(b as u8).wrapping_add(seed));
5031            }
5032            for i in 0..64u8 {
5033                weights.push(i.wrapping_mul(13).wrapping_add(seed).wrapping_add(b as u8));
5034            }
5035            for i in 0..16u8 {
5036                // signed scales in -32..31-ish
5037                weights.push((20i8).wrapping_add(i as i8).wrapping_add(seed as i8) as u8);
5038            }
5039            weights.extend_from_slice(
5040                &f16::from_f32(0.04 + (b as f32 + seed as f32) * 0.008).to_le_bytes(),
5041            );
5042        }
5043        weights
5044    }
5045
5046    fn synth_q4_k_row(n_blocks: usize, seed: u8) -> Vec<u8> {
5047        let mut weights = Vec::with_capacity(n_blocks * Q4_K_BLOCK_BYTES);
5048        for b in 0..n_blocks {
5049            weights.extend_from_slice(
5050                &f16::from_f32(0.05 + (b as f32 + seed as f32) * 0.01).to_le_bytes(),
5051            );
5052            weights.extend_from_slice(
5053                &f16::from_f32(0.01 + (b as f32 + seed as f32) * 0.002).to_le_bytes(),
5054            );
5055            for i in 0..12u8 {
5056                weights.push(20 + i.wrapping_mul(3).wrapping_add(seed));
5057            }
5058            for i in 0..128u8 {
5059                weights.push(i.wrapping_mul(17).wrapping_add(b as u8).wrapping_add(seed));
5060            }
5061        }
5062        weights
5063    }
5064
5065    fn synth_q4_0_row(n_blocks: usize, seed: u8) -> Vec<u8> {
5066        let mut weights = Vec::with_capacity(n_blocks * Q4_0_BLOCK_BYTES);
5067        for b in 0..n_blocks {
5068            weights.extend_from_slice(
5069                &f16::from_f32(0.05 + (b as f32 + seed as f32) * 0.012).to_le_bytes(),
5070            );
5071            for i in 0..16u8 {
5072                weights.push(i.wrapping_mul(23).wrapping_add(b as u8).wrapping_add(seed));
5073            }
5074        }
5075        weights
5076    }
5077
5078    fn synth_q8_0_row(n_blocks: usize, seed: u8) -> Vec<u8> {
5079        let mut weights = Vec::with_capacity(n_blocks * Q8_0_BLOCK_BYTES);
5080        for b in 0..n_blocks {
5081            weights.extend_from_slice(
5082                &f16::from_f32(0.04 + (b as f32 + seed as f32) * 0.008).to_le_bytes(),
5083            );
5084            for i in 0..32u8 {
5085                // signed i8 stored as u8 bytes
5086                let q = ((i as i8)
5087                    .wrapping_mul(3)
5088                    .wrapping_add(seed as i8)
5089                    .wrapping_add(b as i8)) as u8;
5090                weights.push(q);
5091            }
5092        }
5093        weights
5094    }
5095
5096    #[test]
5097    fn pack_and_gemv_matches_scalar_row_dots() {
5098        let n_blocks = 2;
5099        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5100        let rows = 16; // two full groups
5101        let mut matrix = Vec::new();
5102        for r in 0..rows {
5103            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, r as u8));
5104        }
5105        let x: Vec<f32> = (0..cols)
5106            .map(|i| ((i as f32) * 0.017 - 2.1).sin() * 1.8)
5107            .collect();
5108        let act = quantize_activations_q8_k(&x);
5109
5110        let mut reference = vec![0f32; rows];
5111        let row_bytes = n_blocks * Q4_K_BLOCK_BYTES;
5112        for r in 0..rows {
5113            reference[r] = dot_q4_k_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
5114        }
5115
5116        for &interleave in &[4usize, 8] {
5117            let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
5118            let n_groups = rows / Q4_KX8_NROWS;
5119            let mut out = vec![0f32; rows];
5120            gemv_q4_kx8_q8_k(&packed, &act, cols, n_groups, interleave, &mut out);
5121            for r in 0..rows {
5122                let err = (out[r] - reference[r]).abs();
5123                let scale = reference[r].abs().max(1.0);
5124                assert!(
5125                    err / scale < 1e-4 || err < 1e-3,
5126                    "interleave={interleave} row {r}: got {} want {} err={err}",
5127                    out[r],
5128                    reference[r]
5129                );
5130            }
5131        }
5132    }
5133
5134    #[test]
5135    fn q4_0x4_pack_and_gemv_matches_scalar_row_dots() {
5136        let n_blocks = 3;
5137        let cols = n_blocks * Q4_0_BLOCK_ELEMS;
5138        let rows = 12;
5139        let mut matrix = Vec::new();
5140        for r in 0..rows {
5141            matrix.extend_from_slice(&synth_q4_0_row(n_blocks, r as u8));
5142        }
5143        let x: Vec<f32> = (0..cols)
5144            .map(|i| ((i as f32) * 0.019 - 1.2).sin() * 2.1)
5145            .collect();
5146        let act = quantize_activations_q8(&x);
5147
5148        let row_bytes = n_blocks * Q4_0_BLOCK_BYTES;
5149        let mut reference = vec![0f32; rows];
5150        for r in 0..rows {
5151            reference[r] = dot_q4_0_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
5152        }
5153
5154        for &interleave in &[4usize, 8] {
5155            let packed = pack_q4_0_matrix_x4(&matrix, rows, cols, interleave);
5156            let n_groups = rows / Q4_0X4_NROWS;
5157            let mut out = vec![0f32; rows];
5158            gemv_q4_0x4_q8_0(&packed, &act, cols, n_groups, interleave, &mut out);
5159            for r in 0..rows {
5160                let err = (out[r] - reference[r]).abs();
5161                let scale = reference[r].abs().max(1.0);
5162                assert!(
5163                    err / scale < 1e-4 || err < 1e-3,
5164                    "Q4_0x4 interleave={interleave} row {r}: got {} want {} err={err}",
5165                    out[r],
5166                    reference[r]
5167                );
5168            }
5169        }
5170    }
5171
5172    #[test]
5173    fn q4_0x4_gemm_matches_the_gemv_run_once_per_activation() {
5174        let n_blocks = 4;
5175        let cols = n_blocks * Q4_0_BLOCK_ELEMS;
5176        let rows = 8;
5177        let mut matrix = Vec::new();
5178        for r in 0..rows {
5179            matrix.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 2 + 5) as u8));
5180        }
5181        let packed = pack_q4_0_matrix_x4(&matrix, rows, cols, Q4_0X4_INTERLEAVE);
5182
5183        let n_acts = 7;
5184        let acts: Vec<Q8Activations> = (0..n_acts)
5185            .map(|j| {
5186                let x: Vec<f32> = (0..cols)
5187                    .map(|i| (((i + j * 11) as f32) * 0.021 - 0.7).cos() * 1.9)
5188                    .collect();
5189                quantize_activations_q8(&x)
5190            })
5191            .collect();
5192
5193        for group in 0..rows / Q4_0X4_NROWS {
5194            let mut gemm_out = vec![0f32; Q4_0X4_NROWS * n_acts];
5195            gemm_q4_0x4_group(
5196                &packed,
5197                group,
5198                &acts,
5199                cols,
5200                Q4_0X4_INTERLEAVE,
5201                &mut gemm_out,
5202            );
5203
5204            for (j, act) in acts.iter().enumerate() {
5205                let mut gemv_out = [0f32; Q4_0X4_NROWS];
5206                gemv_q4_0x4_group(&packed, group, act, cols, Q4_0X4_INTERLEAVE, &mut gemv_out);
5207                for r in 0..Q4_0X4_NROWS {
5208                    assert_eq!(
5209                        gemm_out[r * n_acts + j],
5210                        gemv_out[r],
5211                        "group {group} row {r} act {j}: Q4_0 GEMM and GEMV disagree"
5212                    );
5213                }
5214            }
5215        }
5216    }
5217
5218    #[test]
5219    fn q4_0x4_gemm_with_no_activations_is_a_no_op() {
5220        let n_blocks = 2;
5221        let cols = n_blocks * Q4_0_BLOCK_ELEMS;
5222        let mut matrix = Vec::new();
5223        for r in 0..Q4_0X4_NROWS {
5224            matrix.extend_from_slice(&synth_q4_0_row(n_blocks, r as u8));
5225        }
5226        let packed = pack_q4_0_matrix_x4(&matrix, Q4_0X4_NROWS, cols, Q4_0X4_INTERLEAVE);
5227        let mut out: Vec<f32> = Vec::new();
5228        gemm_q4_0x4_group(&packed, 0, &[], cols, Q4_0X4_INTERLEAVE, &mut out);
5229        assert!(out.is_empty());
5230    }
5231
5232    #[test]
5233    fn q8_0x4_pack_and_gemv_matches_scalar_row_dots() {
5234        let n_blocks = 3;
5235        let cols = n_blocks * Q8_0_BLOCK_ELEMS;
5236        let rows = 12; // three full groups of 4
5237        let mut matrix = Vec::new();
5238        for r in 0..rows {
5239            matrix.extend_from_slice(&synth_q8_0_row(n_blocks, r as u8));
5240        }
5241        let x: Vec<f32> = (0..cols)
5242            .map(|i| ((i as f32) * 0.023 - 1.4).cos() * 2.2)
5243            .collect();
5244        let act = quantize_activations_q8(&x);
5245
5246        let row_bytes = n_blocks * Q8_0_BLOCK_BYTES;
5247        let mut reference = vec![0f32; rows];
5248        for r in 0..rows {
5249            reference[r] = dot_q8_0_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
5250        }
5251
5252        for &interleave in &[4usize, 8] {
5253            let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, interleave);
5254            let n_groups = rows / Q8_0X4_NROWS;
5255            let mut out = vec![0f32; rows];
5256            gemv_q8_0x4_q8_0(&packed, &act, cols, n_groups, interleave, &mut out);
5257            for r in 0..rows {
5258                let err = (out[r] - reference[r]).abs();
5259                let scale = reference[r].abs().max(1.0);
5260                assert!(
5261                    err / scale < 1e-4 || err < 1e-3,
5262                    "Q8_0x4 interleave={interleave} row {r}: got {} want {} err={err}",
5263                    out[r],
5264                    reference[r]
5265                );
5266            }
5267        }
5268    }
5269
5270    /// The GEMM exists purely to reuse weight loads across activations,
5271    /// so it must produce exactly what the per-activation GEMV produces
5272    /// -- not merely something close. Any divergence would be a
5273    /// batch-size-dependent numeric difference, i.e. prefill and decode
5274    /// disagreeing about the same prompt.
5275    #[test]
5276    fn q8_0x4_gemm_matches_the_gemv_run_once_per_activation() {
5277        let n_blocks = 4;
5278        let cols = n_blocks * Q8_0_BLOCK_ELEMS;
5279        let rows = 8;
5280        let mut matrix = Vec::new();
5281        for r in 0..rows {
5282            matrix.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 3 + 1) as u8));
5283        }
5284        let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, Q8_0X4_INTERLEAVE);
5285
5286        // Deliberately not a multiple of the tile width, so the tail
5287        // path is covered too.
5288        let n_acts = 7;
5289        let acts: Vec<Q8Activations> = (0..n_acts)
5290            .map(|j| {
5291                let x: Vec<f32> = (0..cols)
5292                    .map(|i| (((i + j * 13) as f32) * 0.017 - 0.9).sin() * 1.7)
5293                    .collect();
5294                quantize_activations_q8(&x)
5295            })
5296            .collect();
5297
5298        for group in 0..rows / Q8_0X4_NROWS {
5299            let mut gemm_out = vec![0f32; Q8_0X4_NROWS * n_acts];
5300            gemm_q8_0x4_group(
5301                &packed,
5302                group,
5303                &acts,
5304                cols,
5305                Q8_0X4_INTERLEAVE,
5306                &mut gemm_out,
5307            );
5308
5309            for (j, act) in acts.iter().enumerate() {
5310                let mut gemv_out = [0f32; Q8_0X4_NROWS];
5311                gemv_q8_0x4_group(&packed, group, act, cols, Q8_0X4_INTERLEAVE, &mut gemv_out);
5312                for r in 0..Q8_0X4_NROWS {
5313                    assert_eq!(
5314                        gemm_out[r * n_acts + j],
5315                        gemv_out[r],
5316                        "group {group} row {r} act {j}: GEMM and GEMV disagree"
5317                    );
5318                }
5319            }
5320        }
5321    }
5322
5323    /// The Q4_K GEMM must agree with the GEMV **exactly**, for the same
5324    /// reason as the Q8_0 pair above: the two run on the same prompt in
5325    /// different batch regimes (prefill vs the `< 4` tail vs decode), so
5326    /// any divergence is prefill and decode disagreeing about the same
5327    /// tokens. The GEMM only reorders which loop the weight unpack sits
5328    /// in — every multiply-accumulate happens in the same order and the
5329    /// same precision — so equality is the right assertion, not
5330    /// closeness.
5331    #[test]
5332    fn q4_kx8_gemm_matches_the_gemv_run_once_per_activation() {
5333        let n_blocks = 3;
5334        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5335        let rows = 2 * Q4_KX8_NROWS;
5336        let mut matrix = Vec::new();
5337        for r in 0..rows {
5338            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 5 + 3) as u8));
5339        }
5340        let interleave = q4_kx8_interleave();
5341        let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
5342
5343        // Not a multiple of the tile width, so the ragged tail the
5344        // caller has to chunk around is covered too.
5345        let n_acts = 6;
5346        let acts: Vec<Q8KActivations> = (0..n_acts)
5347            .map(|j| {
5348                let x: Vec<f32> = (0..cols)
5349                    .map(|i| (((i + j * 29) as f32) * 0.011 - 0.4).cos() * 2.3)
5350                    .collect();
5351                quantize_activations_q8_k(&x)
5352            })
5353            .collect();
5354
5355        for group in 0..rows / Q4_KX8_NROWS {
5356            for chunk in acts.chunks(Q4_KX8_GEMM_NC) {
5357                let mut gemm_out = vec![0f32; Q4_KX8_NROWS * chunk.len()];
5358                gemm_q4_kx8_group(&packed, group, chunk, cols, interleave, &mut gemm_out);
5359
5360                for (j, act) in chunk.iter().enumerate() {
5361                    let mut gemv_out = [0f32; Q4_KX8_NROWS];
5362                    gemv_q4_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
5363                    for r in 0..Q4_KX8_NROWS {
5364                        let got = gemm_out[r * chunk.len() + j];
5365                        let want = gemv_out[r];
5366                        if interleave == 4 {
5367                            assert_eq!(
5368                                got, want,
5369                                "group {group} row {r} act {j}: Q4_K GEMM and GEMV disagree"
5370                            );
5371                        } else {
5372                            let err = (got - want).abs();
5373                            let scale = want.abs().max(1.0);
5374                            assert!(
5375                                err / scale < 1e-4 || err < 1e-2,
5376                                "group {group} row {r} act {j}: GEMM {got} vs GEMV {want} (err={err})"
5377                            );
5378                        }
5379                    }
5380                }
5381            }
5382        }
5383    }
5384
5385    #[test]
5386    #[cfg(target_arch = "aarch64")]
5387    fn q4_kx8_gemm_i8mm_matches_scalar_when_available() {
5388        if !std::arch::is_aarch64_feature_detected!("i8mm") {
5389            return;
5390        }
5391        let interleave = q4_kx8_interleave();
5392        assert_eq!(interleave, 8, "i8mm host should pack with interleave 8");
5393
5394        let n_blocks = 3;
5395        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5396        let rows = Q4_KX8_NROWS;
5397        let mut matrix = Vec::new();
5398        for r in 0..rows {
5399            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 7 + 2) as u8));
5400        }
5401        let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
5402
5403        let n_acts = 4;
5404        let acts: Vec<Q8KActivations> = (0..n_acts)
5405            .map(|j| {
5406                let x: Vec<f32> = (0..cols)
5407                    .map(|i| (((i + j * 17) as f32) * 0.013 - 0.6).sin() * 1.9)
5408                    .collect();
5409                quantize_activations_q8_k(&x)
5410            })
5411            .collect();
5412
5413        let mut gemm_out = vec![0f32; Q4_KX8_NROWS * n_acts];
5414        gemm_q4_kx8_group(&packed, 0, &acts, cols, interleave, &mut gemm_out);
5415
5416        for (j, act) in acts.iter().enumerate() {
5417            let mut scalar_out = [0f32; Q4_KX8_NROWS];
5418            gemv_q4_kx8_group(&packed, 0, act, cols, interleave, &mut scalar_out);
5419            for r in 0..Q4_KX8_NROWS {
5420                let got = gemm_out[r * n_acts + j];
5421                let want = scalar_out[r];
5422                let err = (got - want).abs();
5423                let scale = want.abs().max(1.0);
5424                assert!(
5425                    err / scale < 1e-5 || err < 1e-3,
5426                    "row {r} act {j}: i8mm GEMM {got} vs scalar {want} (err={err})"
5427                );
5428            }
5429        }
5430    }
5431
5432    fn synth_q8_k_acts(n: usize, cols: usize) -> Vec<Q8KActivations> {
5433        (0..n)
5434            .map(|j| {
5435                let x: Vec<f32> = (0..cols)
5436                    .map(|i| (((i + j * 17) as f32) * 0.013 - 0.6).sin() * 1.9)
5437                    .collect();
5438                quantize_activations_q8_k(&x)
5439            })
5440            .collect()
5441    }
5442
5443    /// The retired per-block interleave (`pack_q8_k_qs_x4_i8`), kept verbatim
5444    /// as the reference `prepare_q8_k_acts_x4` must reproduce: llama.cpp's
5445    /// `ggml_quantize_mat_q8_K_4x8` qs ordering with 8-byte runs.
5446    fn reference_q8_kx4_block_qs(
5447        acts: &[Q8KActivations],
5448        block: usize,
5449    ) -> [i8; Q4_K_BLOCK_ELEMS * 4] {
5450        const BLCK: usize = 8;
5451        let na = acts.len();
5452        let mut out = [0i8; Q4_K_BLOCK_ELEMS * 4];
5453        for (j, slot) in out.iter_mut().enumerate() {
5454            let src_offset = (j / (4 * BLCK)) * BLCK + (j % BLCK);
5455            let src_id = (j % (4 * BLCK)) / BLCK;
5456            *slot = if src_id < na {
5457                acts[src_id].q[block * Q4_K_BLOCK_ELEMS + src_offset]
5458            } else {
5459                0
5460            };
5461        }
5462        out
5463    }
5464
5465    #[test]
5466    fn prepare_q8_k_acts_x4_matches_block_interleave_reference() {
5467        let n_blocks = 3;
5468        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5469        for na in 1..=Q4_KX8_GEMM_NC {
5470            let acts = synth_q8_k_acts(na, cols);
5471            let tile = prepare_q8_k_acts_x4(&acts, cols);
5472            assert_eq!(tile.na, na);
5473            assert_eq!(tile.n_blocks, n_blocks);
5474            for b in 0..n_blocks {
5475                let want_qs = reference_q8_kx4_block_qs(&acts, b);
5476                assert_eq!(
5477                    &tile.qs[b * Q4_K_BLOCK_ELEMS * 4..][..Q4_K_BLOCK_ELEMS * 4],
5478                    &want_qs[..],
5479                    "qs mismatch, block {b} na {na}"
5480                );
5481                for a in 0..4 {
5482                    let act = acts.get(a);
5483                    for i in 0..8 {
5484                        let want = act.map_or(0, |act| {
5485                            act.bsums[b * 16 + 2 * i] + act.bsums[b * 16 + 2 * i + 1]
5486                        });
5487                        assert_eq!(
5488                            tile.bsums[(b * 4 + a) * 8 + i],
5489                            want,
5490                            "bsums mismatch, block {b} row {a} pair {i} na {na}"
5491                        );
5492                    }
5493                    let want_d = act.map_or(0.0, |act| act.d[b]);
5494                    assert_eq!(tile.d[b * 4 + a], want_d, "d mismatch, block {b} row {a}");
5495                }
5496            }
5497        }
5498    }
5499
5500    #[test]
5501    fn q4_kx8_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
5502        let n_blocks = 3;
5503        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5504        let rows = Q4_KX8_NROWS;
5505        let mut matrix = Vec::new();
5506        for r in 0..rows {
5507            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 5 + 3) as u8));
5508        }
5509        let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, 8);
5510
5511        for na in 1..=Q4_KX8_GEMM_NC {
5512            let acts = synth_q8_k_acts(na, cols);
5513            let tile = prepare_q8_k_acts_x4(&acts, cols);
5514            let mut got = vec![0f32; Q4_KX8_NROWS * na];
5515            gemm_q4_kx8_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
5516
5517            for (j, act) in acts.iter().enumerate() {
5518                let mut want = [0f32; Q4_KX8_NROWS];
5519                gemv_q4_kx8_q8_k_scalar_8(&packed, act, cols, 1, &mut want);
5520                for r in 0..Q4_KX8_NROWS {
5521                    assert_eq!(
5522                        got[r * na + j].to_bits(),
5523                        want[r].to_bits(),
5524                        "row {r} act {j} na {na}: x4 {} vs GEMV {}",
5525                        got[r * na + j],
5526                        want[r]
5527                    );
5528                }
5529            }
5530        }
5531    }
5532
5533    /// The hoisted path must reproduce the in-kernel-interleave behavior it
5534    /// replaced: `gemm_q4_kx8_group` (which now prepares the quad per call)
5535    /// and `gemm_q4_kx8_group_x4` (quad prepared by the caller) share the
5536    /// i8mm kernel, so their outputs must be bit-identical, and both must
5537    /// match the scalar GEMV within the usual tolerance.
5538    #[test]
5539    #[cfg(target_arch = "aarch64")]
5540    fn q4_kx8_gemm_x4_i8mm_matches_group_and_scalar() {
5541        if !std::arch::is_aarch64_feature_detected!("i8mm") {
5542            return;
5543        }
5544        let interleave = q4_kx8_interleave();
5545        assert_eq!(interleave, 8, "i8mm host should pack with interleave 8");
5546
5547        let n_blocks = 3;
5548        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5549        let rows = Q4_KX8_NROWS;
5550        let mut matrix = Vec::new();
5551        for r in 0..rows {
5552            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 7 + 2) as u8));
5553        }
5554        let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, interleave);
5555
5556        assert!(q4_kx8_gemm_uses_acts_x4(interleave));
5557        for na in 1..=Q4_KX8_GEMM_NC {
5558            let acts = synth_q8_k_acts(na, cols);
5559            let tile = prepare_q8_k_acts_x4(&acts, cols);
5560
5561            let mut x4_out = vec![0f32; Q4_KX8_NROWS * na];
5562            gemm_q4_kx8_group_x4(&packed, 0, &tile, cols, interleave, &mut x4_out);
5563
5564            let mut group_out = vec![0f32; Q4_KX8_NROWS * na];
5565            gemm_q4_kx8_group(&packed, 0, &acts, cols, interleave, &mut group_out);
5566            assert_eq!(
5567                x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5568                group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5569                "x4 entry diverged from the compat entry, na {na}"
5570            );
5571
5572            for (j, act) in acts.iter().enumerate() {
5573                let mut scalar_out = [0f32; Q4_KX8_NROWS];
5574                gemv_q4_kx8_group(&packed, 0, act, cols, interleave, &mut scalar_out);
5575                for r in 0..Q4_KX8_NROWS {
5576                    let got = x4_out[r * na + j];
5577                    let want = scalar_out[r];
5578                    let err = (got - want).abs();
5579                    let scale = want.abs().max(1.0);
5580                    assert!(
5581                        err / scale < 1e-5 || err < 1e-3,
5582                        "row {r} act {j} na {na}: i8mm x4 GEMM {got} vs scalar {want} (err={err})"
5583                    );
5584                }
5585            }
5586        }
5587    }
5588
5589    /// `AccelX4` only chooses a kernel; it must never change an answer.
5590    ///
5591    /// Two claims, for all five `×4` GEMMs at once. First, passing the
5592    /// host's own [`AccelX4::detect`] to the `_on` entry point is
5593    /// bit-identical to letting the probing wrapper detect per call -- that
5594    /// is what lets `apply_batch` hoist the probe out of a 10^5-iteration
5595    /// loop. Second, [`AccelX4::Portable`] really does select the portable
5596    /// kernel even on an i8mm host, so the scalar reference stays reachable
5597    /// and the `_portable_is_bit_exact_vs_scalar_gemv` tests keep meaning
5598    /// something on this machine.
5599    #[test]
5600    fn accel_x4_only_picks_a_kernel_it_never_changes_the_answer() {
5601        let na = 3;
5602        let here = AccelX4::detect();
5603
5604        // Q4_K / Q5_K / Q6_K share the Q8_K activation quad.
5605        {
5606            let n_blocks = 3;
5607            let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5608            let acts = synth_q8_k_acts(na, cols);
5609            let tile = prepare_q8_k_acts_x4(&acts, cols);
5610
5611            let mut q4 = Vec::new();
5612            let mut q5 = Vec::new();
5613            let mut q6 = Vec::new();
5614            for r in 0..Q4_KX8_NROWS {
5615                q4.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 7 + 2) as u8));
5616                q5.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 3 + 5) as u8));
5617                q6.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 11 + 1) as u8));
5618            }
5619            let q4 = pack_q4_k_matrix_x8(&q4, Q4_KX8_NROWS, cols, 8);
5620            let q5 = pack_q5_k_matrix_x8(&q5, Q5_KX8_NROWS, cols, 8);
5621            let q6 = pack_q6_k_matrix_x8(&q6, Q6_KX8_NROWS, cols, 8);
5622
5623            let mut wrapper = vec![0f32; Q4_KX8_NROWS * na];
5624            let mut hoisted = vec![0f32; Q4_KX8_NROWS * na];
5625            let mut portable = vec![0f32; Q4_KX8_NROWS * na];
5626            let mut reference = vec![0f32; Q4_KX8_NROWS * na];
5627
5628            gemm_q4_kx8_group_x4(&q4, 0, &tile, cols, 8, &mut wrapper);
5629            gemm_q4_kx8_group_x4_on(&q4, 0, &tile, cols, 8, here, &mut hoisted);
5630            gemm_q4_kx8_group_x4_on(&q4, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
5631            gemm_q4_kx8_acts_x4_scalar_8(&q4, &tile, cols, &mut reference);
5632            assert_bits_eq("q4_k hoisted", &hoisted, &wrapper);
5633            assert_bits_eq("q4_k portable", &portable, &reference);
5634
5635            gemm_q5_kx8_group_x4(&q5, 0, &tile, cols, 8, &mut wrapper);
5636            gemm_q5_kx8_group_x4_on(&q5, 0, &tile, cols, 8, here, &mut hoisted);
5637            gemm_q5_kx8_group_x4_on(&q5, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
5638            gemm_q5_kx8_acts_x4_scalar_8(&q5, &tile, cols, &mut reference);
5639            assert_bits_eq("q5_k hoisted", &hoisted, &wrapper);
5640            assert_bits_eq("q5_k portable", &portable, &reference);
5641
5642            gemm_q6_kx8_group_x4(&q6, 0, &tile, cols, 8, &mut wrapper);
5643            gemm_q6_kx8_group_x4_on(&q6, 0, &tile, cols, 8, here, &mut hoisted);
5644            gemm_q6_kx8_group_x4_on(&q6, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
5645            gemm_q6_kx8_acts_x4_scalar_8(&q6, &tile, cols, &mut reference);
5646            assert_bits_eq("q6_k hoisted", &hoisted, &wrapper);
5647            assert_bits_eq("q6_k portable", &portable, &reference);
5648        }
5649
5650        // Q8_0 / Q4_0 share the Q8_0 activation quad.
5651        {
5652            let n_blocks = 4;
5653            let cols = n_blocks * Q8_0_BLOCK_ELEMS;
5654            let acts = synth_q8_0_acts(na, cols);
5655            let tile = prepare_q8_acts_x4(&acts, cols);
5656
5657            let mut q8 = Vec::new();
5658            let mut q4 = Vec::new();
5659            for r in 0..Q8_0X4_NROWS {
5660                q8.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 9 + 4) as u8));
5661                q4.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 13 + 6) as u8));
5662            }
5663            let q8 = pack_q8_0_matrix_x4(&q8, Q8_0X4_NROWS, cols, 8);
5664            let q4 = pack_q4_0_matrix_x4(&q4, Q4_0X4_NROWS, cols, 8);
5665
5666            let mut wrapper = vec![0f32; Q8_0X4_NROWS * na];
5667            let mut hoisted = vec![0f32; Q8_0X4_NROWS * na];
5668            let mut portable = vec![0f32; Q8_0X4_NROWS * na];
5669            let mut reference = vec![0f32; Q8_0X4_NROWS * na];
5670
5671            gemm_q8_0x4_group_x4(&q8, 0, &tile, cols, 8, &mut wrapper);
5672            gemm_q8_0x4_group_x4_on(&q8, 0, &tile, cols, 8, here, &mut hoisted);
5673            gemm_q8_0x4_group_x4_on(&q8, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
5674            gemm_q8_0x4_acts_x4_scalar_8(&q8, &tile, cols, &mut reference);
5675            assert_bits_eq("q8_0 hoisted", &hoisted, &wrapper);
5676            assert_bits_eq("q8_0 portable", &portable, &reference);
5677
5678            gemm_q4_0x4_group_x4(&q4, 0, &tile, cols, 8, &mut wrapper);
5679            gemm_q4_0x4_group_x4_on(&q4, 0, &tile, cols, 8, here, &mut hoisted);
5680            gemm_q4_0x4_group_x4_on(&q4, 0, &tile, cols, 8, AccelX4::Portable, &mut portable);
5681            gemm_q4_0x4_acts_x4_scalar_8(&q4, &tile, cols, &mut reference);
5682            assert_bits_eq("q4_0 hoisted", &hoisted, &wrapper);
5683            assert_bits_eq("q4_0 portable", &portable, &reference);
5684        }
5685    }
5686
5687    fn assert_bits_eq(what: &str, got: &[f32], want: &[f32]) {
5688        assert_eq!(
5689            got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5690            want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5691            "{what}: {got:?} vs {want:?}"
5692        );
5693    }
5694
5695    #[test]
5696    fn q5_kx8_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
5697        let n_blocks = 3;
5698        let cols = n_blocks * Q5_K_BLOCK_ELEMS;
5699        let rows = Q5_KX8_NROWS;
5700        let mut matrix = Vec::new();
5701        for r in 0..rows {
5702            matrix.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 5 + 3) as u8));
5703        }
5704        let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, 8);
5705
5706        for na in 1..=Q5_KX8_GEMM_NC {
5707            let acts = synth_q8_k_acts(na, cols);
5708            let tile = prepare_q8_k_acts_x4(&acts, cols);
5709            let mut got = vec![0f32; Q5_KX8_NROWS * na];
5710            gemm_q5_kx8_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
5711
5712            for (j, act) in acts.iter().enumerate() {
5713                let mut want = [0f32; Q5_KX8_NROWS];
5714                gemv_q5_kx8_q8_k_scalar_8(&packed, act, cols, 1, &mut want);
5715                for r in 0..Q5_KX8_NROWS {
5716                    assert_eq!(
5717                        got[r * na + j].to_bits(),
5718                        want[r].to_bits(),
5719                        "row {r} act {j} na {na}: x4 {} vs GEMV {}",
5720                        got[r * na + j],
5721                        want[r]
5722                    );
5723                }
5724            }
5725        }
5726    }
5727
5728    #[test]
5729    fn q6_kx8_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
5730        let n_blocks = 2;
5731        let cols = n_blocks * Q6_K_BLOCK_ELEMS;
5732        let rows = Q6_KX8_NROWS;
5733        let mut matrix = Vec::new();
5734        for r in 0..rows {
5735            matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 3 + 2) as u8));
5736        }
5737        let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, 8);
5738
5739        for na in 1..=Q8K_ACTS_X4_NC {
5740            let acts = synth_q8_k_acts(na, cols);
5741            let tile = prepare_q8_k_acts_x4(&acts, cols);
5742            let mut got = vec![0f32; Q6_KX8_NROWS * na];
5743            gemm_q6_kx8_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
5744
5745            for (j, act) in acts.iter().enumerate() {
5746                let mut want = [0f32; Q6_KX8_NROWS];
5747                gemv_q6_kx8_q8_k_scalar(&packed, act, cols, 1, 8, &mut want);
5748                for r in 0..Q6_KX8_NROWS {
5749                    assert_eq!(
5750                        got[r * na + j].to_bits(),
5751                        want[r].to_bits(),
5752                        "row {r} act {j} na {na}: x4 {} vs GEMV {}",
5753                        got[r * na + j],
5754                        want[r]
5755                    );
5756                }
5757            }
5758        }
5759    }
5760
5761    /// The interleave-8 DotProd GEMV against the scalar interleave-8
5762    /// reference (the i8mm GEMM side of Q4_K is covered above).
5763    #[test]
5764    #[cfg(target_arch = "aarch64")]
5765    fn q4_kx8_interleave8_neon_gemv_matches_scalar_when_available() {
5766        if !std::arch::is_aarch64_feature_detected!("dotprod") {
5767            return;
5768        }
5769        let n_blocks = 3;
5770        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
5771        let n_groups = 2;
5772        let rows = n_groups * Q4_KX8_NROWS;
5773        let mut matrix = Vec::new();
5774        for r in 0..rows {
5775            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, (r * 11 + 5) as u8));
5776        }
5777        let packed = pack_q4_k_matrix_x8(&matrix, rows, cols, 8);
5778
5779        let acts = synth_q8_k_acts(4, cols);
5780        for act in &acts {
5781            let mut got = vec![0f32; rows];
5782            gemv_q4_kx8_q8_k(&packed, act, cols, n_groups, 8, &mut got);
5783            let mut want = vec![0f32; rows];
5784            gemv_q4_kx8_q8_k_scalar_8(&packed, act, cols, n_groups, &mut want);
5785            for r in 0..rows {
5786                let err = (got[r] - want[r]).abs();
5787                // Tolerance, not bit equality: the int8 dots are exact in
5788                // i32, but the two paths fold the per-sub-block f32
5789                // scales in a different order, so the f32 accumulation
5790                // rounds differently. Measured on an M2 Pro (the first
5791                // i8mm host this ever ran on): 63 of 64 outputs agree to
5792                // <= 6.3e-6 relative, one to 2.6e-5. 5e-5 keeps a margin
5793                // over that while still being orders of magnitude
5794                // tighter than any indexing error, which moves a result
5795                // by the size of the data. Cross-checked end to end:
5796                // greedy CPU generation with these kernels is
5797                // token-identical to `FERROX_CPU_INT_DOT=0`.
5798                assert!(
5799                    err / want[r].abs().max(1.0) < 5e-5 || err < 1e-3,
5800                    "gemv row {r}: NEON 8x8 {} vs scalar {}",
5801                    got[r],
5802                    want[r]
5803                );
5804            }
5805        }
5806    }
5807
5808    /// The interleave-8 NEON paths (DotProd 8x8 GEMV, i8mm GEMM) against
5809    /// the scalar interleave-8 references, plus the x4 entry against the
5810    /// compat entry (bit-exact -- they share the kernel).
5811    #[test]
5812    #[cfg(target_arch = "aarch64")]
5813    fn q5_kx8_interleave8_neon_matches_references_when_available() {
5814        if !std::arch::is_aarch64_feature_detected!("dotprod") {
5815            return;
5816        }
5817        let n_blocks = 3;
5818        let cols = n_blocks * Q5_K_BLOCK_ELEMS;
5819        let n_groups = 2;
5820        let rows = n_groups * Q5_KX8_NROWS;
5821        let mut matrix = Vec::new();
5822        for r in 0..rows {
5823            matrix.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 7 + 1) as u8));
5824        }
5825        let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, 8);
5826
5827        let acts = synth_q8_k_acts(4, cols);
5828        for act in &acts {
5829            let mut got = vec![0f32; rows];
5830            gemv_q5_kx8_q8_k(&packed, act, cols, n_groups, 8, &mut got);
5831            let mut want = vec![0f32; rows];
5832            gemv_q5_kx8_q8_k_scalar_8(&packed, act, cols, n_groups, &mut want);
5833            for r in 0..rows {
5834                let err = (got[r] - want[r]).abs();
5835                // Tolerance, not bit equality: the int8 dots are exact in
5836                // i32, but the two paths fold the per-sub-block f32
5837                // scales in a different order, so the f32 accumulation
5838                // rounds differently. Measured on an M2 Pro (the first
5839                // i8mm host this ever ran on): 63 of 64 outputs agree to
5840                // <= 6.3e-6 relative, one to 2.6e-5. 5e-5 keeps a margin
5841                // over that while still being orders of magnitude
5842                // tighter than any indexing error, which moves a result
5843                // by the size of the data. Cross-checked end to end:
5844                // greedy CPU generation with these kernels is
5845                // token-identical to `FERROX_CPU_INT_DOT=0`.
5846                assert!(
5847                    err / want[r].abs().max(1.0) < 5e-5 || err < 1e-3,
5848                    "gemv row {r}: NEON 8x8 {} vs scalar {}",
5849                    got[r],
5850                    want[r]
5851                );
5852            }
5853        }
5854
5855        if !std::arch::is_aarch64_feature_detected!("i8mm") {
5856            return;
5857        }
5858        for na in 1..=Q5_KX8_GEMM_NC {
5859            let acts = synth_q8_k_acts(na, cols);
5860            let tile = prepare_q8_k_acts_x4(&acts, cols);
5861            for group in 0..n_groups {
5862                let mut x4_out = vec![0f32; Q5_KX8_NROWS * na];
5863                gemm_q5_kx8_group_x4(&packed, group, &tile, cols, 8, &mut x4_out);
5864
5865                let mut group_out = vec![0f32; Q5_KX8_NROWS * na];
5866                gemm_q5_kx8_group(&packed, group, &acts, cols, 8, &mut group_out);
5867                assert_eq!(
5868                    x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5869                    group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5870                    "x4 entry diverged from the compat entry, group {group} na {na}"
5871                );
5872
5873                let nb = cols / Q5_K_BLOCK_ELEMS;
5874                let slice = &packed[group * nb * Q5_KX8_BLOCK_BYTES..][..nb * Q5_KX8_BLOCK_BYTES];
5875                let mut want = vec![0f32; Q5_KX8_NROWS * na];
5876                gemm_q5_kx8_acts_x4_scalar_8(slice, &tile, cols, &mut want);
5877                for (got, want) in x4_out.iter().zip(want.iter()) {
5878                    let err = (got - want).abs();
5879                    // Same 5e-5 as the GEMV check above, for the same
5880                    // reason: i8mm folds the f32 scales in a different
5881                    // order than the portable path, and the worst
5882                    // deviation measured on an i8mm host is 2.6e-5.
5883                    assert!(
5884                        err / want.abs().max(1.0) < 5e-5 || err < 1e-3,
5885                        "group {group} na {na}: i8mm GEMM {got} vs portable {want}"
5886                    );
5887                }
5888            }
5889        }
5890    }
5891
5892    /// Q6_K twin of the test above.
5893    #[test]
5894    #[cfg(target_arch = "aarch64")]
5895    fn q6_kx8_interleave8_neon_matches_references_when_available() {
5896        if !std::arch::is_aarch64_feature_detected!("dotprod") {
5897            return;
5898        }
5899        let n_blocks = 2;
5900        let cols = n_blocks * Q6_K_BLOCK_ELEMS;
5901        let n_groups = 2;
5902        let rows = n_groups * Q6_KX8_NROWS;
5903        let mut matrix = Vec::new();
5904        for r in 0..rows {
5905            matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 9 + 4) as u8));
5906        }
5907        let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, 8);
5908
5909        let acts = synth_q8_k_acts(4, cols);
5910        for act in &acts {
5911            let mut got = vec![0f32; rows];
5912            gemv_q6_kx8_q8_k(&packed, act, cols, n_groups, 8, &mut got);
5913            let mut want = vec![0f32; rows];
5914            gemv_q6_kx8_q8_k_scalar(&packed, act, cols, n_groups, 8, &mut want);
5915            for r in 0..rows {
5916                let err = (got[r] - want[r]).abs();
5917                // Tolerance, not bit equality: the int8 dots are exact in
5918                // i32, but the two paths fold the per-sub-block f32
5919                // scales in a different order, so the f32 accumulation
5920                // rounds differently. Measured on an M2 Pro (the first
5921                // i8mm host this ever ran on): 63 of 64 outputs agree to
5922                // <= 6.3e-6 relative, one to 2.6e-5. 5e-5 keeps a margin
5923                // over that while still being orders of magnitude
5924                // tighter than any indexing error, which moves a result
5925                // by the size of the data. Cross-checked end to end:
5926                // greedy CPU generation with these kernels is
5927                // token-identical to `FERROX_CPU_INT_DOT=0`.
5928                assert!(
5929                    err / want[r].abs().max(1.0) < 5e-5 || err < 1e-3,
5930                    "gemv row {r}: NEON 8x8 {} vs scalar {}",
5931                    got[r],
5932                    want[r]
5933                );
5934            }
5935        }
5936
5937        if !std::arch::is_aarch64_feature_detected!("i8mm") {
5938            return;
5939        }
5940        for na in 1..=Q8K_ACTS_X4_NC {
5941            let acts = synth_q8_k_acts(na, cols);
5942            let tile = prepare_q8_k_acts_x4(&acts, cols);
5943            for group in 0..n_groups {
5944                let mut x4_out = vec![0f32; Q6_KX8_NROWS * na];
5945                gemm_q6_kx8_group_x4(&packed, group, &tile, cols, 8, &mut x4_out);
5946
5947                let mut group_out = vec![0f32; Q6_KX8_NROWS * na];
5948                gemm_q6_kx8_group(&packed, group, &acts, cols, 8, &mut group_out);
5949                assert_eq!(
5950                    x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5951                    group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
5952                    "x4 entry diverged from the compat entry, group {group} na {na}"
5953                );
5954
5955                let nb = cols / Q6_K_BLOCK_ELEMS;
5956                let slice = &packed[group * nb * Q6_KX8_BLOCK_BYTES..][..nb * Q6_KX8_BLOCK_BYTES];
5957                let mut want = vec![0f32; Q6_KX8_NROWS * na];
5958                gemm_q6_kx8_acts_x4_scalar_8(slice, &tile, cols, &mut want);
5959                for (got, want) in x4_out.iter().zip(want.iter()) {
5960                    let err = (got - want).abs();
5961                    // Same 5e-5 as the GEMV check above, for the same
5962                    // reason: i8mm folds the f32 scales in a different
5963                    // order than the portable path, and the worst
5964                    // deviation measured on an i8mm host is 2.6e-5.
5965                    assert!(
5966                        err / want.abs().max(1.0) < 5e-5 || err < 1e-3,
5967                        "group {group} na {na}: i8mm GEMM {got} vs portable {want}"
5968                    );
5969                }
5970            }
5971        }
5972    }
5973
5974    fn synth_q8_0_acts(n: usize, cols: usize) -> Vec<Q8Activations> {
5975        (0..n)
5976            .map(|j| {
5977                let x: Vec<f32> = (0..cols)
5978                    .map(|i| (((i + j * 13) as f32) * 0.021 - 0.9).sin() * 1.7)
5979                    .collect();
5980                quantize_activations_q8(&x)
5981            })
5982            .collect()
5983    }
5984
5985    /// llama.cpp `ggml_quantize_mat_q8_0_4x8`'s qs ordering, as the
5986    /// reference `prepare_q8_acts_x4` must reproduce.
5987    #[test]
5988    fn prepare_q8_acts_x4_matches_interleave_reference() {
5989        let n_blocks = 3;
5990        let cols = n_blocks * Q8_0_BLOCK_ELEMS;
5991        for na in 1..=Q8K_ACTS_X4_NC {
5992            let acts = synth_q8_0_acts(na, cols);
5993            let tile = prepare_q8_acts_x4(&acts, cols);
5994            assert_eq!(tile.na, na);
5995            assert_eq!(tile.n_blocks, n_blocks);
5996            for b in 0..n_blocks {
5997                for (j, got) in tile.qs[b * 128..(b + 1) * 128].iter().enumerate() {
5998                    let src_offset = (j / 32) * 8 + (j % 8);
5999                    let src_id = (j % 32) / 8;
6000                    let want = if src_id < na {
6001                        acts[src_id].q[b * Q8_0_BLOCK_ELEMS + src_offset]
6002                    } else {
6003                        0
6004                    };
6005                    assert_eq!(*got, want, "qs mismatch, block {b} pos {j} na {na}");
6006                }
6007                for a in 0..4 {
6008                    let want_d = acts.get(a).map_or(0.0, |act| act.d[b]);
6009                    assert_eq!(tile.d[b * 4 + a], want_d, "d mismatch, block {b} row {a}");
6010                }
6011            }
6012        }
6013    }
6014
6015    #[test]
6016    fn q8_0x4_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
6017        let n_blocks = 3;
6018        let cols = n_blocks * Q8_0_BLOCK_ELEMS;
6019        let rows = Q8_0X4_NROWS;
6020        let mut matrix = Vec::new();
6021        for r in 0..rows {
6022            matrix.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 7 + 3) as u8));
6023        }
6024        let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, 8);
6025
6026        for na in 1..=Q8K_ACTS_X4_NC {
6027            let acts = synth_q8_0_acts(na, cols);
6028            let tile = prepare_q8_acts_x4(&acts, cols);
6029            let mut got = vec![0f32; Q8_0X4_NROWS * na];
6030            gemm_q8_0x4_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
6031
6032            for (j, act) in acts.iter().enumerate() {
6033                let mut want = [0f32; Q8_0X4_NROWS];
6034                gemv_q8_0x4_q8_0_scalar(&packed, act, cols, 1, 8, &mut want);
6035                for r in 0..Q8_0X4_NROWS {
6036                    assert_eq!(
6037                        got[r * na + j].to_bits(),
6038                        want[r].to_bits(),
6039                        "row {r} act {j} na {na}: x4 {} vs GEMV {}",
6040                        got[r * na + j],
6041                        want[r]
6042                    );
6043                }
6044            }
6045        }
6046    }
6047
6048    #[test]
6049    fn q4_0x4_gemm_x4_portable_is_bit_exact_vs_scalar_gemv() {
6050        let n_blocks = 3;
6051        let cols = n_blocks * Q4_0_BLOCK_ELEMS;
6052        let rows = Q4_0X4_NROWS;
6053        let mut matrix = Vec::new();
6054        for r in 0..rows {
6055            matrix.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 5 + 1) as u8));
6056        }
6057        let packed = pack_q4_0_matrix_x4(&matrix, rows, cols, 8);
6058
6059        for na in 1..=Q8K_ACTS_X4_NC {
6060            let acts = synth_q8_0_acts(na, cols);
6061            let tile = prepare_q8_acts_x4(&acts, cols);
6062            let mut got = vec![0f32; Q4_0X4_NROWS * na];
6063            gemm_q4_0x4_acts_x4_scalar_8(&packed, &tile, cols, &mut got);
6064
6065            for (j, act) in acts.iter().enumerate() {
6066                let mut want = [0f32; Q4_0X4_NROWS];
6067                gemv_q4_0x4_q8_0_scalar(&packed, act, cols, 1, 8, &mut want);
6068                for r in 0..Q4_0X4_NROWS {
6069                    assert_eq!(
6070                        got[r * na + j].to_bits(),
6071                        want[r].to_bits(),
6072                        "row {r} act {j} na {na}: x4 {} vs GEMV {}",
6073                        got[r * na + j],
6074                        want[r]
6075                    );
6076                }
6077            }
6078        }
6079    }
6080
6081    /// The interleave-8 NEON paths (DotProd `4x8` GEMV, i8mm GEMM) against
6082    /// the scalar interleave-8 references, plus the x4 entries against the
6083    /// compat entries (bit-exact -- they share the kernels).
6084    #[test]
6085    #[cfg(target_arch = "aarch64")]
6086    fn q8_0_q4_0_interleave8_neon_matches_references_when_available() {
6087        if !std::arch::is_aarch64_feature_detected!("dotprod") {
6088            return;
6089        }
6090        let n_blocks = 3;
6091        let n_groups = 2;
6092
6093        // Q8_0
6094        let cols = n_blocks * Q8_0_BLOCK_ELEMS;
6095        let rows = n_groups * Q8_0X4_NROWS;
6096        let mut matrix = Vec::new();
6097        for r in 0..rows {
6098            matrix.extend_from_slice(&synth_q8_0_row(n_blocks, (r * 3 + 2) as u8));
6099        }
6100        let packed = pack_q8_0_matrix_x4(&matrix, rows, cols, 8);
6101        let acts = synth_q8_0_acts(4, cols);
6102        for act in &acts {
6103            let mut got = vec![0f32; rows];
6104            gemv_q8_0x4_q8_0(&packed, act, cols, n_groups, 8, &mut got);
6105            let mut want = vec![0f32; rows];
6106            gemv_q8_0x4_q8_0_scalar(&packed, act, cols, n_groups, 8, &mut want);
6107            for r in 0..rows {
6108                let err = (got[r] - want[r]).abs();
6109                assert!(
6110                    err / want[r].abs().max(1.0) < 1e-5 || err < 1e-3,
6111                    "q8_0 gemv row {r}: NEON 4x8 {} vs scalar {}",
6112                    got[r],
6113                    want[r]
6114                );
6115            }
6116        }
6117        // Q4_0
6118        let q4_cols = n_blocks * Q4_0_BLOCK_ELEMS;
6119        let q4_rows = n_groups * Q4_0X4_NROWS;
6120        let mut q4_matrix = Vec::new();
6121        for r in 0..q4_rows {
6122            q4_matrix.extend_from_slice(&synth_q4_0_row(n_blocks, (r * 9 + 1) as u8));
6123        }
6124        let q4_packed = pack_q4_0_matrix_x4(&q4_matrix, q4_rows, q4_cols, 8);
6125        let q4_acts = synth_q8_0_acts(4, q4_cols);
6126        for act in &q4_acts {
6127            let mut got = vec![0f32; q4_rows];
6128            gemv_q4_0x4_q8_0(&q4_packed, act, q4_cols, n_groups, 8, &mut got);
6129            let mut want = vec![0f32; q4_rows];
6130            gemv_q4_0x4_q8_0_scalar(&q4_packed, act, q4_cols, n_groups, 8, &mut want);
6131            for r in 0..q4_rows {
6132                let err = (got[r] - want[r]).abs();
6133                assert!(
6134                    err / want[r].abs().max(1.0) < 1e-5 || err < 1e-3,
6135                    "q4_0 gemv row {r}: NEON 4x8 {} vs scalar {}",
6136                    got[r],
6137                    want[r]
6138                );
6139            }
6140        }
6141
6142        if !std::arch::is_aarch64_feature_detected!("i8mm") {
6143            return;
6144        }
6145        for na in 1..=Q8K_ACTS_X4_NC {
6146            let acts = synth_q8_0_acts(na, cols);
6147            let tile = prepare_q8_acts_x4(&acts, cols);
6148            for group in 0..n_groups {
6149                let mut x4_out = vec![0f32; Q8_0X4_NROWS * na];
6150                gemm_q8_0x4_group_x4(&packed, group, &tile, cols, 8, &mut x4_out);
6151
6152                let mut group_out = vec![0f32; Q8_0X4_NROWS * na];
6153                gemm_q8_0x4_group(&packed, group, &acts, cols, 8, &mut group_out);
6154                assert_eq!(
6155                    x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
6156                    group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
6157                    "q8_0 x4 entry diverged from the compat entry, group {group} na {na}"
6158                );
6159
6160                let slice = &packed[group * n_blocks * Q8_0X4_BLOCK_BYTES..]
6161                    [..n_blocks * Q8_0X4_BLOCK_BYTES];
6162                let mut want = vec![0f32; Q8_0X4_NROWS * na];
6163                gemm_q8_0x4_acts_x4_scalar_8(slice, &tile, cols, &mut want);
6164                for (got, want) in x4_out.iter().zip(want.iter()) {
6165                    let err = (got - want).abs();
6166                    assert!(
6167                        err / want.abs().max(1.0) < 1e-5 || err < 1e-3,
6168                        "q8_0 group {group} na {na}: i8mm GEMM {got} vs portable {want}"
6169                    );
6170                }
6171            }
6172
6173            let q4_acts = synth_q8_0_acts(na, q4_cols);
6174            let q4_tile = prepare_q8_acts_x4(&q4_acts, q4_cols);
6175            for group in 0..n_groups {
6176                let mut x4_out = vec![0f32; Q4_0X4_NROWS * na];
6177                gemm_q4_0x4_group_x4(&q4_packed, group, &q4_tile, q4_cols, 8, &mut x4_out);
6178
6179                let mut group_out = vec![0f32; Q4_0X4_NROWS * na];
6180                gemm_q4_0x4_group(&q4_packed, group, &q4_acts, q4_cols, 8, &mut group_out);
6181                assert_eq!(
6182                    x4_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
6183                    group_out.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
6184                    "q4_0 x4 entry diverged from the compat entry, group {group} na {na}"
6185                );
6186
6187                let slice = &q4_packed[group * n_blocks * Q4_0X4_BLOCK_BYTES..]
6188                    [..n_blocks * Q4_0X4_BLOCK_BYTES];
6189                let mut want = vec![0f32; Q4_0X4_NROWS * na];
6190                gemm_q4_0x4_acts_x4_scalar_8(slice, &q4_tile, q4_cols, &mut want);
6191                for (got, want) in x4_out.iter().zip(want.iter()) {
6192                    let err = (got - want).abs();
6193                    assert!(
6194                        err / want.abs().max(1.0) < 1e-5 || err < 1e-3,
6195                        "q4_0 group {group} na {na}: i8mm GEMM {got} vs portable {want}"
6196                    );
6197                }
6198            }
6199        }
6200    }
6201
6202    #[test]
6203    fn q4_kx8_gemm_with_no_activations_is_a_no_op() {
6204        let n_blocks = 2;
6205        let cols = n_blocks * Q4_K_BLOCK_ELEMS;
6206        let mut matrix = Vec::new();
6207        for r in 0..Q4_KX8_NROWS {
6208            matrix.extend_from_slice(&synth_q4_k_row(n_blocks, r as u8));
6209        }
6210        let packed = pack_q4_k_matrix_x8(&matrix, Q4_KX8_NROWS, cols, 4);
6211        let mut out: Vec<f32> = Vec::new();
6212        gemm_q4_kx8_group(&packed, 0, &[], cols, 4, &mut out);
6213        assert!(out.is_empty());
6214    }
6215
6216    #[test]
6217    fn q8_0x4_gemm_with_no_activations_is_a_no_op() {
6218        let n_blocks = 2;
6219        let cols = n_blocks * Q8_0_BLOCK_ELEMS;
6220        let mut matrix = Vec::new();
6221        for r in 0..Q8_0X4_NROWS {
6222            matrix.extend_from_slice(&synth_q8_0_row(n_blocks, r as u8));
6223        }
6224        let packed = pack_q8_0_matrix_x4(&matrix, Q8_0X4_NROWS, cols, Q8_0X4_INTERLEAVE);
6225        let mut out: Vec<f32> = Vec::new();
6226        gemm_q8_0x4_group(&packed, 0, &[], cols, Q8_0X4_INTERLEAVE, &mut out);
6227        assert!(out.is_empty());
6228    }
6229
6230    #[test]
6231    fn q5_kx8_pack_and_gemv_matches_scalar_row_dots() {
6232        let n_blocks = 2;
6233        let cols = n_blocks * Q5_K_BLOCK_ELEMS;
6234        let rows = 16;
6235        let mut matrix = Vec::new();
6236        for r in 0..rows {
6237            matrix.extend_from_slice(&synth_q5_k_row(n_blocks, r as u8));
6238        }
6239        let x: Vec<f32> = (0..cols)
6240            .map(|i| ((i as f32) * 0.019 - 1.8).sin() * 1.6)
6241            .collect();
6242        let act = quantize_activations_q8_k(&x);
6243
6244        let row_bytes = n_blocks * Q5_K_BLOCK_BYTES;
6245        let mut reference = vec![0f32; rows];
6246        for r in 0..rows {
6247            reference[r] = dot_q5_k_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
6248        }
6249
6250        for &interleave in &[4usize, 8] {
6251            let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, interleave);
6252            let n_groups = rows / Q5_KX8_NROWS;
6253            let mut out = vec![0f32; rows];
6254            gemv_q5_kx8_q8_k(&packed, &act, cols, n_groups, interleave, &mut out);
6255            for r in 0..rows {
6256                let err = (out[r] - reference[r]).abs();
6257                let scale = reference[r].abs().max(1.0);
6258                assert!(
6259                    err / scale < 1e-4 || err < 1e-3,
6260                    "interleave={interleave} row {r}: got {} want {} err={err}",
6261                    out[r],
6262                    reference[r]
6263                );
6264            }
6265        }
6266    }
6267
6268    #[test]
6269    fn q5_kx8_gemm_matches_the_gemv_run_once_per_activation() {
6270        let n_blocks = 3;
6271        let cols = n_blocks * Q5_K_BLOCK_ELEMS;
6272        let rows = 2 * Q5_KX8_NROWS;
6273        let mut matrix = Vec::new();
6274        for r in 0..rows {
6275            matrix.extend_from_slice(&synth_q5_k_row(n_blocks, (r * 5 + 3) as u8));
6276        }
6277        let interleave = q5_kx8_interleave();
6278        let packed = pack_q5_k_matrix_x8(&matrix, rows, cols, interleave);
6279
6280        let n_acts = 6;
6281        let acts: Vec<Q8KActivations> = (0..n_acts)
6282            .map(|j| {
6283                let x: Vec<f32> = (0..cols)
6284                    .map(|i| (((i + j * 29) as f32) * 0.011 - 0.4).cos() * 2.3)
6285                    .collect();
6286                quantize_activations_q8_k(&x)
6287            })
6288            .collect();
6289
6290        let row_bytes = n_blocks * Q5_K_BLOCK_BYTES;
6291        for group in 0..rows / Q5_KX8_NROWS {
6292            for chunk in acts.chunks(Q5_KX8_GEMM_NC) {
6293                let mut gemm_out = vec![0f32; Q5_KX8_NROWS * chunk.len()];
6294                gemm_q5_kx8_group(&packed, group, chunk, cols, interleave, &mut gemm_out);
6295
6296                for (j, act) in chunk.iter().enumerate() {
6297                    let mut gemv_out = [0f32; Q5_KX8_NROWS];
6298                    gemv_q5_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
6299                    for r in 0..Q5_KX8_NROWS {
6300                        let got = gemm_out[r * chunk.len() + j];
6301                        let want = gemv_out[r];
6302                        // Tolerance, not bit equality: on i8mm hosts the
6303                        // GEMM (i8mm, per-block f32 accumulation) and the
6304                        // GEMV (DotProd, per-sub-block) round differently.
6305                        // Measured on an M2 Pro: worst observed relative
6306                        // deviation 2.9e-5, everything else <= 4.8e-6.
6307                        let err = (got - want).abs();
6308                        let scale = want.abs().max(1.0);
6309                        assert!(
6310                            err / scale < 5e-5 || err < 1e-3,
6311                            "group {group} row {r} act {j}: Q5_K GEMM {got} vs GEMV {want}"
6312                        );
6313                    }
6314                }
6315            }
6316
6317            // Also check against per-row dot_q5_k_q8 for each activation.
6318            for (j, act) in acts.iter().enumerate() {
6319                for r in 0..Q5_KX8_NROWS {
6320                    let row_idx = group * Q5_KX8_NROWS + r;
6321                    let row = &matrix[row_idx * row_bytes..(row_idx + 1) * row_bytes];
6322                    let want = dot_q5_k_q8_scalar(row, act);
6323                    let mut gemv_out = [0f32; Q5_KX8_NROWS];
6324                    gemv_q5_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
6325                    let err = (gemv_out[r] - want).abs();
6326                    let scale = want.abs().max(1.0);
6327                    assert!(
6328                        err / scale < 1e-4 || err < 1e-3,
6329                        "group {group} row {r} act {j}: packed gemv {} vs dot {want}",
6330                        gemv_out[r]
6331                    );
6332                }
6333            }
6334        }
6335    }
6336
6337    #[test]
6338    fn q5_kx8_gemm_with_no_activations_is_a_no_op() {
6339        let n_blocks = 2;
6340        let cols = n_blocks * Q5_K_BLOCK_ELEMS;
6341        let mut matrix = Vec::new();
6342        for r in 0..Q5_KX8_NROWS {
6343            matrix.extend_from_slice(&synth_q5_k_row(n_blocks, r as u8));
6344        }
6345        let packed = pack_q5_k_matrix_x8(&matrix, Q5_KX8_NROWS, cols, 4);
6346        let mut out: Vec<f32> = Vec::new();
6347        gemm_q5_kx8_group(&packed, 0, &[], cols, 4, &mut out);
6348        assert!(out.is_empty());
6349    }
6350
6351    #[test]
6352    fn q6_kx8_pack_and_gemv_matches_scalar_row_dots() {
6353        let n_blocks = 3;
6354        let cols = n_blocks * Q6_K_BLOCK_ELEMS;
6355        let rows = 2 * Q6_KX8_NROWS;
6356        let mut matrix = Vec::new();
6357        for r in 0..rows {
6358            matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 5 + 1) as u8));
6359        }
6360        let x: Vec<f32> = (0..cols)
6361            .map(|i| ((i as f32) * 0.017 - 0.8).cos() * 1.8)
6362            .collect();
6363        let act = quantize_activations_q8_k(&x);
6364        let row_bytes = n_blocks * Q6_K_BLOCK_BYTES;
6365        let mut reference = vec![0f32; rows];
6366        for r in 0..rows {
6367            reference[r] = dot_q6_k_q8_scalar(&matrix[r * row_bytes..(r + 1) * row_bytes], &act);
6368        }
6369        for interleave in [4usize, 8] {
6370            let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, interleave);
6371            let n_groups = rows / Q6_KX8_NROWS;
6372            let mut out = vec![0f32; rows];
6373            gemv_q6_kx8_q8_k(&packed, &act, cols, n_groups, interleave, &mut out);
6374            for r in 0..rows {
6375                let err = (out[r] - reference[r]).abs();
6376                let scale = reference[r].abs().max(1.0);
6377                assert!(
6378                    err / scale < 1e-4 || err < 1e-3,
6379                    "interleave={interleave} row {r}: got {} want {} err={err}",
6380                    out[r],
6381                    reference[r]
6382                );
6383            }
6384        }
6385    }
6386
6387    #[test]
6388    fn q6_kx8_gemm_matches_the_gemv_run_once_per_activation() {
6389        let n_blocks = 2;
6390        let cols = n_blocks * Q6_K_BLOCK_ELEMS;
6391        let rows = 2 * Q6_KX8_NROWS;
6392        let mut matrix = Vec::new();
6393        for r in 0..rows {
6394            matrix.extend_from_slice(&synth_q6_k_row(n_blocks, (r * 3 + 2) as u8));
6395        }
6396        let interleave = q6_kx8_interleave();
6397        let packed = pack_q6_k_matrix_x8(&matrix, rows, cols, interleave);
6398        let acts: Vec<_> = (0..Q6_KX8_GEMM_NC)
6399            .map(|j| {
6400                let x: Vec<f32> = (0..cols)
6401                    .map(|i| (((i + j * 11) as f32) * 0.015 - 0.7).sin() * 2.0)
6402                    .collect();
6403                quantize_activations_q8_k(&x)
6404            })
6405            .collect();
6406        let row_bytes = n_blocks * Q6_K_BLOCK_BYTES;
6407        for group in 0..rows / Q6_KX8_NROWS {
6408            for chunk in acts.chunks(Q6_KX8_GEMM_NC) {
6409                let mut gemm_out = vec![0f32; Q6_KX8_NROWS * chunk.len()];
6410                gemm_q6_kx8_group(&packed, group, chunk, cols, interleave, &mut gemm_out);
6411                for (j, act) in chunk.iter().enumerate() {
6412                    let mut gemv_out = [0f32; Q6_KX8_NROWS];
6413                    gemv_q6_kx8_group(&packed, group, act, cols, interleave, &mut gemv_out);
6414                    for r in 0..Q6_KX8_NROWS {
6415                        let got = gemm_out[r * chunk.len() + j];
6416                        let want = gemv_out[r];
6417                        let err = (got - want).abs();
6418                        let scale = want.abs().max(1.0);
6419                        assert!(
6420                            err / scale < 1e-4 || err < 1e-3,
6421                            "group {group} row {r} act {j}: gemm {got} vs gemv {want}"
6422                        );
6423                        let row_idx = group * Q6_KX8_NROWS + r;
6424                        let row = &matrix[row_idx * row_bytes..(row_idx + 1) * row_bytes];
6425                        let dot = dot_q6_k_q8_scalar(row, act);
6426                        let err2 = (got - dot).abs();
6427                        let scale2 = dot.abs().max(1.0);
6428                        assert!(
6429                            err2 / scale2 < 1e-4 || err2 < 1e-3,
6430                            "group {group} row {r} act {j}: gemm {got} vs dot {dot}"
6431                        );
6432                    }
6433                }
6434            }
6435        }
6436    }
6437
6438    #[test]
6439    fn block_size_matches_ggml() {
6440        assert_eq!(Q4_KX8_BLOCK_BYTES, 16 + 16 + 96 + 1024);
6441        assert_eq!(Q5_KX8_BLOCK_BYTES, 16 + 16 + 96 + 256 + 1024);
6442        assert_eq!(Q6_KX8_BLOCK_BYTES, 16 + 128 + 1024 + 512);
6443        assert_eq!(Q8_0X4_BLOCK_BYTES, 4 * 2 + Q8_0_BLOCK_ELEMS * Q8_0X4_NROWS);
6444        assert_eq!(Q4_0X4_BLOCK_BYTES, 4 * 2 + Q4_0_BLOCK_ELEMS * 2);
6445    }
6446}