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