1use 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
21pub const Q4_KX8_BLOCK_BYTES: usize = 1152;
23pub const Q4_KX8_NROWS: usize = 8;
25
26const KMASK1: u32 = 0x3f3f_3f3f;
27const KMASK2: u32 = 0x0f0f_0f0f;
28const KMASK3: u32 = 0x0303_0303;
29
30#[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
51pub 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 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; 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 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
132pub 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#[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
174fn 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); 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
243fn 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); 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
311pub 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#[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
378pub const Q8_0X4_BLOCK_BYTES: usize = 136;
384pub const Q8_0X4_NROWS: usize = 4;
386pub const Q8_0X4_INTERLEAVE: usize = 4;
389
390#[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
404pub 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
432pub 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
454fn 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
494pub 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
535pub const Q8_0X4_GEMM_NC: usize = 8;
540
541pub 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 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 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
615pub struct Q8ActsX4 {
621 pub na: usize,
623 pub n_blocks: usize,
625 pub qs: Vec<i8>,
629 pub d: Vec<f32>,
631}
632
633pub 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.chunks_exact(8).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#[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
679pub 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
714fn 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 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
752pub const Q4_KX8_GEMM_NC: usize = 4;
759
760pub 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 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 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
819pub const Q8K_ACTS_X4_NC: usize = 4;
831
832pub struct Q8KActsX4 {
833 pub na: usize,
835 pub n_blocks: usize,
837 pub qs: Vec<i8>,
841 pub bsums: Vec<i16>,
844 pub d: Vec<f32>,
846}
847
848pub 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.chunks_exact(8).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.chunks_exact(2)) {
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#[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
903pub 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
944fn 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); 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 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
1009pub const Q5_KX8_BLOCK_BYTES: usize = 1408;
1016pub const Q5_KX8_NROWS: usize = 8;
1018
1019#[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
1035pub 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 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 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
1126pub 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
1148fn 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); 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
1227fn 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); 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
1305pub 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#[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
1362pub const Q5_KX8_GEMM_NC: usize = 4;
1364
1365pub 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 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#[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
1430pub 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
1466fn 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); 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 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
1679pub const Q6_KX8_BLOCK_BYTES: usize = 1680;
1685pub const Q6_KX8_NROWS: usize = 8;
1687
1688#[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
1704pub 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 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 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
1884pub 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 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 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#[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
1998pub 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
2035fn 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 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
2107pub const Q4_0X4_BLOCK_BYTES: usize = 72;
2113pub const Q4_0X4_NROWS: usize = 4;
2115pub const Q4_0X4_INTERLEAVE: usize = 4;
2118
2119#[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
2136pub 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
2180pub 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
2209fn 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
2252pub 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
2293pub const Q4_0X4_GEMM_NC: usize = 4;
2295
2296pub 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 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#[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
2373pub 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
2408fn 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 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#[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#[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 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 #[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 #[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 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 #[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 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 #[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 #[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 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 #[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 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 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; 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 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 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 #[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 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 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 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 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 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 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 #[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 #[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 #[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 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 #[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 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; 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 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 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 #[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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 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 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 #[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 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 #[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 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 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; 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; 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 #[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 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 #[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 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 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 #[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 #[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 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 #[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 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 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 #[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 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 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 #[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 #[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 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 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 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 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}