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