1use crate::pool::{Pool, matvec_rows, matvec_rows2};
15use cortiq_core::quant::{
16 GROUP_SIZE, Q1_TILE, Q2TP_CHUNK, Q4_TILE, Q4TP_NIB, f16_to_f32, q2tp_ladder, q2tp_sections,
17 q4tp_code, q4tp_ladder, q4tp_sections,
18};
19use cortiq_core::{CmfModel, TensorDtype};
20use std::cell::UnsafeCell;
21use std::sync::Arc;
22
23pub enum QTensor {
24 F32 {
25 data: Vec<f32>,
26 rows: usize,
27 cols: usize,
28 },
29 Mapped {
30 model: Arc<CmfModel>,
31 idx: usize,
33 dtype: TensorDtype,
34 rows: usize,
35 cols: usize,
36 row_scale: Vec<f32>,
38 col_field: Vec<f32>,
40 vbit_offsets: Vec<usize>,
44 repack: Vec<u8>,
52 },
53}
54
55fn repack_enabled() -> bool {
62 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
63 *ON.get_or_init(|| {
64 std::env::var("CMF_REPACK")
65 .map(|v| v == "1")
66 .unwrap_or(cfg!(target_os = "android"))
67 })
68}
69
70fn q8_repack(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
74 #[cfg(target_arch = "aarch64")]
75 let arch_ok = sdot_enabled();
76 #[cfg(not(target_arch = "aarch64"))]
77 let arch_ok = false;
78 if !arch_ok || !repack_enabled() || rows < 256 || cols % 16 != 0 {
79 return Vec::new();
80 }
81 q8_repack_layout(bytes, rows, cols)
82}
83
84fn q8_repack_layout(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
87 let groups = rows / 4;
88 let mut rep = vec![0u8; groups * 4 * cols];
89 for g in 0..groups {
90 let dst = &mut rep[g * 4 * cols..(g + 1) * 4 * cols];
91 for c in 0..cols / 16 {
92 for lane in 0..4 {
93 let src = (g * 4 + lane) * cols + c * 16;
94 dst[c * 64 + lane * 16..c * 64 + lane * 16 + 16]
95 .copy_from_slice(&bytes[src..src + 16]);
96 }
97 }
98 }
99 rep
100}
101
102fn vbit_row_offsets(bytes: &[u8], rows: usize, cols: usize) -> Vec<usize> {
105 let ng = cols / GROUP_SIZE;
106 let bits = &bytes[..rows];
107 let mut offsets = Vec::with_capacity(rows + 1);
108 let mut off = rows + rows * ng * 2;
109 for r in 0..rows {
110 offsets.push(off);
111 off += (cols * bits[r] as usize).div_ceil(8);
112 }
113 offsets.push(off);
114 offsets
115}
116
117fn blocked_enabled() -> bool {
123 use std::sync::atomic::Ordering::Relaxed;
124 match BLOCKED_OVERRIDE.load(Relaxed) {
125 1 => false,
126 2 => true,
127 _ => {
128 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
129 *ON.get_or_init(|| {
130 std::env::var("CMF_X86_BLOCKED")
131 .map(|v| v != "0")
132 .unwrap_or(true)
133 })
134 }
135 }
136}
137
138static BLOCKED_OVERRIDE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
139
140pub fn set_blocked_override(on: Option<bool>) {
147 let v = match on {
148 None => 0,
149 Some(false) => 1,
150 Some(true) => 2,
151 };
152 BLOCKED_OVERRIDE.store(v, std::sync::atomic::Ordering::Relaxed);
153}
154
155fn gpu_lmhead_enabled() -> bool {
156 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
157 *ON.get_or_init(|| {
158 std::env::var("CMF_GPU_LMHEAD")
159 .map(|v| v != "0")
160 .unwrap_or(true)
161 })
162}
163
164fn gpu_split_frac() -> f32 {
165 if FULL_GPU_Q8.get() {
168 return 1.0;
169 }
170 static F: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
171 *F.get_or_init(|| {
172 std::env::var("CMF_GPU_SPLIT")
173 .ok()
174 .and_then(|v| v.parse::<f32>().ok())
175 .unwrap_or(0.5)
176 .clamp(0.0, 1.0)
177 })
178}
179
180impl QTensor {
181 pub fn from_f32(data: Vec<f32>, rows: usize, cols: usize) -> Self {
182 debug_assert_eq!(data.len(), rows * cols);
183 Self::F32 { data, rows, cols }
184 }
185
186 pub fn from_model(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
189 let idx = model
192 .tensor_index(name)
193 .ok_or_else(|| format!("tensor '{name}' not found in CMF directory"))?;
194 let entry = &model.tensors[idx];
195 if entry.shape.len() != 2 {
196 return Err(format!("QTensor::from_model needs 2-D, got '{name}'"));
197 }
198 let (rows, cols) = (entry.shape[0], entry.shape[1]);
199 let bytes = model.entry_bytes(entry);
200
201 match entry.dtype {
202 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
203 let n = rows * cols;
204 let scales_off = n;
205 let row_scale: Vec<f32> = (0..rows)
206 .map(|o| {
207 f16_to_f32(u16::from_le_bytes([
208 bytes[scales_off + o * 2],
209 bytes[scales_off + o * 2 + 1],
210 ]))
211 })
212 .collect();
213 let col_field: Vec<f32> = if entry.dtype == TensorDtype::Q8_2f {
214 let col_off = n + rows * 2;
215 (0..cols)
216 .map(|i| {
217 f16_to_f32(u16::from_le_bytes([
218 bytes[col_off + i * 2],
219 bytes[col_off + i * 2 + 1],
220 ]))
221 })
222 .collect()
223 } else {
224 Vec::new()
225 };
226 Ok(Self::Mapped {
227 model: model.clone(),
228 idx,
229 dtype: entry.dtype,
230 rows,
231 cols,
232 row_scale,
233 col_field,
234 vbit_offsets: Vec::new(),
235 repack: q8_repack(bytes, rows, cols),
236 })
237 }
238 TensorDtype::Vbit if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
240 model: model.clone(),
241 idx,
242 dtype: entry.dtype,
243 rows,
244 cols,
245 row_scale: Vec::new(),
246 col_field: Vec::new(),
247 vbit_offsets: vbit_row_offsets(bytes, rows, cols),
248 repack: Vec::new(),
249 }),
250 TensorDtype::VbitRo if cols % GROUP_SIZE == 0 => {
254 let (_, off_off, packed_off) = cortiq_core::quant::vbit_ro_sections(rows, cols);
255 let offsets: Vec<usize> = (0..=rows)
256 .map(|r| packed_off + cortiq_core::quant::vbit_ro_offset(bytes, off_off, r))
257 .collect();
258 Ok(Self::Mapped {
259 model: model.clone(),
260 idx,
261 dtype: entry.dtype,
262 rows,
263 cols,
264 row_scale: Vec::new(),
265 col_field: Vec::new(),
266 vbit_offsets: offsets,
267 repack: Vec::new(),
268 })
269 }
270 TensorDtype::Q4Tiled if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
276 model: model.clone(),
277 idx,
278 dtype: entry.dtype,
279 rows,
280 cols,
281 row_scale: Vec::new(),
282 col_field: Vec::new(),
283 vbit_offsets: Vec::new(),
284 repack: Vec::new(),
285 }),
286 TensorDtype::Q4TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
289 model: model.clone(),
290 idx,
291 dtype: entry.dtype,
292 rows,
293 cols,
294 row_scale: Vec::new(),
295 col_field: Vec::new(),
296 vbit_offsets: Vec::new(),
297 repack: Vec::new(),
298 }),
299 TensorDtype::Q2TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
301 model: model.clone(),
302 idx,
303 dtype: entry.dtype,
304 rows,
305 cols,
306 row_scale: Vec::new(),
307 col_field: Vec::new(),
308 vbit_offsets: Vec::new(),
309 repack: Vec::new(),
310 }),
311 TensorDtype::Q4Block if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
312 model: model.clone(),
313 idx,
314 dtype: entry.dtype,
315 rows,
316 cols,
317 row_scale: Vec::new(),
318 col_field: Vec::new(),
319 vbit_offsets: Vec::new(),
320 repack: Vec::new(),
321 }),
322 TensorDtype::Q1 if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
324 model: model.clone(),
325 idx,
326 dtype: entry.dtype,
327 rows,
328 cols,
329 row_scale: Vec::new(),
330 col_field: Vec::new(),
331 vbit_offsets: Vec::new(),
332 repack: Vec::new(),
333 }),
334 TensorDtype::Q1T if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
338 model: model.clone(),
339 idx,
340 dtype: entry.dtype,
341 rows,
342 cols,
343 row_scale: Vec::new(),
344 col_field: Vec::new(),
345 vbit_offsets: Vec::new(),
346 repack: Vec::new(),
347 }),
348 _ => {
350 let mut data = vec![0.0f32; rows * cols];
351 cortiq_core::quant::dequant_tensor(entry, bytes, &mut data)?;
352 Ok(Self::from_f32(data, rows, cols))
353 }
354 }
355 }
356
357 pub(crate) fn is_q1(&self) -> bool {
360 matches!(
361 self,
362 Self::Mapped {
363 dtype: TensorDtype::Q1,
364 ..
365 }
366 )
367 }
368
369 pub(crate) fn f32_parts(&self) -> Option<(&[f32], usize, usize)> {
372 match self {
373 Self::F32 { data, rows, cols } => Some((data, *rows, *cols)),
374 _ => None,
375 }
376 }
377
378 pub(crate) fn q1_parts(&self) -> Option<(usize, usize, usize)> {
385 if self.has_prism_contract() {
386 return None;
387 }
388 match self {
389 #[cfg(target_os = "macos")]
390 Self::Mapped {
391 dtype: TensorDtype::Q1T,
392 ..
393 } if !crate::gpu::metal_q1t_enabled() => None,
394 Self::Mapped {
395 idx,
396 dtype:
397 TensorDtype::Q1
398 | TensorDtype::Q1T
399 | TensorDtype::Q4Block
400 | TensorDtype::Q4Tiled
401 | TensorDtype::Q4TiledP
405 | TensorDtype::Q8Row
406 | TensorDtype::Q8_2f,
407 rows,
408 cols,
409 ..
410 } => Some((*idx, *rows, *cols)),
411 _ => None,
412 }
413 }
414
415 #[cfg(target_os = "macos")]
422 pub(crate) fn metal_graph_parts(&self) -> Option<(usize, usize, usize)> {
423 if let Some((model, idx, kind, _)) = self.graph_weight_descriptor() {
424 let name = &model.tensors[idx].name;
425 let forward = kind == 9 && crate::prism::is_forward_weight(model, name);
426 let affine = kind == 9 && crate::prism::is_affine_target(model, name);
427 if forward && affine {
428 let e = model.tensors.get(idx)?;
429 return Some((idx, *e.shape.first()?, *e.shape.get(1)?));
430 }
431 }
432 self.q1_parts()
433 }
434
435 pub(crate) fn q4t_parts(&self) -> Option<(usize, usize, usize)> {
441 if self.has_prism_contract() {
442 return None;
443 }
444 match self {
445 Self::Mapped {
446 idx,
447 dtype: TensorDtype::Q4Tiled,
448 rows,
449 cols,
450 ..
451 } => Some((*idx, *rows, *cols)),
452 _ => None,
453 }
454 }
455
456 pub(crate) fn q4tp_parts(&self) -> Option<(usize, usize, usize)> {
460 if self.has_prism_contract() {
461 return None;
462 }
463 match self {
464 Self::Mapped {
465 idx,
466 dtype: TensorDtype::Q4TiledP,
467 rows,
468 cols,
469 ..
470 } => Some((*idx, *rows, *cols)),
471 _ => None,
472 }
473 }
474
475 pub(crate) fn q8_row_parts(&self) -> Option<(usize, usize, usize, &[f32])> {
480 if self.has_prism_contract() {
481 return None;
482 }
483 match self {
484 Self::Mapped {
485 idx,
486 dtype: TensorDtype::Q8Row,
487 rows,
488 cols,
489 row_scale,
490 col_field,
491 ..
492 } if col_field.is_empty() => Some((*idx, *rows, *cols, row_scale)),
493 _ => None,
494 }
495 }
496
497 pub fn model_dtype(&self) -> Option<cortiq_core::TensorDtype> {
501 match self {
502 Self::Mapped { dtype, .. } => Some(*dtype),
503 _ => None,
504 }
505 }
506
507 pub fn model_idx(&self) -> Option<usize> {
512 match self {
513 Self::Mapped { idx, .. } => Some(*idx),
514 _ => None,
515 }
516 }
517
518 pub fn model_arc(&self) -> Option<std::sync::Arc<cortiq_core::CmfModel>> {
523 match self {
524 Self::Mapped { model, .. } => Some(model.clone()),
525 _ => None,
526 }
527 }
528
529 pub(crate) fn has_prism_contract(&self) -> bool {
534 matches!(self, Self::Mapped { model, .. } if crate::prism::has_contract(model))
535 }
536
537 pub fn rows(&self) -> usize {
538 match self {
539 Self::F32 { rows, .. } | Self::Mapped { rows, .. } => *rows,
540 }
541 }
542
543 pub(crate) fn mapped_q4t(&self) -> Option<(&Arc<CmfModel>, usize)> {
546 if self.has_prism_contract() {
547 return None;
548 }
549 match self {
550 Self::Mapped {
551 model,
552 idx,
553 dtype: TensorDtype::Q4Tiled,
554 ..
555 } => Some((model, *idx)),
556 _ => None,
557 }
558 }
559
560 pub fn mapped_q4tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
563 if self.has_prism_contract() {
564 return None;
565 }
566 match self {
567 Self::Mapped {
568 model,
569 idx,
570 dtype: TensorDtype::Q4TiledP,
571 ..
572 } => Some((model, *idx)),
573 _ => None,
574 }
575 }
576
577 pub fn mapped_device_gemm(&self) -> Option<(&Arc<CmfModel>, usize)> {
585 if self.has_prism_contract() {
586 return None;
587 }
588 match self {
589 Self::Mapped {
590 model,
591 idx,
592 dtype: TensorDtype::Q4TiledP | TensorDtype::Q8Row | TensorDtype::Q8_2f,
593 ..
594 } => Some((model, *idx)),
595 _ => None,
596 }
597 }
598
599 pub fn mapped_q2tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
602 if self.has_prism_contract() {
603 return None;
604 }
605 match self {
606 Self::Mapped {
607 model,
608 idx,
609 dtype: TensorDtype::Q2TiledP,
610 ..
611 } => Some((model, *idx)),
612 _ => None,
613 }
614 }
615
616 pub fn cols(&self) -> usize {
617 match self {
618 Self::F32 { cols, .. } | Self::Mapped { cols, .. } => *cols,
619 }
620 }
621
622 pub fn mapped_q1(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
625 if self.has_prism_contract() {
626 return None;
627 }
628 match self {
629 Self::Mapped {
630 model,
631 idx,
632 dtype: TensorDtype::Q1,
633 ..
634 } => Some((model, *idx)),
635 _ => None,
636 }
637 }
638
639 pub fn graph_weight(&self) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
649 if self.has_prism_contract() {
650 return None;
651 }
652 self.graph_weight_descriptor()
653 }
654
655 pub(crate) fn graph_weight_descriptor(
660 &self,
661 ) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
662 match self {
663 Self::Mapped {
664 model,
665 idx,
666 dtype: TensorDtype::Q8Row,
667 row_scale,
668 ..
669 } => Some((model, *idx, 0, row_scale.as_slice())),
670 Self::Mapped {
671 model,
672 idx,
673 dtype: TensorDtype::Q1,
674 ..
675 } => Some((model, *idx, 1, &[])),
676 Self::Mapped {
681 model,
682 idx,
683 dtype: TensorDtype::Q4Tiled,
684 ..
685 } => Some((model, *idx, 5, &[])),
686 Self::Mapped {
690 model,
691 idx,
692 dtype: TensorDtype::Q4TiledP,
693 ..
694 } => Some((model, *idx, 6, &[])),
695 Self::Mapped {
696 model,
697 idx,
698 dtype: TensorDtype::Q4Block,
699 ..
700 } => Some((model, *idx, 2, &[])),
701 Self::Mapped {
706 model,
707 idx,
708 dtype: TensorDtype::Q8_2f,
709 ..
710 } => Some((model, *idx, 7, &[])),
711 Self::Mapped {
712 model,
713 idx,
714 dtype: TensorDtype::Q1T,
715 ..
716 } => Some((model, *idx, 3, &[])),
717 Self::Mapped {
721 model,
722 idx,
723 dtype: TensorDtype::Q2TiledP,
724 ..
725 } => Some((model, *idx, 9, &[])),
726 _ => None,
727 }
728 }
729
730 pub fn as_f32(&self) -> Option<&[f32]> {
733 match self {
734 Self::F32 { data, .. } => Some(data),
735 Self::Mapped { .. } => None,
736 }
737 }
738
739 fn quant_bytes(&self) -> &[u8] {
740 match self {
741 Self::Mapped { model, idx, .. } => model.entry_bytes(&model.tensors[*idx]),
742 Self::F32 { .. } => unreachable!("quant_bytes on F32"),
743 }
744 }
745
746 pub fn row_f32(&self, r: usize, dst: &mut [f32]) {
748 let cols = self.cols();
749 debug_assert_eq!(dst.len(), cols);
750 match self {
751 Self::F32 { data, .. } => dst.copy_from_slice(&data[r * cols..(r + 1) * cols]),
752 Self::Mapped {
753 model,
754 idx,
755 dtype,
756 row_scale,
757 col_field,
758 vbit_offsets,
759 ..
760 } => {
761 if *dtype == TensorDtype::Q4Tiled {
762 let bytes = self.quant_bytes();
763 let gpr = cols / GROUP_SIZE;
764 for gi in 0..gpr {
765 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
766 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
767 for (k, &b) in tile[2..].iter().enumerate() {
768 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
769 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
770 }
771 }
772 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
773 crate::prism::inverse_embedding(model, dst);
774 }
775 return;
776 }
777 if *dtype == TensorDtype::Q4TiledP {
778 let bytes = self.quant_bytes();
779 let gpr = cols / GROUP_SIZE;
780 let v = Q4tpView::new(bytes, self.rows(), cols);
781 let mut sc = vec![0f32; gpr];
782 v.scales_into(r, gpr, &mut sc);
783 for gi in 0..gpr {
784 let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
785 let s = sc[gi];
786 for (k, &b) in tile.iter().enumerate() {
787 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
788 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
789 }
790 }
791 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
792 crate::prism::inverse_embedding(model, dst);
793 }
794 return;
795 }
796 if *dtype == TensorDtype::Q2TiledP {
797 let bytes = self.quant_bytes();
798 let gpr = cols / GROUP_SIZE;
799 let v = Q4tpView::new_q2(bytes, self.rows(), cols);
800 let mut sc = vec![0f32; gpr];
801 v.scales_into(r, gpr, &mut sc);
802 for gi in 0..gpr {
803 let ch =
804 &v.nib[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
805 let s = sc[gi];
806 for (k, &b) in ch.iter().enumerate() {
807 for j in 0..4 {
808 let center = if crate::prism::is_affine_target(
809 model,
810 &model.tensors[*idx].name,
811 ) {
812 1.0
813 } else {
814 1.5
815 };
816 dst[gi * GROUP_SIZE + k * 4 + j] =
817 (((b >> (2 * j)) & 3) as f32 - center) * s;
818 }
819 }
820 }
821 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
822 crate::prism::inverse_embedding(model, dst);
823 }
824 return;
825 }
826 if *dtype == TensorDtype::Q4Block {
827 let (packed, scales) = q4_split(self.quant_bytes(), self.rows(), cols);
828 let gpr = cols / GROUP_SIZE;
829 for gi in 0..gpr {
830 let g = r * gpr + gi;
831 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
832 for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
833 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
834 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
835 }
836 }
837 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
838 crate::prism::inverse_embedding(model, dst);
839 }
840 return;
841 }
842 if *dtype == TensorDtype::Q1 {
843 let bytes = self.quant_bytes();
844 let gpr = cols / GROUP_SIZE;
845 for gi in 0..gpr {
846 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
847 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
848 for (j, &b) in tile[2..].iter().enumerate() {
849 for k in 0..8 {
850 dst[gi * GROUP_SIZE + j * 8 + k] =
851 (((b >> k) & 1) as f32 * 2.0 - 1.0) * s;
852 }
853 }
854 }
855 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
856 crate::prism::inverse_embedding(model, dst);
857 }
858 return;
859 }
860 if *dtype == TensorDtype::Q1T {
861 let bytes = self.quant_bytes();
862 let gpr = cols / GROUP_SIZE;
863 let base_len = self.rows() * gpr * cortiq_core::quant::Q1T_TILE;
864 for gi in 0..gpr {
865 let off = (r * gpr + gi) * cortiq_core::quant::Q1T_TILE;
866 let s = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
867 bytes[off],
868 bytes[off + 1],
869 ]));
870 let codes = &bytes[off + 2..off + cortiq_core::quant::Q1T_TILE];
871 for k in 0..GROUP_SIZE {
872 dst[gi * GROUP_SIZE + k] = match cortiq_core::quant::q1t_code(codes, k)
873 {
874 1 => s,
875 2 => -s,
876 _ => 0.0,
877 };
878 }
879 }
880 let rows = self.rows();
882 let entries = base_len + (rows + 1) * 4;
883 if entries <= bytes.len() {
884 let ptrs = &bytes[base_len..base_len + (rows + 1) * 4];
885 let r0 = u32::from_le_bytes([
886 ptrs[r * 4],
887 ptrs[r * 4 + 1],
888 ptrs[r * 4 + 2],
889 ptrs[r * 4 + 3],
890 ]) as usize;
891 let r1 = u32::from_le_bytes([
892 ptrs[(r + 1) * 4],
893 ptrs[(r + 1) * 4 + 1],
894 ptrs[(r + 1) * 4 + 2],
895 ptrs[(r + 1) * 4 + 3],
896 ]) as usize;
897 let off = entries + r0 * 4;
898 for i in 0..r1 - r0 {
899 let item = &bytes[off + i * 4..off + i * 4 + 4];
900 let c = u16::from_le_bytes([item[0], item[1]]) as usize;
901 let v = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
902 item[2], item[3],
903 ]));
904 if c < cols {
905 dst[c] = v;
906 }
907 }
908 }
909 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
910 crate::prism::inverse_embedding(model, dst);
911 }
912 return;
913 }
914 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
915 let bytes = self.quant_bytes();
916 let rows = self.rows();
917 let ng = cols / GROUP_SIZE;
918 let bits = &bytes[..rows];
919 let sc_off = rows;
920 let off = vbit_offsets[r];
923 let b = bits[r] as usize;
924 let l = ((1usize << (b - 1)) - 1) as f32;
925 let data = &bytes[off..];
926 let (mut acc, mut nbits, mut byte_idx) = (0u64, 0usize, 0usize);
927 for (i, d) in dst.iter_mut().enumerate() {
928 while nbits < b {
929 acc = (acc << 8) | data[byte_idx] as u64;
930 byte_idx += 1;
931 nbits += 8;
932 }
933 let u = ((acc >> (nbits - b)) & ((1u64 << b) - 1)) as f32;
934 nbits -= b;
935 let so = (r * ng + i / GROUP_SIZE) * 2;
936 let sv = f16_to_f32(u16::from_le_bytes([
937 bytes[sc_off + so],
938 bytes[sc_off + so + 1],
939 ]));
940 *d = (u - l) * sv;
941 }
942 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
943 crate::prism::inverse_embedding(model, dst);
944 }
945 return;
946 }
947 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
948 let s = row_scale[r];
949 match dtype {
950 TensorDtype::Q8Row => {
951 for (d, &b) in dst.iter_mut().zip(q) {
952 *d = (b as i8) as f32 * s;
953 }
954 }
955 TensorDtype::Q8_2f => {
956 for (i, (d, &b)) in dst.iter_mut().zip(q).enumerate() {
957 *d = (b as i8) as f32 * s * col_field[i];
958 }
959 }
960 _ => unreachable!(),
961 }
962 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
963 crate::prism::inverse_embedding(model, dst);
964 }
965 }
966 }
967 }
968
969 pub fn sparse_col_ok(&self) -> bool {
974 match self {
975 Self::F32 { .. } => true,
976 Self::Mapped { dtype, .. } => {
977 matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
978 }
979 }
980 }
981
982 pub fn add_col_scaled(&self, c: usize, w: f32, out: &mut [f32]) {
986 let inter = self.cols();
987 let hidden = self.rows();
988 debug_assert_eq!(out.len(), hidden);
989 match self {
990 Self::F32 { data, .. } => {
991 for (k, o) in out.iter_mut().enumerate() {
992 *o += w * data[k * inter + c];
993 }
994 }
995 Self::Mapped {
996 dtype,
997 row_scale,
998 col_field,
999 ..
1000 } => {
1001 let q = self.quant_bytes();
1002 let colf = if *dtype == TensorDtype::Q8_2f {
1003 col_field[c]
1004 } else {
1005 1.0
1006 };
1007 let wc = w * colf;
1008 for (k, o) in out.iter_mut().enumerate() {
1009 let b = q[k * inter + c] as i8 as f32;
1010 *o += wc * b * row_scale[k];
1011 }
1012 }
1013 }
1014 }
1015
1016 #[inline]
1025 pub fn prefetch_row(&self, r: usize) {
1026 let Self::Mapped { dtype, .. } = self else {
1027 return;
1028 };
1029 if !matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f) {
1030 return;
1031 }
1032 let cols = self.cols();
1033 let q = self.quant_bytes();
1034 let (a, b) = (r * cols, (r + 1) * cols);
1035 if b > q.len() {
1036 return;
1037 }
1038 let mut j = a;
1039 while j < b {
1040 unsafe { std::ptr::read_volatile(q.as_ptr().add(j)) };
1041 j += 512;
1042 }
1043 }
1044
1045 pub fn add_row_scaled(&self, r: usize, w: f32, out: &mut [f32], scratch: &mut [f32]) {
1054 let cols = self.cols();
1055 debug_assert_eq!(out.len(), cols);
1056 match self {
1057 Self::F32 { data, .. } => {
1058 let row = &data[r * cols..(r + 1) * cols];
1059 for (o, v) in out.iter_mut().zip(row) {
1060 *o += w * v;
1061 }
1062 }
1063 Self::Mapped {
1064 dtype,
1065 row_scale,
1066 col_field,
1067 ..
1068 } => match dtype {
1069 TensorDtype::Q8Row => {
1070 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1071 let ws = w * row_scale[r];
1072 let row: &[i8] =
1073 unsafe { std::slice::from_raw_parts(q.as_ptr() as *const i8, q.len()) };
1074 axpy_i8_f32(out, row, ws);
1075 }
1076 TensorDtype::Q8_2f => {
1077 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1078 let ws = w * row_scale[r];
1079 for ((o, b), c) in out.iter_mut().zip(q).zip(col_field) {
1080 *o += ws * c * (*b as i8 as f32);
1081 }
1082 }
1083 _ => {
1084 self.row_f32(r, scratch);
1085 for (o, v) in out.iter_mut().zip(scratch.iter()) {
1086 *o += w * v;
1087 }
1088 }
1089 },
1090 }
1091 }
1092
1093 pub fn row_dot(&self, r: usize, x: &[f32], scratch: &mut [f32]) -> f32 {
1097 let cols = self.cols();
1098 match self {
1099 Self::F32 { data, .. } => {
1100 let row = &data[r * cols..(r + 1) * cols];
1101 row.iter().zip(x).map(|(w, v)| w * v).sum()
1102 }
1103 Self::Mapped {
1104 model,
1105 idx,
1106 dtype,
1107 row_scale,
1108 col_field,
1109 ..
1110 } => {
1111 let prism_forward =
1112 crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1113 if prism_forward {
1114 let transformed = crate::prism::forward(model, &x[..cols]);
1115 let gpr = cols / GROUP_SIZE;
1116 match dtype {
1117 TensorDtype::Q2TiledP => {
1118 let v = Q4tpView::new_q2(self.quant_bytes(), self.rows(), cols);
1119 let mut sc = vec![0f32; gpr];
1120 v.scales_into(r, gpr, &mut sc);
1121 if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1122 return q2tp_affine_row_exact(v.nib, r, gpr, &transformed, &sc);
1123 }
1124 return q2tp_row_exact(v.nib, r, gpr, &transformed, &sc);
1125 }
1126 _ => {
1127 self.row_f32(r, scratch);
1128 return scratch.iter().zip(&transformed).map(|(w, v)| w * v).sum();
1129 }
1130 }
1131 }
1132 match dtype {
1133 TensorDtype::Q8Row => {
1134 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1135 dot_i8_f32(q, x) * row_scale[r]
1136 }
1137 TensorDtype::Q8_2f => {
1138 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1139 dot_i8_col_f32(q, x, col_field) * row_scale[r]
1140 }
1141 _ => {
1142 self.row_f32(r, scratch);
1143 scratch.iter().zip(x).map(|(w, v)| w * v).sum()
1144 }
1145 }
1146 }
1147 }
1148 }
1149
1150 pub fn matvec(&self, x: &[f32], out: &mut [f32], pool: Option<&Pool>) {
1153 match self {
1154 Self::F32 { data, .. } => {
1158 if !crate::f32_backend::matvec(data, x, out) {
1159 matvec_rows(pool, data, x, out);
1160 }
1161 }
1162 Self::Mapped {
1163 model,
1164 idx,
1165 dtype,
1166 rows,
1167 cols,
1168 row_scale,
1169 col_field,
1170 vbit_offsets,
1171 repack,
1172 } => {
1173 let _ = (model, idx);
1174 assert!(
1184 out.len() >= *rows && x.len() >= *cols,
1185 "matvec {rows}x{cols}: out {} (need {rows}), x {} (need {cols})",
1186 out.len(),
1187 x.len(),
1188 );
1189 let prism_forward =
1190 crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1191 if *dtype == TensorDtype::Q2TiledP
1192 && std::env::var("CMF_Q2TP_TRACE").as_deref() == Ok("1")
1193 {
1194 use std::sync::atomic::{AtomicUsize, Ordering};
1195 static N: AtomicUsize = AtomicUsize::new(0);
1196 let n = N.fetch_add(1, Ordering::Relaxed);
1197 if n < 128 {
1198 eprintln!(
1199 "q2tp-dispatch #{n} name={} prism={} rows={} cols={} gpu={} optin={} layer={}",
1200 model.tensors[*idx].name,
1201 prism_forward,
1202 rows,
1203 cols,
1204 crate::gpu::enabled_here(),
1205 crate::gpu::q2tp_gpu_opt_in(),
1206 crate::gpu::cur_layer(),
1207 );
1208 }
1209 }
1210 if prism_forward {
1215 let transformed = crate::prism::forward(model, &x[..*cols]);
1216 match dtype {
1217 TensorDtype::Q4Block => {
1218 q4matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1219 }
1220 TensorDtype::Q4Tiled => {
1221 q4t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1222 }
1223 TensorDtype::Q4TiledP => {
1224 q4tp_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1225 }
1226 TensorDtype::Q2TiledP => {
1227 let affine =
1228 crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1229 if *rows * *cols >= 8_388_608
1230 && crate::gpu::enabled_here()
1231 && crate::gpu::q2tp_gpu_opt_in()
1232 {
1233 let gpu_ok = if affine {
1234 crate::gpu::q2tp_affine_matvec(
1235 model,
1236 *idx,
1237 &transformed,
1238 *rows,
1239 *cols,
1240 out,
1241 )
1242 } else {
1243 crate::gpu::q2tp_matvec(
1244 model,
1245 *idx,
1246 &transformed,
1247 *rows,
1248 *cols,
1249 out,
1250 )
1251 };
1252 if gpu_ok {
1253 return;
1254 }
1255 }
1256 if affine {
1257 q2tp_affine_matvec(
1258 self.quant_bytes(),
1259 &transformed,
1260 *rows,
1261 *cols,
1262 out,
1263 pool,
1264 )
1265 } else {
1266 q2tp_matvec(
1267 self.quant_bytes(),
1268 &transformed,
1269 *rows,
1270 *cols,
1271 out,
1272 pool,
1273 )
1274 }
1275 }
1276 TensorDtype::Q1 => {
1277 q1_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1278 }
1279 TensorDtype::Q1T => {
1280 q1t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1281 }
1282 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec(
1283 self.quant_bytes(),
1284 vbit_offsets,
1285 &transformed,
1286 *rows,
1287 *cols,
1288 out,
1289 pool,
1290 ),
1291 TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec(
1292 self.quant_bytes(),
1293 repack,
1294 row_scale,
1295 &transformed,
1296 col_field,
1297 *dtype,
1298 *rows,
1299 *cols,
1300 out,
1301 pool,
1302 ),
1303 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1304 }
1305 return;
1306 }
1307 if *dtype == TensorDtype::Q4Block {
1308 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1312 let t0 = std::time::Instant::now();
1313 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1314 crate::gpu::ProbeArm::Gpu => {
1315 if crate::gpu::q4b_matvec(model, *idx, x, *rows, *cols, out) {
1316 crate::gpu::probe_record(
1317 crate::gpu::OpClass::Matvec,
1318 true,
1319 t0.elapsed(),
1320 );
1321 return;
1322 }
1323 }
1324 crate::gpu::ProbeArm::CpuTimed => {
1325 q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1326 crate::gpu::probe_record(
1327 crate::gpu::OpClass::Matvec,
1328 false,
1329 t0.elapsed(),
1330 );
1331 return;
1332 }
1333 crate::gpu::ProbeArm::Cpu => {}
1334 }
1335 }
1336 q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1337 return;
1338 }
1339 if *dtype == TensorDtype::Q4Tiled {
1340 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1345 let t0 = std::time::Instant::now();
1346 let cls = crate::gpu::matvec_class(*rows, *cols);
1347 match crate::gpu::probe_arm(cls) {
1348 crate::gpu::ProbeArm::Gpu => {
1349 if crate::gpu::q4t_matvec(model, *idx, x, *rows, *cols, out) {
1350 crate::gpu::probe_record(cls, true, t0.elapsed());
1351 return;
1352 }
1353 }
1354 crate::gpu::ProbeArm::CpuTimed => {
1355 q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1356 crate::gpu::probe_record(cls, false, t0.elapsed());
1357 return;
1358 }
1359 crate::gpu::ProbeArm::Cpu => {}
1360 }
1361 }
1362 q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1363 return;
1364 }
1365 if *dtype == TensorDtype::Q4TiledP {
1366 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1372 let t0 = std::time::Instant::now();
1373 let cls = crate::gpu::matvec_class(*rows, *cols);
1374 match crate::gpu::probe_arm(cls) {
1375 crate::gpu::ProbeArm::Gpu => {
1376 if crate::gpu::q4tp_matvec(model, *idx, x, *rows, *cols, out) {
1377 crate::gpu::probe_record(cls, true, t0.elapsed());
1378 return;
1379 }
1380 }
1381 crate::gpu::ProbeArm::CpuTimed => {
1382 q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1383 crate::gpu::probe_record(cls, false, t0.elapsed());
1384 return;
1385 }
1386 crate::gpu::ProbeArm::Cpu => {}
1387 }
1388 }
1389 q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1390 return;
1391 }
1392 if *dtype == TensorDtype::Q2TiledP {
1393 q2tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1394 return;
1395 }
1396 if *dtype == TensorDtype::Q1 {
1397 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1402 let t0 = std::time::Instant::now();
1403 let arm = if crate::gpu::q1_force() {
1404 crate::gpu::ProbeArm::Gpu
1405 } else {
1406 crate::gpu::probe_arm(crate::gpu::OpClass::Matvec)
1407 };
1408 match arm {
1409 crate::gpu::ProbeArm::Gpu => {
1410 if crate::gpu::q1_matvec(model, *idx, x, *rows, *cols, out) {
1411 crate::gpu::probe_record(
1412 crate::gpu::OpClass::Matvec,
1413 true,
1414 t0.elapsed(),
1415 );
1416 return;
1417 }
1418 }
1419 crate::gpu::ProbeArm::CpuTimed => {
1420 q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1421 crate::gpu::probe_record(
1422 crate::gpu::OpClass::Matvec,
1423 false,
1424 t0.elapsed(),
1425 );
1426 return;
1427 }
1428 crate::gpu::ProbeArm::Cpu => {}
1429 }
1430 }
1431 q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1432 return;
1433 }
1434 if *dtype == TensorDtype::Q1T {
1435 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1439 let t0 = std::time::Instant::now();
1440 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1441 crate::gpu::ProbeArm::Gpu => {
1442 if crate::gpu::q1t_matvec(model, *idx, x, *rows, *cols, out) {
1443 q1t_add_overlay(self.quant_bytes(), x, *rows, *cols, out, pool);
1444 crate::gpu::probe_record(
1445 crate::gpu::OpClass::Matvec,
1446 true,
1447 t0.elapsed(),
1448 );
1449 return;
1450 }
1451 }
1452 crate::gpu::ProbeArm::CpuTimed => {
1453 q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1454 crate::gpu::probe_record(
1455 crate::gpu::OpClass::Matvec,
1456 false,
1457 t0.elapsed(),
1458 );
1459 return;
1460 }
1461 crate::gpu::ProbeArm::Cpu => {}
1462 }
1463 }
1464 q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1465 return;
1466 }
1467 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1468 vbitmatvec(self.quant_bytes(), vbit_offsets, x, *rows, *cols, out, pool);
1469 return;
1470 }
1471 let xs = prescale(x, col_field, *dtype);
1472 if *rows >= crate::gpu::min_rows()
1477 && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
1478 && gpu_lmhead_enabled()
1479 && crate::gpu::enabled_here()
1480 {
1481 let t0 = std::time::Instant::now();
1484 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1485 crate::gpu::ProbeArm::Gpu => {}
1486 crate::gpu::ProbeArm::CpuTimed => {
1487 qmatvec(
1488 self.quant_bytes(),
1489 repack,
1490 row_scale,
1491 x,
1492 col_field,
1493 *dtype,
1494 *rows,
1495 *cols,
1496 out,
1497 pool,
1498 );
1499 crate::gpu::probe_record(
1500 crate::gpu::OpClass::Matvec,
1501 false,
1502 t0.elapsed(),
1503 );
1504 return;
1505 }
1506 crate::gpu::ProbeArm::Cpu => {
1507 qmatvec(
1508 self.quant_bytes(),
1509 repack,
1510 row_scale,
1511 x,
1512 col_field,
1513 *dtype,
1514 *rows,
1515 *cols,
1516 out,
1517 pool,
1518 );
1519 return;
1520 }
1521 }
1522 let frac = gpu_split_frac();
1523 let cpu_rows = ((*rows as f32) * (1.0 - frac)) as usize;
1524 let (out_cpu, out_gpu) = out.split_at_mut(cpu_rows);
1525 let bytes = self.quant_bytes();
1526 let ok = std::thread::scope(|sc| {
1527 let g = sc.spawn(|| {
1528 crate::gpu::q8_matvec_range(
1529 model,
1530 *idx,
1531 cpu_rows,
1532 &row_scale[cpu_rows..],
1533 &xs,
1534 *rows - cpu_rows,
1535 *cols,
1536 out_gpu,
1537 )
1538 });
1539 if cpu_rows > 0 {
1540 let rep_cpu = if repack.is_empty() {
1543 &[][..]
1544 } else {
1545 &repack[..(cpu_rows / 4) * 4 * *cols]
1546 };
1547 qmatvec(
1548 &bytes[..cpu_rows * *cols],
1549 rep_cpu,
1550 &row_scale[..cpu_rows],
1551 x,
1552 col_field,
1553 *dtype,
1554 cpu_rows,
1555 *cols,
1556 out_cpu,
1557 pool,
1558 );
1559 }
1560 g.join().unwrap_or(false)
1561 });
1562 if ok {
1563 crate::gpu::probe_record(crate::gpu::OpClass::Matvec, true, t0.elapsed());
1564 return;
1565 }
1566 qmatvec(
1569 &bytes[cpu_rows * *cols..(*rows) * *cols],
1570 &[],
1571 &row_scale[cpu_rows..],
1572 x,
1573 col_field,
1574 *dtype,
1575 *rows - cpu_rows,
1576 *cols,
1577 out_gpu,
1578 pool,
1579 );
1580 return;
1581 }
1582 qmatvec(
1583 self.quant_bytes(),
1584 repack,
1585 row_scale,
1586 x,
1587 col_field,
1588 *dtype,
1589 *rows,
1590 *cols,
1591 out,
1592 pool,
1593 );
1594 }
1595 }
1596 }
1597
1598 pub fn matvec2(
1600 &self,
1601 x1: &[f32],
1602 x2: &[f32],
1603 o1: &mut [f32],
1604 o2: &mut [f32],
1605 pool: Option<&Pool>,
1606 ) {
1607 match self {
1608 Self::F32 { data, .. } => matvec_rows2(pool, data, x1, x2, o1, o2),
1609 Self::Mapped {
1610 model,
1611 idx,
1612 dtype,
1613 rows,
1614 cols,
1615 row_scale,
1616 col_field,
1617 vbit_offsets,
1618 ..
1619 } => {
1620 if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1621 let tx1 = crate::prism::forward(model, &x1[..*cols]);
1622 let tx2 = crate::prism::forward(model, &x2[..*cols]);
1623 match dtype {
1624 TensorDtype::Q4Block => {
1625 q4matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1626 }
1627 TensorDtype::Q4Tiled => {
1628 q4t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1629 }
1630 TensorDtype::Q4TiledP => {
1631 q4tp_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1632 }
1633 TensorDtype::Q2TiledP => {
1634 if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1635 q2tp_affine_matvec2(
1636 self.quant_bytes(),
1637 &tx1,
1638 &tx2,
1639 *rows,
1640 *cols,
1641 o1,
1642 o2,
1643 pool,
1644 )
1645 } else {
1646 q2tp_matvec2(
1647 self.quant_bytes(),
1648 &tx1,
1649 &tx2,
1650 *rows,
1651 *cols,
1652 o1,
1653 o2,
1654 pool,
1655 )
1656 }
1657 }
1658 TensorDtype::Q1 => {
1659 q1_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1660 }
1661 TensorDtype::Q1T => {
1662 q1t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1663 }
1664 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec2(
1665 self.quant_bytes(),
1666 vbit_offsets,
1667 &tx1,
1668 &tx2,
1669 *rows,
1670 *cols,
1671 o1,
1672 o2,
1673 pool,
1674 ),
1675 TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec2(
1676 self.quant_bytes(),
1677 row_scale,
1678 &tx1,
1679 &tx2,
1680 col_field,
1681 *dtype,
1682 *rows,
1683 *cols,
1684 o1,
1685 o2,
1686 pool,
1687 ),
1688 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1689 }
1690 return;
1691 }
1692 if *dtype == TensorDtype::Q4Block {
1693 q4matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1694 return;
1695 }
1696 if *dtype == TensorDtype::Q4Tiled {
1697 q4t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1698 return;
1699 }
1700 if *dtype == TensorDtype::Q4TiledP {
1701 q4tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1702 return;
1703 }
1704 if *dtype == TensorDtype::Q2TiledP {
1705 q2tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1706 return;
1707 }
1708 if *dtype == TensorDtype::Q1 {
1709 q1_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1710 return;
1711 }
1712 if *dtype == TensorDtype::Q1T {
1713 q1t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1719 return;
1720 }
1721 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1722 vbitmatvec2(
1723 self.quant_bytes(),
1724 vbit_offsets,
1725 x1,
1726 x2,
1727 *rows,
1728 *cols,
1729 o1,
1730 o2,
1731 pool,
1732 );
1733 return;
1734 }
1735 qmatvec2(
1736 self.quant_bytes(),
1737 row_scale,
1738 x1,
1739 x2,
1740 col_field,
1741 *dtype,
1742 *rows,
1743 *cols,
1744 o1,
1745 o2,
1746 pool,
1747 );
1748 }
1749 }
1750 }
1751}
1752
1753impl QTensor {
1754 pub fn q4tp_mapped(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
1762 if self.has_prism_contract() {
1763 return None;
1764 }
1765 match self {
1766 Self::Mapped {
1767 model, idx, dtype, ..
1768 } if *dtype == TensorDtype::Q4TiledP => Some((model, *idx)),
1769 _ => None,
1770 }
1771 }
1772
1773 pub fn matmat(&self, xs_all: &[f32], b: usize, out: &mut [f32], pool: Option<&Pool>) {
1774 let cols = self.cols();
1775 let rows = self.rows();
1776 debug_assert_eq!(xs_all.len(), b * cols);
1777 debug_assert_eq!(out.len(), b * rows);
1778 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Matmat);
1779 if crate::gptq_capture::capturing() {
1783 if let Self::Mapped { model, idx, .. } = self {
1784 crate::gptq_capture::accumulate(&model.tensors[*idx].name, xs_all, b, cols);
1785 }
1786 }
1787 match self {
1788 Self::F32 { data, .. } => {
1789 if crate::f32_backend::matmat(data, xs_all, b, rows, cols, out) {
1790 return;
1791 }
1792 let out_addr = SendMut(out.as_mut_ptr());
1793 let run = |start: usize, end: usize| {
1800 let mut bi = start;
1801 while bi < end {
1802 let n = (end - bi).min(4);
1803 for o in 0..rows {
1804 let row = &data[o * cols..(o + 1) * cols];
1805 if n == 4 {
1806 let x = |k: usize| &xs_all[(bi + k) * cols..(bi + k + 1) * cols];
1807 let (x0, x1, x2, x3) = (x(0), x(1), x(2), x(3));
1808 let (mut a0, mut a1, mut a2, mut a3) = (0f32, 0f32, 0f32, 0f32);
1809 for j in 0..cols {
1810 let w = row[j];
1811 a0 += w * x0[j];
1812 a1 += w * x1[j];
1813 a2 += w * x2[j];
1814 a3 += w * x3[j];
1815 }
1816 for (k, acc) in [a0, a1, a2, a3].into_iter().enumerate() {
1817 unsafe { *out_addr.at((bi + k) * rows + o) = acc };
1818 }
1819 } else {
1820 for k in 0..n {
1821 let x = &xs_all[(bi + k) * cols..(bi + k + 1) * cols];
1822 let mut acc = 0f32;
1823 for j in 0..cols {
1824 acc += row[j] * x[j];
1825 }
1826 unsafe { *out_addr.at((bi + k) * rows + o) = acc };
1827 }
1828 }
1829 }
1830 bi += n;
1831 }
1832 };
1833 dispatch_rows(pool, b, &run);
1834 }
1835 Self::Mapped {
1836 model,
1837 idx,
1838 dtype,
1839 row_scale,
1840 col_field,
1841 vbit_offsets,
1842 ..
1843 } => {
1844 if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1845 let mut transformed = Vec::with_capacity(xs_all.len());
1846 for bi in 0..b {
1847 transformed.extend_from_slice(&crate::prism::forward(
1848 model,
1849 &xs_all[bi * cols..(bi + 1) * cols],
1850 ));
1851 }
1852 match dtype {
1853 TensorDtype::Q4Block => {
1854 q4matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1855 }
1856 TensorDtype::Q4Tiled => {
1857 q4t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1858 }
1859 TensorDtype::Q4TiledP => {
1860 q4tp_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1861 }
1862 TensorDtype::Q2TiledP => {
1863 let affine =
1864 crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1865 let gpu_batch_ok = if affine {
1871 b >= 2
1872 } else {
1873 b >= 32 && b * rows * cols >= 128_000_000
1874 };
1875 if gpu_batch_ok
1876 && cols % 32 == 0
1877 && crate::gpu::enabled_here()
1878 && crate::gpu::q2tp_gpu_opt_in()
1879 {
1880 let gpu_ok = if affine {
1881 crate::gpu::q2tp_affine_matmat(
1882 model,
1883 *idx,
1884 &transformed,
1885 b,
1886 rows,
1887 cols,
1888 out,
1889 )
1890 } else {
1891 crate::gpu::q2tp_matmat(
1892 model,
1893 *idx,
1894 &transformed,
1895 b,
1896 rows,
1897 cols,
1898 out,
1899 )
1900 };
1901 if gpu_ok {
1902 return;
1903 }
1904 }
1905 if affine
1909 && b == 1
1910 && cols % 32 == 0
1911 && crate::gpu::enabled_here()
1912 && crate::gpu::q2tp_gpu_opt_in()
1913 && crate::gpu::q2tp_affine_matvec(
1914 model,
1915 *idx,
1916 &transformed[..cols],
1917 rows,
1918 cols,
1919 &mut out[..rows],
1920 )
1921 {
1922 return;
1923 }
1924 if affine {
1925 q2tp_affine_matmat(
1926 self.quant_bytes(),
1927 &transformed,
1928 b,
1929 rows,
1930 cols,
1931 out,
1932 pool,
1933 )
1934 } else {
1935 q2tp_matmat(
1936 self.quant_bytes(),
1937 &transformed,
1938 b,
1939 rows,
1940 cols,
1941 out,
1942 pool,
1943 )
1944 }
1945 }
1946 TensorDtype::Q1 => {
1947 q1_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1948 }
1949 TensorDtype::Q1T => {
1950 q1t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1951 }
1952 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatmat(
1953 self.quant_bytes(),
1954 vbit_offsets,
1955 &transformed,
1956 b,
1957 rows,
1958 cols,
1959 out,
1960 pool,
1961 ),
1962 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
1963 let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1964 .map(|bi| {
1965 prescale(
1966 &transformed[bi * cols..(bi + 1) * cols],
1967 col_field,
1968 *dtype,
1969 )
1970 })
1971 .collect();
1972 qmatmat(self.quant_bytes(), row_scale, &pre, rows, cols, out, pool)
1973 }
1974 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1975 }
1976 return;
1977 }
1978 if *dtype == TensorDtype::Q4Block {
1979 q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1980 return;
1981 }
1982 if *dtype == TensorDtype::Q4TiledP {
1983 if b >= 32
1997 && b * rows * cols >= 128_000_000
1998 && cols % 32 == 0
1999 && !row_exact()
2000 && !crate::gpu::mm_killed()
2001 && crate::gpu::enabled_here()
2002 {
2003 let class = if b >= 128 {
2004 crate::gpu::OpClass::MatmatWide
2005 } else {
2006 crate::gpu::OpClass::Matmat
2007 };
2008 if let Self::Mapped { model, idx, .. } = self {
2009 if crate::mm_ab::on() {
2021 let mut g = vec![0f32; b * rows];
2022 let t = std::time::Instant::now();
2023 let took = crate::gpu::q4tp_matmat(
2024 model, *idx, xs_all, b, rows, cols, &mut g,
2025 );
2026 let dg = t.elapsed();
2027 let t = std::time::Instant::now();
2028 q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2029 let dc = t.elapsed();
2030 crate::mm_ab::record(b, rows, cols, took, dg, dc, &g, out);
2031 return;
2032 }
2033 let t0 = std::time::Instant::now();
2034 let resident = crate::gpu::weight_is_resident(model, *idx);
2038 match crate::gpu::probe_arm_cold_prefers_gpu(class, resident) {
2039 crate::gpu::ProbeArm::Gpu => {
2040 if crate::gpu::q4tp_matmat(
2041 model, *idx, xs_all, b, rows, cols, out,
2042 ) {
2043 let el = t0.elapsed();
2044 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2054 let budget = std::time::Duration::from_secs_f64(
2055 flops / 1.5e12 * 8.0 + 0.020,
2056 );
2057 crate::gpu::mm_budget_check(
2058 "q4tp matmat",
2059 el,
2060 budget,
2061 crate::gpu::probe_was_cold() || !resident,
2062 );
2063 crate::gpu::probe_record(class, true, el);
2064 return;
2065 }
2066 }
2067 crate::gpu::ProbeArm::CpuTimed => {
2068 q4tp_matmat(
2069 self.quant_bytes(),
2070 xs_all,
2071 b,
2072 rows,
2073 cols,
2074 out,
2075 pool,
2076 );
2077 crate::gpu::probe_record(class, false, t0.elapsed());
2078 return;
2079 }
2080 crate::gpu::ProbeArm::Cpu => {}
2081 }
2082 }
2083 }
2084 q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2085 return;
2086 }
2087 if *dtype == TensorDtype::Q2TiledP {
2088 if b >= 32
2094 && b * rows * cols >= 128_000_000
2095 && cols % 32 == 0
2096 && !crate::gpu::mm_killed()
2097 && crate::gpu::enabled_here()
2098 {
2099 let class = if b >= 128 {
2100 crate::gpu::OpClass::MatmatWide
2101 } else {
2102 crate::gpu::OpClass::Matmat
2103 };
2104 if let Self::Mapped { model, idx, .. } = self {
2105 let t0 = std::time::Instant::now();
2106 match crate::gpu::probe_arm(class) {
2107 crate::gpu::ProbeArm::Gpu => {
2108 if crate::gpu::q2tp_matmat(
2109 model, *idx, xs_all, b, rows, cols, out,
2110 ) {
2111 crate::gpu::probe_record(class, true, t0.elapsed());
2112 return;
2113 }
2114 }
2115 crate::gpu::ProbeArm::CpuTimed => {
2116 q2tp_matmat(
2117 self.quant_bytes(),
2118 xs_all,
2119 b,
2120 rows,
2121 cols,
2122 out,
2123 pool,
2124 );
2125 crate::gpu::probe_record(class, false, t0.elapsed());
2126 return;
2127 }
2128 crate::gpu::ProbeArm::Cpu => {}
2129 }
2130 }
2131 }
2132 q2tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2137 return;
2138 }
2139 if *dtype == TensorDtype::Q4Tiled {
2140 if b >= 32
2152 && b * rows * cols >= 128_000_000
2153 && cols % 32 == 0
2154 && !crate::gpu::mm_killed()
2155 && crate::gpu::enabled_here()
2156 {
2157 let class = if b >= 128 {
2158 crate::gpu::OpClass::MatmatWide
2159 } else {
2160 crate::gpu::OpClass::Matmat
2161 };
2162 if let Self::Mapped { model, idx, .. } = self {
2163 let t0 = std::time::Instant::now();
2164 match crate::gpu::probe_arm(class) {
2165 crate::gpu::ProbeArm::Gpu => {
2166 if crate::gpu::q4t_matmat(
2167 model, *idx, xs_all, b, rows, cols, out,
2168 ) {
2169 let el = t0.elapsed();
2170 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2180 let budget = std::time::Duration::from_secs_f64(
2181 flops / 1.5e12 * 8.0 + 0.020,
2182 );
2183 crate::gpu::mm_budget_check(
2184 "q4t matmat",
2185 el,
2186 budget,
2187 crate::gpu::probe_was_cold(),
2188 );
2189 crate::gpu::probe_record(class, true, el);
2190 return;
2191 }
2192 }
2193 crate::gpu::ProbeArm::CpuTimed => {
2194 q4t_matmat(
2195 self.quant_bytes(),
2196 xs_all,
2197 b,
2198 rows,
2199 cols,
2200 out,
2201 pool,
2202 );
2203 crate::gpu::probe_record(class, false, t0.elapsed());
2204 return;
2205 }
2206 crate::gpu::ProbeArm::Cpu => {}
2207 }
2208 }
2209 }
2210 q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2211 return;
2212 }
2213 if *dtype == TensorDtype::Q1 {
2214 if b >= 32
2217 && b * rows * cols >= 128_000_000
2218 && cols % 64 == 0
2219 && crate::gpu::enabled_here()
2220 {
2221 if let Self::Mapped { model, idx, .. } = self {
2222 let t0 = std::time::Instant::now();
2223 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2224 crate::gpu::ProbeArm::Gpu => {
2225 if crate::gpu::q1_matmat(
2226 model, *idx, xs_all, b, rows, cols, out,
2227 ) {
2228 crate::gpu::probe_record(
2229 crate::gpu::OpClass::Matmat,
2230 true,
2231 t0.elapsed(),
2232 );
2233 return;
2234 }
2235 }
2236 crate::gpu::ProbeArm::CpuTimed => {
2237 q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2238 crate::gpu::probe_record(
2239 crate::gpu::OpClass::Matmat,
2240 false,
2241 t0.elapsed(),
2242 );
2243 return;
2244 }
2245 crate::gpu::ProbeArm::Cpu => {}
2246 }
2247 }
2248 }
2249 q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2250 return;
2251 }
2252 if *dtype == TensorDtype::Q1T {
2253 if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2256 if let Self::Mapped { model, idx, .. } = self {
2257 let t0 = std::time::Instant::now();
2258 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2259 crate::gpu::ProbeArm::Gpu => {
2260 if crate::gpu::q1t_matmat(
2261 model, *idx, xs_all, b, rows, cols, out,
2262 ) {
2263 crate::gpu::probe_record(
2264 crate::gpu::OpClass::Matmat,
2265 true,
2266 t0.elapsed(),
2267 );
2268 return;
2269 }
2270 }
2271 crate::gpu::ProbeArm::CpuTimed => {
2272 q1t_matmat(
2273 self.quant_bytes(),
2274 xs_all,
2275 b,
2276 rows,
2277 cols,
2278 out,
2279 pool,
2280 );
2281 crate::gpu::probe_record(
2282 crate::gpu::OpClass::Matmat,
2283 false,
2284 t0.elapsed(),
2285 );
2286 return;
2287 }
2288 crate::gpu::ProbeArm::Cpu => {}
2289 }
2290 }
2291 }
2292 q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2293 return;
2294 }
2295 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
2296 vbitmatmat(
2297 self.quant_bytes(),
2298 vbit_offsets,
2299 xs_all,
2300 b,
2301 rows,
2302 cols,
2303 out,
2304 pool,
2305 );
2306 return;
2307 }
2308 let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
2309 .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
2310 .collect();
2311 if row_exact()
2317 && (1..=4).contains(&b)
2318 && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
2319 && crate::gpu::enabled_here()
2320 && crate::gpu::wgpu_active()
2321 {
2322 let flat: Vec<f32> = pre.iter().flat_map(|v| v.iter().copied()).collect();
2323 if crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out) {
2324 return;
2325 }
2326 }
2327 if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2333 if let Self::Mapped { model, idx, .. } = self {
2334 let t0 = std::time::Instant::now();
2335 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2336 crate::gpu::ProbeArm::Gpu
2337 if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
2338 && !crate::gpu::q8_resident_or_upload(model, *idx) =>
2339 {
2340 let q = self.quant_bytes();
2344 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2345 return;
2346 }
2347 crate::gpu::ProbeArm::Gpu => {
2348 if *dtype == TensorDtype::Q8_2f
2356 && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2357 && crate::gpu::q8_matmat_2f(
2358 model, *idx, row_scale, col_field, xs_all, b, rows,
2359 cols, out,
2360 )
2361 {
2362 crate::gpu::probe_record(
2363 crate::gpu::OpClass::Matmat,
2364 true,
2365 t0.elapsed(),
2366 );
2367 return;
2368 }
2369 let flat: Vec<f32> =
2370 pre.iter().flat_map(|v| v.iter().copied()).collect();
2371 if crate::gpu::q8_matmat(
2372 model, *idx, row_scale, &flat, b, rows, cols, out,
2373 ) {
2374 crate::gpu::probe_record(
2375 crate::gpu::OpClass::Matmat,
2376 true,
2377 t0.elapsed(),
2378 );
2379 return;
2380 }
2381 }
2382 crate::gpu::ProbeArm::CpuTimed => {
2383 let q = self.quant_bytes();
2384 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2385 crate::gpu::probe_record(
2386 crate::gpu::OpClass::Matmat,
2387 false,
2388 t0.elapsed(),
2389 );
2390 return;
2391 }
2392 crate::gpu::ProbeArm::Cpu => {}
2393 }
2394 }
2395 }
2396 let q = self.quant_bytes();
2397 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2398 }
2399 }
2400 }
2401}
2402
2403impl QTensor {
2404 pub fn device_matmat(&self, xs: &[f32], b: usize, out: &mut [f32]) -> bool {
2414 let (rows, cols) = (self.rows(), self.cols());
2415 let Self::Mapped {
2416 model,
2417 idx,
2418 dtype,
2419 row_scale,
2420 col_field,
2421 ..
2422 } = self
2423 else {
2424 return false;
2425 };
2426 if crate::prism::has_contract(model) {
2427 return false;
2428 }
2429 match *dtype {
2430 TensorDtype::Q4TiledP => crate::gpu::q4tp_matmat(model, *idx, xs, b, rows, cols, out),
2431 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
2435 if *dtype == TensorDtype::Q8_2f
2438 && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2439 && crate::gpu::q8_matmat_2f(
2440 model, *idx, row_scale, col_field, xs, b, rows, cols, out,
2441 )
2442 {
2443 return true;
2444 }
2445 let flat: Vec<f32> = (0..b)
2446 .flat_map(|bi| {
2447 prescale(&xs[bi * cols..(bi + 1) * cols], col_field, *dtype).into_owned()
2448 })
2449 .collect();
2450 crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out)
2451 }
2452 _ => false,
2453 }
2454 }
2455
2456 pub fn matvec_many<const N: usize>(
2463 ts: [&QTensor; N],
2464 x: &[f32],
2465 mut outs: [&mut [f32]; N],
2466 pool: Option<&Pool>,
2467 ) {
2468 let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2469 if ts.iter().any(|t| t.has_prism_contract()) {
2470 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2474 t.matvec(x, o, pool);
2475 }
2476 return;
2477 }
2478 let uniform_q8 = ts.iter().all(|t| {
2479 matches!(
2480 t,
2481 Self::Mapped {
2482 dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2483 ..
2484 }
2485 )
2486 });
2487 let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2488 if uniform_f32 && crate::f32_backend::active() {
2489 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2490 t.matvec(x, o, pool);
2491 }
2492 return;
2493 }
2494 let uniform_q4 = ts.iter().all(|t| {
2495 matches!(
2496 t,
2497 Self::Mapped {
2498 dtype: TensorDtype::Q4Block,
2499 ..
2500 }
2501 )
2502 });
2503 let uniform_vbit = ts.iter().all(|t| {
2504 matches!(
2505 t,
2506 Self::Mapped {
2507 dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2508 ..
2509 }
2510 )
2511 });
2512 let uniform_q1 = ts.iter().all(|t| {
2513 matches!(
2514 t,
2515 Self::Mapped {
2516 dtype: TensorDtype::Q1,
2517 ..
2518 }
2519 )
2520 });
2521 let uniform_q1t = ts.iter().all(|t| {
2522 matches!(
2523 t,
2524 Self::Mapped {
2525 dtype: TensorDtype::Q1T,
2526 ..
2527 }
2528 )
2529 });
2530 let uniform_q4tp = ts.iter().all(|t| {
2535 matches!(
2536 t,
2537 Self::Mapped {
2538 dtype: TensorDtype::Q4TiledP,
2539 ..
2540 }
2541 )
2542 }) && ts
2543 .iter()
2544 .all(|t| t.cols() == ts[0].cols() && t.cols() % GROUP_SIZE == 0);
2545 let Some(pool) = pool else {
2546 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2547 t.matvec(x, o, None);
2548 }
2549 return;
2550 };
2551 if total_rows < 256
2552 || !(uniform_q8
2553 || uniform_f32
2554 || uniform_q4
2555 || uniform_vbit
2556 || uniform_q1
2557 || uniform_q1t
2558 || uniform_q4tp)
2559 {
2560 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2561 t.matvec(x, o, Some(pool));
2562 }
2563 return;
2564 }
2565
2566 if uniform_q4tp {
2567 let cols = ts[0].cols();
2573 let gpr = cols / GROUP_SIZE;
2574 let views: [Q4tpView; N] =
2575 std::array::from_fn(|i| Q4tpView::new(ts[i].quant_bytes(), ts[i].rows(), cols));
2576 let rows_of: [usize; N] = std::array::from_fn(|i| ts[i].rows());
2577 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2578 let locate = |flat: usize| -> (usize, usize) {
2580 let mut acc = 0;
2581 for (i, &r) in rows_of.iter().enumerate() {
2582 if flat < acc + r {
2583 return (i, flat - acc);
2584 }
2585 acc += r;
2586 }
2587 (rows_of.len() - 1, 0)
2588 };
2589 let (views, outs_addr) = (&views, &outs_addr);
2590 if a8w8_enabled() {
2591 let act = split_act(x);
2592 let act = &act;
2593 let run = |start: usize, end: usize| {
2594 with_krow(gpr, |sc| {
2595 for flat in start..end {
2596 let (t, r) = locate(flat);
2597 let v = &views[t];
2598 v.scales_into(r, gpr, sc);
2599 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
2600 for &(j, xv) in &act.outliers {
2601 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
2602 acc += w * s * xv;
2603 }
2604 unsafe { *outs_addr[t].at(r) = acc };
2606 }
2607 });
2608 };
2609 pool.run_rows(total_rows, &run);
2610 } else {
2611 let run = |start: usize, end: usize| {
2612 with_krow(gpr, |sc| {
2613 for flat in start..end {
2614 let (t, r) = locate(flat);
2615 let v = &views[t];
2616 v.scales_into(r, gpr, sc);
2617 unsafe { *outs_addr[t].at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
2619 }
2620 });
2621 };
2622 pool.run_rows(total_rows, &run);
2623 }
2624 return;
2625 }
2626
2627 if uniform_q1 {
2628 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2631 if a8w8_enabled() {
2632 let act = split_act(x);
2633 let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
2634 let (act, gsum) = (&act, &gsum);
2635 let closures: [_; N] = std::array::from_fn(|i| {
2636 let (bytes, gpr, out) =
2637 (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2638 move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
2639 });
2640 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2641 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2642 pool.run_many(&parts);
2643 } else {
2644 let closures: [_; N] = std::array::from_fn(|i| {
2645 let (bytes, gpr, out) =
2646 (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2647 move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
2648 });
2649 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2650 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2651 pool.run_many(&parts);
2652 }
2653 return;
2654 }
2655
2656 if uniform_q1t {
2657 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2661 const TILE: usize = cortiq_core::quant::Q1T_TILE;
2662 if a8w8_enabled() {
2663 let act = split_act(x);
2664 let act = &act;
2665 let x_ref = x;
2666 let closures: [_; N] = std::array::from_fn(|i| {
2667 let bytes = ts[i].quant_bytes();
2668 let (rows, cols) = (ts[i].rows(), ts[i].cols());
2669 let gpr = cols / GROUP_SIZE;
2670 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2671 let out = outs_addr[i];
2672 move |s: usize, e: usize| {
2673 q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
2674 }
2675 });
2676 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2677 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2678 pool.run_many(&parts);
2679 } else {
2680 let x_ref = x;
2681 let closures: [_; N] = std::array::from_fn(|i| {
2682 let bytes = ts[i].quant_bytes();
2683 let (rows, cols) = (ts[i].rows(), ts[i].cols());
2684 let gpr = cols / GROUP_SIZE;
2685 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2686 let out = outs_addr[i];
2687 move |s: usize, e: usize| {
2688 q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
2689 }
2690 });
2691 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2692 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2693 pool.run_many(&parts);
2694 }
2695 return;
2696 }
2697
2698 if uniform_q4 || uniform_vbit {
2699 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2700 if a8w8_enabled() {
2702 let act = split_act(x);
2703 let act = &act;
2704 if uniform_q4 {
2705 let closures: [_; N] = std::array::from_fn(|i| {
2706 let (packed, scales) =
2707 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2708 let (gpr, cols, out) =
2709 (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
2710 move |s: usize, e: usize| {
2711 q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
2712 }
2713 });
2714 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2715 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2716 pool.run_many(&parts);
2717 } else {
2718 let closures: [_; N] = std::array::from_fn(|i| {
2719 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2720 unreachable!()
2721 };
2722 let (bytes, rows, cols, out) = (
2723 ts[i].quant_bytes(),
2724 ts[i].rows(),
2725 ts[i].cols(),
2726 outs_addr[i],
2727 );
2728 move |s: usize, e: usize| {
2729 vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
2730 }
2731 });
2732 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2733 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2734 pool.run_many(&parts);
2735 }
2736 return;
2737 }
2738 if uniform_q4 {
2739 let closures: [_; N] = std::array::from_fn(|i| {
2740 let (packed, scales) =
2741 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2742 let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2743 move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
2744 });
2745 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2746 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2747 pool.run_many(&parts);
2748 } else {
2749 let closures: [_; N] = std::array::from_fn(|i| {
2750 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2751 unreachable!()
2752 };
2753 let (bytes, rows, cols, out) = (
2754 ts[i].quant_bytes(),
2755 ts[i].rows(),
2756 ts[i].cols(),
2757 outs_addr[i],
2758 );
2759 move |s: usize, e: usize| {
2760 vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
2761 }
2762 });
2763 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2764 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2765 pool.run_many(&parts);
2766 }
2767 return;
2768 }
2769
2770 if uniform_f32 {
2771 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2772 let closures: [_; N] = std::array::from_fn(|i| {
2773 let Self::F32 { data, cols, .. } = ts[i] else {
2774 unreachable!()
2775 };
2776 let out = outs_addr[i];
2777 move |start: usize, end: usize| {
2778 for o in start..end {
2779 let row = &data[o * cols..(o + 1) * cols];
2780 let mut sum = 0.0f32;
2781 for j in 0..*cols {
2782 sum += row[j] * x[j];
2783 }
2784 unsafe { *out.at(o) = sum };
2786 }
2787 }
2788 });
2789 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2790 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2791 pool.run_many(&parts);
2792 return;
2793 }
2794
2795 struct Ctx<'a> {
2798 bytes: &'a [u8],
2799 #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
2800 rep: &'a [u8],
2801 row_scale: &'a [f32],
2802 cols: usize,
2803 xs: std::borrow::Cow<'a, [f32]>,
2804 }
2805 let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2806 let Self::Mapped {
2807 dtype,
2808 cols,
2809 row_scale,
2810 col_field,
2811 repack,
2812 ..
2813 } = ts[i]
2814 else {
2815 unreachable!()
2816 };
2817 Ctx {
2818 bytes: ts[i].quant_bytes(),
2819 rep: repack,
2820 row_scale,
2821 cols: *cols,
2822 xs: prescale(x, col_field, *dtype),
2823 }
2824 });
2825 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2826 #[cfg(target_arch = "aarch64")]
2827 if sdot_enabled() {
2828 let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2829 let closures: [_; N] = std::array::from_fn(|i| {
2830 let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2831 move |start: usize, end: usize| {
2832 q8_range_sdot(c.bytes, c.rep, c.row_scale, act, c.cols, out, start, end)
2833 }
2834 });
2835 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2836 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2837 pool.run_many(&parts);
2838 return;
2839 }
2840 #[cfg(target_arch = "x86_64")]
2841 if avx2_a8w8_enabled() {
2842 let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2843 let closures: [_; N] = std::array::from_fn(|i| {
2844 let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2845 move |start: usize, end: usize| {
2846 q8_range_avx2(c.bytes, c.row_scale, act, c.cols, out, start, end)
2847 }
2848 });
2849 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2850 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2851 pool.run_many(&parts);
2852 return;
2853 }
2854 let closures: [_; N] = std::array::from_fn(|i| {
2855 let (c, out) = (&ctxs[i], outs_addr[i]);
2856 move |start: usize, end: usize| {
2857 q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
2858 }
2859 });
2860 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2861 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2862 pool.run_many(&parts);
2863 }
2864}
2865
2866impl QTensor {
2867 #[allow(clippy::needless_range_loop)]
2872 pub fn matvec2_many<const N: usize>(
2873 ts: [&QTensor; N],
2874 x1: &[f32],
2875 x2: &[f32],
2876 mut o1s: [&mut [f32]; N],
2877 mut o2s: [&mut [f32]; N],
2878 pool: Option<&Pool>,
2879 ) {
2880 let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2881 if ts.iter().any(|t| t.has_prism_contract()) {
2882 for i in 0..N {
2883 ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2884 }
2885 return;
2886 }
2887 let uniform_q8 = ts.iter().all(|t| {
2888 matches!(
2889 t,
2890 Self::Mapped {
2891 dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2892 ..
2893 }
2894 )
2895 });
2896 let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2897 let uniform_q4 = ts.iter().all(|t| {
2898 matches!(
2899 t,
2900 Self::Mapped {
2901 dtype: TensorDtype::Q4Block,
2902 ..
2903 }
2904 )
2905 });
2906 let uniform_vbit = ts.iter().all(|t| {
2907 matches!(
2908 t,
2909 Self::Mapped {
2910 dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2911 ..
2912 }
2913 )
2914 });
2915 let fusable = pool.is_some()
2916 && total_rows >= 256
2917 && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
2918 if !fusable {
2919 for i in 0..N {
2920 ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2921 }
2922 return;
2923 }
2924 let pool = pool.unwrap();
2925
2926 if uniform_q4 || uniform_vbit {
2927 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2928 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2929 if a8w8_enabled() {
2931 let a1 = split_act(x1);
2932 let a2 = split_act(x2);
2933 let (a1, a2) = (&a1, &a2);
2934 if uniform_q4 {
2935 let closures: [_; N] = std::array::from_fn(|i| {
2936 let (packed, scales) =
2937 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2938 let (gpr, cols, o1, o2) =
2939 (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
2940 move |s: usize, e: usize| {
2941 q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
2942 }
2943 });
2944 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2945 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2946 pool.run_many(&parts);
2947 } else {
2948 let closures: [_; N] = std::array::from_fn(|i| {
2949 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2950 unreachable!()
2951 };
2952 let (bytes, rows, cols, o1, o2) = (
2953 ts[i].quant_bytes(),
2954 ts[i].rows(),
2955 ts[i].cols(),
2956 p1[i],
2957 p2[i],
2958 );
2959 move |s: usize, e: usize| {
2960 vbit_range2_a8w8(
2961 bytes,
2962 vbit_offsets,
2963 x1,
2964 x2,
2965 a1,
2966 a2,
2967 rows,
2968 cols,
2969 o1,
2970 o2,
2971 s,
2972 e,
2973 )
2974 }
2975 });
2976 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2977 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2978 pool.run_many(&parts);
2979 }
2980 return;
2981 }
2982 if uniform_q4 {
2983 let closures: [_; N] = std::array::from_fn(|i| {
2984 let (packed, scales) =
2985 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2986 let (gpr, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
2987 move |s: usize, e: usize| {
2988 q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
2989 }
2990 });
2991 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2992 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2993 pool.run_many(&parts);
2994 } else {
2995 let closures: [_; N] = std::array::from_fn(|i| {
2996 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2997 unreachable!()
2998 };
2999 let (bytes, rows, cols, o1, o2) = (
3000 ts[i].quant_bytes(),
3001 ts[i].rows(),
3002 ts[i].cols(),
3003 p1[i],
3004 p2[i],
3005 );
3006 move |s: usize, e: usize| {
3007 vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
3008 }
3009 });
3010 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3011 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3012 pool.run_many(&parts);
3013 }
3014 return;
3015 }
3016
3017 if uniform_f32 {
3018 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3019 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3020 let closures: [_; N] = std::array::from_fn(|i| {
3021 let Self::F32 { data, cols, .. } = ts[i] else {
3022 unreachable!()
3023 };
3024 let (o1, o2) = (p1[i], p2[i]);
3025 move |start: usize, end: usize| {
3026 for o in start..end {
3027 let row = &data[o * cols..(o + 1) * cols];
3028 let (mut s1, mut s2) = (0.0f32, 0.0f32);
3029 for j in 0..*cols {
3030 s1 += row[j] * x1[j];
3031 s2 += row[j] * x2[j];
3032 }
3033 unsafe {
3035 *o1.at(o) = s1;
3036 *o2.at(o) = s2;
3037 }
3038 }
3039 }
3040 });
3041 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3042 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3043 pool.run_many(&parts);
3044 return;
3045 }
3046
3047 struct Ctx<'a> {
3048 bytes: &'a [u8],
3049 row_scale: &'a [f32],
3050 cols: usize,
3051 xs1: std::borrow::Cow<'a, [f32]>,
3052 xs2: std::borrow::Cow<'a, [f32]>,
3053 }
3054 let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
3055 let Self::Mapped {
3056 dtype,
3057 cols,
3058 row_scale,
3059 col_field,
3060 ..
3061 } = ts[i]
3062 else {
3063 unreachable!()
3064 };
3065 Ctx {
3066 bytes: ts[i].quant_bytes(),
3067 row_scale,
3068 cols: *cols,
3069 xs1: prescale(x1, col_field, *dtype),
3070 xs2: prescale(x2, col_field, *dtype),
3071 }
3072 });
3073 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3074 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3075 #[cfg(target_arch = "aarch64")]
3076 if sdot_enabled() {
3077 let acts: [(SplitAct, SplitAct); N] =
3078 std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3079 let closures: [_; N] = std::array::from_fn(|i| {
3080 let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3081 move |start: usize, end: usize| {
3082 q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3083 }
3084 });
3085 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3086 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3087 pool.run_many(&parts);
3088 return;
3089 }
3090 #[cfg(target_arch = "x86_64")]
3091 if avx2_a8w8_enabled() {
3092 let acts: [(SplitAct, SplitAct); N] =
3093 std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3094 let closures: [_; N] = std::array::from_fn(|i| {
3095 let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3096 move |start: usize, end: usize| {
3097 q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3098 }
3099 });
3100 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3101 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3102 pool.run_many(&parts);
3103 return;
3104 }
3105 let closures: [_; N] = std::array::from_fn(|i| {
3106 let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
3107 move |start: usize, end: usize| {
3108 q8_range2_f32(
3109 c.bytes,
3110 c.row_scale,
3111 &c.xs1,
3112 &c.xs2,
3113 c.cols,
3114 o1,
3115 o2,
3116 start,
3117 end,
3118 )
3119 }
3120 });
3121 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3122 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3123 pool.run_many(&parts);
3124 }
3125
3126 pub fn matvec_silu_mul(
3131 gate: &QTensor,
3132 up: &QTensor,
3133 x: &[f32],
3134 out: &mut [f32],
3135 pool: Option<&Pool>,
3136 ) -> bool {
3137 Self::matvec_silu_mul_limited(gate, up, x, out, 0.0, pool)
3138 }
3139
3140 pub fn matvec_silu_mul_limited(
3146 gate: &QTensor,
3147 up: &QTensor,
3148 x: &[f32],
3149 out: &mut [f32],
3150 limit: f32,
3151 pool: Option<&Pool>,
3152 ) -> bool {
3153 if gate.has_prism_contract() || up.has_prism_contract() {
3154 return false;
3158 }
3159 let inter = gate.rows();
3160 debug_assert_eq!(up.rows(), inter);
3161 debug_assert_eq!(out.len(), inter);
3162 debug_assert_eq!(gate.cols(), up.cols());
3163 if !a8w8_enabled() {
3164 return false;
3165 }
3166 let act = split_act(x);
3167 let act = &act;
3168 let x_ref = x;
3169 let out_addr = SendMut(out.as_mut_ptr());
3170
3171 match (gate, up) {
3172 (
3174 Self::Mapped {
3175 dtype: TensorDtype::Q4Block,
3176 ..
3177 },
3178 Self::Mapped {
3179 dtype: TensorDtype::Q4Block,
3180 ..
3181 },
3182 ) => {
3183 let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
3184 let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
3185 let gpr = gate.cols() / GROUP_SIZE;
3186 let cols = gate.cols();
3187 let run = move |start: usize, end: usize| {
3188 for r in start..end {
3189 let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
3190 let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
3191 for &(j, xv) in &act.outliers {
3192 let flat = r * cols + j;
3193 let gb = gp[flat / 2];
3194 let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
3195 let gsc = f16_to_f32(u16::from_le_bytes([
3196 gs[(flat / GROUP_SIZE) * 2],
3197 gs[(flat / GROUP_SIZE) * 2 + 1],
3198 ]));
3199 gv += ((gn as i32 - 8) as f32) * gsc * xv;
3200 let ub = up_p[flat / 2];
3201 let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
3202 let usc = f16_to_f32(u16::from_le_bytes([
3203 up_s[(flat / GROUP_SIZE) * 2],
3204 up_s[(flat / GROUP_SIZE) * 2 + 1],
3205 ]));
3206 uv += ((un as i32 - 8) as f32) * usc * xv;
3207 }
3208 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3210 }
3211 };
3212 dispatch_rows(pool, inter, &run);
3213 true
3214 }
3215 (
3219 Self::Mapped {
3220 dtype: TensorDtype::Q4Tiled,
3221 ..
3222 },
3223 Self::Mapped {
3224 dtype: TensorDtype::Q4Tiled,
3225 ..
3226 },
3227 ) => {
3228 let g_bytes = gate.quant_bytes();
3229 let u_bytes = up.quant_bytes();
3230 let gpr = gate.cols() / GROUP_SIZE;
3231 let run = move |start: usize, end: usize| {
3232 for r in start..end {
3233 let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3234 let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3235 for &(j, xv) in &act.outliers {
3236 let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
3237 gv += w * s * xv;
3238 let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
3239 uv += w * s * xv;
3240 }
3241 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3243 }
3244 };
3245 dispatch_rows(pool, inter, &run);
3246 true
3247 }
3248 (
3251 Self::Mapped {
3252 dtype: TensorDtype::Q4TiledP,
3253 ..
3254 },
3255 Self::Mapped {
3256 dtype: TensorDtype::Q4TiledP,
3257 ..
3258 },
3259 ) => {
3260 let cols = gate.cols();
3261 let gpr = cols / GROUP_SIZE;
3262 let gv_view = Q4tpView::new(gate.quant_bytes(), inter, cols);
3263 let uv_view = Q4tpView::new(up.quant_bytes(), inter, cols);
3264 let run = |start: usize, end: usize| {
3265 with_krows(gpr, |gsc, usc| {
3266 for r in start..end {
3267 gv_view.scales_into(r, gpr, gsc);
3268 uv_view.scales_into(r, gpr, usc);
3269 let mut gv =
3270 dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx;
3271 let mut uv =
3272 dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx;
3273 for &(j, xv) in &act.outliers {
3274 let (w, s) = q4tp_outlier(gv_view.nib, r, gpr, j, gsc);
3275 gv += w * s * xv;
3276 let (w, s) = q4tp_outlier(uv_view.nib, r, gpr, j, usc);
3277 uv += w * s * xv;
3278 }
3279 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3281 }
3282 });
3283 };
3284 dispatch_rows(pool, inter, &run);
3285 true
3286 }
3287 (
3293 Self::Mapped {
3294 dtype: TensorDtype::Q1,
3295 ..
3296 },
3297 Self::Mapped {
3298 dtype: TensorDtype::Q1,
3299 ..
3300 },
3301 ) => {
3302 let g_bytes = gate.quant_bytes();
3303 let u_bytes = up.quant_bytes();
3304 let gpr = gate.cols() / GROUP_SIZE;
3305 let gsum = q1_group_sums(&act.xq, gpr);
3306 let gsum = &gsum;
3307 let run = move |start: usize, end: usize| {
3308 for r in start..end {
3309 let mut gv = dot_q1_row_i8(g_bytes, r, gpr, &act.xq, gsum) * act.sx;
3310 let mut uv = dot_q1_row_i8(u_bytes, r, gpr, &act.xq, gsum) * act.sx;
3311 for &(j, xv) in &act.outliers {
3312 let (w, s) = q1_outlier(g_bytes, r, gpr, j);
3313 gv += w * s * xv;
3314 let (w, s) = q1_outlier(u_bytes, r, gpr, j);
3315 uv += w * s * xv;
3316 }
3317 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3319 }
3320 };
3321 dispatch_rows(pool, inter, &run);
3322 true
3323 }
3324 (
3328 Self::Mapped {
3329 dtype: TensorDtype::Q2TiledP,
3330 ..
3331 },
3332 Self::Mapped {
3333 dtype: TensorDtype::Q2TiledP,
3334 ..
3335 },
3336 ) => {
3337 let cols = gate.cols();
3338 let gpr = cols / GROUP_SIZE;
3339 let gv_view = Q4tpView::new_q2(gate.quant_bytes(), inter, cols);
3340 let uv_view = Q4tpView::new_q2(up.quant_bytes(), inter, cols);
3341 let gsum = q1_group_sums(&act.xq, gpr);
3342 let gsum = &gsum;
3343 let run = move |start: usize, end: usize| {
3344 with_krows(gpr, |gsc, usc| {
3345 for r in start..end {
3346 gv_view.scales_into(r, gpr, gsc);
3347 uv_view.scales_into(r, gpr, usc);
3348 let mut gv =
3349 dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx;
3350 let mut uv =
3351 dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx;
3352 for &(j, xv) in &act.outliers {
3353 let (w, s) = q2tp_outlier(gv_view.nib, r, gpr, j, gsc);
3354 gv += w * s * xv;
3355 let (w, s) = q2tp_outlier(uv_view.nib, r, gpr, j, usc);
3356 uv += w * s * xv;
3357 }
3358 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3360 }
3361 });
3362 };
3363 dispatch_rows(pool, inter, &run);
3364 true
3365 }
3366 (
3371 Self::Mapped {
3372 dtype: TensorDtype::Q8Row,
3373 row_scale: g_rs,
3374 ..
3375 },
3376 Self::Mapped {
3377 dtype: TensorDtype::Q8Row,
3378 row_scale: u_rs,
3379 ..
3380 },
3381 ) => {
3382 let g_bytes = gate.quant_bytes();
3383 let u_bytes = up.quant_bytes();
3384 let cols = gate.cols();
3385 let run = move |start: usize, end: usize| {
3386 for r in start..end {
3387 let gv = q8_row_dot(&g_bytes[r * cols..(r + 1) * cols], act) * g_rs[r];
3388 let uv = q8_row_dot(&u_bytes[r * cols..(r + 1) * cols], act) * u_rs[r];
3389 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3391 }
3392 };
3393 dispatch_rows(pool, inter, &run);
3394 true
3395 }
3396 (
3398 Self::Mapped {
3399 dtype: TensorDtype::Q1T,
3400 ..
3401 },
3402 Self::Mapped {
3403 dtype: TensorDtype::Q1T,
3404 ..
3405 },
3406 ) => {
3407 const TILE: usize = cortiq_core::quant::Q1T_TILE;
3408 let g_bytes = gate.quant_bytes();
3409 let u_bytes = up.quant_bytes();
3410 let gpr = gate.cols() / GROUP_SIZE;
3411 let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
3412 let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
3413 let run = move |start: usize, end: usize| {
3414 for r in start..end {
3415 let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3416 let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3417 for &(j, xv) in &act.outliers {
3418 gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
3419 uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
3420 }
3421 gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
3422 uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
3423 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3425 }
3426 };
3427 dispatch_rows(pool, inter, &run);
3428 true
3429 }
3430 _ => false,
3431 }
3432 }
3433
3434 pub fn moe_gate_up_many(
3449 pairs: &[(&QTensor, &QTensor)],
3450 x: &[f32],
3451 outs: &mut [Vec<f32>],
3452 pool: Option<&Pool>,
3453 ) -> bool {
3454 Self::moe_gate_up_many_limited(pairs, x, outs, 0.0, pool)
3455 }
3456
3457 pub fn moe_gate_up_many_limited(
3461 pairs: &[(&QTensor, &QTensor)],
3462 x: &[f32],
3463 outs: &mut [Vec<f32>],
3464 limit: f32,
3465 pool: Option<&Pool>,
3466 ) -> bool {
3467 if pairs.is_empty() || pairs.len() != outs.len() {
3468 return false;
3469 }
3470 if !a8w8_enabled() {
3471 if limit > 0.0 {
3472 return false;
3475 }
3476 let groups = vec![vec![0]; pairs.len()];
3477 return Self::moe_gate_up_rows(pairs, &groups, x, outs, pool);
3478 }
3479 let inter = pairs[0].0.rows();
3480 let cols = pairs[0].0.cols();
3481 if cols % GROUP_SIZE != 0 {
3482 return false;
3483 }
3484 let gpr = cols / GROUP_SIZE;
3485 let q2 = matches!(
3488 pairs[0].0,
3489 Self::Mapped {
3490 dtype: TensorDtype::Q2TiledP,
3491 ..
3492 }
3493 );
3494 let want = if q2 {
3495 TensorDtype::Q2TiledP
3496 } else {
3497 TensorDtype::Q4TiledP
3498 };
3499 let mut views = Vec::with_capacity(pairs.len() * 2);
3500 for ((g, u), o) in pairs.iter().zip(outs.iter()) {
3501 let both = matches!(g, Self::Mapped { dtype, .. } if *dtype == want)
3502 && matches!(u, Self::Mapped { dtype, .. } if *dtype == want);
3503 if !both
3504 || g.rows() != inter
3505 || u.rows() != inter
3506 || g.cols() != cols
3507 || u.cols() != cols
3508 || o.len() != inter
3509 {
3510 return false;
3511 }
3512 let mk = if q2 { Q4tpView::new_q2 } else { Q4tpView::new };
3513 views.push(mk(g.quant_bytes(), inter, cols));
3514 views.push(mk(u.quant_bytes(), inter, cols));
3515 }
3516 let act = split_act(x);
3517 let gsum = if q2 {
3518 q1_group_sums(&act.xq, gpr)
3519 } else {
3520 Vec::new()
3521 };
3522 let (act, gsum) = (&act, &gsum);
3523 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3524 let (views, ptrs) = (&views, &ptrs);
3525 let run = |start: usize, end: usize| {
3526 with_krows(gpr, |gsc, usc| {
3527 for flat in start..end {
3528 let (e, r) = (flat / inter, flat % inter);
3529 let gv_view = &views[e * 2];
3530 let uv_view = &views[e * 2 + 1];
3531 gv_view.scales_into(r, gpr, gsc);
3532 uv_view.scales_into(r, gpr, usc);
3533 let (mut gv, mut uv) = if q2 {
3534 (
3535 dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx,
3536 dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx,
3537 )
3538 } else {
3539 (
3540 dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx,
3541 dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx,
3542 )
3543 };
3544 for &(j, xv) in &act.outliers {
3545 let (og, ou) = if q2 {
3546 (
3547 q2tp_outlier(gv_view.nib, r, gpr, j, gsc),
3548 q2tp_outlier(uv_view.nib, r, gpr, j, usc),
3549 )
3550 } else {
3551 (
3552 q4tp_outlier(gv_view.nib, r, gpr, j, gsc),
3553 q4tp_outlier(uv_view.nib, r, gpr, j, usc),
3554 )
3555 };
3556 gv += og.0 * og.1 * xv;
3557 uv += ou.0 * ou.1 * xv;
3558 }
3559 unsafe { *ptrs[e].at(r) = silu_mul_limited(gv, uv, limit) };
3561 }
3562 });
3563 };
3564 dispatch_rows(pool, pairs.len() * inter, &run);
3565 true
3566 }
3567
3568 pub fn moe_down_many(
3577 downs: &[&QTensor],
3578 gs: &[Vec<f32>],
3579 weights: &[f32],
3580 out: &mut [f32],
3581 pool: Option<&Pool>,
3582 ) -> bool {
3583 if downs.is_empty() || downs.len() != gs.len() || downs.len() != weights.len() {
3584 return false;
3585 }
3586 if !a8w8_enabled() {
3587 let mut terms = vec![vec![0.0; out.len()]; downs.len()];
3588 if !Self::moe_down_rows(downs, &vec![1; downs.len()], gs, &mut terms, pool) {
3589 return false;
3590 }
3591 out.fill(0.0);
3592 for (row, &w) in terms.iter().zip(weights) {
3593 for (o, &v) in out.iter_mut().zip(row) {
3594 *o += w * v;
3595 }
3596 }
3597 return true;
3598 }
3599 let rows = out.len();
3600 let cols = downs[0].cols();
3601 if cols % GROUP_SIZE != 0 {
3602 return false;
3603 }
3604 let gpr = cols / GROUP_SIZE;
3605 let mut views = Vec::with_capacity(downs.len());
3606 for (d, g) in downs.iter().zip(gs.iter()) {
3607 if !matches!(
3608 d,
3609 Self::Mapped {
3610 dtype: TensorDtype::Q4TiledP,
3611 ..
3612 }
3613 ) || d.rows() != rows
3614 || d.cols() != cols
3615 || g.len() != cols
3616 {
3617 return false;
3618 }
3619 views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3620 }
3621 let acts: Vec<SplitAct> = gs.iter().map(|g| split_act(g)).collect();
3623 let out_addr = SendMut(out.as_mut_ptr());
3631 let (views, acts, weights) = (&views, &acts, &weights);
3632 let run = |start: usize, end: usize| {
3633 with_krow(gpr, |sc| {
3634 for r in start..end {
3635 let mut acc = 0f32;
3636 for (e, v) in views.iter().enumerate() {
3637 v.scales_into(r, gpr, sc);
3638 let a = &acts[e];
3639 let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, sc) * a.sx;
3640 for &(j, xv) in &a.outliers {
3641 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
3642 d += w * s * xv;
3643 }
3644 acc += weights[e] * d;
3645 }
3646 unsafe { *out_addr.at(r) = acc };
3648 }
3649 });
3650 };
3651 dispatch_rows(pool, rows, &run);
3652 true
3653 }
3654
3655 pub fn moe_gate_up_rows(
3667 pairs: &[(&QTensor, &QTensor)],
3668 groups: &[Vec<usize>],
3669 xs: &[f32],
3670 outs: &mut [Vec<f32>],
3671 pool: Option<&Pool>,
3672 ) -> bool {
3673 if pairs.is_empty() || pairs.len() != groups.len() {
3674 return false;
3675 }
3676 let inter = pairs[0].0.rows();
3677 let cols = pairs[0].0.cols();
3678 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
3679 if cols == 0 || cols % GROUP_SIZE != 0 || outs.len() != n_pairs || xs.len() % cols != 0 {
3680 return false;
3681 }
3682 let b = xs.len() / cols;
3683 let gpr = cols / GROUP_SIZE;
3684 let mut views = Vec::with_capacity(pairs.len() * 2);
3685 for (g, u) in pairs {
3686 let q4tp = |t: &QTensor| {
3687 matches!(
3688 t,
3689 Self::Mapped {
3690 dtype: TensorDtype::Q4TiledP,
3691 ..
3692 }
3693 )
3694 };
3695 if g.has_prism_contract()
3696 || u.has_prism_contract()
3697 || !q4tp(g)
3698 || !q4tp(u)
3699 || g.rows() != inter
3700 || u.rows() != inter
3701 || g.cols() != cols
3702 || u.cols() != cols
3703 {
3704 return false;
3705 }
3706 views.push(Q4tpView::new(g.quant_bytes(), inter, cols));
3707 views.push(Q4tpView::new(u.quant_bytes(), inter, cols));
3708 }
3709 if outs.iter().any(|o| o.len() != inter) || groups.iter().flatten().any(|&t| t >= b) {
3710 return false;
3711 }
3712 let quantized = a8w8_enabled();
3713 let acts: Vec<SplitAct> = if quantized {
3714 (0..b)
3715 .map(|t| split_act(&xs[t * cols..(t + 1) * cols]))
3716 .collect()
3717 } else {
3718 Vec::new()
3719 };
3720 let mut offs = Vec::with_capacity(groups.len());
3721 let mut o = 0usize;
3722 for g in groups {
3723 offs.push(o);
3724 o += g.len();
3725 }
3726 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3727 let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3728 let run = |start: usize, end: usize| {
3729 let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3730 for flat in start..end {
3731 let (e, r) = (flat / inter, flat % inter);
3732 let (gv_view, uv_view) = (&views[e * 2], &views[e * 2 + 1]);
3733 gv_view.scales_into(r, gpr, &mut gsc);
3734 uv_view.scales_into(r, gpr, &mut usc);
3735 for (k, &t) in groups[e].iter().enumerate() {
3736 if !quantized {
3737 let x = &xs[t * cols..(t + 1) * cols];
3738 let gv = q4tp_row_exact(gv_view.nib, r, gpr, x, &gsc);
3739 let uv = q4tp_row_exact(uv_view.nib, r, gpr, x, &usc);
3740 unsafe { *ptrs[offs[e] + k].at(r) = (gv / (1.0 + (-gv).exp())) * uv };
3741 continue;
3742 }
3743 let act = &acts[t];
3744 let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3745 let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3746 for &(j, xv) in &act.outliers {
3747 let og = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3748 let ou = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3749 gv += og.0 * og.1 * xv;
3750 uv += ou.0 * ou.1 * xv;
3751 }
3752 let silu_g = gv / (1.0 + (-gv).exp());
3753 unsafe { *ptrs[offs[e] + k].at(r) = silu_g * uv };
3756 }
3757 }
3758 };
3759 dispatch_rows(pool, pairs.len() * inter, &run);
3760 true
3761 }
3762
3763 pub fn moe_down_rows(
3770 downs: &[&QTensor],
3771 group_lens: &[usize],
3772 gs: &[Vec<f32>],
3773 outs: &mut [Vec<f32>],
3774 pool: Option<&Pool>,
3775 ) -> bool {
3776 if downs.is_empty() || downs.len() != group_lens.len() {
3777 return false;
3778 }
3779 let rows = downs[0].rows();
3780 let cols = downs[0].cols();
3781 let n_pairs: usize = group_lens.iter().sum();
3782 if cols == 0 || cols % GROUP_SIZE != 0 || gs.len() != n_pairs || outs.len() != n_pairs {
3783 return false;
3784 }
3785 let gpr = cols / GROUP_SIZE;
3786 let mut views = Vec::with_capacity(downs.len());
3787 for d in downs {
3788 if d.has_prism_contract()
3789 || !matches!(
3790 d,
3791 Self::Mapped {
3792 dtype: TensorDtype::Q4TiledP,
3793 ..
3794 }
3795 )
3796 || d.rows() != rows
3797 || d.cols() != cols
3798 {
3799 return false;
3800 }
3801 views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3802 }
3803 if gs.iter().any(|g| g.len() != cols) || outs.iter().any(|o| o.len() != rows) {
3804 return false;
3805 }
3806 let quantized = a8w8_enabled();
3807 let acts: Vec<SplitAct> = if quantized {
3808 gs.iter().map(|g| split_act(g)).collect()
3809 } else {
3810 Vec::new()
3811 };
3812 let mut offs = Vec::with_capacity(group_lens.len());
3813 let mut o = 0usize;
3814 for &l in group_lens {
3815 offs.push(o);
3816 o += l;
3817 }
3818 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3819 let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3820 let run = |start: usize, end: usize| {
3821 let mut sc = vec![0f32; gpr];
3822 for flat in start..end {
3823 let (e, r) = (flat / rows, flat % rows);
3824 let v = &views[e];
3825 v.scales_into(r, gpr, &mut sc);
3826 for k in 0..group_lens[e] {
3827 if !quantized {
3828 let d = q4tp_row_exact(v.nib, r, gpr, &gs[offs[e] + k], &sc);
3829 unsafe { *ptrs[offs[e] + k].at(r) = d };
3830 continue;
3831 }
3832 let a = &acts[offs[e] + k];
3833 let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3834 for &(j, xv) in &a.outliers {
3835 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3836 d += w * s * xv;
3837 }
3838 unsafe { *ptrs[offs[e] + k].at(r) = d };
3840 }
3841 }
3842 };
3843 dispatch_rows(pool, downs.len() * rows, &run);
3844 true
3845 }
3846}
3847
3848#[cfg(target_os = "macos")]
3853mod accel_blas {
3854 #[link(name = "Accelerate", kind = "framework")]
3855 unsafe extern "C" {
3856 pub fn cblas_sgemm(
3857 order: i32,
3858 trans_a: i32,
3859 trans_b: i32,
3860 m: i32,
3861 n: i32,
3862 k: i32,
3863 alpha: f32,
3864 a: *const f32,
3865 lda: i32,
3866 b: *const f32,
3867 ldb: i32,
3868 beta: f32,
3869 c: *mut f32,
3870 ldc: i32,
3871 );
3872 }
3873}
3874
3875#[cfg(target_os = "macos")]
3876pub(crate) fn accel_gemm_enabled() -> bool {
3877 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3878 *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
3879}
3880
3881#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3884pub(crate) fn accel_gemm_enabled() -> bool {
3885 true
3886}
3887
3888#[cfg(target_arch = "aarch64")]
3895#[allow(clippy::too_many_arguments)]
3896pub(crate) fn neon_gemm_rm(
3897 m: usize,
3898 n: usize,
3899 k: usize,
3900 alpha: f32,
3901 a: &[f32],
3902 lda: usize,
3903 b_mat: &[f32],
3904 ldb: usize,
3905 b_rows_are_n: bool,
3906 c: &mut [f32],
3907 ldc: usize,
3908) {
3909 debug_assert!(a.len() >= (m - 1) * lda + k);
3910 debug_assert!(c.len() >= (m - 1) * ldc + n);
3911 unsafe {
3913 use core::arch::aarch64::*;
3914 let mut i = 0usize;
3915 while i < m {
3916 let mi = (m - i).min(4);
3917 let mut j = 0usize;
3918 while j < n {
3919 let nj = (n - j).min(8);
3920 if mi == 4 && nj == 8 {
3921 let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3922 let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3923 let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3924 let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3925 for p in 0..k {
3926 let (b0, b1) = if b_rows_are_n {
3927 let base = b_mat.as_ptr().add(j * ldb + p);
3930 let g = |o: usize| *base.add(o * ldb);
3931 ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
3932 } else {
3933 let base = b_mat.as_ptr().add(p * ldb + j);
3934 (
3935 [*base, *base.add(1), *base.add(2), *base.add(3)],
3936 [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
3937 )
3938 };
3939 let bv0 = vld1q_f32(b0.as_ptr());
3940 let bv1 = vld1q_f32(b1.as_ptr());
3941 let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
3942 let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
3943 let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
3944 let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
3945 c0a = vfmaq_f32(c0a, a0, bv0);
3946 c0b = vfmaq_f32(c0b, a0, bv1);
3947 c1a = vfmaq_f32(c1a, a1, bv0);
3948 c1b = vfmaq_f32(c1b, a1, bv1);
3949 c2a = vfmaq_f32(c2a, a2, bv0);
3950 c2b = vfmaq_f32(c2b, a2, bv1);
3951 c3a = vfmaq_f32(c3a, a3, bv0);
3952 c3b = vfmaq_f32(c3b, a3, bv1);
3953 }
3954 let al = vdupq_n_f32(alpha);
3955 for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
3956 .iter()
3957 .enumerate()
3958 {
3959 let dst = c.as_mut_ptr().add((i + r) * ldc + j);
3960 vst1q_f32(dst, vmulq_f32(*ca, al));
3961 vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
3962 }
3963 } else {
3964 for r in 0..mi {
3965 for q in 0..nj {
3966 let mut acc = 0f32;
3967 for p in 0..k {
3968 let bv = if b_rows_are_n {
3969 b_mat[(j + q) * ldb + p]
3970 } else {
3971 b_mat[p * ldb + j + q]
3972 };
3973 acc += a[(i + r) * lda + p] * bv;
3974 }
3975 c[(i + r) * ldc + j + q] = acc * alpha;
3976 }
3977 }
3978 }
3979 j += nj;
3980 }
3981 i += mi;
3982 }
3983 }
3984}
3985
3986#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3988#[allow(clippy::too_many_arguments)]
3989pub(crate) fn sgemm_rm(
3990 m: usize,
3991 n: usize,
3992 k: usize,
3993 alpha: f32,
3994 a: &[f32],
3995 lda: usize,
3996 b_mat: &[f32],
3997 ldb: usize,
3998 b_rows_are_n: bool,
3999 c: &mut [f32],
4000 ldc: usize,
4001) {
4002 neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4003}
4004
4005#[allow(clippy::too_many_arguments)]
4009pub fn sgemm_public(
4010 m: usize,
4011 n: usize,
4012 k: usize,
4013 alpha: f32,
4014 a: &[f32],
4015 lda: usize,
4016 b_mat: &[f32],
4017 ldb: usize,
4018 b_rows_are_n: bool,
4019 c: &mut [f32],
4020 ldc: usize,
4021) {
4022 #[cfg(any(target_os = "macos", target_arch = "aarch64"))]
4023 {
4024 sgemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4025 }
4026 #[cfg(not(any(target_os = "macos", target_arch = "aarch64")))]
4031 {
4032 for i in 0..m {
4033 for j in 0..n {
4034 let mut acc = 0f32;
4035 for p in 0..k {
4036 let bv = if b_rows_are_n {
4037 b_mat[j * ldb + p]
4038 } else {
4039 b_mat[p * ldb + j]
4040 };
4041 acc += a[i * lda + p] * bv;
4042 }
4043 c[i * ldc + j] = alpha * acc;
4044 }
4045 }
4046 }
4047}
4048
4049#[cfg(target_os = "macos")]
4052#[allow(clippy::too_many_arguments)]
4053pub(crate) fn sgemm_rm(
4054 m: usize,
4055 n: usize,
4056 k: usize,
4057 alpha: f32,
4058 a: &[f32],
4059 lda: usize,
4060 b_mat: &[f32],
4061 ldb: usize,
4062 b_rows_are_n: bool,
4063 c: &mut [f32],
4064 ldc: usize,
4065) {
4066 debug_assert!(a.len() >= (m - 1) * lda + k);
4067 debug_assert!(c.len() >= (m - 1) * ldc + n);
4068 #[cfg(target_arch = "aarch64")]
4073 if std::env::var("CMF_FORCE_NEON_GEMM")
4074 .map(|v| v == "1")
4075 .unwrap_or(false)
4076 {
4077 return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4078 }
4079 unsafe {
4080 accel_blas::cblas_sgemm(
4081 101, 111, if b_rows_are_n { 112 } else { 111 },
4084 m as i32,
4085 n as i32,
4086 k as i32,
4087 alpha,
4088 a.as_ptr(),
4089 lda as i32,
4090 b_mat.as_ptr(),
4091 ldb as i32,
4092 0.0,
4093 c.as_mut_ptr(),
4094 ldc as i32,
4095 );
4096 }
4097}
4098
4099#[cfg(target_os = "macos")]
4106fn qmatmat_accel(
4107 q: &[u8],
4108 row_scale: &[f32],
4109 pre: &[std::borrow::Cow<'_, [f32]>],
4110 rows: usize,
4111 cols: usize,
4112 out: &mut [f32],
4113 pool: Option<&Pool>,
4114) {
4115 const TR: usize = 2048;
4120 let b = pre.len();
4121 thread_local! {
4122 static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4123 static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4124 }
4125 XPANEL.with(|xp| {
4126 WTILE.with(|wt| {
4127 let mut xpanel = xp.borrow_mut();
4128 xpanel.clear();
4129 for x in pre {
4130 xpanel.extend_from_slice(x);
4131 }
4132 let mut wtile = wt.borrow_mut();
4133 wtile.resize(TR * cols, 0.0);
4134 let mut r0 = 0usize;
4135 while r0 < rows {
4136 let tr = TR.min(rows - r0);
4137 let wt_addr = SendMut(wtile.as_mut_ptr());
4139 let run = |start: usize, end: usize| {
4140 for r in start..end {
4141 let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
4142 let s = row_scale[r0 + r];
4143 let dst =
4145 unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
4146 for (d, &v) in dst.iter_mut().zip(row) {
4147 *d = (v as i8) as f32 * s;
4148 }
4149 }
4150 };
4151 dispatch_rows(pool, tr, &run);
4152 unsafe {
4154 accel_blas::cblas_sgemm(
4155 101, 111, 112, b as i32,
4159 tr as i32,
4160 cols as i32,
4161 1.0,
4162 xpanel.as_ptr(),
4163 cols as i32,
4164 wtile.as_ptr(),
4165 cols as i32,
4166 0.0,
4167 out.as_mut_ptr().add(r0),
4168 rows as i32,
4169 );
4170 }
4171 r0 += tr;
4172 }
4173 })
4174 });
4175}
4176
4177fn qmatmat(
4178 q: &[u8],
4179 row_scale: &[f32],
4180 pre: &[std::borrow::Cow<'_, [f32]>],
4181 rows: usize,
4182 cols: usize,
4183 out: &mut [f32],
4184 pool: Option<&Pool>,
4185) {
4186 let b = pre.len();
4187 debug_assert_eq!(out.len(), b * rows);
4188 #[cfg(target_os = "macos")]
4193 if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
4194 qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
4195 return;
4196 }
4197 #[cfg(target_arch = "aarch64")]
4198 if sdot_enabled() {
4199 let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4200 let out_addr = SendMut(out.as_mut_ptr());
4201 let blocked_ok = blocked_enabled();
4204 let use_i8mm = i8mm_enabled();
4205 if blocked_ok {
4206 let run = |start: usize, end: usize| {
4207 let mut o = start;
4208 while o < end {
4209 if o + 2 <= end {
4210 let r0 = &q[o * cols..(o + 1) * cols];
4211 let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4212 let mut bi = 0usize;
4213 while bi + 4 <= acts.len() {
4214 let xs = [
4215 acts[bi].xq.as_slice(),
4216 acts[bi + 1].xq.as_slice(),
4217 acts[bi + 2].xq.as_slice(),
4218 acts[bi + 3].xq.as_slice(),
4219 ];
4220 let d = if use_i8mm {
4221 unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
4222 } else {
4223 unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
4224 };
4225 for (r, row) in [r0, r1].into_iter().enumerate() {
4226 for k in 0..4 {
4227 let act = &acts[bi + k];
4228 let mut v = d[r][k] as f32 * act.sx;
4229 for &(j, xv) in &act.outliers {
4230 v += (row[j] as i8) as f32 * xv;
4231 }
4232 unsafe {
4233 *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4234 };
4235 }
4236 }
4237 bi += 4;
4238 }
4239 while bi < acts.len() {
4240 for (r, row) in [r0, r1].into_iter().enumerate() {
4241 let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
4242 unsafe { *out_addr.at(bi * rows + o + r) = v };
4243 }
4244 bi += 1;
4245 }
4246 o += 2;
4247 } else {
4248 let row = &q[o * cols..(o + 1) * cols];
4249 for (bi, act) in acts.iter().enumerate() {
4250 let v = row_dot_sdot(row, act) * row_scale[o];
4251 unsafe { *out_addr.at(bi * rows + o) = v };
4252 }
4253 o += 1;
4254 }
4255 }
4256 };
4257 dispatch_rows(pool, rows, &run);
4258 return;
4259 }
4260 let run = |start: usize, end: usize| {
4261 for o in start..end {
4262 let row = &q[o * cols..(o + 1) * cols];
4263 for (bi, act) in acts.iter().enumerate() {
4264 let v = row_dot_sdot(row, act) * row_scale[o];
4265 unsafe { *out_addr.at(bi * rows + o) = v };
4266 }
4267 }
4268 };
4269 dispatch_rows(pool, rows, &run);
4270 return;
4271 }
4272 #[cfg(target_arch = "x86_64")]
4277 if avx2_a8w8_enabled() {
4278 let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4279 let out_addr = SendMut(out.as_mut_ptr());
4280 let blocked_ok = blocked_enabled();
4283 if !avx512vnni_enabled() && blocked_ok && !row_exact() {
4284 let run = |start: usize, end: usize| {
4285 let mut o = start;
4286 while o < end {
4287 if o + 2 <= end {
4288 let r0 = &q[o * cols..(o + 1) * cols];
4289 let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4290 let mut bi = 0usize;
4291 while bi + 4 <= acts.len() {
4292 let xs = [
4293 acts[bi].xq.as_slice(),
4294 acts[bi + 1].xq.as_slice(),
4295 acts[bi + 2].xq.as_slice(),
4296 acts[bi + 3].xq.as_slice(),
4297 ];
4298 let d = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
4299 for (r, row) in [r0, r1].into_iter().enumerate() {
4300 for k in 0..4 {
4301 let act = &acts[bi + k];
4302 let mut v = d[r][k] as f32 * act.sx;
4303 for &(j, xv) in &act.outliers {
4304 v += (row[j] as i8) as f32 * xv;
4305 }
4306 unsafe {
4307 *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4308 };
4309 }
4310 }
4311 bi += 4;
4312 }
4313 while bi < acts.len() {
4314 for (r, row) in [r0, r1].into_iter().enumerate() {
4315 let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
4316 unsafe { *out_addr.at(bi * rows + o + r) = v };
4317 }
4318 bi += 1;
4319 }
4320 o += 2;
4321 } else {
4322 let row = &q[o * cols..(o + 1) * cols];
4323 for (bi, act) in acts.iter().enumerate() {
4324 let v = row_dot_avx2(row, act) * row_scale[o];
4325 unsafe { *out_addr.at(bi * rows + o) = v };
4326 }
4327 o += 1;
4328 }
4329 }
4330 };
4331 dispatch_rows(pool, rows, &run);
4332 return;
4333 }
4334 let run = |start: usize, end: usize| {
4335 for o in start..end {
4336 let row = &q[o * cols..(o + 1) * cols];
4337 for (bi, act) in acts.iter().enumerate() {
4338 let v = row_dot_avx2(row, act) * row_scale[o];
4339 unsafe { *out_addr.at(bi * rows + o) = v };
4340 }
4341 }
4342 };
4343 dispatch_rows(pool, rows, &run);
4344 return;
4345 }
4346 let out_addr = SendMut(out.as_mut_ptr());
4347 let run = |start: usize, end: usize| {
4348 for o in start..end {
4349 let row = &q[o * cols..(o + 1) * cols];
4350 for (bi, x) in pre.iter().enumerate() {
4351 let mut acc = 0f32;
4352 for j in 0..cols {
4353 acc += (row[j] as i8) as f32 * x[j];
4354 }
4355 unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
4356 }
4357 }
4358 };
4359 dispatch_rows(pool, rows, &run);
4360}
4361
4362fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
4365 match pool {
4366 Some(pool) if rows >= 256 => pool.run_rows(rows, run),
4367 _ => run(0, rows),
4368 }
4369}
4370
4371fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
4373 let groups = rows * cols / GROUP_SIZE;
4374 bytes.split_at(groups * 16)
4375}
4376
4377#[inline]
4382fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
4383 #[cfg(target_arch = "aarch64")]
4384 unsafe {
4385 return vbit_fill4_neon(data, buf);
4386 }
4387 #[cfg(target_arch = "x86_64")]
4388 if avx2_enabled() {
4389 return unsafe { vbit_fill4_avx2(data, buf) };
4390 }
4391 #[allow(unreachable_code)]
4392 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4393 let u = unpack8::<4>(&data[blk * 4..]);
4394 for k in 0..8 {
4395 chunk[k] = (u[k] - 7) as i8 as u8;
4396 }
4397 }
4398}
4399
4400#[cfg(target_arch = "aarch64")]
4401#[target_feature(enable = "neon")]
4402unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
4403 unsafe {
4406 use core::arch::aarch64::*;
4407 let n = buf.len();
4408 let mask = vdupq_n_u8(0x0F);
4409 let seven = vdupq_n_s8(7);
4410 let mut g = 0usize;
4411 while g * 32 + 32 <= n {
4412 let b = vld1q_u8(data.as_ptr().add(g * 16));
4413 let hi = vshrq_n_u8::<4>(b);
4414 let lo = vandq_u8(b, mask);
4415 let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
4416 let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
4417 vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
4418 vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
4419 g += 1;
4420 }
4421 }
4422}
4423
4424#[cfg(target_arch = "x86_64")]
4425#[target_feature(enable = "avx2")]
4426unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
4427 unsafe {
4429 use core::arch::x86_64::*;
4430 let n = buf.len();
4431 let mask = _mm_set1_epi8(0x0F);
4432 let seven = _mm256_set1_epi8(7);
4433 let mut g = 0usize;
4434 while g * 32 + 32 <= n {
4435 let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
4436 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
4437 let lo = _mm_and_si128(b, mask);
4438 let z = _mm256_sub_epi8(
4439 _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
4440 seven,
4441 );
4442 _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
4443 g += 1;
4444 }
4445 }
4446}
4447
4448#[inline(always)]
4453fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
4454 let mut acc = 0u64;
4455 for i in 0..B {
4456 acc = (acc << 8) | data[i] as u64;
4457 }
4458 let mask = (1u64 << B) - 1;
4459 let mut out = [0i32; 8];
4460 for (k, o) in out.iter_mut().enumerate() {
4461 *o = ((acc >> ((7 - k) * B)) & mask) as i32;
4462 }
4463 out
4464}
4465
4466#[allow(clippy::too_many_arguments)]
4472fn vbitmatvec(
4473 bytes: &[u8],
4474 offsets: &[usize],
4475 x: &[f32],
4476 rows: usize,
4477 cols: usize,
4478 out: &mut [f32],
4479 pool: Option<&Pool>,
4480) {
4481 debug_assert_eq!(out.len(), rows);
4482 debug_assert_eq!(offsets.len(), rows + 1);
4483
4484 if a8w8_enabled() {
4488 let act = split_act(x);
4489 let out_addr = SendMut(out.as_mut_ptr());
4490 let run = move |start: usize, end: usize| {
4491 vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
4492 };
4493 dispatch_rows(pool, rows, &run);
4494 return;
4495 }
4496
4497 let out_addr = SendMut(out.as_mut_ptr());
4498 let run = move |start: usize, end: usize| {
4499 vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
4500 };
4501 dispatch_rows(pool, rows, &run);
4502}
4503
4504#[allow(clippy::too_many_arguments)]
4508fn vbit_range_a8w8(
4509 bytes: &[u8],
4510 offsets: &[usize],
4511 x: &[f32],
4512 act: &SplitAct,
4513 rows: usize,
4514 cols: usize,
4515 out: SendMut,
4516 start: usize,
4517 end: usize,
4518) {
4519 let ng = cols / GROUP_SIZE;
4520 let bits = &bytes[..rows];
4521 let sc_off = rows;
4522 let row_dot = |r: usize| -> f32 {
4523 let b = bits[r] as usize;
4524 let l = (1i32 << (b - 1)) - 1;
4525 let mask = (1u64 << b) - 1;
4526 let data = &bytes[offsets[r]..offsets[r + 1]];
4527 if b == 8 {
4528 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4530 let mut dot = 0f32;
4531 for g in 0..ng {
4532 let so = (r * ng + g) * 2;
4533 let sgf = f16_to_f32(u16::from_le_bytes([
4534 bytes[sc_off + so],
4535 bytes[sc_off + so + 1],
4536 ]));
4537 let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4538 let mut gd = 0f32;
4539 for &xv in xg.iter() {
4540 if nbits < 8 {
4541 acc = (acc << 8) | data[idx] as u64;
4542 idx += 1;
4543 nbits += 8;
4544 }
4545 let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4546 nbits -= 8;
4547 gd += (u - l) as f32 * xv;
4548 }
4549 dot += gd * sgf;
4550 }
4551 return dot;
4552 }
4553 thread_local! {
4557 static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
4558 const { std::cell::RefCell::new(Vec::new()) };
4559 }
4560 #[inline(always)]
4561 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4562 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4563 let u = unpack8::<B>(&data[blk * B..]);
4564 for k in 0..8 {
4565 chunk[k] = (u[k] - l) as i8 as u8;
4566 }
4567 }
4568 }
4569 let _ = mask;
4570 VBIT_SCRATCH.with(|scratch| {
4571 let mut buf = scratch.borrow_mut();
4572 buf.resize(cols, 0);
4573 match b {
4574 3 => fill::<3>(data, l, &mut buf),
4575 4 => vbit_fill4(data, &mut buf),
4576 5 => fill::<5>(data, l, &mut buf),
4577 6 => fill::<6>(data, l, &mut buf),
4578 _ => unreachable!(),
4579 }
4580 let mut dot = 0f32;
4581 for g in 0..ng {
4582 let so = (r * ng + g) * 2;
4583 let s = f16_to_f32(u16::from_le_bytes([
4584 bytes[sc_off + so],
4585 bytes[sc_off + so + 1],
4586 ]));
4587 let d = dot_i8_i8(
4588 &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4589 &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4590 ) as f32
4591 * act.sx;
4592 dot += d * s;
4593 }
4594 for &(j, xv) in &act.outliers {
4595 let so = (r * ng + j / GROUP_SIZE) * 2;
4596 let s = f16_to_f32(u16::from_le_bytes([
4597 bytes[sc_off + so],
4598 bytes[sc_off + so + 1],
4599 ]));
4600 dot += (buf[j] as i8) as f32 * s * xv;
4602 }
4603 dot
4604 })
4605 };
4606 for r in start..end {
4607 unsafe { *out.at(r) = row_dot(r) };
4609 }
4610}
4611
4612#[allow(clippy::too_many_arguments)]
4614fn vbit_range_f32(
4615 bytes: &[u8],
4616 offsets: &[usize],
4617 x: &[f32],
4618 rows: usize,
4619 cols: usize,
4620 out: SendMut,
4621 start: usize,
4622 end: usize,
4623) {
4624 let ng = cols / GROUP_SIZE;
4625 let bits = &bytes[..rows];
4626 let sc_off = rows;
4627 #[inline(always)]
4631 fn dot_row<const B: usize>(
4632 data: &[u8],
4633 bytes: &[u8],
4634 sc_off: usize,
4635 r: usize,
4636 ng: usize,
4637 x: &[f32],
4638 ) -> f32 {
4639 let l = ((1i32 << (B - 1)) - 1) as f32;
4640 let gbytes = GROUP_SIZE * B / 8;
4641 let mut dot = 0f32;
4642 for g in 0..ng {
4643 let so = (r * ng + g) * 2;
4644 let s = f16_to_f32(u16::from_le_bytes([
4645 bytes[sc_off + so],
4646 bytes[sc_off + so + 1],
4647 ]));
4648 let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4649 let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4650 let mut gd = 0f32;
4651 for blk in 0..GROUP_SIZE / 8 {
4652 let u = unpack8::<B>(&gd0[blk * B..]);
4653 let xb = &xg[blk * 8..blk * 8 + 8];
4654 for k in 0..8 {
4655 gd += (u[k] as f32 - l) * xb[k];
4656 }
4657 }
4658 dot += gd * s;
4659 }
4660 dot
4661 }
4662 for r in start..end {
4663 let data = &bytes[offsets[r]..offsets[r + 1]];
4664 let v = match bits[r] {
4665 3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
4666 4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
4667 5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
4668 6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
4669 8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
4670 b => unreachable!("vbit bit-width {b} (validated at load)"),
4671 };
4672 unsafe { *out.at(r) = v };
4674 }
4675}
4676
4677#[allow(clippy::too_many_arguments)]
4682fn vbitmatvec2(
4683 bytes: &[u8],
4684 offsets: &[usize],
4685 x1: &[f32],
4686 x2: &[f32],
4687 rows: usize,
4688 cols: usize,
4689 o1: &mut [f32],
4690 o2: &mut [f32],
4691 pool: Option<&Pool>,
4692) {
4693 debug_assert_eq!(o1.len(), rows);
4694 debug_assert_eq!(o2.len(), rows);
4695
4696 if a8w8_enabled() {
4697 let a1 = split_act(x1);
4698 let a2 = split_act(x2);
4699 let p1 = SendMut(o1.as_mut_ptr());
4700 let p2 = SendMut(o2.as_mut_ptr());
4701 let run = move |start: usize, end: usize| {
4702 vbit_range2_a8w8(
4703 bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
4704 )
4705 };
4706 dispatch_rows(pool, rows, &run);
4707 return;
4708 }
4709
4710 let p1 = SendMut(o1.as_mut_ptr());
4711 let p2 = SendMut(o2.as_mut_ptr());
4712 let run = move |start: usize, end: usize| {
4713 vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
4714 };
4715 dispatch_rows(pool, rows, &run);
4716}
4717
4718#[allow(clippy::too_many_arguments)]
4722fn vbit_range2_a8w8(
4723 bytes: &[u8],
4724 offsets: &[usize],
4725 x1: &[f32],
4726 x2: &[f32],
4727 a1: &SplitAct,
4728 a2: &SplitAct,
4729 rows: usize,
4730 cols: usize,
4731 p1: SendMut,
4732 p2: SendMut,
4733 start: usize,
4734 end: usize,
4735) {
4736 let ng = cols / GROUP_SIZE;
4737 let bits = &bytes[..rows];
4738 let sc_off = rows;
4739 let row_dots = |r: usize| -> (f32, f32) {
4740 let b = bits[r] as usize;
4741 let l = (1i32 << (b - 1)) - 1;
4742 let data = &bytes[offsets[r]..offsets[r + 1]];
4743 if b == 8 {
4744 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4747 let (mut d1, mut d2) = (0f32, 0f32);
4748 for g in 0..ng {
4749 let so = (r * ng + g) * 2;
4750 let sgf = f16_to_f32(u16::from_le_bytes([
4751 bytes[sc_off + so],
4752 bytes[sc_off + so + 1],
4753 ]));
4754 let (mut g1, mut g2) = (0f32, 0f32);
4755 for k in 0..GROUP_SIZE {
4756 if nbits < 8 {
4757 acc = (acc << 8) | data[idx] as u64;
4758 idx += 1;
4759 nbits += 8;
4760 }
4761 let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4762 nbits -= 8;
4763 let w = (u - l) as f32;
4764 g1 += w * x1[g * GROUP_SIZE + k];
4765 g2 += w * x2[g * GROUP_SIZE + k];
4766 }
4767 d1 += g1 * sgf;
4768 d2 += g2 * sgf;
4769 }
4770 return (d1, d2);
4771 }
4772 thread_local! {
4773 static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
4774 const { std::cell::RefCell::new(Vec::new()) };
4775 }
4776 #[inline(always)]
4777 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4778 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4779 let u = unpack8::<B>(&data[blk * B..]);
4780 for k in 0..8 {
4781 chunk[k] = (u[k] - l) as i8 as u8;
4782 }
4783 }
4784 }
4785 VBIT_SCRATCH2.with(|scratch| {
4786 let mut buf = scratch.borrow_mut();
4787 buf.resize(cols, 0);
4788 match b {
4789 3 => fill::<3>(data, l, &mut buf),
4790 4 => vbit_fill4(data, &mut buf),
4791 5 => fill::<5>(data, l, &mut buf),
4792 6 => fill::<6>(data, l, &mut buf),
4793 _ => unreachable!(),
4794 }
4795 let (mut d1, mut d2) = (0f32, 0f32);
4796 for g in 0..ng {
4797 let so = (r * ng + g) * 2;
4798 let s = f16_to_f32(u16::from_le_bytes([
4799 bytes[sc_off + so],
4800 bytes[sc_off + so + 1],
4801 ]));
4802 let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4803 let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
4804 let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
4805 d1 += v1 * s;
4806 d2 += v2 * s;
4807 }
4808 for &(j, xv) in &a1.outliers {
4809 let so = (r * ng + j / GROUP_SIZE) * 2;
4810 let s = f16_to_f32(u16::from_le_bytes([
4811 bytes[sc_off + so],
4812 bytes[sc_off + so + 1],
4813 ]));
4814 d1 += (buf[j] as i8) as f32 * s * xv;
4815 }
4816 for &(j, xv) in &a2.outliers {
4817 let so = (r * ng + j / GROUP_SIZE) * 2;
4818 let s = f16_to_f32(u16::from_le_bytes([
4819 bytes[sc_off + so],
4820 bytes[sc_off + so + 1],
4821 ]));
4822 d2 += (buf[j] as i8) as f32 * s * xv;
4823 }
4824 (d1, d2)
4825 })
4826 };
4827 for r in start..end {
4828 let (v1, v2) = row_dots(r);
4829 unsafe {
4831 *p1.at(r) = v1;
4832 *p2.at(r) = v2;
4833 }
4834 }
4835}
4836
4837#[allow(clippy::too_many_arguments)]
4841fn vbit_range2_f32(
4842 bytes: &[u8],
4843 offsets: &[usize],
4844 x1: &[f32],
4845 x2: &[f32],
4846 rows: usize,
4847 cols: usize,
4848 p1: SendMut,
4849 p2: SendMut,
4850 start: usize,
4851 end: usize,
4852) {
4853 let ng = cols / GROUP_SIZE;
4854 let bits = &bytes[..rows];
4855 let sc_off = rows;
4856 #[inline(always)]
4857 #[allow(clippy::too_many_arguments)]
4858 fn dot_row2<const B: usize>(
4859 data: &[u8],
4860 bytes: &[u8],
4861 sc_off: usize,
4862 r: usize,
4863 ng: usize,
4864 x1: &[f32],
4865 x2: &[f32],
4866 ) -> (f32, f32) {
4867 let l = ((1i32 << (B - 1)) - 1) as f32;
4868 let gbytes = GROUP_SIZE * B / 8;
4869 let (mut d1, mut d2) = (0f32, 0f32);
4870 for g in 0..ng {
4871 let so = (r * ng + g) * 2;
4872 let s = f16_to_f32(u16::from_le_bytes([
4873 bytes[sc_off + so],
4874 bytes[sc_off + so + 1],
4875 ]));
4876 let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4877 let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4878 let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4879 let (mut g1, mut g2) = (0f32, 0f32);
4880 for blk in 0..GROUP_SIZE / 8 {
4881 let u = unpack8::<B>(&gd0[blk * B..]);
4882 for k in 0..8 {
4883 let w = u[k] as f32 - l;
4884 g1 += w * x1g[blk * 8 + k];
4885 g2 += w * x2g[blk * 8 + k];
4886 }
4887 }
4888 d1 += g1 * s;
4889 d2 += g2 * s;
4890 }
4891 (d1, d2)
4892 }
4893 for r in start..end {
4894 let data = &bytes[offsets[r]..offsets[r + 1]];
4895 let (v1, v2) = match bits[r] {
4896 3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
4897 4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
4898 5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
4899 6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
4900 8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
4901 b => unreachable!("vbit bit-width {b} (validated at load)"),
4902 };
4903 unsafe {
4905 *p1.at(r) = v1;
4906 *p2.at(r) = v2;
4907 }
4908 }
4909}
4910
4911#[inline]
4918#[allow(unreachable_code)]
4919fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4920 #[cfg(target_arch = "aarch64")]
4921 unsafe {
4922 return dot_q4t_row_sdot(bytes, r, gpr, xq);
4923 }
4924 #[cfg(target_arch = "x86_64")]
4925 unsafe {
4926 if vnni_tiles_enabled() {
4927 return dot_q4t_row_vnni(bytes, r, gpr, xq);
4928 }
4929 return dot_q4t_row_avx2(bytes, r, gpr, xq);
4930 }
4931 let mut acc = 0f32;
4932 for gi in 0..gpr {
4933 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
4934 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
4935 let mut d = 0i32;
4936 for (k, &b) in tile[2..].iter().enumerate() {
4937 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
4938 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
4939 }
4940 acc += d as f32 * s;
4941 }
4942 acc
4943}
4944
4945#[cfg(target_arch = "aarch64")]
4946#[target_feature(enable = "neon,dotprod")]
4947unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4948 unsafe {
4951 use core::arch::aarch64::*;
4952 use core::arch::asm;
4953 let lomask = vdupq_n_u8(0x0F);
4954 let eight = vdupq_n_s8(8);
4955 let mut acc = 0f32;
4956 for gi in 0..gpr {
4957 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4958 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4959 let b = vld1q_u8(t.add(2));
4960 let lo = vandq_u8(b, lomask);
4961 let hi = vshrq_n_u8::<4>(b);
4962 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
4963 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
4964 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
4965 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
4966 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4967 asm!(
4968 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
4969 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
4970 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4971 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
4972 options(pure, nomem, nostack),
4973 );
4974 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
4975 }
4976 acc
4977 }
4978}
4979
4980#[cfg(target_arch = "x86_64")]
4981#[target_feature(enable = "avx2")]
4982unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4983 unsafe {
4985 use core::arch::x86_64::*;
4986 let lomask = _mm_set1_epi8(0x0F);
4987 let eight = _mm256_set1_epi8(8);
4988 let ones = _mm256_set1_epi16(1);
4989 let mut acc = 0f32;
4990 for gi in 0..gpr {
4991 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4992 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4993 let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4994 let lo = _mm_and_si128(b, lomask);
4995 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4996 let w = _mm256_sub_epi8(
4997 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4998 eight,
4999 );
5000 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5001 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5002 let d = _mm256_madd_epi16(p16, ones);
5003 let hi128 = _mm256_extracti128_si256::<1>(d);
5004 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5005 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5006 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5007 acc += _mm_cvtsi128_si32(s32) as f32 * s;
5008 }
5009 acc
5010 }
5011}
5012
5013#[cfg(target_arch = "x86_64")]
5017#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5018unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
5019 unsafe {
5021 use core::arch::x86_64::*;
5022 let lomask = _mm_set1_epi8(0x0F);
5023 let eight = _mm256_set1_epi8(8);
5024 let mut acc = 0f32;
5025 for gi in 0..gpr {
5026 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5027 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5028 let b = _mm_loadu_si128(t.add(2) as *const __m128i);
5029 let lo = _mm_and_si128(b, lomask);
5030 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5031 let w = _mm256_sub_epi8(
5032 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5033 eight,
5034 );
5035 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5036 let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5037 acc += d as f32 * s;
5038 }
5039 acc
5040 }
5041}
5042
5043#[cfg(target_arch = "x86_64")]
5048#[target_feature(enable = "avx2,fma")]
5053unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5054 unsafe {
5056 use core::arch::x86_64::*;
5057 let lomask = _mm_set1_epi8(0x0F);
5058 let eight = _mm256_set1_epi8(8);
5059 let ones = _mm256_set1_epi16(1);
5060 let mut f0 = _mm256_setzero_ps();
5073 let mut f1 = _mm256_setzero_ps();
5074 let mut f2 = _mm256_setzero_ps();
5075 let mut f3 = _mm256_setzero_ps();
5076 for gi in 0..gpr {
5077 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5078 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5079 let sv = _mm256_set1_ps(s);
5080 let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5081 let lo = _mm_and_si128(bb, lomask);
5082 let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5083 let w = _mm256_sub_epi8(
5084 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5085 eight,
5086 );
5087 let aw = _mm256_abs_epi8(w);
5088 let off = gi * GROUP_SIZE;
5089 let dot = |xq: &[i8]| {
5090 let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5091 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5092 _mm256_cvtepi32_ps(_mm256_madd_epi16(p16, ones))
5093 };
5094 f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5095 f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5096 f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5097 f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5098 }
5099 [
5100 hsum256_ps(f0),
5101 hsum256_ps(f1),
5102 hsum256_ps(f2),
5103 hsum256_ps(f3),
5104 ]
5105 }
5106}
5107
5108#[cfg(target_arch = "x86_64")]
5111#[target_feature(enable = "avx2")]
5112#[inline]
5113unsafe fn hsum256_ps(v: core::arch::x86_64::__m256) -> f32 {
5114 unsafe {
5116 use core::arch::x86_64::*;
5117 let hi = _mm256_extractf128_ps::<1>(v);
5118 let s = _mm_add_ps(_mm256_castps256_ps128(v), hi);
5119 let s = _mm_add_ps(s, _mm_movehl_ps(s, s));
5120 let s = _mm_add_ss(s, _mm_shuffle_ps::<0x55>(s, s));
5121 _mm_cvtss_f32(s)
5122 }
5123}
5124
5125#[cfg(target_arch = "x86_64")]
5127#[target_feature(enable = "avx2,fma,avx512f,avx512bw,avx512vl,avx512vnni")]
5128unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5129 unsafe {
5131 use core::arch::x86_64::*;
5132 let lomask = _mm_set1_epi8(0x0F);
5133 let eight = _mm256_set1_epi8(8);
5134 let mut f0 = _mm256_setzero_ps();
5137 let mut f1 = _mm256_setzero_ps();
5138 let mut f2 = _mm256_setzero_ps();
5139 let mut f3 = _mm256_setzero_ps();
5140 for gi in 0..gpr {
5141 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5142 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5143 let sv = _mm256_set1_ps(s);
5144 let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5145 let lo = _mm_and_si128(bb, lomask);
5146 let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5147 let w = _mm256_sub_epi8(
5148 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5149 eight,
5150 );
5151 let aw = _mm256_abs_epi8(w);
5152 let off = gi * GROUP_SIZE;
5153 let dot = |xq: &[i8]| {
5154 let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5155 _mm256_cvtepi32_ps(_mm256_dpbusd_epi32(
5156 _mm256_setzero_si256(),
5157 aw,
5158 _mm256_sign_epi8(x, w),
5159 ))
5160 };
5161 f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5162 f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5163 f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5164 f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5165 }
5166 let acc = [
5167 hsum256_ps(f0),
5168 hsum256_ps(f1),
5169 hsum256_ps(f2),
5170 hsum256_ps(f3),
5171 ];
5172 acc
5173 }
5174}
5175
5176#[cfg(target_arch = "aarch64")]
5181#[target_feature(enable = "neon,dotprod")]
5182unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5183 unsafe {
5185 use core::arch::aarch64::*;
5186 use core::arch::asm;
5187 let lomask = vdupq_n_u8(0x0F);
5188 let eight = vdupq_n_s8(8);
5189 let mut acc = [0f32; 4];
5190 for gi in 0..gpr {
5191 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5192 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5193 let b = vld1q_u8(t.add(2));
5194 let lo = vandq_u8(b, lomask);
5195 let hi = vshrq_n_u8::<4>(b);
5196 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5197 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5198 for (k, xq) in xs.iter().enumerate() {
5199 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5200 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5201 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5202 asm!(
5203 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5204 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5205 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5206 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5207 options(pure, nomem, nostack),
5208 );
5209 acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5210 }
5211 }
5212 acc
5213 }
5214}
5215
5216#[inline]
5218fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
5219 let gi = j / GROUP_SIZE;
5220 let k = j % GROUP_SIZE;
5221 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5222 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5223 let byte = tile[2 + k / 2];
5224 let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5225 ((nib as i32 - 8) as f32, s)
5226}
5227
5228#[inline]
5231fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
5232 let mut acc = 0f32;
5233 for gi in 0..gpr {
5234 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5235 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5236 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5237 let mut ga = 0f32;
5238 for (k, &b) in tile[2..].iter().enumerate() {
5239 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5240 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5241 }
5242 acc += ga * s;
5243 }
5244 acc
5245}
5246
5247struct Q4tpView<'a> {
5251 nib: &'a [u8],
5252 params: &'a [u8],
5253 codes: &'a [u8],
5254 stride: usize,
5255 zero_rung: bool,
5257}
5258
5259impl<'a> Q4tpView<'a> {
5260 fn new(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5261 let (params_off, codes_off, stride) = q4tp_sections(rows, cols);
5262 Self {
5263 nib: &bytes[..params_off],
5264 params: &bytes[params_off..codes_off],
5265 codes: &bytes[codes_off..],
5266 stride,
5267 zero_rung: false,
5268 }
5269 }
5270
5271 fn new_q2(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5273 let (params_off, codes_off, stride) = q2tp_sections(rows, cols);
5274 Self {
5275 nib: &bytes[..params_off],
5276 params: &bytes[params_off..codes_off],
5277 codes: &bytes[codes_off..],
5278 stride,
5279 zero_rung: true,
5280 }
5281 }
5282
5283 #[inline]
5299 fn scales_into(&self, r: usize, gpr: usize, out: &mut [f32]) {
5300 let tab = if self.zero_rung {
5301 q2tp_ladder(self.params, r)
5302 } else {
5303 q4tp_ladder(self.params, r)
5304 };
5305 let codes = &self.codes[r * self.stride..(r + 1) * self.stride];
5306 let out = &mut out[..gpr];
5307 let mut chunks = out.chunks_exact_mut(8);
5308 let mut ci = 0usize;
5309 for c in &mut chunks {
5310 let w = u64::from(codes[ci])
5311 | u64::from(codes[ci + 1]) << 8
5312 | u64::from(codes[ci + 2]) << 16
5313 | u64::from(codes[ci + 3]) << 24
5314 | u64::from(codes[ci + 4]) << 32;
5315 for (k, o) in c.iter_mut().enumerate() {
5316 *o = tab[((w >> (5 * k)) & 31) as usize];
5317 }
5318 ci += 5;
5319 }
5320 let tail = &codes[ci..];
5323 for (k, o) in chunks.into_remainder().iter_mut().enumerate() {
5324 *o = tab[q4tp_code(tail, k)];
5325 }
5326 }
5327}
5328
5329#[inline]
5330fn dot_q4tp_row_i8(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5331 #[cfg(target_arch = "aarch64")]
5332 unsafe {
5333 return dot_q4tp_row_sdot(nib, r, gpr, xq, scales);
5334 }
5335 #[cfg(target_arch = "x86_64")]
5336 unsafe {
5337 if vnni_tiles_enabled() {
5338 return dot_q4tp_row_vnni(nib, r, gpr, xq, scales);
5339 }
5340 return dot_q4tp_row_avx2(nib, r, gpr, xq, scales);
5341 }
5342 #[allow(unreachable_code)]
5343 {
5344 let mut acc = 0f32;
5345 for gi in 0..gpr {
5346 let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5347 let s = scales[gi];
5348 let mut d = 0i32;
5349 for (k, &b) in tile.iter().enumerate() {
5350 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5351 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5352 }
5353 acc += d as f32 * s;
5354 }
5355 acc
5356 }
5357}
5358
5359#[cfg(target_arch = "aarch64")]
5362#[target_feature(enable = "neon,dotprod")]
5363unsafe fn dot_q4tp_row_sdot(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5364 unsafe {
5367 use core::arch::aarch64::*;
5368 use core::arch::asm;
5369 let lomask = vdupq_n_u8(0x0F);
5370 let eight = vdupq_n_s8(8);
5371 let mut acc = 0f32;
5372 for gi in 0..gpr {
5373 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5374 let s = *scales.get_unchecked(gi);
5375 let b = vld1q_u8(t);
5376 let lo = vandq_u8(b, lomask);
5377 let hi = vshrq_n_u8::<4>(b);
5378 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5379 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5380 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5381 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5382 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5383 asm!(
5384 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5385 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5386 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5387 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5388 options(pure, nomem, nostack),
5389 );
5390 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5391 }
5392 acc
5393 }
5394}
5395
5396#[cfg(target_arch = "x86_64")]
5397#[target_feature(enable = "avx2")]
5398unsafe fn dot_q4tp_row_avx2(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5399 unsafe {
5401 use core::arch::x86_64::*;
5402 let lomask = _mm_set1_epi8(0x0F);
5403 let eight = _mm256_set1_epi8(8);
5404 let ones = _mm256_set1_epi16(1);
5405 let mut acc = 0f32;
5406 for gi in 0..gpr {
5407 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5408 let s = *scales.get_unchecked(gi);
5409 let b = _mm_loadu_si128(t as *const __m128i);
5410 let lo = _mm_and_si128(b, lomask);
5411 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5412 let w = _mm256_sub_epi8(
5413 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5414 eight,
5415 );
5416 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5417 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5418 let d = _mm256_madd_epi16(p16, ones);
5419 let hi128 = _mm256_extracti128_si256::<1>(d);
5420 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5421 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5422 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5423 acc += _mm_cvtsi128_si32(s32) as f32 * s;
5424 }
5425 acc
5426 }
5427}
5428
5429#[cfg(target_arch = "x86_64")]
5432#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5433unsafe fn dot_q4tp_row_vnni(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5434 unsafe {
5436 use core::arch::x86_64::*;
5437 let lomask = _mm_set1_epi8(0x0F);
5438 let eight = _mm256_set1_epi8(8);
5439 let mut acc = 0f32;
5440 for gi in 0..gpr {
5441 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5442 let s = *scales.get_unchecked(gi);
5443 let b = _mm_loadu_si128(t as *const __m128i);
5444 let lo = _mm_and_si128(b, lomask);
5445 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5446 let w = _mm256_sub_epi8(
5447 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5448 eight,
5449 );
5450 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5451 acc += dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w)) as f32 * s;
5452 }
5453 acc
5454 }
5455}
5456
5457#[inline]
5460fn q4tp_row_exact(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5461 #[cfg(target_arch = "x86_64")]
5462 if avx2_enabled() {
5463 return unsafe { q4tp_row_float_avx2(nib, r, gpr, x, scales) };
5465 }
5466 q4tp_row_float_scalar(nib, r, gpr, x, scales)
5467}
5468
5469#[inline]
5470fn q4tp_row_float_scalar(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5471 let mut acc = 0f32;
5472 for gi in 0..gpr {
5473 let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5474 let s = scales[gi];
5475 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5476 let mut ga = 0f32;
5477 for (k, &b) in tile.iter().enumerate() {
5478 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5479 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5480 }
5481 acc += ga * s;
5482 }
5483 acc
5484}
5485
5486#[cfg(target_arch = "x86_64")]
5490#[target_feature(enable = "avx2")]
5491unsafe fn q4tp_row_float_avx2(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5492 unsafe {
5495 use core::arch::x86_64::*;
5496 let mask = _mm_set1_epi8(15);
5497 let eight = _mm_set1_epi8(8);
5498 let order = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
5499 let mut acc = 0.0f32;
5500 for gi in 0..gpr {
5501 let packed = _mm_loadu_si128(nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB).cast());
5502 let lo = _mm_and_si128(packed, mask);
5503 let hi = _mm_and_si128(_mm_srli_epi16::<4>(packed), mask);
5504 let w0 = _mm_sub_epi8(_mm_unpacklo_epi8(lo, hi), eight);
5505 let w1 = _mm_sub_epi8(_mm_unpackhi_epi8(lo, hi), eight);
5506 let xp = x.as_ptr().add(gi * GROUP_SIZE);
5507 let a = _mm256_mul_ps(
5508 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w0)),
5509 _mm256_loadu_ps(xp),
5510 );
5511 let b = _mm256_mul_ps(
5512 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w0))),
5513 _mm256_loadu_ps(xp.add(8)),
5514 );
5515 let c = _mm256_mul_ps(
5516 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w1)),
5517 _mm256_loadu_ps(xp.add(16)),
5518 );
5519 let d = _mm256_mul_ps(
5520 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w1))),
5521 _mm256_loadu_ps(xp.add(24)),
5522 );
5523 let mut pairs = [0.0f32; 16];
5524 _mm256_storeu_ps(
5525 pairs.as_mut_ptr(),
5526 _mm256_permutevar8x32_ps(_mm256_hadd_ps(a, b), order),
5527 );
5528 _mm256_storeu_ps(
5529 pairs.as_mut_ptr().add(8),
5530 _mm256_permutevar8x32_ps(_mm256_hadd_ps(c, d), order),
5531 );
5532 let mut ga = 0.0f32;
5533 for v in pairs {
5534 ga += v;
5535 }
5536 acc += ga * scales[gi];
5537 }
5538 acc
5539 }
5540}
5541
5542#[inline]
5545fn q4tp_outlier(nib: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5546 let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5547 let byte = nib[(r * gpr + gi) * Q4TP_NIB + k / 2];
5548 let n = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5549 ((n as i32 - 8) as f32, scales[gi])
5550}
5551
5552fn q4tp_matvec(
5554 bytes: &[u8],
5555 x: &[f32],
5556 rows: usize,
5557 cols: usize,
5558 out: &mut [f32],
5559 pool: Option<&Pool>,
5560) {
5561 debug_assert_eq!(out.len(), rows);
5562 let gpr = cols / GROUP_SIZE;
5563 let v = Q4tpView::new(bytes, rows, cols);
5564 let out_addr = SendMut(out.as_mut_ptr());
5565 if a8w8_enabled() {
5566 let act = split_act(x);
5567 let run = |start: usize, end: usize| {
5568 with_krow(gpr, |sc| {
5570 for r in start..end {
5571 v.scales_into(r, gpr, sc);
5572 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
5573 for &(j, xv) in &act.outliers {
5574 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
5575 acc += w * s * xv;
5576 }
5577 unsafe { *out_addr.at(r) = acc };
5579 }
5580 })
5581 };
5582 dispatch_rows(pool, rows, &run);
5583 return;
5584 }
5585 let run = |start: usize, end: usize| {
5586 with_krow(gpr, |sc| {
5587 for r in start..end {
5588 v.scales_into(r, gpr, sc);
5589 unsafe { *out_addr.at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
5591 }
5592 })
5593 };
5594 dispatch_rows(pool, rows, &run);
5595}
5596
5597#[allow(clippy::too_many_arguments)]
5600fn q4tp_matvec2(
5601 bytes: &[u8],
5602 x1: &[f32],
5603 x2: &[f32],
5604 rows: usize,
5605 cols: usize,
5606 o1: &mut [f32],
5607 o2: &mut [f32],
5608 pool: Option<&Pool>,
5609) {
5610 let gpr = cols / GROUP_SIZE;
5611 let v = Q4tpView::new(bytes, rows, cols);
5612 let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5613 let run = |start: usize, end: usize| {
5614 let mut sc = vec![0f32; gpr];
5615 for r in start..end {
5616 v.scales_into(r, gpr, &mut sc);
5617 unsafe {
5619 *p1.at(r) = q4tp_row_exact(v.nib, r, gpr, x1, &sc);
5620 *p2.at(r) = q4tp_row_exact(v.nib, r, gpr, x2, &sc);
5621 }
5622 }
5623 };
5624 dispatch_rows(pool, rows, &run);
5625}
5626
5627#[inline]
5630fn q2tp_outlier(chunks: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5631 let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5632 let byte = chunks[(r * gpr + gi) * Q2TP_CHUNK + k / 4];
5633 let c = (byte >> (2 * (k % 4))) & 3;
5634 (c as f32 - 1.5, scales[gi])
5635}
5636
5637#[cfg(target_arch = "x86_64")]
5638const Q2TP_DECODE_U32: [u32; 256] = {
5639 let mut tab = [0u32; 256];
5640 let mut b = 0usize;
5641 while b < 256 {
5642 tab[b] = ((b as u32) & 3)
5643 | ((((b as u32) >> 2) & 3) << 8)
5644 | ((((b as u32) >> 4) & 3) << 16)
5645 | ((((b as u32) >> 6) & 3) << 24);
5646 b += 1;
5647 }
5648 tab
5649};
5650
5651#[cfg(target_arch = "x86_64")]
5655#[target_feature(enable = "avx2")]
5656unsafe fn q2tp_code_dot_avx2(ch: &[u8], x: &[i8]) -> i32 {
5657 use core::arch::x86_64::*;
5658 debug_assert!(ch.len() >= Q2TP_CHUNK && x.len() >= GROUP_SIZE);
5659 let codes = _mm256_setr_epi32(
5660 Q2TP_DECODE_U32[ch[0] as usize] as i32,
5661 Q2TP_DECODE_U32[ch[1] as usize] as i32,
5662 Q2TP_DECODE_U32[ch[2] as usize] as i32,
5663 Q2TP_DECODE_U32[ch[3] as usize] as i32,
5664 Q2TP_DECODE_U32[ch[4] as usize] as i32,
5665 Q2TP_DECODE_U32[ch[5] as usize] as i32,
5666 Q2TP_DECODE_U32[ch[6] as usize] as i32,
5667 Q2TP_DECODE_U32[ch[7] as usize] as i32,
5668 );
5669 let xv = unsafe { _mm256_loadu_si256(x.as_ptr().cast()) };
5670 let pair = _mm256_maddubs_epi16(codes, xv);
5671 let quad = _mm256_madd_epi16(pair, _mm256_set1_epi16(1));
5672 let sum128 = _mm_add_epi32(
5673 _mm256_castsi256_si128(quad),
5674 _mm256_extracti128_si256(quad, 1),
5675 );
5676 let sum64 = _mm_hadd_epi32(sum128, sum128);
5677 _mm_cvtsi128_si32(_mm_hadd_epi32(sum64, sum64))
5678}
5679
5680#[cfg(target_arch = "x86_64")]
5681#[inline]
5682fn q2tp_avx2_enabled() -> bool {
5683 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5684 *ON.get_or_init(|| std::arch::is_x86_feature_detected!("avx2"))
5685}
5686
5687#[inline]
5694fn dot_q2tp_row_i8(
5695 chunks: &[u8],
5696 r: usize,
5697 gpr: usize,
5698 xq: &[i8],
5699 gsum: &[i32],
5700 scales: &[f32],
5701) -> f32 {
5702 let mut acc = 0f32;
5703 let base = r * gpr * Q2TP_CHUNK;
5704 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5705 let mut codes = [0i8; GROUP_SIZE];
5706 #[cfg(target_arch = "x86_64")]
5707 let avx2 = q2tp_avx2_enabled();
5712 for gi in 0..gpr {
5713 let ch = &chunks[base + gi * Q2TP_CHUNK..base + (gi + 1) * Q2TP_CHUNK];
5714 let xg = &xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5715 #[cfg(target_arch = "aarch64")]
5716 let dot = unsafe {
5722 use core::arch::aarch64::*;
5723 let b = vld1_u8(ch.as_ptr());
5724 let three = vdup_n_u8(3);
5725 let c0 = vreinterpret_s8_u8(vand_u8(b, three));
5726 let c1 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 2), three));
5727 let c2 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 4), three));
5728 let c3 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 6), three));
5729 let x4 = vld4_s8(xg.as_ptr());
5730 let mut acc4 = vdupq_n_s32(0);
5731 acc4 = vpadalq_s16(acc4, vmull_s8(c0, x4.0));
5732 acc4 = vpadalq_s16(acc4, vmull_s8(c1, x4.1));
5733 acc4 = vpadalq_s16(acc4, vmull_s8(c2, x4.2));
5734 acc4 = vpadalq_s16(acc4, vmull_s8(c3, x4.3));
5735 vaddvq_s32(acc4)
5736 };
5737 #[cfg(target_arch = "x86_64")]
5738 let dot: i32 = if avx2 {
5739 unsafe { q2tp_code_dot_avx2(ch, xg) }
5742 } else {
5743 ch.iter()
5744 .enumerate()
5745 .map(|(k, &b)| {
5746 ((b & 3) as i32) * xg[k * 4] as i32
5747 + (((b >> 2) & 3) as i32) * xg[k * 4 + 1] as i32
5748 + (((b >> 4) & 3) as i32) * xg[k * 4 + 2] as i32
5749 + (((b >> 6) & 3) as i32) * xg[k * 4 + 3] as i32
5750 })
5751 .sum()
5752 };
5753 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5754 let dot: i32 = {
5755 for (k, &b) in ch.iter().enumerate() {
5756 codes[k * 4] = (b & 3) as i8;
5757 codes[k * 4 + 1] = ((b >> 2) & 3) as i8;
5758 codes[k * 4 + 2] = ((b >> 4) & 3) as i8;
5759 codes[k * 4 + 3] = ((b >> 6) & 3) as i8;
5760 }
5761 codes
5762 .iter()
5763 .zip(xg)
5764 .map(|(&c, &x)| c as i32 * x as i32)
5765 .sum()
5766 };
5767 acc += scales[gi] * (dot as f32 - 1.5 * gsum[gi] as f32);
5768 }
5769 acc
5770}
5771
5772fn q2tp_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5776 q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.5)
5777}
5778
5779#[inline]
5783fn q2tp_affine_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5784 q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.0)
5785}
5786
5787#[inline]
5788fn q2tp_row_exact_center(
5789 chunks: &[u8],
5790 r: usize,
5791 gpr: usize,
5792 x: &[f32],
5793 scales: &[f32],
5794 center: f32,
5795) -> f32 {
5796 let mut acc = 0f32;
5797 for gi in 0..gpr {
5798 let ch = &chunks[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
5799 let s = scales[gi];
5800 let xb = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5801 let mut g = 0f32;
5802 for (k, &b) in ch.iter().enumerate() {
5803 g += ((b & 3) as f32 - center) * xb[k * 4]
5804 + (((b >> 2) & 3) as f32 - center) * xb[k * 4 + 1]
5805 + (((b >> 4) & 3) as f32 - center) * xb[k * 4 + 2]
5806 + (((b >> 6) & 3) as f32 - center) * xb[k * 4 + 3];
5807 }
5808 acc += s * g;
5809 }
5810 acc
5811}
5812
5813fn q2tp_matvec(
5814 bytes: &[u8],
5815 x: &[f32],
5816 rows: usize,
5817 cols: usize,
5818 out: &mut [f32],
5819 pool: Option<&Pool>,
5820) {
5821 q2tp_matvec_mode(bytes, x, rows, cols, out, pool, false);
5822}
5823
5824fn q2tp_affine_matvec(
5825 bytes: &[u8],
5826 x: &[f32],
5827 rows: usize,
5828 cols: usize,
5829 out: &mut [f32],
5830 pool: Option<&Pool>,
5831) {
5832 q2tp_matvec_mode(bytes, x, rows, cols, out, pool, true);
5833}
5834
5835fn q2tp_matvec_mode(
5836 bytes: &[u8],
5837 x: &[f32],
5838 rows: usize,
5839 cols: usize,
5840 out: &mut [f32],
5841 pool: Option<&Pool>,
5842 affine: bool,
5843) {
5844 debug_assert_eq!(out.len(), rows);
5845 let gpr = cols / GROUP_SIZE;
5846 let v = Q4tpView::new_q2(bytes, rows, cols);
5847 let out_addr = SendMut(out.as_mut_ptr());
5848 if !affine && a8w8_enabled() {
5853 let act = split_act(x);
5854 let gsum = q1_group_sums(&act.xq, gpr);
5855 let (act, gsum) = (&act, &gsum);
5856 let run = move |start: usize, end: usize| {
5857 with_krow(gpr, |sc| {
5858 for r in start..end {
5859 v.scales_into(r, gpr, sc);
5860 let mut acc = dot_q2tp_row_i8(v.nib, r, gpr, &act.xq, gsum, sc) * act.sx;
5861 for &(j, xv) in &act.outliers {
5862 let (w, s) = q2tp_outlier(v.nib, r, gpr, j, sc);
5863 acc += w * s * xv;
5864 }
5865 unsafe { *out_addr.at(r) = acc };
5867 }
5868 })
5869 };
5870 dispatch_rows(pool, rows, &run);
5871 return;
5872 }
5873 let run = |start: usize, end: usize| {
5874 with_krow(gpr, |sc| {
5875 for r in start..end {
5876 v.scales_into(r, gpr, sc);
5877 unsafe {
5879 *out_addr.at(r) = if affine {
5880 q2tp_affine_row_exact(v.nib, r, gpr, x, sc)
5881 } else {
5882 q2tp_row_exact(v.nib, r, gpr, x, sc)
5883 }
5884 };
5885 }
5886 })
5887 };
5888 dispatch_rows(pool, rows, &run);
5889}
5890
5891#[allow(clippy::too_many_arguments)]
5893fn q2tp_matvec2(
5894 bytes: &[u8],
5895 x1: &[f32],
5896 x2: &[f32],
5897 rows: usize,
5898 cols: usize,
5899 o1: &mut [f32],
5900 o2: &mut [f32],
5901 pool: Option<&Pool>,
5902) {
5903 q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, false);
5904}
5905
5906#[allow(clippy::too_many_arguments)]
5907fn q2tp_affine_matvec2(
5908 bytes: &[u8],
5909 x1: &[f32],
5910 x2: &[f32],
5911 rows: usize,
5912 cols: usize,
5913 o1: &mut [f32],
5914 o2: &mut [f32],
5915 pool: Option<&Pool>,
5916) {
5917 q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, true);
5918}
5919
5920#[allow(clippy::too_many_arguments)]
5921fn q2tp_matvec2_mode(
5922 bytes: &[u8],
5923 x1: &[f32],
5924 x2: &[f32],
5925 rows: usize,
5926 cols: usize,
5927 o1: &mut [f32],
5928 o2: &mut [f32],
5929 pool: Option<&Pool>,
5930 affine: bool,
5931) {
5932 let gpr = cols / GROUP_SIZE;
5933 let v = Q4tpView::new_q2(bytes, rows, cols);
5934 let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5935 let run = |start: usize, end: usize| {
5936 let mut sc = vec![0f32; gpr];
5937 for r in start..end {
5938 v.scales_into(r, gpr, &mut sc);
5939 unsafe {
5941 *p1.at(r) = if affine {
5942 q2tp_affine_row_exact(v.nib, r, gpr, x1, &sc)
5943 } else {
5944 q2tp_row_exact(v.nib, r, gpr, x1, &sc)
5945 };
5946 *p2.at(r) = if affine {
5947 q2tp_affine_row_exact(v.nib, r, gpr, x2, &sc)
5948 } else {
5949 q2tp_row_exact(v.nib, r, gpr, x2, &sc)
5950 };
5951 }
5952 }
5953 };
5954 dispatch_rows(pool, rows, &run);
5955}
5956
5957pub fn q2tp_matvec_for_test(bytes: &[u8], x: &[f32], rows: usize, cols: usize, out: &mut [f32]) {
5964 let gpr = cols / GROUP_SIZE;
5969 let v = Q4tpView::new_q2(bytes, rows, cols);
5970 with_krow(gpr, |sc| {
5971 for r in 0..rows {
5972 v.scales_into(r, gpr, sc);
5973 out[r] = q2tp_row_exact(v.nib, r, gpr, x, sc);
5974 }
5975 });
5976}
5977
5978pub fn q2tp_affine_matvec_for_test(
5981 bytes: &[u8],
5982 x: &[f32],
5983 rows: usize,
5984 cols: usize,
5985 out: &mut [f32],
5986) {
5987 q2tp_affine_matvec(bytes, x, rows, cols, out, None);
5988}
5989
5990pub fn q2tp_matmat_for_test(
5991 bytes: &[u8],
5992 xs_all: &[f32],
5993 b: usize,
5994 rows: usize,
5995 cols: usize,
5996 out: &mut [f32],
5997) {
5998 q2tp_matmat(bytes, xs_all, b, rows, cols, out, None);
5999}
6000
6001fn q2tp_matmat(
6002 bytes: &[u8],
6003 xs_all: &[f32],
6004 b: usize,
6005 rows: usize,
6006 cols: usize,
6007 out: &mut [f32],
6008 pool: Option<&Pool>,
6009) {
6010 q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, false);
6011}
6012
6013fn q2tp_affine_matmat(
6014 bytes: &[u8],
6015 xs_all: &[f32],
6016 b: usize,
6017 rows: usize,
6018 cols: usize,
6019 out: &mut [f32],
6020 pool: Option<&Pool>,
6021) {
6022 q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, true);
6023}
6024
6025fn q2tp_matmat_mode(
6026 bytes: &[u8],
6027 xs_all: &[f32],
6028 b: usize,
6029 rows: usize,
6030 cols: usize,
6031 out: &mut [f32],
6032 pool: Option<&Pool>,
6033 affine: bool,
6034) {
6035 debug_assert_eq!(out.len(), b * rows);
6036 let gpr = cols / GROUP_SIZE;
6037 let v = Q4tpView::new_q2(bytes, rows, cols);
6038 let out_addr = SendMut(out.as_mut_ptr());
6039 let run = |start: usize, end: usize| {
6040 let mut sc = vec![0f32; gpr];
6041 for r in start..end {
6042 v.scales_into(r, gpr, &mut sc);
6043 for bi in 0..b {
6044 let x = &xs_all[bi * cols..(bi + 1) * cols];
6045 unsafe {
6047 *out_addr.at(bi * rows + r) = if affine {
6048 q2tp_affine_row_exact(v.nib, r, gpr, x, &sc)
6049 } else {
6050 q2tp_row_exact(v.nib, r, gpr, x, &sc)
6051 }
6052 };
6053 }
6054 }
6055 };
6056 dispatch_rows(pool, rows, &run);
6057}
6058
6059#[cfg(target_arch = "aarch64")]
6067#[target_feature(enable = "neon,dotprod")]
6068unsafe fn dot_q4tp_row_1x4_sdot_v1(
6069 nib: &[u8],
6070 r: usize,
6071 gpr: usize,
6072 xs: [&[i8]; 4],
6073 scales: &[f32],
6074) -> [f32; 4] {
6075 unsafe {
6076 use core::arch::aarch64::*;
6077 use core::arch::asm;
6078 let lomask = vdupq_n_u8(0x0F);
6079 let eight = vdupq_n_s8(8);
6080 let (mut f0, mut f1, mut f2, mut f3) = (0f32, 0f32, 0f32, 0f32);
6081 for gi in 0..gpr {
6082 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6083 let s = *scales.get_unchecked(gi);
6084 let bb = vld1q_u8(t);
6085 let lo = vandq_u8(bb, lomask);
6086 let hi = vshrq_n_u8::<4>(bb);
6087 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6088 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6089 let mut d = [0f32; 4];
6090 for (k, dk) in d.iter_mut().enumerate() {
6091 let x0 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE));
6092 let x1 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE + 16));
6093 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6094 asm!(
6095 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6096 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6097 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6098 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6099 options(pure, nomem, nostack),
6100 );
6101 *dk = vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
6102 }
6103 f0 += d[0];
6104 f1 += d[1];
6105 f2 += d[2];
6106 f3 += d[3];
6107 }
6108 [f0, f1, f2, f3]
6109 }
6110}
6111
6112#[allow(dead_code)]
6119static Q4TP_ALT: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
6120
6121#[cfg(test)]
6124static Q4TP_ALT_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
6125
6126#[cfg(target_arch = "x86_64")]
6132fn q4tp_blocked_x86() -> bool {
6133 match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6134 1 => false,
6135 2 => avx512vnni_enabled(),
6140 _ => blocked_enabled() && avx512vnni_enabled(),
6144 }
6145}
6146
6147#[cfg(target_arch = "aarch64")]
6149#[allow(dead_code)]
6150fn q4tp_v1() -> bool {
6151 match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6152 1 => true,
6153 2 => false,
6154 _ => {
6155 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6156 *ON.get_or_init(|| std::env::var("CMF_Q4TP_V1").is_ok_and(|v| v != "0"))
6157 }
6158 }
6159}
6160
6161#[cfg(target_arch = "x86_64")]
6173#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6174unsafe fn dot_q4tp_2x8_avx512(
6175 nib: &[u8],
6176 r0: usize,
6177 gpr: usize,
6178 xs: [&[i8]; 8],
6179 sc0: &[f32],
6180 sc1: &[f32],
6181) -> [[f32; 8]; 2] {
6182 unsafe {
6185 use core::arch::x86_64::*;
6186 let lomask = _mm256_set1_epi8(0x0F);
6187 let eight = _mm256_set1_epi8(8);
6188 let zero = _mm512_setzero_si512();
6189 let mut v0 = [_mm512_setzero_ps(); 8];
6190 let mut v1 = [_mm512_setzero_ps(); 8];
6191 let pairs = gpr / 2;
6192 let unpack = |r: usize, gi: usize| -> (__m512i, __mmask64) {
6193 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6194 let bb = _mm256_loadu_si256(t as *const __m256i);
6195 let lo = _mm256_and_si256(bb, lomask);
6196 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6197 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6198 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6199 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6200 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6201 (_mm512_abs_epi8(w), _mm512_movepi8_mask(w))
6202 };
6203 for gp in 0..pairs {
6204 let gi = gp * 2;
6205 let (wa0, neg0) = unpack(r0, gi);
6206 let (wa1, neg1) = unpack(r0 + 1, gi);
6207 let off = gi * GROUP_SIZE;
6208 let sv = |sc: &[f32]| {
6209 _mm512_insertf32x8::<1>(
6210 _mm512_castps256_ps512(_mm256_set1_ps(*sc.get_unchecked(gi))),
6211 _mm256_set1_ps(*sc.get_unchecked(gi + 1)),
6212 )
6213 };
6214 let s0 = sv(sc0);
6215 let s1 = sv(sc1);
6216 for k in 0..8 {
6217 let xv = _mm512_loadu_si512(xs[k].as_ptr().add(off) as *const __m512i);
6218 let d0 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6219 zero,
6220 wa0,
6221 _mm512_mask_sub_epi8(xv, neg0, zero, xv),
6222 ));
6223 let d1 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6224 zero,
6225 wa1,
6226 _mm512_mask_sub_epi8(xv, neg1, zero, xv),
6227 ));
6228 v0[k] = _mm512_fmadd_ps(d0, s0, v0[k]);
6229 v1[k] = _mm512_fmadd_ps(d1, s1, v1[k]);
6230 }
6231 }
6232 let mut acc = [[0f32; 8]; 2];
6233 for k in 0..8 {
6234 acc[0][k] = _mm512_reduce_add_ps(v0[k]);
6235 acc[1][k] = _mm512_reduce_add_ps(v1[k]);
6236 }
6237 if gpr % 2 == 1 {
6238 let off = (gpr - 1) * GROUP_SIZE;
6239 for j in off..off + GROUP_SIZE {
6240 let (w0, sa) = q4tp_outlier(nib, r0, gpr, j, sc0);
6241 let (w1, sb) = q4tp_outlier(nib, r0 + 1, gpr, j, sc1);
6242 for k in 0..8 {
6243 let x = *xs[k].get_unchecked(j) as f32;
6244 acc[0][k] += w0 * sa * x;
6245 acc[1][k] += w1 * sb * x;
6246 }
6247 }
6248 }
6249 acc
6250 }
6251}
6252
6253#[cfg(target_arch = "x86_64")]
6258#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6259unsafe fn dot_q4tp_row_1x8_avx512(
6260 nib: &[u8],
6261 r: usize,
6262 gpr: usize,
6263 xs: [&[i8]; 8],
6264 scales: &[f32],
6265) -> [f32; 8] {
6266 unsafe {
6268 use core::arch::x86_64::*;
6269 let lomask = _mm256_set1_epi8(0x0F);
6270 let eight = _mm256_set1_epi8(8);
6271 let zero = _mm512_setzero_si512();
6272 let (mut v0, mut v1, mut v2, mut v3) = (
6273 _mm512_setzero_ps(),
6274 _mm512_setzero_ps(),
6275 _mm512_setzero_ps(),
6276 _mm512_setzero_ps(),
6277 );
6278 let (mut v4, mut v5, mut v6, mut v7) = (
6279 _mm512_setzero_ps(),
6280 _mm512_setzero_ps(),
6281 _mm512_setzero_ps(),
6282 _mm512_setzero_ps(),
6283 );
6284 let pairs = gpr / 2;
6285 for gp in 0..pairs {
6286 let gi = gp * 2;
6287 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6288 let bb = _mm256_loadu_si256(t as *const __m256i);
6289 let lo = _mm256_and_si256(bb, lomask);
6290 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6291 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6296 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6297 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6298 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6299 let wabs = _mm512_abs_epi8(w);
6300 let neg = _mm512_movepi8_mask(w);
6301 let off = gi * GROUP_SIZE;
6302 let sv = _mm512_insertf32x8::<1>(
6303 _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6304 _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6305 );
6306 let dot = |x: &[i8]| -> __m512 {
6307 let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6308 let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6309 _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6310 };
6311 v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6312 v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6313 v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6314 v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6315 v4 = _mm512_fmadd_ps(dot(xs[4]), sv, v4);
6316 v5 = _mm512_fmadd_ps(dot(xs[5]), sv, v5);
6317 v6 = _mm512_fmadd_ps(dot(xs[6]), sv, v6);
6318 v7 = _mm512_fmadd_ps(dot(xs[7]), sv, v7);
6319 }
6320 let mut acc = [
6321 _mm512_reduce_add_ps(v0),
6322 _mm512_reduce_add_ps(v1),
6323 _mm512_reduce_add_ps(v2),
6324 _mm512_reduce_add_ps(v3),
6325 _mm512_reduce_add_ps(v4),
6326 _mm512_reduce_add_ps(v5),
6327 _mm512_reduce_add_ps(v6),
6328 _mm512_reduce_add_ps(v7),
6329 ];
6330 if gpr % 2 == 1 {
6333 let off = (gpr - 1) * GROUP_SIZE;
6334 for j in off..off + GROUP_SIZE {
6335 let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6336 let ws = w * s;
6337 for k in 0..8 {
6338 acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6339 }
6340 }
6341 }
6342 acc
6343 }
6344}
6345
6346#[cfg(target_arch = "x86_64")]
6359#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6360unsafe fn dot_q4tp_row_1x4_avx512(
6361 nib: &[u8],
6362 r: usize,
6363 gpr: usize,
6364 xs: [&[i8]; 4],
6365 scales: &[f32],
6366) -> [f32; 4] {
6367 unsafe {
6369 use core::arch::x86_64::*;
6370 let lomask = _mm256_set1_epi8(0x0F);
6371 let eight = _mm256_set1_epi8(8);
6372 let zero = _mm512_setzero_si512();
6373 let (mut v0, mut v1, mut v2, mut v3) = (
6374 _mm512_setzero_ps(),
6375 _mm512_setzero_ps(),
6376 _mm512_setzero_ps(),
6377 _mm512_setzero_ps(),
6378 );
6379 let pairs = gpr / 2;
6380 for gp in 0..pairs {
6381 let gi = gp * 2;
6382 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6383 let bb = _mm256_loadu_si256(t as *const __m256i);
6384 let lo = _mm256_and_si256(bb, lomask);
6385 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6386 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6391 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6392 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6393 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6394 let wabs = _mm512_abs_epi8(w);
6395 let neg = _mm512_movepi8_mask(w);
6396 let off = gi * GROUP_SIZE;
6397 let sv = _mm512_insertf32x8::<1>(
6398 _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6399 _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6400 );
6401 let dot = |x: &[i8]| -> __m512 {
6402 let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6403 let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6404 _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6405 };
6406 v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6407 v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6408 v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6409 v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6410 }
6411 let mut acc = [
6412 _mm512_reduce_add_ps(v0),
6413 _mm512_reduce_add_ps(v1),
6414 _mm512_reduce_add_ps(v2),
6415 _mm512_reduce_add_ps(v3),
6416 ];
6417 if gpr % 2 == 1 {
6420 let off = (gpr - 1) * GROUP_SIZE;
6421 for j in off..off + GROUP_SIZE {
6422 let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6423 let ws = w * s;
6424 for k in 0..4 {
6425 acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6426 }
6427 }
6428 }
6429 acc
6430 }
6431}
6432
6433#[cfg(target_arch = "aarch64")]
6437#[target_feature(enable = "neon,dotprod")]
6438unsafe fn dot_q4tp_row_1x4_sdot(
6439 nib: &[u8],
6440 r: usize,
6441 gpr: usize,
6442 xs: [&[i8]; 4],
6443 scales: &[f32],
6444) -> [f32; 4] {
6445 unsafe {
6447 use core::arch::aarch64::*;
6448 use core::arch::asm;
6449 let lomask = vdupq_n_u8(0x0F);
6450 let eight = vdupq_n_s8(8);
6451 let (mut v0, mut v1, mut v2, mut v3) = (
6467 vdupq_n_f32(0.0),
6468 vdupq_n_f32(0.0),
6469 vdupq_n_f32(0.0),
6470 vdupq_n_f32(0.0),
6471 );
6472 for gi in 0..gpr {
6473 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6474 let s = *scales.get_unchecked(gi);
6475 let bb = vld1q_u8(t);
6476 let lo = vandq_u8(bb, lomask);
6477 let hi = vshrq_n_u8::<4>(bb);
6478 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6479 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6480 let off = gi * GROUP_SIZE;
6481 let dot4 = |x: &[i8]| -> int32x4_t {
6482 let x0 = vld1q_s8(x.as_ptr().add(off));
6483 let x1 = vld1q_s8(x.as_ptr().add(off + 16));
6484 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6485 asm!(
6486 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6487 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6488 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6489 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6490 options(pure, nomem, nostack),
6491 );
6492 vaddq_s32(a0, a1)
6493 };
6494 v0 = vfmaq_n_f32(v0, vcvtq_f32_s32(dot4(xs[0])), s);
6495 v1 = vfmaq_n_f32(v1, vcvtq_f32_s32(dot4(xs[1])), s);
6496 v2 = vfmaq_n_f32(v2, vcvtq_f32_s32(dot4(xs[2])), s);
6497 v3 = vfmaq_n_f32(v3, vcvtq_f32_s32(dot4(xs[3])), s);
6498 }
6499 [
6500 vaddvq_f32(v0),
6501 vaddvq_f32(v1),
6502 vaddvq_f32(v2),
6503 vaddvq_f32(v3),
6504 ]
6505 }
6506}
6507
6508fn q4tp_matmat(
6516 bytes: &[u8],
6517 xs_all: &[f32],
6518 b: usize,
6519 rows: usize,
6520 cols: usize,
6521 out: &mut [f32],
6522 pool: Option<&Pool>,
6523) {
6524 q4tp_matmat_with(bytes, xs_all, b, rows, cols, out, pool, row_exact())
6525}
6526
6527#[allow(clippy::too_many_arguments)]
6531fn q4tp_matmat_with(
6532 bytes: &[u8],
6533 xs_all: &[f32],
6534 b: usize,
6535 rows: usize,
6536 cols: usize,
6537 out: &mut [f32],
6538 pool: Option<&Pool>,
6539 exact: bool,
6540) {
6541 debug_assert_eq!(out.len(), b * rows);
6542 let gpr = cols / GROUP_SIZE;
6543 let v = Q4tpView::new(bytes, rows, cols);
6544
6545 #[cfg(target_os = "macos")]
6549 if !exact && b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6550 dequant_matmat_accel(
6551 &|r, dst| {
6552 let mut sc = [0f32; 32];
6553 let mut scv;
6554 let s: &[f32] = if gpr <= 32 {
6555 v.scales_into(r, gpr, &mut sc);
6556 &sc[..gpr]
6557 } else {
6558 scv = vec![0f32; gpr];
6559 v.scales_into(r, gpr, &mut scv);
6560 &scv
6561 };
6562 for gi in 0..gpr {
6563 let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
6564 for (k, &bb) in tile.iter().enumerate() {
6565 dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s[gi];
6566 dst[gi * GROUP_SIZE + k * 2 + 1] =
6567 (((bb >> 4) & 0x0F) as f32 - 8.0) * s[gi];
6568 }
6569 }
6570 },
6571 xs_all,
6572 b,
6573 rows,
6574 cols,
6575 out,
6576 pool,
6577 );
6578 return;
6579 }
6580
6581 let out_addr = SendMut(out.as_mut_ptr());
6582 if a8w8_enabled() {
6583 let acts: Vec<SplitAct> = (0..b)
6584 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6585 .collect();
6586 let acts = &acts;
6587 #[cfg(target_arch = "aarch64")]
6594 let blocked_ok = sdot_enabled() && blocked_enabled();
6595 #[cfg(target_arch = "x86_64")]
6602 let blocked_ok = q4tp_blocked_x86() && !exact;
6603 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
6604 let blocked_ok = {
6605 let _ = exact;
6606 false
6607 };
6608 let panel_cols: usize = std::env::var("CMF_Q4TP_PANEL")
6618 .ok()
6619 .and_then(|v| v.parse().ok())
6620 .filter(|v| *v > 0)
6621 .unwrap_or(256);
6622 let run = |start: usize, end: usize| {
6623 for abase in (0..acts.len()).step_by(panel_cols) {
6624 let alen = (acts.len() - abase).min(panel_cols);
6625 let mut sc = vec![0f32; gpr];
6626 #[cfg(target_arch = "x86_64")]
6627 let mut r_lo = start;
6628 #[cfg(target_arch = "x86_64")]
6629 if blocked_ok && alen >= 8 {
6630 let mut sc1 = vec![0f32; gpr];
6631 while r_lo + 2 <= end {
6632 v.scales_into(r_lo, gpr, &mut sc);
6633 v.scales_into(r_lo + 1, gpr, &mut sc1);
6634 let mut bi = 0usize;
6635 while bi + 8 <= alen {
6636 let xs = [
6637 acts[abase + bi].xq.as_slice(),
6638 acts[abase + bi + 1].xq.as_slice(),
6639 acts[abase + bi + 2].xq.as_slice(),
6640 acts[abase + bi + 3].xq.as_slice(),
6641 acts[abase + bi + 4].xq.as_slice(),
6642 acts[abase + bi + 5].xq.as_slice(),
6643 acts[abase + bi + 6].xq.as_slice(),
6644 acts[abase + bi + 7].xq.as_slice(),
6645 ];
6646 let d = unsafe { dot_q4tp_2x8_avx512(v.nib, r_lo, gpr, xs, &sc, &sc1) };
6647 for (row, dr, scr) in [(r_lo, &d[0], &sc), (r_lo + 1, &d[1], &sc1)] {
6648 for k in 0..8 {
6649 let act = &acts[abase + bi + k];
6650 let mut acc = dr[k] * act.sx;
6651 for &(j, xv) in &act.outliers {
6652 let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6653 acc += w * s * xv;
6654 }
6655 unsafe { *out_addr.at((abase + bi + k) * rows + row) = acc };
6657 }
6658 }
6659 bi += 8;
6660 }
6661 for row in [r_lo, r_lo + 1] {
6664 let scr: &[f32] = if row == r_lo { &sc } else { &sc1 };
6665 for b2 in bi..alen {
6666 let act = &acts[abase + b2];
6667 let xs4 = [
6668 act.xq.as_slice(),
6669 act.xq.as_slice(),
6670 act.xq.as_slice(),
6671 act.xq.as_slice(),
6672 ];
6673 let d =
6674 unsafe { dot_q4tp_row_1x4_avx512(v.nib, row, gpr, xs4, scr) };
6675 let mut acc = d[0] * act.sx;
6676 for &(j, xv) in &act.outliers {
6677 let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6678 acc += w * s * xv;
6679 }
6680 unsafe { *out_addr.at((abase + b2) * rows + row) = acc };
6682 }
6683 }
6684 r_lo += 2;
6685 }
6686 }
6687 #[cfg(target_arch = "x86_64")]
6688 let row_start = r_lo;
6689 #[cfg(not(target_arch = "x86_64"))]
6690 let row_start = start;
6691 for r in row_start..end {
6692 v.scales_into(r, gpr, &mut sc);
6693 let mut bi = 0usize;
6694 #[cfg(target_arch = "x86_64")]
6695 if blocked_ok {
6696 while bi + 8 <= alen {
6697 let xs = [
6698 acts[abase + bi].xq.as_slice(),
6699 acts[abase + bi + 1].xq.as_slice(),
6700 acts[abase + bi + 2].xq.as_slice(),
6701 acts[abase + bi + 3].xq.as_slice(),
6702 acts[abase + bi + 4].xq.as_slice(),
6703 acts[abase + bi + 5].xq.as_slice(),
6704 acts[abase + bi + 6].xq.as_slice(),
6705 acts[abase + bi + 7].xq.as_slice(),
6706 ];
6707 let d = unsafe { dot_q4tp_row_1x8_avx512(v.nib, r, gpr, xs, &sc) };
6708 for k in 0..8 {
6709 let act = &acts[abase + bi + k];
6710 let mut acc = d[k] * act.sx;
6711 for &(j, xv) in &act.outliers {
6712 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6713 acc += w * s * xv;
6714 }
6715 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6717 }
6718 bi += 8;
6719 }
6720 while bi + 4 <= alen {
6721 let xs = [
6722 acts[abase + bi].xq.as_slice(),
6723 acts[abase + bi + 1].xq.as_slice(),
6724 acts[abase + bi + 2].xq.as_slice(),
6725 acts[abase + bi + 3].xq.as_slice(),
6726 ];
6727 let d = unsafe { dot_q4tp_row_1x4_avx512(v.nib, r, gpr, xs, &sc) };
6728 for k in 0..4 {
6729 let act = &acts[abase + bi + k];
6730 let mut acc = d[k] * act.sx;
6731 for &(j, xv) in &act.outliers {
6732 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6733 acc += w * s * xv;
6734 }
6735 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6737 }
6738 bi += 4;
6739 }
6740 }
6741 #[cfg(target_arch = "aarch64")]
6742 if blocked_ok {
6743 while bi + 4 <= alen {
6744 let xs = [
6745 acts[abase + bi].xq.as_slice(),
6746 acts[abase + bi + 1].xq.as_slice(),
6747 acts[abase + bi + 2].xq.as_slice(),
6748 acts[abase + bi + 3].xq.as_slice(),
6749 ];
6750 let d = unsafe {
6751 if exact || q4tp_v1() {
6752 dot_q4tp_row_1x4_sdot_v1(v.nib, r, gpr, xs, &sc)
6753 } else {
6754 dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs, &sc)
6755 }
6756 };
6757 for k in 0..4 {
6758 let act = &acts[abase + bi + k];
6759 let mut acc = d[k] * act.sx;
6760 for &(j, xv) in &act.outliers {
6761 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6762 acc += w * s * xv;
6763 }
6764 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6766 }
6767 bi += 4;
6768 }
6769 }
6770 let _ = blocked_ok;
6771 while bi < alen {
6772 let act = &acts[abase + bi];
6773 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
6774 for &(j, xv) in &act.outliers {
6775 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6776 acc += w * s * xv;
6777 }
6778 unsafe { *out_addr.at((abase + bi) * rows + r) = acc };
6780 bi += 1;
6781 }
6782 }
6783 }
6784 };
6785 dispatch_rows(pool, rows, &run);
6786 return;
6787 }
6788
6789 let run = |start: usize, end: usize| {
6790 let mut sc = vec![0f32; gpr];
6791 for r in start..end {
6792 v.scales_into(r, gpr, &mut sc);
6793 for bi in 0..b {
6794 let x = &xs_all[bi * cols..(bi + 1) * cols];
6795 unsafe { *out_addr.at(bi * rows + r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
6797 }
6798 }
6799 };
6800 dispatch_rows(pool, rows, &run);
6801}
6802
6803fn q4t_matvec(
6805 bytes: &[u8],
6806 x: &[f32],
6807 rows: usize,
6808 cols: usize,
6809 out: &mut [f32],
6810 pool: Option<&Pool>,
6811) {
6812 debug_assert_eq!(out.len(), rows);
6813 let gpr = cols / GROUP_SIZE;
6814 let out_addr = SendMut(out.as_mut_ptr());
6815 if a8w8_enabled() {
6816 let act = split_act(x);
6817 let run = move |start: usize, end: usize| {
6818 for r in start..end {
6819 let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6820 for &(j, xv) in &act.outliers {
6821 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6822 acc += w * s * xv;
6823 }
6824 unsafe { *out_addr.at(r) = acc };
6826 }
6827 };
6828 dispatch_rows(pool, rows, &run);
6829 return;
6830 }
6831 let run = move |start: usize, end: usize| {
6832 for r in start..end {
6833 unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
6835 }
6836 };
6837 dispatch_rows(pool, rows, &run);
6838}
6839
6840#[allow(clippy::too_many_arguments)]
6842fn q4t_matvec2(
6843 bytes: &[u8],
6844 x1: &[f32],
6845 x2: &[f32],
6846 rows: usize,
6847 cols: usize,
6848 o1: &mut [f32],
6849 o2: &mut [f32],
6850 pool: Option<&Pool>,
6851) {
6852 let gpr = cols / GROUP_SIZE;
6853 let p1 = SendMut(o1.as_mut_ptr());
6854 let p2 = SendMut(o2.as_mut_ptr());
6855 if a8w8_enabled() {
6856 let a1 = split_act(x1);
6857 let a2 = split_act(x2);
6858 let run = move |start: usize, end: usize| {
6859 for r in start..end {
6860 let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
6861 let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
6862 for &(j, xv) in &a1.outliers {
6863 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6864 v1 += w * s * xv;
6865 }
6866 for &(j, xv) in &a2.outliers {
6867 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6868 v2 += w * s * xv;
6869 }
6870 unsafe {
6872 *p1.at(r) = v1;
6873 *p2.at(r) = v2;
6874 }
6875 }
6876 };
6877 dispatch_rows(pool, rows, &run);
6878 return;
6879 }
6880 let run = move |start: usize, end: usize| {
6881 for r in start..end {
6882 unsafe {
6884 *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
6885 *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
6886 }
6887 }
6888 };
6889 dispatch_rows(pool, rows, &run);
6890}
6891
6892#[allow(clippy::too_many_arguments)]
6894#[cfg(target_os = "macos")]
6900fn dequant_matmat_accel(
6901 dequant_row: &(dyn Fn(usize, &mut [f32]) + Sync),
6902 xs_all: &[f32],
6903 b: usize,
6904 rows: usize,
6905 cols: usize,
6906 out: &mut [f32],
6907 pool: Option<&Pool>,
6908) {
6909 const TR: usize = 2048;
6910 thread_local! {
6911 static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
6912 }
6913 WTILE.with(|wt| {
6914 let mut wtile = wt.borrow_mut();
6915 wtile.resize(TR * cols, 0.0);
6916 let mut r0 = 0usize;
6917 while r0 < rows {
6918 let tr = TR.min(rows - r0);
6919 let wt_addr = SendMut(wtile.as_mut_ptr());
6920 let run = |start: usize, end: usize| {
6921 for r in start..end {
6922 let dst = unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
6924 dequant_row(r0 + r, dst);
6925 }
6926 };
6927 dispatch_rows(pool, tr, &run);
6928 unsafe {
6929 accel_blas::cblas_sgemm(
6930 101, 111, 112, b as i32,
6934 tr as i32,
6935 cols as i32,
6936 1.0,
6937 xs_all.as_ptr(),
6938 cols as i32,
6939 wtile.as_ptr(),
6940 cols as i32,
6941 0.0,
6942 out.as_mut_ptr().add(r0),
6943 rows as i32,
6944 );
6945 }
6946 r0 += tr;
6947 }
6948 });
6949}
6950
6951fn q4t_matmat(
6952 bytes: &[u8],
6953 xs_all: &[f32],
6954 b: usize,
6955 rows: usize,
6956 cols: usize,
6957 out: &mut [f32],
6958 pool: Option<&Pool>,
6959) {
6960 debug_assert_eq!(out.len(), b * rows);
6961 let gpr = cols / GROUP_SIZE;
6962 #[cfg(target_os = "macos")]
6966 if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6967 dequant_matmat_accel(
6968 &|r, dst| {
6969 for gi in 0..gpr {
6970 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
6971 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
6972 for (k, &bb) in tile[2..].iter().enumerate() {
6973 dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
6974 dst[gi * GROUP_SIZE + k * 2 + 1] = (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
6975 }
6976 }
6977 },
6978 xs_all,
6979 b,
6980 rows,
6981 cols,
6982 out,
6983 pool,
6984 );
6985 return;
6986 }
6987 let out_addr = SendMut(out.as_mut_ptr());
6988 if a8w8_enabled() {
6989 let acts: Vec<SplitAct> = (0..b)
6990 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6991 .collect();
6992 let acts = &acts;
6993 #[cfg(target_arch = "x86_64")]
6994 let blocked_ok = avx2_enabled() && blocked_enabled();
6995 #[cfg(target_arch = "aarch64")]
6996 let blocked_ok = sdot_enabled() && blocked_enabled();
6997 #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
6998 let blocked_ok = false;
6999 let run = move |start: usize, end: usize| {
7000 for r in start..end {
7001 let mut bi = 0usize;
7002 #[cfg(target_arch = "aarch64")]
7003 if blocked_ok {
7004 while bi + 4 <= acts.len() {
7005 let xs = [
7006 acts[bi].xq.as_slice(),
7007 acts[bi + 1].xq.as_slice(),
7008 acts[bi + 2].xq.as_slice(),
7009 acts[bi + 3].xq.as_slice(),
7010 ];
7011 let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
7012 for k in 0..4 {
7013 let act = &acts[bi + k];
7014 let mut acc = d[k] * act.sx;
7015 for &(j, xv) in &act.outliers {
7016 let (w, sc) = q4t_outlier(bytes, r, gpr, j);
7017 acc += w * sc * xv;
7018 }
7019 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
7021 }
7022 bi += 4;
7023 }
7024 }
7025 #[cfg(target_arch = "x86_64")]
7026 if blocked_ok {
7027 while bi + 4 <= acts.len() {
7028 let xs = [
7029 acts[bi].xq.as_slice(),
7030 acts[bi + 1].xq.as_slice(),
7031 acts[bi + 2].xq.as_slice(),
7032 acts[bi + 3].xq.as_slice(),
7033 ];
7034 let d = unsafe {
7035 if vnni_tiles_enabled() {
7036 dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
7037 } else {
7038 dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
7039 }
7040 };
7041 for k in 0..4 {
7042 let act = &acts[bi + k];
7043 let mut acc = d[k] * act.sx;
7044 for &(j, xv) in &act.outliers {
7045 let (w, sc) = q4t_outlier(bytes, r, gpr, j);
7046 acc += w * sc * xv;
7047 }
7048 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
7050 }
7051 bi += 4;
7052 }
7053 }
7054 let _ = blocked_ok;
7055 while bi < acts.len() {
7056 let act = &acts[bi];
7057 let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
7058 for &(j, xv) in &act.outliers {
7059 let (w, s) = q4t_outlier(bytes, r, gpr, j);
7060 acc += w * s * xv;
7061 }
7062 unsafe { *out_addr.at(bi * rows + r) = acc };
7064 bi += 1;
7065 }
7066 }
7067 };
7068 dispatch_rows(pool, rows, &run);
7069 return;
7070 }
7071 let run = move |start: usize, end: usize| {
7072 for r in start..end {
7073 for bi in 0..b {
7074 let x = &xs_all[bi * cols..(bi + 1) * cols];
7075 unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
7077 }
7078 }
7079 };
7080 dispatch_rows(pool, rows, &run);
7081}
7082
7083fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
7092 (0..gpr)
7093 .map(|gi| {
7094 xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
7095 .iter()
7096 .map(|&v| v as i32)
7097 .sum()
7098 })
7099 .collect()
7100}
7101
7102#[inline]
7106#[allow(unreachable_code)]
7107#[cfg(target_arch = "x86_64")]
7112#[target_feature(enable = "avx2")]
7113unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7114 unsafe {
7116 use core::arch::x86_64::*;
7117 let expand = _mm256_setr_epi8(
7119 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7120 3, 3, 3,
7121 );
7122 let bitsel = _mm256_setr_epi8(
7123 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7124 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7125 );
7126 let ones8 = _mm256_set1_epi8(1);
7127 let ones16 = _mm256_set1_epi16(1);
7128 let mut acc = 0f32;
7129 for gi in 0..gpr {
7130 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7131 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7132 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7133 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7134 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7135 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7136 let sel = _mm256_and_si256(x, mask);
7137 let p16 = _mm256_maddubs_epi16(ones8, sel);
7139 let d32 = _mm256_madd_epi16(p16, ones16);
7140 let hi128 = _mm256_extracti128_si256::<1>(d32);
7141 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7142 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7143 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7144 let msum = _mm_cvtsi128_si32(s32);
7145 let d = 2 * msum - gsum[gi];
7148 acc += d as f32 * s;
7149 }
7150 acc
7151 }
7152}
7153
7154#[cfg(target_arch = "x86_64")]
7157#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7158unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7159 unsafe {
7161 use core::arch::x86_64::*;
7162 let expand = _mm256_setr_epi8(
7163 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7164 3, 3, 3,
7165 );
7166 let bitsel = _mm256_setr_epi8(
7167 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7168 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7169 );
7170 let ones8 = _mm256_set1_epi8(1);
7171 let mut acc = 0f32;
7172 for gi in 0..gpr {
7173 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7174 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7175 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7176 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7177 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7178 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7179 let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7180 let d = 2 * msum - gsum[gi];
7181 acc += d as f32 * s;
7182 }
7183 acc
7184 }
7185}
7186
7187#[cfg(target_arch = "x86_64")]
7189#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7190unsafe fn dot_q1_row_1x4_vnni(
7191 bytes: &[u8],
7192 r: usize,
7193 gpr: usize,
7194 xs: [&[i8]; 4],
7195 gsums: [&[i32]; 4],
7196) -> [f32; 4] {
7197 unsafe {
7199 use core::arch::x86_64::*;
7200 let expand = _mm256_setr_epi8(
7201 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7202 3, 3, 3,
7203 );
7204 let bitsel = _mm256_setr_epi8(
7205 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7206 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7207 );
7208 let ones8 = _mm256_set1_epi8(1);
7209 let mut acc = [0f32; 4];
7210 for gi in 0..gpr {
7211 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7212 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7213 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7214 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7215 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7216 for (k, xq) in xs.iter().enumerate() {
7217 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7218 let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7219 let d = 2 * msum - gsums[k][gi];
7220 acc[k] += d as f32 * s;
7221 }
7222 }
7223 acc
7224 }
7225}
7226
7227#[cfg(target_arch = "x86_64")]
7230#[target_feature(enable = "avx2")]
7231unsafe fn dot_q1_row_1x4_avx2(
7232 bytes: &[u8],
7233 r: usize,
7234 gpr: usize,
7235 xs: [&[i8]; 4],
7236 gsums: [&[i32]; 4],
7237) -> [f32; 4] {
7238 unsafe {
7240 use core::arch::x86_64::*;
7241 let expand = _mm256_setr_epi8(
7242 0, 0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1, 1, 2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3,
7243 3, 3, 3,
7244 );
7245 let bitsel = _mm256_setr_epi8(
7246 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7247 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7248 );
7249 let ones8 = _mm256_set1_epi8(1);
7250 let ones16 = _mm256_set1_epi16(1);
7251 let mut acc = [0f32; 4];
7252 for gi in 0..gpr {
7253 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7254 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7255 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7256 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7257 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7258 for (k, xq) in xs.iter().enumerate() {
7259 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7260 let sel = _mm256_and_si256(x, mask);
7261 let p16 = _mm256_maddubs_epi16(ones8, sel);
7262 let d32 = _mm256_madd_epi16(p16, ones16);
7263 let hi128 = _mm256_extracti128_si256::<1>(d32);
7264 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7265 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7266 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7267 let msum = _mm_cvtsi128_si32(s32);
7268 let d = 2 * msum - gsums[k][gi];
7269 acc[k] += d as f32 * s;
7270 }
7271 }
7272 acc
7273 }
7274}
7275
7276#[allow(unreachable_code)]
7277fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7278 #[cfg(target_arch = "aarch64")]
7279 unsafe {
7280 return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
7281 }
7282 #[cfg(target_arch = "x86_64")]
7283 if avx2_enabled() {
7284 unsafe {
7285 if vnni_tiles_enabled() {
7286 return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
7287 }
7288 return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
7289 }
7290 }
7291 let _ = gsum;
7292 let mut acc = 0f32;
7293 for gi in 0..gpr {
7294 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7295 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7296 let mut d = 0i32;
7297 for (j, &b) in tile[2..].iter().enumerate() {
7298 for k in 0..8 {
7299 let w = ((b >> k) & 1) as i32 * 2 - 1;
7300 d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
7301 }
7302 }
7303 acc += d as f32 * s;
7304 }
7305 acc
7306}
7307
7308#[cfg(target_arch = "aarch64")]
7317#[target_feature(enable = "neon,dotprod")]
7318unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7319 unsafe {
7322 use core::arch::aarch64::*;
7323 use core::arch::asm;
7324 const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7325 let m = vld1q_u8(MASKS.as_ptr());
7326 macro_rules! tile_dot {
7328 ($t:expr, $x:expr) => {{
7329 let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
7330 let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
7331 let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7332 let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7333 let x0 = vld1q_s8($x);
7334 let x1 = vld1q_s8($x.add(16));
7335 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7336 asm!(
7337 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7338 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7339 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7340 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7341 options(pure, nomem, nostack),
7342 );
7343 vaddq_s32(a0, a1)
7344 }};
7345 }
7346 const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7355 const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7356 const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7357 const IW11: [u8; 16] = [
7358 10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7359 ];
7360 const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7361 let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7362 let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7363 let isc = vld1_u8(ISC.as_ptr());
7364 macro_rules! tile_dot_tbl {
7366 ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
7367 let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
7368 let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
7369 let x0 = vld1q_s8($x);
7370 let x1 = vld1q_s8($x.add(16));
7371 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7372 asm!(
7373 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7374 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7375 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7376 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7377 options(pure, nomem, nostack),
7378 );
7379 vaddq_s32(a0, a1)
7380 }};
7381 }
7382 let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7383 let row_base = r * gpr * Q1_TILE;
7384 let abs_end = bytes.len();
7385 let xp = xq.as_ptr();
7386 let gp = gsum.as_ptr();
7387 let mut accv = vdupq_n_f32(0.0);
7388 let mut gi = 0;
7389 while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7392 let t0 = base.add(gi * Q1_TILE);
7393 let ld_a = vld1q_u8(t0);
7394 let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7395 let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
7396 let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
7397 let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
7398 let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
7399 let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7401 let g = vld1q_s32(gp.add(gi));
7402 let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7403 let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7404 let scf: float32x4_t;
7405 asm!(
7406 "fcvtl {o:v}.4s, {i:v}.4h",
7407 o = out(vreg) scf, i = in(vreg) sc16,
7408 options(pure, nomem, nostack),
7409 );
7410 accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
7411 gi += 4;
7412 }
7413 let mut acc = vaddvq_f32(accv);
7414 while gi < gpr {
7415 let t = base.add(gi * Q1_TILE);
7416 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7417 let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
7418 acc += (-(2 * d + *gp.add(gi))) as f32 * s;
7419 gi += 1;
7420 }
7421 acc
7422 }
7423}
7424
7425#[cfg(target_arch = "aarch64")]
7430#[target_feature(enable = "neon,dotprod")]
7431unsafe fn dot_q1_row_1x4_sdot(
7432 bytes: &[u8],
7433 r: usize,
7434 gpr: usize,
7435 xs: [&[i8]; 4],
7436 gs: [&[i32]; 4],
7437) -> [f32; 4] {
7438 unsafe {
7440 use core::arch::aarch64::*;
7441 use core::arch::asm;
7442 const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7443 const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7444 const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7445 const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7446 const IW11: [u8; 16] = [
7447 10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7448 ];
7449 const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7450 let m = vld1q_u8(MASKS.as_ptr());
7451 let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7452 let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7453 let isc = vld1_u8(ISC.as_ptr());
7454 macro_rules! sdot2 {
7455 ($w0:expr, $w1:expr, $x:expr) => {{
7456 let x0 = vld1q_s8($x);
7457 let x1 = vld1q_s8($x.add(16));
7458 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7459 asm!(
7460 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7461 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7462 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7463 w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
7464 options(pure, nomem, nostack),
7465 );
7466 vaddq_s32(a0, a1)
7467 }};
7468 }
7469 let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7470 let row_base = r * gpr * Q1_TILE;
7471 let abs_end = bytes.len();
7472 let mut accv = [vdupq_n_f32(0.0); 4];
7473 let mut gi = 0;
7474 while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7475 let t0 = base.add(gi * Q1_TILE);
7476 let ld_a = vld1q_u8(t0);
7477 let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7478 let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
7480 let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
7481 let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
7482 let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
7483 let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
7484 let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
7485 let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
7486 let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
7487 let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7488 let scf: float32x4_t;
7489 asm!(
7490 "fcvtl {o:v}.4s, {i:v}.4h",
7491 o = out(vreg) scf, i = in(vreg) sc16,
7492 options(pure, nomem, nostack),
7493 );
7494 for k in 0..4 {
7495 let xp = xs[k].as_ptr();
7496 let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
7497 let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
7498 let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
7499 let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
7500 let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7501 let g = vld1q_s32(gs[k].as_ptr().add(gi));
7502 let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7503 accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
7504 }
7505 gi += 4;
7506 }
7507 let mut acc = [
7508 vaddvq_f32(accv[0]),
7509 vaddvq_f32(accv[1]),
7510 vaddvq_f32(accv[2]),
7511 vaddvq_f32(accv[3]),
7512 ];
7513 while gi < gpr {
7514 let t = base.add(gi * Q1_TILE);
7515 let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7516 let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
7517 let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
7518 let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7519 let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7520 for k in 0..4 {
7521 let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
7522 acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
7523 }
7524 gi += 1;
7525 }
7526 acc
7527 }
7528}
7529
7530#[inline]
7532fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
7533 let gi = j / GROUP_SIZE;
7534 let k = j % GROUP_SIZE;
7535 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7536 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7537 let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
7538 ((bit as i32 * 2 - 1) as f32, s)
7539}
7540
7541#[inline]
7543fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
7544 let mut acc = 0f32;
7545 for gi in 0..gpr {
7546 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7547 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7548 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
7549 let mut ga = 0f32;
7550 for (j, &b) in tile[2..].iter().enumerate() {
7551 for k in 0..8 {
7552 ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
7553 }
7554 }
7555 acc += ga * s;
7556 }
7557 acc
7558}
7559
7560#[allow(clippy::too_many_arguments)]
7563fn q1_range_a8w8(
7564 bytes: &[u8],
7565 gpr: usize,
7566 act: &SplitAct,
7567 gsum: &[i32],
7568 out: SendMut,
7569 start: usize,
7570 end: usize,
7571) {
7572 for r in start..end {
7573 let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
7574 for &(j, xv) in &act.outliers {
7575 let (w, s) = q1_outlier(bytes, r, gpr, j);
7576 acc += w * s * xv;
7577 }
7578 unsafe { *out.at(r) = acc };
7580 }
7581}
7582
7583fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
7585 for r in start..end {
7586 unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
7588 }
7589}
7590
7591fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
7596 let entries = base_len + (rows + 1) * 4;
7597 (base_len, entries, entries <= bytes.len())
7598}
7599
7600#[inline]
7602fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
7603 let o = rp_off + r * 4;
7604 u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
7605}
7606
7607const SIGN5: [[f32; 5]; 256] = {
7611 let mut lut = [[0.0f32; 5]; 256];
7612 let pow3 = [1u16, 3, 9, 27, 81];
7613 let mut byte = 0usize;
7614 while byte < 256 {
7615 let mut i = 0usize;
7616 while i < 5 {
7617 let code = (byte as u16 / pow3[i]) % 3;
7618 lut[byte][i] = if code == 1 {
7619 1.0
7620 } else if code == 2 {
7621 -1.0
7622 } else {
7623 0.0
7624 };
7625 i += 1;
7626 }
7627 byte += 1;
7628 }
7629 lut
7630};
7631
7632const SIGN5_I8: [[i8; 5]; 256] = {
7634 let mut lut = [[0i8; 5]; 256];
7635 let pow3 = [1u16, 3, 9, 27, 81];
7636 let mut byte = 0usize;
7637 while byte < 256 {
7638 let mut i = 0usize;
7639 while i < 5 {
7640 let code = (byte as u16 / pow3[i]) % 3;
7641 lut[byte][i] = if code == 1 {
7642 1
7643 } else if code == 2 {
7644 -1
7645 } else {
7646 0
7647 };
7648 i += 1;
7649 }
7650 byte += 1;
7651 }
7652 lut
7653};
7654
7655const SIGN5_U64: [u64; 256] = {
7661 let mut lut = [0u64; 256];
7662 let pow3 = [1u16, 3, 9, 27, 81];
7663 let mut byte = 0usize;
7664 while byte < 256 {
7665 let mut v = 0u64;
7666 let mut i = 0usize;
7667 while i < 5 {
7668 let code = (byte as u16 / pow3[i]) % 3;
7669 let s: u8 = if code == 1 {
7670 1
7671 } else if code == 2 {
7672 0xFF
7673 } else {
7674 0
7675 };
7676 v |= (s as u64) << (i * 8);
7677 i += 1;
7678 }
7679 lut[byte] = v;
7680 byte += 1;
7681 }
7682 lut
7683};
7684
7685#[inline]
7690fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
7691 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7692 let off = (r * gpr + j / GROUP_SIZE) * TILE;
7693 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7694 let within = j % GROUP_SIZE;
7695 SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
7696}
7697
7698#[cfg(target_arch = "aarch64")]
7701#[target_feature(enable = "neon,dotprod")]
7702#[inline]
7703unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
7704 unsafe {
7706 use core::arch::aarch64::*;
7707 use core::arch::asm;
7708 let w0 = vld1q_s8(w);
7709 let w1 = vld1q_s8(w.add(16));
7710 let x0 = vld1q_s8(x);
7711 let x1 = vld1q_s8(x.add(16));
7712 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7713 asm!(
7714 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7715 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7716 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7717 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7718 options(pure, nomem, nostack),
7719 );
7720 vaddvq_s32(vaddq_s32(a0, a1))
7721 }
7722}
7723
7724#[cfg(target_arch = "x86_64")]
7727#[target_feature(enable = "avx2")]
7728#[inline]
7729unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
7730 unsafe {
7732 use core::arch::x86_64::*;
7733 let wv = _mm256_loadu_si256(w as *const __m256i);
7734 let xv = _mm256_loadu_si256(x as *const __m256i);
7735 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7736 let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
7737 let hi128 = _mm256_extracti128_si256::<1>(d);
7738 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
7739 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7740 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7741 _mm_cvtsi128_si32(s32)
7742 }
7743}
7744
7745#[inline]
7750fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
7751 debug_assert!(dst.len() >= 40);
7752 unsafe {
7755 let p = dst.as_mut_ptr();
7756 for bi in 0..7 {
7757 core::ptr::write_unaligned(
7758 p.add(bi * 5) as *mut u64,
7759 SIGN5_U64[*codes.add(bi) as usize],
7760 );
7761 }
7762 }
7763}
7764
7765#[inline]
7770fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
7771 #[cfg(target_arch = "aarch64")]
7772 unsafe {
7773 return sdot32_i8(w, x);
7774 }
7775 #[cfg(target_arch = "x86_64")]
7776 unsafe {
7777 return i8dot32_avx2(w, x);
7778 }
7779 #[allow(unreachable_code)]
7780 unsafe {
7781 let mut s = 0i32;
7782 for k in 0..GROUP_SIZE {
7783 s += *w.add(k) as i32 * *x.add(k) as i32;
7784 }
7785 s
7786 }
7787}
7788
7789#[inline]
7790unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
7791 let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
7792 (
7793 SIGN5_U64[*codes as usize],
7794 SIGN5_U64[*codes.add(1) as usize],
7795 SIGN5_U64[*codes.add(2) as usize],
7796 SIGN5_U64[*codes.add(3) as usize],
7797 SIGN5_U64[*codes.add(4) as usize],
7798 SIGN5_U64[*codes.add(5) as usize],
7799 SIGN5_U64[*codes.add(6) as usize],
7800 )
7801 };
7802
7803 let u0 = s0 | (s1 << 40);
7804 let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
7805 let u2 = (s3 >> 8) | (s4 << 32);
7806 let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
7807
7808 (u0, u1, u2, u3)
7809}
7810
7811#[cfg(target_arch = "aarch64")]
7815#[target_feature(enable = "neon,dotprod")]
7816unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7817 use core::arch::aarch64::*;
7818 use core::arch::asm;
7819 unsafe {
7820 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7821 let mut acc = 0f32;
7822 let bytes_ptr = bytes.as_ptr();
7823 let xq_ptr = xq.as_ptr();
7824 let row_off = r * gpr * TILE;
7825
7826 let gpr2 = gpr & !1;
7827 let mut gi = 0;
7828 while gi < gpr2 {
7829 let off0 = row_off + gi * TILE;
7830 let off1 = off0 + TILE;
7831 let s0 = f16_to_f32(u16::from_le_bytes([
7832 *bytes_ptr.add(off0),
7833 *bytes_ptr.add(off0 + 1),
7834 ]));
7835 let s1 = f16_to_f32(u16::from_le_bytes([
7836 *bytes_ptr.add(off1),
7837 *bytes_ptr.add(off1 + 1),
7838 ]));
7839
7840 let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
7841 let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
7842
7843 let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
7844 let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
7845 let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
7846 let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
7847
7848 let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7849 let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7850 let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
7851 let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
7852
7853 let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
7854 let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7855 asm!(
7856 "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
7857 "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
7858 "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
7859 "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
7860 a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
7861 a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
7862 w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
7863 w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
7864 options(pure, nomem, nostack),
7865 );
7866 let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
7867 let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
7868 acc += d0 as f32 * s0 + d1 as f32 * s1;
7869 gi += 2;
7870 }
7871
7872 if gi < gpr {
7873 let off = row_off + gi * TILE;
7874 let s = f16_to_f32(u16::from_le_bytes([
7875 *bytes_ptr.add(off),
7876 *bytes_ptr.add(off + 1),
7877 ]));
7878 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7879 let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
7880 let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
7881 let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7882 let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7883 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7884 asm!(
7885 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7886 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7887 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7888 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7889 options(pure, nomem, nostack),
7890 );
7891 let d = vaddvq_s32(vaddq_s32(a0, a1));
7892 acc += d as f32 * s;
7893 }
7894 acc
7895 }
7896}
7897
7898#[cfg(target_arch = "x86_64")]
7900#[target_feature(enable = "avx2")]
7901unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7902 use core::arch::x86_64::*;
7903 unsafe {
7904 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7905 let mut acc = 0f32;
7906 let bytes_ptr = bytes.as_ptr();
7907 let xq_ptr = xq.as_ptr();
7908 let row_off = r * gpr * TILE;
7909
7910 let ones = _mm256_set1_epi16(1);
7911 for gi in 0..gpr {
7912 let off = row_off + gi * TILE;
7913 let s = f16_to_f32(u16::from_le_bytes([
7914 *bytes_ptr.add(off),
7915 *bytes_ptr.add(off + 1),
7916 ]));
7917 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7918 let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7919 let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7920 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7921 let d256 = _mm256_madd_epi16(p16, ones);
7922 let d128 = _mm_add_epi32(
7923 _mm256_castsi256_si128(d256),
7924 _mm256_extracti128_si256(d256, 1),
7925 );
7926 let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
7927 let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
7928 acc += d32 as f32 * s;
7929 }
7930 acc
7931 }
7932}
7933
7934#[cfg(target_arch = "x86_64")]
7936#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7937unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7938 use core::arch::x86_64::*;
7939 unsafe {
7941 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7942 let mut acc = 0f32;
7943 let bytes_ptr = bytes.as_ptr();
7944 let xq_ptr = xq.as_ptr();
7945 let row_off = r * gpr * TILE;
7946 for gi in 0..gpr {
7947 let off = row_off + gi * TILE;
7948 let s = f16_to_f32(u16::from_le_bytes([
7949 *bytes_ptr.add(off),
7950 *bytes_ptr.add(off + 1),
7951 ]));
7952 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7953 let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7954 let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7955 let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7956 acc += d as f32 * s;
7957 }
7958 acc
7959 }
7960}
7961
7962#[inline]
7966fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7967 #[cfg(target_arch = "aarch64")]
7968 unsafe {
7969 return q1t_dot_row_sdot(bytes, r, gpr, xq);
7970 }
7971 #[cfg(target_arch = "x86_64")]
7972 unsafe {
7973 if vnni_tiles_enabled() {
7974 return q1t_dot_row_vnni(bytes, r, gpr, xq);
7975 }
7976 return q1t_dot_row_avx2(bytes, r, gpr, xq);
7977 }
7978 #[allow(unreachable_code)]
7979 {
7980 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7981 let mut acc = 0f32;
7982 let mut sg = [0i8; GROUP_SIZE + 8]; for gi in 0..gpr {
7984 let off = (r * gpr + gi) * TILE;
7985 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7986 q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
7987 let mut d = 0i32;
7988 for k in 0..GROUP_SIZE {
7989 d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
7990 }
7991 acc += d as f32 * s;
7992 }
7993 acc
7994 }
7995}
7996
7997fn q1t_row_outlier_correction(
8004 bytes: &[u8],
8005 r: usize,
8006 rp_off: usize,
8007 entries_off: usize,
8008 has_ov: bool,
8009 x: &[f32],
8010) -> f32 {
8011 if !has_ov {
8012 return 0.0;
8013 }
8014 let (c0, c1) = (
8015 q1t_rowptr(bytes, rp_off, r),
8016 q1t_rowptr(bytes, rp_off, r + 1),
8017 );
8018 let mut corr = 0f32;
8019 for p in c0..c1 {
8020 let e = entries_off + p * 4;
8021 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8022 let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8023 corr += val * x[col];
8024 }
8025 corr
8026}
8027
8028fn q1t_dequant_row(
8032 bytes: &[u8],
8033 r: usize,
8034 gpr: usize,
8035 rp_off: usize,
8036 entries_off: usize,
8037 has_ov: bool,
8038 buf: &mut [f32],
8039) {
8040 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8041 for g in 0..gpr {
8042 let off = (r * gpr + g) * TILE;
8043 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8044 let codes = &bytes[off + 2..off + TILE];
8045 let bc = g * GROUP_SIZE;
8046 for bi in 0..6 {
8048 let lut = &SIGN5[codes[bi] as usize];
8049 let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
8050 for i in 0..5 {
8051 d[i] = lut[i] * s;
8052 }
8053 }
8054 let lut = &SIGN5[codes[6] as usize];
8055 buf[bc + 30] = lut[0] * s;
8056 buf[bc + 31] = lut[1] * s;
8057 }
8058 if !has_ov {
8059 return;
8060 }
8061 let (c0, c1) = (
8062 q1t_rowptr(bytes, rp_off, r),
8063 q1t_rowptr(bytes, rp_off, r + 1),
8064 );
8065 for p in c0..c1 {
8066 let e = entries_off + p * 4;
8067 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8068 buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8069 }
8070}
8071
8072fn q1t_add_overlay(
8076 bytes: &[u8],
8077 x: &[f32],
8078 rows: usize,
8079 cols: usize,
8080 out: &mut [f32],
8081 pool: Option<&Pool>,
8082) {
8083 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8084 let gpr = cols / GROUP_SIZE;
8085 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8086 if !has_ov {
8087 return;
8088 }
8089 let out_addr = SendMut(out.as_mut_ptr());
8090 let run = move |start: usize, end: usize| {
8091 for r in start..end {
8092 let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8093 unsafe { *out_addr.at(r) += corr };
8095 }
8096 };
8097 dispatch_rows(pool, rows, &run);
8098}
8099
8100#[allow(clippy::too_many_arguments)]
8103fn q1t_range_a8w8(
8104 bytes: &[u8],
8105 gpr: usize,
8106 rp_off: usize,
8107 ent_off: usize,
8108 has_ov: bool,
8109 act: &SplitAct,
8110 x: &[f32],
8111 out: SendMut,
8112 start: usize,
8113 end: usize,
8114) {
8115 for r in start..end {
8116 let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8117 for &(j, xv) in &act.outliers {
8118 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8119 }
8120 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8121 unsafe { *out.at(r) = acc };
8123 }
8124}
8125
8126#[allow(clippy::too_many_arguments)]
8129fn q1t_range_f32_batch(
8130 bytes: &[u8],
8131 gpr: usize,
8132 rp_off: usize,
8133 ent_off: usize,
8134 has_ov: bool,
8135 x: &[f32],
8136 out: SendMut,
8137 start: usize,
8138 end: usize,
8139) {
8140 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8141 let mut sg = [0f32; GROUP_SIZE];
8142 for r in start..end {
8143 let mut acc = 0f32;
8144 for g in 0..gpr {
8145 let off = (r * gpr + g) * TILE;
8146 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8147 let codes = &bytes[off + 2..off + TILE];
8148 let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8149 for bi in 0..6 {
8150 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8151 }
8152 let lut = &SIGN5[codes[6] as usize];
8153 sg[30] = lut[0];
8154 sg[31] = lut[1];
8155 let mut gsum = 0f32;
8156 for k in 0..GROUP_SIZE {
8157 gsum += sg[k] * xg[k];
8158 }
8159 acc += s * gsum;
8160 }
8161 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8162 unsafe { *out.at(r) = acc };
8164 }
8165}
8166
8167fn q1t_matvec(
8171 bytes: &[u8],
8172 x: &[f32],
8173 rows: usize,
8174 cols: usize,
8175 out: &mut [f32],
8176 pool: Option<&Pool>,
8177) {
8178 debug_assert_eq!(out.len(), rows);
8179 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8180 let gpr = cols / GROUP_SIZE;
8181 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8182 let out_addr = SendMut(out.as_mut_ptr());
8183 if a8w8_enabled() {
8187 let act = split_act(x);
8188 let act = &act;
8189 let run = move |start: usize, end: usize| {
8190 for r in start..end {
8191 let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8192 for &(j, xv) in &act.outliers {
8193 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8194 }
8195 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8196 unsafe { *out_addr.at(r) = acc };
8198 }
8199 };
8200 dispatch_rows(pool, rows, &run);
8201 return;
8202 }
8203 let run = move |start: usize, end: usize| {
8204 let mut sg = [0f32; GROUP_SIZE];
8208 for r in start..end {
8209 let mut acc = 0f32;
8210 for g in 0..gpr {
8211 let off = (r * gpr + g) * TILE;
8212 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8213 let codes = &bytes[off + 2..off + TILE];
8214 let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8215 for bi in 0..6 {
8216 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8217 }
8218 let lut = &SIGN5[codes[6] as usize];
8219 sg[30] = lut[0];
8220 sg[31] = lut[1];
8221 let mut gsum = 0f32;
8222 for k in 0..GROUP_SIZE {
8223 gsum += sg[k] * xg[k];
8224 }
8225 acc += s * gsum;
8226 }
8227 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8228 unsafe { *out_addr.at(r) = acc };
8229 }
8230 };
8231 dispatch_rows(pool, rows, &run);
8232}
8233
8234#[cfg(target_arch = "aarch64")]
8240#[target_feature(enable = "neon,dotprod")]
8241unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
8242 use core::arch::aarch64::*;
8243 use core::arch::asm;
8244 unsafe {
8246 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8247 let bytes_ptr = bytes.as_ptr();
8248 let row_off = r * gpr * TILE;
8249 let xp = [xa.as_ptr(), xb.as_ptr()];
8250 let mut acc = [0f32; 2];
8251 macro_rules! sdot2 {
8252 ($w0:expr, $w1:expr, $x:expr) => {{
8253 let x0 = vld1q_s8($x);
8254 let x1 = vld1q_s8($x.add(16));
8255 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
8256 asm!(
8257 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
8258 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
8259 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
8260 w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
8261 options(pure, nomem, nostack),
8262 );
8263 vaddvq_s32(vaddq_s32(a0, a1))
8264 }};
8265 }
8266 let gpr2 = gpr & !1;
8267 let mut gi = 0;
8268 while gi < gpr2 {
8269 let off0 = row_off + gi * TILE;
8270 let off1 = off0 + TILE;
8271 let s0 = f16_to_f32(u16::from_le_bytes([
8272 *bytes_ptr.add(off0),
8273 *bytes_ptr.add(off0 + 1),
8274 ]));
8275 let s1 = f16_to_f32(u16::from_le_bytes([
8276 *bytes_ptr.add(off1),
8277 *bytes_ptr.add(off1 + 1),
8278 ]));
8279 let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
8280 let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
8281 let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
8282 let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
8283 let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
8284 let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
8285 for k in 0..2 {
8286 let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
8287 let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
8288 acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
8289 }
8290 gi += 2;
8291 }
8292 if gi < gpr {
8293 let off = row_off + gi * TILE;
8294 let s = f16_to_f32(u16::from_le_bytes([
8295 *bytes_ptr.add(off),
8296 *bytes_ptr.add(off + 1),
8297 ]));
8298 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
8299 let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
8300 let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
8301 for k in 0..2 {
8302 let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
8303 acc[k] += d as f32 * s;
8304 }
8305 }
8306 acc
8307 }
8308}
8309
8310fn q1t_matvec2(
8316 bytes: &[u8],
8317 x1: &[f32],
8318 x2: &[f32],
8319 rows: usize,
8320 cols: usize,
8321 o1: &mut [f32],
8322 o2: &mut [f32],
8323 pool: Option<&Pool>,
8324) {
8325 debug_assert_eq!(o1.len(), rows);
8326 debug_assert_eq!(o2.len(), rows);
8327 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8328 let gpr = cols / GROUP_SIZE;
8329 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8330 let out1 = SendMut(o1.as_mut_ptr());
8331 let out2 = SendMut(o2.as_mut_ptr());
8332 if a8w8_enabled() {
8333 let a1 = split_act(x1);
8334 let a2 = split_act(x2);
8335 let (a1, a2) = (&a1, &a2);
8336 let run = move |start: usize, end: usize| {
8337 for r in start..end {
8338 #[cfg(target_arch = "aarch64")]
8339 let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
8342 #[cfg(not(target_arch = "aarch64"))]
8343 let ds = [
8344 q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
8345 q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
8346 ];
8347 let mut acc1 = ds[0] * a1.sx;
8348 for &(j, xv) in &a1.outliers {
8349 acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
8350 }
8351 acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8352 let mut acc2 = ds[1] * a2.sx;
8353 for &(j, xv) in &a2.outliers {
8354 acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
8355 }
8356 acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8357 unsafe {
8359 *out1.at(r) = acc1;
8360 *out2.at(r) = acc2;
8361 }
8362 }
8363 };
8364 dispatch_rows(pool, rows, &run);
8365 return;
8366 }
8367 let run = move |start: usize, end: usize| {
8368 let mut sg = [0f32; GROUP_SIZE];
8371 for r in start..end {
8372 let mut acc1 = 0f32;
8373 let mut acc2 = 0f32;
8374 for g in 0..gpr {
8375 let off = (r * gpr + g) * TILE;
8376 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8377 let codes = &bytes[off + 2..off + TILE];
8378 for bi in 0..6 {
8379 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8380 }
8381 let lut = &SIGN5[codes[6] as usize];
8382 sg[30] = lut[0];
8383 sg[31] = lut[1];
8384 let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8385 let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8386 let mut gsum1 = 0f32;
8387 for k in 0..GROUP_SIZE {
8388 gsum1 += sg[k] * xg1[k];
8389 }
8390 acc1 += s * gsum1;
8391 let mut gsum2 = 0f32;
8392 for k in 0..GROUP_SIZE {
8393 gsum2 += sg[k] * xg2[k];
8394 }
8395 acc2 += s * gsum2;
8396 }
8397 acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8398 acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8399 unsafe {
8401 *out1.at(r) = acc1;
8402 *out2.at(r) = acc2;
8403 }
8404 }
8405 };
8406 dispatch_rows(pool, rows, &run);
8407}
8408
8409fn q1t_matmat(
8412 bytes: &[u8],
8413 xs: &[f32],
8414 b: usize,
8415 rows: usize,
8416 cols: usize,
8417 out: &mut [f32],
8418 pool: Option<&Pool>,
8419) {
8420 debug_assert_eq!(out.len(), b * rows);
8421 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8422 let gpr = cols / GROUP_SIZE;
8423 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8424 let out_addr = SendMut(out.as_mut_ptr());
8425 if a8w8_enabled() {
8429 let acts: Vec<SplitAct> = (0..b)
8430 .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
8431 .collect();
8432 let acts = &acts;
8433 let run = move |start: usize, end: usize| {
8434 let mut sg = vec![0i8; cols + 8]; let mut sc = vec![0f32; gpr]; let mut accs = vec![0f32; b]; for r in start..end {
8438 for g in 0..gpr {
8439 let off = (r * gpr + g) * TILE;
8440 sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8441 q1t_unpack_group_i8(
8442 bytes.as_ptr().wrapping_add(off + 2),
8443 &mut sg[g * GROUP_SIZE..],
8444 );
8445 }
8446 for bi in 0..b {
8447 let act = &acts[bi];
8448 let mut isum = 0f32;
8449 for g in 0..gpr {
8450 let d = q1t_i8dot32(
8451 sg.as_ptr().wrapping_add(g * GROUP_SIZE),
8452 act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
8453 );
8454 isum += d as f32 * sc[g];
8455 }
8456 let mut acc = isum * act.sx;
8457 for &(j, xv) in &act.outliers {
8458 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8459 }
8460 accs[bi] = acc;
8461 }
8462 if has_ov {
8466 let (c0, c1) = (
8467 q1t_rowptr(bytes, rp_off, r),
8468 q1t_rowptr(bytes, rp_off, r + 1),
8469 );
8470 for p in c0..c1 {
8471 let e = ent_off + p * 4;
8472 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8473 let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8474 for bi in 0..b {
8475 accs[bi] += val * xs[bi * cols + col];
8476 }
8477 }
8478 }
8479 for bi in 0..b {
8480 unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
8481 }
8482 }
8483 };
8484 dispatch_rows(pool, rows, &run);
8485 return;
8486 }
8487 let run = move |start: usize, end: usize| {
8488 let mut buf = vec![0f32; cols];
8489 for r in start..end {
8490 q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
8491 for bi in 0..b {
8492 let xr = &xs[bi * cols..(bi + 1) * cols];
8493 let mut acc = 0f32;
8494 for j in 0..cols {
8495 acc += buf[j] * xr[j];
8496 }
8497 unsafe { *out_addr.at(bi * rows + r) = acc };
8498 }
8499 }
8500 };
8501 dispatch_rows(pool, rows, &run);
8502}
8503
8504fn q1_matvec(
8505 bytes: &[u8],
8506 x: &[f32],
8507 rows: usize,
8508 cols: usize,
8509 out: &mut [f32],
8510 pool: Option<&Pool>,
8511) {
8512 debug_assert_eq!(out.len(), rows);
8513 let gpr = cols / GROUP_SIZE;
8514 let out_addr = SendMut(out.as_mut_ptr());
8515 if a8w8_enabled() {
8516 let act = split_act(x);
8517 let gsum = q1_group_sums(&act.xq, gpr);
8518 let (act, gsum) = (&act, &gsum);
8519 let run = move |start: usize, end: usize| {
8520 q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
8521 };
8522 dispatch_rows(pool, rows, &run);
8523 return;
8524 }
8525 let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
8526 dispatch_rows(pool, rows, &run);
8527}
8528
8529#[allow(clippy::too_many_arguments)]
8531fn q1_matvec2(
8532 bytes: &[u8],
8533 x1: &[f32],
8534 x2: &[f32],
8535 rows: usize,
8536 cols: usize,
8537 o1: &mut [f32],
8538 o2: &mut [f32],
8539 pool: Option<&Pool>,
8540) {
8541 let gpr = cols / GROUP_SIZE;
8542 let p1 = SendMut(o1.as_mut_ptr());
8543 let p2 = SendMut(o2.as_mut_ptr());
8544 if a8w8_enabled() {
8545 let a1 = split_act(x1);
8546 let a2 = split_act(x2);
8547 let g1 = q1_group_sums(&a1.xq, gpr);
8548 let g2 = q1_group_sums(&a2.xq, gpr);
8549 let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
8550 let run = move |start: usize, end: usize| {
8551 for r in start..end {
8552 let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
8553 let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
8554 for &(j, xv) in &a1.outliers {
8555 let (w, s) = q1_outlier(bytes, r, gpr, j);
8556 v1 += w * s * xv;
8557 }
8558 for &(j, xv) in &a2.outliers {
8559 let (w, s) = q1_outlier(bytes, r, gpr, j);
8560 v2 += w * s * xv;
8561 }
8562 unsafe {
8564 *p1.at(r) = v1;
8565 *p2.at(r) = v2;
8566 }
8567 }
8568 };
8569 dispatch_rows(pool, rows, &run);
8570 return;
8571 }
8572 let run = move |start: usize, end: usize| {
8573 for r in start..end {
8574 unsafe {
8576 *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
8577 *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
8578 }
8579 }
8580 };
8581 dispatch_rows(pool, rows, &run);
8582}
8583
8584#[allow(clippy::too_many_arguments)]
8586fn q1_matmat(
8587 bytes: &[u8],
8588 xs_all: &[f32],
8589 b: usize,
8590 rows: usize,
8591 cols: usize,
8592 out: &mut [f32],
8593 pool: Option<&Pool>,
8594) {
8595 debug_assert_eq!(out.len(), b * rows);
8596 let gpr = cols / GROUP_SIZE;
8597 let out_addr = SendMut(out.as_mut_ptr());
8598 if a8w8_enabled() {
8599 let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
8600 .map(|bi| {
8601 let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
8602 let gsum = q1_group_sums(&act.xq, gpr);
8603 (act, gsum)
8604 })
8605 .collect();
8606 let acts = &acts;
8607 #[cfg(target_arch = "x86_64")]
8608 let blocked_ok = avx2_enabled() && blocked_enabled();
8609 #[cfg(target_arch = "aarch64")]
8610 let blocked_ok = sdot_enabled() && blocked_enabled();
8611 let run = move |start: usize, end: usize| {
8612 for r in start..end {
8613 let mut bi = 0usize;
8614 #[cfg(target_arch = "aarch64")]
8617 if blocked_ok {
8618 while bi + 4 <= acts.len() {
8619 let xs = [
8620 acts[bi].0.xq.as_slice(),
8621 acts[bi + 1].0.xq.as_slice(),
8622 acts[bi + 2].0.xq.as_slice(),
8623 acts[bi + 3].0.xq.as_slice(),
8624 ];
8625 let gs = [
8626 acts[bi].1.as_slice(),
8627 acts[bi + 1].1.as_slice(),
8628 acts[bi + 2].1.as_slice(),
8629 acts[bi + 3].1.as_slice(),
8630 ];
8631 let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
8632 for k in 0..4 {
8633 let (act, _) = &acts[bi + k];
8634 let mut acc = d[k] * act.sx;
8635 for &(j, xv) in &act.outliers {
8636 let (w, sc) = q1_outlier(bytes, r, gpr, j);
8637 acc += w * sc * xv;
8638 }
8639 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8641 }
8642 bi += 4;
8643 }
8644 }
8645 #[cfg(target_arch = "x86_64")]
8646 if blocked_ok {
8647 while bi + 4 <= acts.len() {
8648 let xs = [
8649 acts[bi].0.xq.as_slice(),
8650 acts[bi + 1].0.xq.as_slice(),
8651 acts[bi + 2].0.xq.as_slice(),
8652 acts[bi + 3].0.xq.as_slice(),
8653 ];
8654 let gs = [
8655 acts[bi].1.as_slice(),
8656 acts[bi + 1].1.as_slice(),
8657 acts[bi + 2].1.as_slice(),
8658 acts[bi + 3].1.as_slice(),
8659 ];
8660 let d = unsafe {
8661 if vnni_tiles_enabled() {
8662 dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
8663 } else {
8664 dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
8665 }
8666 };
8667 for k in 0..4 {
8668 let (act, _) = &acts[bi + k];
8669 let mut acc = d[k] * act.sx;
8670 for &(j, xv) in &act.outliers {
8671 let (w, sc) = q1_outlier(bytes, r, gpr, j);
8672 acc += w * sc * xv;
8673 }
8674 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8676 }
8677 bi += 4;
8678 }
8679 }
8680 while bi < acts.len() {
8681 let (act, gsum) = &acts[bi];
8682 let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
8683 for &(j, xv) in &act.outliers {
8684 let (w, s) = q1_outlier(bytes, r, gpr, j);
8685 acc += w * s * xv;
8686 }
8687 unsafe { *out_addr.at(bi * rows + r) = acc };
8689 bi += 1;
8690 }
8691 }
8692 };
8693 dispatch_rows(pool, rows, &run);
8694 return;
8695 }
8696 let run = move |start: usize, end: usize| {
8697 for r in start..end {
8698 for bi in 0..b {
8699 let x = &xs_all[bi * cols..(bi + 1) * cols];
8700 unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
8702 }
8703 }
8704 };
8705 dispatch_rows(pool, rows, &run);
8706}
8707
8708fn q4matvec(
8714 bytes: &[u8],
8715 x: &[f32],
8716 rows: usize,
8717 cols: usize,
8718 out: &mut [f32],
8719 pool: Option<&Pool>,
8720) {
8721 debug_assert_eq!(out.len(), rows);
8722 let (packed, scales) = q4_split(bytes, rows, cols);
8723 let gpr = cols / GROUP_SIZE;
8724 let out_addr = SendMut(out.as_mut_ptr());
8725
8726 if a8w8_enabled() {
8727 let act = split_act(x);
8728 let run = move |start: usize, end: usize| {
8729 q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
8730 };
8731 dispatch_rows(pool, rows, &run);
8732 return;
8733 }
8734
8735 let run =
8736 move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
8737 dispatch_rows(pool, rows, &run);
8738}
8739
8740#[inline]
8743#[allow(unreachable_code)]
8744#[cfg(target_arch = "x86_64")]
8749#[target_feature(enable = "avx2")]
8750unsafe fn dot_q4b_row_1x4_avx2(
8751 buf: &[u8],
8752 scales: &[u8],
8753 g0: usize,
8754 gpr: usize,
8755 xs: [&[i8]; 4],
8756) -> [f32; 4] {
8757 unsafe {
8759 use core::arch::x86_64::*;
8760 let ones = _mm256_set1_epi16(1);
8761 let mut acc = [0f32; 4];
8762 for gi in 0..gpr {
8763 let s = f16_to_f32(u16::from_le_bytes([
8764 scales[(g0 + gi) * 2],
8765 scales[(g0 + gi) * 2 + 1],
8766 ]));
8767 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8768 let aw = _mm256_abs_epi8(w);
8769 for (k, xq) in xs.iter().enumerate() {
8770 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8771 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8772 let d = _mm256_madd_epi16(p16, ones);
8773 let hi128 = _mm256_extracti128_si256::<1>(d);
8774 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8775 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8776 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8777 acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
8778 }
8779 }
8780 acc
8781 }
8782}
8783
8784#[cfg(target_arch = "x86_64")]
8786#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8787unsafe fn dot_q4b_row_1x4_vnni(
8788 buf: &[u8],
8789 scales: &[u8],
8790 g0: usize,
8791 gpr: usize,
8792 xs: [&[i8]; 4],
8793) -> [f32; 4] {
8794 unsafe {
8796 use core::arch::x86_64::*;
8797 let mut acc = [0f32; 4];
8798 for gi in 0..gpr {
8799 let s = f16_to_f32(u16::from_le_bytes([
8800 scales[(g0 + gi) * 2],
8801 scales[(g0 + gi) * 2 + 1],
8802 ]));
8803 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8804 let aw = _mm256_abs_epi8(w);
8805 for (k, xq) in xs.iter().enumerate() {
8806 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8807 let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8808 acc[k] += d as f32 * s;
8809 }
8810 }
8811 acc
8812 }
8813}
8814
8815#[cfg(target_arch = "x86_64")]
8821#[target_feature(enable = "avx2")]
8822unsafe fn dot_q4b_row_1x4_sx_avx2(
8823 buf: &[u8],
8824 scales: &[u8],
8825 g0: usize,
8826 gpr: usize,
8827 xs: [&[i8]; 4],
8828 sxs: [f32; 4],
8829) -> [f32; 4] {
8830 unsafe {
8832 use core::arch::x86_64::*;
8833 let ones = _mm256_set1_epi16(1);
8834 let mut acc = [0f32; 4];
8835 for gi in 0..gpr {
8836 let s = f16_to_f32(u16::from_le_bytes([
8837 scales[(g0 + gi) * 2],
8838 scales[(g0 + gi) * 2 + 1],
8839 ]));
8840 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8841 let aw = _mm256_abs_epi8(w);
8842 for (k, xq) in xs.iter().enumerate() {
8843 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8844 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8845 let d = _mm256_madd_epi16(p16, ones);
8846 let hi128 = _mm256_extracti128_si256::<1>(d);
8847 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8848 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8849 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8850 acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
8851 }
8852 }
8853 acc
8854 }
8855}
8856
8857#[cfg(target_arch = "x86_64")]
8860#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8861unsafe fn dot_q4b_row_1x4_sx_vnni(
8862 buf: &[u8],
8863 scales: &[u8],
8864 g0: usize,
8865 gpr: usize,
8866 xs: [&[i8]; 4],
8867 sxs: [f32; 4],
8868) -> [f32; 4] {
8869 unsafe {
8871 use core::arch::x86_64::*;
8872 let mut acc = [0f32; 4];
8873 for gi in 0..gpr {
8874 let s = f16_to_f32(u16::from_le_bytes([
8875 scales[(g0 + gi) * 2],
8876 scales[(g0 + gi) * 2 + 1],
8877 ]));
8878 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8879 let aw = _mm256_abs_epi8(w);
8880 for (k, xq) in xs.iter().enumerate() {
8881 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8882 let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8883 acc[k] += (d as f32 * sxs[k]) * s;
8884 }
8885 }
8886 acc
8887 }
8888}
8889
8890#[allow(unreachable_code)]
8891fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
8892 #[cfg(target_arch = "aarch64")]
8893 unsafe {
8894 return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
8895 }
8896 #[cfg(target_arch = "x86_64")]
8897 unsafe {
8898 return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
8899 }
8900 let mut acc = 0f32;
8901 for gi in 0..gpr {
8902 let g = g0 + gi;
8903 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8904 let mut d = 0i32;
8905 for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
8906 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
8907 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
8908 }
8909 acc += d as f32 * s;
8910 }
8911 acc
8912}
8913
8914#[inline]
8916#[allow(unreachable_code)]
8917fn dot_q4_row_i8_2(
8918 packed: &[u8],
8919 scales: &[u8],
8920 g0: usize,
8921 gpr: usize,
8922 xq1: &[i8],
8923 xq2: &[i8],
8924) -> (f32, f32) {
8925 #[cfg(target_arch = "aarch64")]
8926 unsafe {
8927 return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
8928 }
8929 #[cfg(target_arch = "x86_64")]
8930 unsafe {
8931 return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
8932 }
8933 (
8934 dot_q4_row_i8(packed, scales, g0, gpr, xq1),
8935 dot_q4_row_i8(packed, scales, g0, gpr, xq2),
8936 )
8937}
8938
8939#[allow(clippy::too_many_arguments)]
8942fn q4_range_a8w8(
8943 packed: &[u8],
8944 scales: &[u8],
8945 gpr: usize,
8946 cols: usize,
8947 act: &SplitAct,
8948 out: SendMut,
8949 start: usize,
8950 end: usize,
8951) {
8952 for r in start..end {
8953 let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
8954 for &(j, xv) in &act.outliers {
8956 let flat = r * cols + j;
8957 let byte = packed[flat / 2];
8958 let nib = if flat & 1 == 0 {
8959 byte & 0x0F
8960 } else {
8961 byte >> 4
8962 };
8963 let s = f16_to_f32(u16::from_le_bytes([
8964 scales[(flat / GROUP_SIZE) * 2],
8965 scales[(flat / GROUP_SIZE) * 2 + 1],
8966 ]));
8967 acc += ((nib as i32 - 8) as f32) * s * xv;
8968 }
8969 unsafe { *out.at(r) = acc };
8971 }
8972}
8973
8974#[allow(clippy::too_many_arguments)]
8977fn q4_range2_a8w8(
8978 packed: &[u8],
8979 scales: &[u8],
8980 gpr: usize,
8981 cols: usize,
8982 a1: &SplitAct,
8983 a2: &SplitAct,
8984 p1: SendMut,
8985 p2: SendMut,
8986 start: usize,
8987 end: usize,
8988) {
8989 for r in start..end {
8990 let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
8991 let mut acc1 = s1 * a1.sx;
8992 let mut acc2 = s2 * a2.sx;
8993 let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
8995 for &(j, xv) in outliers {
8996 let flat = r * cols + j;
8997 let byte = packed[flat / 2];
8998 let nib = if flat & 1 == 0 {
8999 byte & 0x0F
9000 } else {
9001 byte >> 4
9002 };
9003 let s = f16_to_f32(u16::from_le_bytes([
9004 scales[(flat / GROUP_SIZE) * 2],
9005 scales[(flat / GROUP_SIZE) * 2 + 1],
9006 ]));
9007 *acc += ((nib as i32 - 8) as f32) * s * xv;
9008 }
9009 };
9010 fix(&a1.outliers, &mut acc1);
9011 fix(&a2.outliers, &mut acc2);
9012 unsafe {
9014 *p1.at(r) = acc1;
9015 *p2.at(r) = acc2;
9016 }
9017 }
9018}
9019
9020fn q4_range_f32(
9022 packed: &[u8],
9023 scales: &[u8],
9024 gpr: usize,
9025 x: &[f32],
9026 out: SendMut,
9027 start: usize,
9028 end: usize,
9029) {
9030 for r in start..end {
9031 let mut acc = 0f32;
9032 for gi in 0..gpr {
9033 let g = r * gpr + gi;
9034 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9035 let pk = &packed[g * 16..(g + 1) * 16];
9036 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9037 let mut ga = 0f32;
9038 for (k, &b) in pk.iter().enumerate() {
9039 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
9040 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
9041 }
9042 acc += ga * s;
9043 }
9044 unsafe { *out.at(r) = acc };
9046 }
9047}
9048
9049#[allow(clippy::too_many_arguments)]
9053fn q4matvec2(
9054 bytes: &[u8],
9055 x1: &[f32],
9056 x2: &[f32],
9057 rows: usize,
9058 cols: usize,
9059 o1: &mut [f32],
9060 o2: &mut [f32],
9061 pool: Option<&Pool>,
9062) {
9063 debug_assert_eq!(o1.len(), rows);
9064 debug_assert_eq!(o2.len(), rows);
9065 let (packed, scales) = q4_split(bytes, rows, cols);
9066 let gpr = cols / GROUP_SIZE;
9067
9068 if a8w8_enabled() {
9069 let a1 = split_act(x1);
9070 let a2 = split_act(x2);
9071 let p1 = SendMut(o1.as_mut_ptr());
9072 let p2 = SendMut(o2.as_mut_ptr());
9073 let run = move |start: usize, end: usize| {
9074 q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
9075 };
9076 dispatch_rows(pool, rows, &run);
9077 return;
9078 }
9079
9080 let p1 = SendMut(o1.as_mut_ptr());
9081 let p2 = SendMut(o2.as_mut_ptr());
9082 let run = move |start: usize, end: usize| {
9083 q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
9084 };
9085 dispatch_rows(pool, rows, &run);
9086}
9087
9088#[allow(clippy::too_many_arguments)]
9090fn q4_range2_f32(
9091 packed: &[u8],
9092 scales: &[u8],
9093 gpr: usize,
9094 x1: &[f32],
9095 x2: &[f32],
9096 p1: SendMut,
9097 p2: SendMut,
9098 start: usize,
9099 end: usize,
9100) {
9101 for r in start..end {
9102 let (mut acc1, mut acc2) = (0f32, 0f32);
9103 for gi in 0..gpr {
9104 let g = r * gpr + gi;
9105 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9106 let pk = &packed[g * 16..(g + 1) * 16];
9107 let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9108 let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9109 let (mut g1, mut g2) = (0f32, 0f32);
9110 for (k, &b) in pk.iter().enumerate() {
9111 let wl = (b & 0x0F) as f32 - 8.0;
9112 let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
9113 g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
9114 g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
9115 }
9116 acc1 += g1 * s;
9117 acc2 += g2 * s;
9118 }
9119 unsafe {
9121 *p1.at(r) = acc1;
9122 *p2.at(r) = acc2;
9123 }
9124 }
9125}
9126
9127thread_local! {
9128 static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
9131 static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9132}
9133
9134#[allow(clippy::too_many_arguments)]
9140fn q4matmat(
9141 bytes: &[u8],
9142 xs_all: &[f32],
9143 b: usize,
9144 rows: usize,
9145 cols: usize,
9146 out: &mut [f32],
9147 pool: Option<&Pool>,
9148) {
9149 debug_assert_eq!(xs_all.len(), b * cols);
9150 debug_assert_eq!(out.len(), b * rows);
9151 let (packed, scales) = q4_split(bytes, rows, cols);
9152 let gpr = cols / GROUP_SIZE;
9153 let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9154
9155 if a8w8_enabled() {
9156 let acts: Vec<SplitAct> = (0..b)
9157 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9158 .collect();
9159 let acts = &acts;
9160 let out_addr = SendMut(out.as_mut_ptr());
9161 let run = move |start: usize, end: usize| {
9162 ROW_I8.with(|rb| {
9163 let mut buf = rb.borrow_mut();
9164 buf.resize(cols, 0);
9165 for r in start..end {
9166 for gi in 0..gpr {
9170 let g = r * gpr + gi;
9171 for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9172 buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
9173 buf[gi * GROUP_SIZE + k * 2 + 1] =
9174 (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
9175 }
9176 }
9177 let mut bi = 0usize;
9178 #[cfg(target_arch = "x86_64")]
9179 if avx2_enabled() && blocked_enabled() {
9180 while bi + 4 <= acts.len() {
9181 let xs = [
9182 acts[bi].xq.as_slice(),
9183 acts[bi + 1].xq.as_slice(),
9184 acts[bi + 2].xq.as_slice(),
9185 acts[bi + 3].xq.as_slice(),
9186 ];
9187 let d = unsafe {
9188 if vnni_tiles_enabled() {
9189 dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
9190 } else {
9191 dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
9192 }
9193 };
9194 for k in 0..4 {
9195 let act = &acts[bi + k];
9196 let mut acc = d[k] * act.sx;
9197 for &(j, xv) in &act.outliers {
9198 acc += (buf[j] as i8) as f32
9199 * gscale((r * cols + j) / GROUP_SIZE)
9200 * xv;
9201 }
9202 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
9204 }
9205 bi += 4;
9206 }
9207 }
9208 while bi < acts.len() {
9209 let act = &acts[bi];
9210 let mut acc = 0f32;
9211 for gi in 0..gpr {
9212 let d = dot_i8_i8(
9213 &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9214 &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9215 );
9216 acc += d as f32 * gscale(r * gpr + gi);
9217 }
9218 acc *= act.sx;
9219 for &(j, xv) in &act.outliers {
9221 acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
9222 }
9223 unsafe { *out_addr.at(bi * rows + r) = acc };
9225 bi += 1;
9226 }
9227 }
9228 })
9229 };
9230 dispatch_rows(pool, rows, &run);
9231 return;
9232 }
9233
9234 let out_addr = SendMut(out.as_mut_ptr());
9235 let run = move |start: usize, end: usize| {
9236 ROW_F32.with(|rb| {
9237 let mut buf = rb.borrow_mut();
9238 buf.resize(cols, 0.0);
9239 for r in start..end {
9240 for gi in 0..gpr {
9243 let g = r * gpr + gi;
9244 for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9245 buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
9246 buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
9247 }
9248 }
9249 for bi in 0..b {
9250 let x = &xs_all[bi * cols..(bi + 1) * cols];
9251 let mut acc = 0f32;
9252 for gi in 0..gpr {
9253 let mut ga = 0f32;
9254 for k in 0..GROUP_SIZE / 2 {
9259 let e = gi * GROUP_SIZE + k * 2;
9260 ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
9261 }
9262 acc += ga * gscale(r * gpr + gi);
9263 }
9264 unsafe { *out_addr.at(bi * rows + r) = acc };
9266 }
9267 }
9268 })
9269 };
9270 dispatch_rows(pool, rows, &run);
9271}
9272
9273#[allow(clippy::too_many_arguments)]
9278fn vbitmatmat(
9279 bytes: &[u8],
9280 offsets: &[usize],
9281 xs_all: &[f32],
9282 b: usize,
9283 rows: usize,
9284 cols: usize,
9285 out: &mut [f32],
9286 pool: Option<&Pool>,
9287) {
9288 debug_assert_eq!(xs_all.len(), b * cols);
9289 debug_assert_eq!(out.len(), b * rows);
9290 debug_assert_eq!(offsets.len(), rows + 1);
9291 let ng = cols / GROUP_SIZE;
9292 let bits = &bytes[..rows];
9293 let sc_off = rows;
9294 let gscale = |r: usize, g: usize| {
9295 let so = (r * ng + g) * 2;
9296 f16_to_f32(u16::from_le_bytes([
9297 bytes[sc_off + so],
9298 bytes[sc_off + so + 1],
9299 ]))
9300 };
9301
9302 let decode_f32 = |r: usize, dst: &mut [f32]| {
9304 let bw = bits[r] as usize;
9305 let l = ((1i32 << (bw - 1)) - 1) as f32;
9306 let data = &bytes[offsets[r]..offsets[r + 1]];
9307 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
9308 for d in dst.iter_mut() {
9309 while nbits < bw {
9310 acc = (acc << 8) | data[idx] as u64;
9311 idx += 1;
9312 nbits += 8;
9313 }
9314 let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
9315 nbits -= bw;
9316 *d = u - l;
9317 }
9318 };
9319
9320 if a8w8_enabled() {
9321 let acts: Vec<SplitAct> = (0..b)
9322 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9323 .collect();
9324 let acts = &acts;
9325 let out_addr = SendMut(out.as_mut_ptr());
9326 let run = move |start: usize, end: usize| {
9327 for r in start..end {
9328 let bw = bits[r] as usize;
9329 if bw == 8 {
9330 ROW_F32.with(|rb| {
9333 let mut buf = rb.borrow_mut();
9334 buf.resize(cols, 0.0);
9335 decode_f32(r, &mut buf);
9336 for bi in 0..b {
9337 let x = &xs_all[bi * cols..(bi + 1) * cols];
9338 let mut dot = 0f32;
9339 for g in 0..ng {
9340 let mut gd = 0f32;
9341 for k in 0..GROUP_SIZE {
9342 gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9343 }
9344 dot += gd * gscale(r, g);
9345 }
9346 unsafe { *out_addr.at(bi * rows + r) = dot };
9348 }
9349 });
9350 continue;
9351 }
9352 let l = (1i32 << (bw - 1)) - 1;
9353 let data = &bytes[offsets[r]..offsets[r + 1]];
9354 ROW_I8.with(|rb| {
9355 let mut buf = rb.borrow_mut();
9356 buf.resize(cols, 0);
9357 #[inline(always)]
9358 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
9359 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
9360 let u = unpack8::<B>(&data[blk * B..]);
9361 for k in 0..8 {
9362 chunk[k] = (u[k] - l) as i8 as u8;
9363 }
9364 }
9365 }
9366 match bw {
9367 3 => fill::<3>(data, l, &mut buf),
9368 4 => vbit_fill4(data, &mut buf),
9369 5 => fill::<5>(data, l, &mut buf),
9370 6 => fill::<6>(data, l, &mut buf),
9371 _ => unreachable!("vbit bit-width {bw} (validated at load)"),
9372 }
9373 let mut bi = 0usize;
9374 #[cfg(target_arch = "x86_64")]
9378 if avx2_enabled() && blocked_enabled() {
9379 while bi + 4 <= acts.len() {
9380 let xs = [
9381 acts[bi].xq.as_slice(),
9382 acts[bi + 1].xq.as_slice(),
9383 acts[bi + 2].xq.as_slice(),
9384 acts[bi + 3].xq.as_slice(),
9385 ];
9386 let sxs = [
9387 acts[bi].sx,
9388 acts[bi + 1].sx,
9389 acts[bi + 2].sx,
9390 acts[bi + 3].sx,
9391 ];
9392 let d = unsafe {
9393 if vnni_tiles_enabled() {
9394 dot_q4b_row_1x4_sx_vnni(
9395 &buf,
9396 &bytes[sc_off..],
9397 r * ng,
9398 ng,
9399 xs,
9400 sxs,
9401 )
9402 } else {
9403 dot_q4b_row_1x4_sx_avx2(
9404 &buf,
9405 &bytes[sc_off..],
9406 r * ng,
9407 ng,
9408 xs,
9409 sxs,
9410 )
9411 }
9412 };
9413 for k in 0..4 {
9414 let act = &acts[bi + k];
9415 let mut dot = d[k];
9416 for &(j, xv) in &act.outliers {
9417 dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9418 }
9419 unsafe { *out_addr.at((bi + k) * rows + r) = dot };
9421 }
9422 bi += 4;
9423 }
9424 }
9425 while bi < acts.len() {
9426 let act = &acts[bi];
9427 let mut dot = 0f32;
9428 for g in 0..ng {
9429 let d = dot_i8_i8(
9430 &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9431 &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9432 ) as f32
9433 * act.sx;
9434 dot += d * gscale(r, g);
9435 }
9436 for &(j, xv) in &act.outliers {
9437 dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9438 }
9439 unsafe { *out_addr.at(bi * rows + r) = dot };
9441 bi += 1;
9442 }
9443 });
9444 }
9445 };
9446 dispatch_rows(pool, rows, &run);
9447 return;
9448 }
9449
9450 let out_addr = SendMut(out.as_mut_ptr());
9451 let run = move |start: usize, end: usize| {
9452 ROW_F32.with(|rb| {
9453 let mut buf = rb.borrow_mut();
9454 buf.resize(cols, 0.0);
9455 for r in start..end {
9456 decode_f32(r, &mut buf);
9457 for bi in 0..b {
9458 let x = &xs_all[bi * cols..(bi + 1) * cols];
9459 let mut dot = 0f32;
9460 for g in 0..ng {
9461 let mut gd = 0f32;
9462 for k in 0..GROUP_SIZE {
9463 gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9464 }
9465 dot += gd * gscale(r, g);
9466 }
9467 unsafe { *out_addr.at(bi * rows + r) = dot };
9469 }
9470 }
9471 })
9472 };
9473 dispatch_rows(pool, rows, &run);
9474}
9475
9476pub(crate) fn gpu_batch_job<'a>(
9480 t: &'a QTensor,
9481 x: &[f32],
9482) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
9483 match t {
9484 QTensor::Mapped {
9485 model,
9486 idx,
9487 dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
9488 rows,
9489 cols,
9490 row_scale,
9491 col_field,
9492 ..
9493 } => Some((
9494 model.clone(),
9495 crate::gpu::BatchJob {
9496 idx: *idx,
9497 rows: *rows,
9498 cols: *cols,
9499 row_scale,
9500 xs: prescale(x, col_field, *dt).into_owned(),
9501 layout: crate::gpu::BatchLayout::Q8,
9502 },
9503 )),
9504 QTensor::Mapped {
9506 model,
9507 idx,
9508 dtype: TensorDtype::Q1,
9509 rows,
9510 cols,
9511 ..
9512 } => Some((
9513 model.clone(),
9514 crate::gpu::BatchJob {
9515 idx: *idx,
9516 rows: *rows,
9517 cols: *cols,
9518 row_scale: &[],
9519 xs: x.to_vec(),
9520 layout: crate::gpu::BatchLayout::Q1,
9521 },
9522 )),
9523 QTensor::Mapped {
9528 model,
9529 idx,
9530 dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
9531 rows,
9532 cols,
9533 ..
9534 } => Some((
9535 model.clone(),
9536 crate::gpu::BatchJob {
9537 idx: *idx,
9538 rows: *rows,
9539 cols: *cols,
9540 row_scale: &[],
9541 xs: x.to_vec(),
9542 layout: if *dt == TensorDtype::Q4Tiled {
9543 crate::gpu::BatchLayout::Q4t
9544 } else {
9545 crate::gpu::BatchLayout::Q4tp
9546 },
9547 },
9548 )),
9549 _ => None,
9550 }
9551}
9552
9553thread_local! {
9554 static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9555 static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9556}
9557
9558pub(crate) fn prescale<'a>(
9559 x: &'a [f32],
9560 col_field: &[f32],
9561 dtype: TensorDtype,
9562) -> std::borrow::Cow<'a, [f32]> {
9563 if dtype == TensorDtype::Q8_2f {
9564 x.iter().zip(col_field).map(|(a, c)| a * c).collect()
9565 } else {
9566 std::borrow::Cow::Borrowed(x)
9567 }
9568}
9569
9570pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
9573 x: &[f32],
9574 col_field: &[f32],
9575 dtype: TensorDtype,
9576 buf_id: u8,
9577 f: F,
9578) -> R {
9579 if dtype == TensorDtype::Q8_2f {
9580 if buf_id == 1 {
9581 PRESCALE_BUF1.with(|b| {
9582 let mut buf = b.borrow_mut();
9583 buf.clear();
9584 buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9585 f(&buf)
9586 })
9587 } else {
9588 PRESCALE_BUF2.with(|b| {
9589 let mut buf = b.borrow_mut();
9590 buf.clear();
9591 buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9592 f(&buf)
9593 })
9594 }
9595 } else {
9596 f(x)
9597 }
9598}
9599
9600#[cfg(target_arch = "x86_64")]
9605pub(crate) fn avx2_enabled() -> bool {
9606 use std::sync::OnceLock;
9607 static ON: OnceLock<bool> = OnceLock::new();
9608 *ON.get_or_init(|| {
9609 std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
9610 && std::arch::is_x86_feature_detected!("avx2")
9611 && std::arch::is_x86_feature_detected!("fma")
9612 })
9613}
9614
9615#[cfg(target_arch = "x86_64")]
9620fn avx2_a8w8_enabled() -> bool {
9621 if FLOAT_ACTIVATIONS.get() {
9622 return false;
9623 }
9624 use std::sync::OnceLock;
9625 static ON: OnceLock<bool> = OnceLock::new();
9626 *ON.get_or_init(|| {
9627 avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
9628 })
9629}
9630
9631thread_local! {
9632 static FULL_GPU_Q8: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9633}
9634
9635pub(crate) fn enter_full_gpu_q8_scope() -> impl Drop {
9638 struct Restore(bool, std::marker::PhantomData<std::rc::Rc<()>>);
9639 impl Drop for Restore {
9640 fn drop(&mut self) {
9641 FULL_GPU_Q8.set(self.0);
9642 }
9643 }
9644 Restore(FULL_GPU_Q8.replace(true), std::marker::PhantomData)
9645}
9646
9647thread_local! {
9651 static FLOAT_ACTIVATIONS: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9652}
9653
9654pub(crate) fn float_activations_scope<R>(f: impl FnOnce() -> R) -> R {
9655 struct Restore(bool);
9656 impl Drop for Restore {
9657 fn drop(&mut self) {
9658 FLOAT_ACTIVATIONS.set(self.0);
9659 }
9660 }
9661 let _restore = Restore(FLOAT_ACTIVATIONS.replace(true));
9662 f()
9663}
9664
9665static ROW_EXACT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
9677
9678pub(crate) fn row_exact() -> bool {
9679 ROW_EXACT.load(std::sync::atomic::Ordering::Acquire) != 0
9680}
9681
9682fn counted_row_exact_scope<R>(active: &std::sync::atomic::AtomicUsize, f: impl FnOnce() -> R) -> R {
9683 struct Restore<'a>(&'a std::sync::atomic::AtomicUsize);
9684 impl Drop for Restore<'_> {
9685 fn drop(&mut self) {
9686 self.0.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
9687 }
9688 }
9689 active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
9690 let _restore = Restore(active);
9691 f()
9692}
9693
9694pub(crate) fn row_exact_scope<R>(f: impl FnOnce() -> R) -> R {
9696 counted_row_exact_scope(&ROW_EXACT, f)
9697}
9698
9699#[inline]
9703pub(crate) fn a8w8_enabled() -> bool {
9704 #[cfg(target_arch = "aarch64")]
9705 {
9706 sdot_enabled()
9707 }
9708 #[cfg(target_arch = "x86_64")]
9709 {
9710 avx2_a8w8_enabled()
9711 }
9712 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
9713 {
9714 false
9715 }
9716}
9717
9718#[inline]
9721#[allow(unreachable_code)]
9722fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
9723 #[cfg(target_arch = "aarch64")]
9724 unsafe {
9725 return dot_i8_sdot(w, xq);
9726 }
9727 #[cfg(target_arch = "x86_64")]
9728 unsafe {
9729 if avx512vnni_enabled() {
9730 return dot_i8_i8_vnni(w, xq);
9731 }
9732 return dot_i8_i8_avx2(w, xq);
9733 }
9734 w.iter()
9735 .zip(xq)
9736 .map(|(&a, &b)| (a as i8) as i32 * b as i32)
9737 .sum()
9738}
9739
9740#[cfg(target_arch = "x86_64")]
9744fn avx512vnni_enabled() -> bool {
9745 use std::sync::OnceLock;
9746 static ON: OnceLock<bool> = OnceLock::new();
9747 *ON.get_or_init(|| {
9748 std::env::var("CMF_AVX512")
9749 .map(|v| v != "0")
9750 .unwrap_or(true)
9751 && std::arch::is_x86_feature_detected!("avx512f")
9752 && std::arch::is_x86_feature_detected!("avx512bw")
9753 && std::arch::is_x86_feature_detected!("avx512vl")
9754 && std::arch::is_x86_feature_detected!("avx512vnni")
9755 })
9756}
9757
9758#[cfg(target_arch = "x86_64")]
9766fn vnni_tiles_enabled() -> bool {
9767 use std::sync::OnceLock;
9768 static ON: OnceLock<bool> = OnceLock::new();
9769 *ON.get_or_init(|| {
9770 std::env::var("CMF_VNNI_TILES")
9771 .map(|v| v != "0")
9772 .unwrap_or(true)
9773 && avx512vnni_enabled()
9774 })
9775}
9776
9777#[cfg(target_arch = "x86_64")]
9782#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9783#[inline]
9784unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
9785 unsafe {
9787 use core::arch::x86_64::*;
9788 let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
9789 let hi128 = _mm256_extracti128_si256::<1>(d);
9790 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9791 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9792 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9793 _mm_cvtsi128_si32(s32)
9794 }
9795}
9796
9797#[cfg(target_arch = "x86_64")]
9802#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9803unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
9804 unsafe {
9806 use core::arch::x86_64::*;
9807 let n = w.len();
9808 let mut j = 0usize;
9809 let mut total: i32;
9810 {
9815 #[inline(always)]
9816 unsafe fn step(
9817 w: *const u8,
9818 x: *const i8,
9819 acc: core::arch::x86_64::__m512i,
9820 ) -> core::arch::x86_64::__m512i {
9821 unsafe {
9822 use core::arch::x86_64::*;
9823 let wv = _mm512_loadu_si512(w as *const _);
9824 let xv = _mm512_loadu_si512(x as *const _);
9825 let aw = _mm512_abs_epi8(wv);
9826 let neg = _mm512_movepi8_mask(wv);
9827 let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
9828 _mm512_dpbusd_epi32(acc, aw, sx)
9829 }
9830 }
9831 let (mut a0, mut a1, mut a2, mut a3) = (
9832 _mm512_setzero_si512(),
9833 _mm512_setzero_si512(),
9834 _mm512_setzero_si512(),
9835 _mm512_setzero_si512(),
9836 );
9837 while j + 256 <= n {
9838 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9839 a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
9840 a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
9841 a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
9842 j += 256;
9843 }
9844 while j + 64 <= n {
9845 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9846 j += 64;
9847 }
9848 let s01 = _mm512_add_epi32(a0, a1);
9849 let s23 = _mm512_add_epi32(a2, a3);
9850 total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
9851 }
9852 if j + 32 <= n {
9854 let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9855 let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9856 let d = _mm256_dpbusd_epi32(
9857 _mm256_setzero_si256(),
9858 _mm256_abs_epi8(wv),
9859 _mm256_sign_epi8(xv, wv),
9860 );
9861 let hi128 = _mm256_extracti128_si256::<1>(d);
9862 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9863 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9864 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9865 total += _mm_cvtsi128_si32(s32);
9866 j += 32;
9867 }
9868 while j < n {
9869 total += (w[j] as i8) as i32 * xq[j] as i32;
9870 j += 1;
9871 }
9872 total
9873 }
9874}
9875
9876#[cfg(target_arch = "x86_64")]
9878#[target_feature(enable = "avx2,fma")]
9879unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
9880 unsafe {
9882 use core::arch::x86_64::*;
9883 let n = x.len();
9884 let wp = w.as_ptr();
9885 let xp = x.as_ptr();
9886 let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
9887 let mut j = 0usize;
9888 while j + 16 <= n {
9889 let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
9890 let lo = _mm256_cvtepi8_epi32(wb);
9891 let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
9892 a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
9893 a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
9894 j += 16;
9895 }
9896 let acc = _mm256_add_ps(a0, a1);
9897 let hi128 = _mm256_extractf128_ps::<1>(acc);
9898 let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
9899 let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
9900 let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
9901 let mut sum = _mm_cvtss_f32(s32);
9902 while j < n {
9903 sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
9904 j += 1;
9905 }
9906 sum
9907 }
9908}
9909
9910#[cfg(target_arch = "x86_64")]
9915#[target_feature(enable = "avx2")]
9916unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
9917 unsafe {
9919 use core::arch::x86_64::*;
9920 let n = w.len();
9921 let ones = _mm256_set1_epi16(1);
9922 let mut acc = _mm256_setzero_si256();
9923 let mut j = 0usize;
9924 while j + 32 <= n {
9925 let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9926 let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9927 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
9928 acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
9929 j += 32;
9930 }
9931 let hi128 = _mm256_extracti128_si256::<1>(acc);
9932 let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
9933 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9934 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9935 let mut s = _mm_cvtsi128_si32(s32);
9936 while j < n {
9937 s += (w[j] as i8) as i32 * xq[j] as i32;
9938 j += 1;
9939 }
9940 s
9941 }
9942}
9943
9944#[cfg(target_arch = "aarch64")]
9948#[target_feature(enable = "neon,i8mm")]
9949unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9950 unsafe {
9952 use core::arch::aarch64::*;
9953 use core::arch::asm;
9954 let n = w0.len();
9955 let w0p = w0.as_ptr() as *const i8;
9956 let w1p = w1.as_ptr() as *const i8;
9957 let mut acc01 = vdupq_n_s32(0);
9960 let mut acc23 = vdupq_n_s32(0);
9961 let mut i = 0usize;
9962 while i + 8 <= n {
9963 let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
9964 let xb01 = vcombine_s8(
9965 vld1_s8(xs[0].as_ptr().add(i)),
9966 vld1_s8(xs[1].as_ptr().add(i)),
9967 );
9968 let xb23 = vcombine_s8(
9969 vld1_s8(xs[2].as_ptr().add(i)),
9970 vld1_s8(xs[3].as_ptr().add(i)),
9971 );
9972 asm!(
9973 "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
9974 "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
9975 a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
9976 w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
9977 options(pure, nomem, nostack),
9978 );
9979 i += 8;
9980 }
9981 let mut out = [[0i32; 4]; 2];
9982 let a01: [i32; 4] = core::mem::transmute(acc01);
9983 let a23: [i32; 4] = core::mem::transmute(acc23);
9984 out[0][0] = a01[0];
9985 out[0][1] = a01[1];
9986 out[1][0] = a01[2];
9987 out[1][1] = a01[3];
9988 out[0][2] = a23[0];
9989 out[0][3] = a23[1];
9990 out[1][2] = a23[2];
9991 out[1][3] = a23[3];
9992 if i < n {
9993 for (k, x) in xs.iter().enumerate() {
9994 for j in i..n {
9995 out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9996 out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9997 }
9998 }
9999 }
10000 out
10001 }
10002}
10003
10004#[cfg(target_arch = "aarch64")]
10008#[target_feature(enable = "neon,dotprod")]
10009unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
10010 unsafe {
10012 use core::arch::aarch64::*;
10013 use core::arch::asm;
10014 let n = w0.len();
10015 let w0p = w0.as_ptr() as *const i8;
10016 let w1p = w1.as_ptr() as *const i8;
10017 let mut acc = [[vdupq_n_s32(0); 4]; 2];
10018 let mut i = 0usize;
10019 while i + 16 <= n {
10020 let wv0 = vld1q_s8(w0p.add(i));
10021 let wv1 = vld1q_s8(w1p.add(i));
10022 for (k, x) in xs.iter().enumerate() {
10023 let xv = vld1q_s8(x.as_ptr().add(i));
10024 let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
10025 asm!(
10026 "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
10027 "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
10028 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10029 w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
10030 options(pure, nomem, nostack),
10031 );
10032 acc[0][k] = a0;
10033 acc[1][k] = a1;
10034 }
10035 i += 16;
10036 }
10037 let mut out = [[0i32; 4]; 2];
10038 for r in 0..2 {
10039 for k in 0..4 {
10040 out[r][k] = vaddvq_s32(acc[r][k]);
10041 }
10042 }
10043 if i < n {
10044 for (k, x) in xs.iter().enumerate() {
10045 for j in i..n {
10046 out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
10047 out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
10048 }
10049 }
10050 }
10051 out
10052 }
10053}
10054
10055#[cfg(target_arch = "x86_64")]
10061#[target_feature(enable = "avx2")]
10062unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
10063 unsafe {
10065 use core::arch::x86_64::*;
10066 let n = w0.len();
10067 let ones = _mm256_set1_epi16(1);
10068 let mut acc = [[_mm256_setzero_si256(); 4]; 2];
10069 let mut j = 0usize;
10070 while j + 32 <= n {
10071 let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
10072 let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
10073 let aw0 = _mm256_abs_epi8(wv0);
10074 let aw1 = _mm256_abs_epi8(wv1);
10075 for (k, x) in xs.iter().enumerate() {
10076 let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
10077 let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
10078 acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
10079 let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
10080 acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
10081 }
10082 j += 32;
10083 }
10084 let mut out = [[0i32; 4]; 2];
10085 for r in 0..2 {
10086 for k in 0..4 {
10087 let a = acc[r][k];
10088 let hi128 = _mm256_extracti128_si256::<1>(a);
10089 let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
10090 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10091 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10092 out[r][k] = _mm_cvtsi128_si32(s32);
10093 }
10094 }
10095 if j < n {
10096 for (k, x) in xs.iter().enumerate() {
10097 for i in j..n {
10098 out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
10099 out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
10100 }
10101 }
10102 }
10103 out
10104 }
10105}
10106
10107#[cfg(target_arch = "x86_64")]
10112#[inline]
10113fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
10114 let dot = if avx512vnni_enabled() && row.len() >= 64 {
10115 (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
10116 } else {
10117 unsafe { dot_i8_i8_avx2(row, &act.xq) }
10118 };
10119 let mut acc = dot as f32 * act.sx;
10120 for &(j, xv) in &act.outliers {
10121 acc += (row[j] as i8) as f32 * xv;
10122 }
10123 acc
10124}
10125
10126#[cfg(target_arch = "x86_64")]
10130#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
10131unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
10132 unsafe {
10134 use core::arch::x86_64::*;
10135 let n = w.len();
10136 let flip = _mm512_set1_epi8(-128); #[inline(always)]
10138 unsafe fn step(
10139 w: *const u8,
10140 x: *const i8,
10141 flip: core::arch::x86_64::__m512i,
10142 acc: core::arch::x86_64::__m512i,
10143 ) -> core::arch::x86_64::__m512i {
10144 unsafe {
10145 use core::arch::x86_64::*;
10146 let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
10147 _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
10148 }
10149 }
10150 let (mut a0, mut a1, mut a2, mut a3) = (
10151 _mm512_setzero_si512(),
10152 _mm512_setzero_si512(),
10153 _mm512_setzero_si512(),
10154 _mm512_setzero_si512(),
10155 );
10156 let mut j = 0usize;
10157 while j + 256 <= n {
10158 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10159 a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
10160 a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
10161 a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
10162 j += 256;
10163 }
10164 while j + 64 <= n {
10165 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10166 j += 64;
10167 }
10168 let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
10169 _mm512_add_epi32(a0, a1),
10170 _mm512_add_epi32(a2, a3),
10171 ));
10172 while j < n {
10174 total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
10175 j += 1;
10176 }
10177 total
10178 }
10179}
10180
10181#[cfg(target_arch = "x86_64")]
10187#[target_feature(enable = "avx2")]
10188unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10189 unsafe {
10192 use core::arch::x86_64::*;
10193 let lomask = _mm_set1_epi8(0x0F);
10194 let eight = _mm256_set1_epi8(8);
10195 let ones = _mm256_set1_epi16(1);
10196 let mut acc = 0f32;
10197 for gi in 0..gpr {
10198 let g = g0 + gi;
10199 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10200 let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10201 let lo = _mm_and_si128(b, lomask);
10202 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10203 let w = _mm256_sub_epi8(
10204 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10205 eight,
10206 );
10207 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10208 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
10209 let d = _mm256_madd_epi16(p16, ones);
10210 let hi128 = _mm256_extracti128_si256::<1>(d);
10211 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10212 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10213 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10214 acc += _mm_cvtsi128_si32(s32) as f32 * s;
10215 }
10216 acc
10217 }
10218}
10219
10220#[cfg(target_arch = "x86_64")]
10223#[target_feature(enable = "avx2")]
10224unsafe fn dot_q4_row_avx2_2(
10225 packed: &[u8],
10226 scales: &[u8],
10227 g0: usize,
10228 gpr: usize,
10229 xq1: &[i8],
10230 xq2: &[i8],
10231) -> (f32, f32) {
10232 unsafe {
10234 use core::arch::x86_64::*;
10235 let lomask = _mm_set1_epi8(0x0F);
10236 let eight = _mm256_set1_epi8(8);
10237 let ones = _mm256_set1_epi16(1);
10238 let (mut acc1, mut acc2) = (0f32, 0f32);
10239 #[inline(always)]
10240 unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
10241 unsafe {
10242 use core::arch::x86_64::*;
10243 let hi128 = _mm256_extracti128_si256::<1>(d);
10244 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10245 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10246 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10247 _mm_cvtsi128_si32(s32)
10248 }
10249 }
10250 for gi in 0..gpr {
10251 let g = g0 + gi;
10252 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10253 let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10254 let lo = _mm_and_si128(b, lomask);
10255 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10256 let w = _mm256_sub_epi8(
10257 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10258 eight,
10259 );
10260 let aw = _mm256_abs_epi8(w);
10261 let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10262 let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10263 let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
10264 let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
10265 acc1 += hsum(d1) as f32 * s;
10266 acc2 += hsum(d2) as f32 * s;
10267 }
10268 (acc1, acc2)
10269 }
10270}
10271
10272#[cfg(target_arch = "x86_64")]
10274fn q8_range_avx2(
10275 q: &[u8],
10276 row_scale: &[f32],
10277 act: &SplitAct,
10278 cols: usize,
10279 out_addr: SendMut,
10280 start: usize,
10281 end: usize,
10282) {
10283 for o in start..end {
10284 let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10285 unsafe { *out_addr.at(o) = v };
10287 }
10288}
10289
10290#[cfg(target_arch = "x86_64")]
10292#[allow(clippy::too_many_arguments)]
10293fn q8_range2_avx2(
10294 q: &[u8],
10295 row_scale: &[f32],
10296 a1: &SplitAct,
10297 a2: &SplitAct,
10298 cols: usize,
10299 p1: SendMut,
10300 p2: SendMut,
10301 start: usize,
10302 end: usize,
10303) {
10304 for o in start..end {
10305 let row = &q[o * cols..(o + 1) * cols];
10306 unsafe {
10308 *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
10309 *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
10310 }
10311 }
10312}
10313
10314#[cfg(target_arch = "aarch64")]
10325fn i8mm_enabled() -> bool {
10326 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10327 *ON.get_or_init(|| {
10328 std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
10329 && std::arch::is_aarch64_feature_detected!("i8mm")
10330 })
10331}
10332
10333#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
10337fn sdot_enabled() -> bool {
10338 if FLOAT_ACTIVATIONS.get() {
10339 return false;
10340 }
10341 use std::sync::OnceLock;
10342 static ON: OnceLock<bool> = OnceLock::new();
10343 *ON.get_or_init(|| {
10344 let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
10345 if !want {
10346 return false;
10347 }
10348
10349 #[cfg(target_arch = "aarch64")]
10350 {
10351 if std::arch::is_aarch64_feature_detected!("dotprod") {
10352 return true;
10353 }
10354 #[cfg(target_os = "android")]
10355 {
10356 if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
10357 if cpuinfo.lines().any(|l| {
10358 (l.starts_with("Features") || l.starts_with("features"))
10359 && l.contains("asimddp")
10360 }) {
10361 return true;
10362 }
10363 }
10364 }
10365 false
10366 }
10367 #[cfg(not(target_arch = "aarch64"))]
10368 {
10369 false
10370 }
10371 })
10372}
10373
10374struct SplitAct {
10379 xq: Vec<i8>,
10380 sx: f32,
10381 outliers: Vec<(usize, f32)>,
10382 #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
10385 xsum: i32,
10386}
10387
10388thread_local! {
10389 static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
10392 const { std::cell::RefCell::new(Vec::new()) };
10393}
10394
10395impl Drop for SplitAct {
10396 fn drop(&mut self) {
10397 let buf = std::mem::take(&mut self.xq);
10398 if buf.capacity() > 0 {
10399 XQ_FREE.with(|f| {
10400 let mut f = f.borrow_mut();
10401 if f.len() < 16 {
10402 f.push(buf);
10403 }
10404 });
10405 }
10406 }
10407}
10408
10409thread_local! {
10410 static KROW: UnsafeCell<Vec<f32>> = const { UnsafeCell::new(Vec::new()) };
10417 static KROWS: UnsafeCell<[Vec<f32>; 2]> =
10422 const { UnsafeCell::new([Vec::new(), Vec::new()]) };
10423}
10424
10425#[inline]
10428fn with_krow<R>(n: usize, f: impl FnOnce(&mut [f32]) -> R) -> R {
10429 KROW.with(|s| {
10430 let b = unsafe { &mut *s.get() };
10433 if b.len() < n {
10434 b.resize(n, 0.0);
10435 }
10436 f(&mut b[..n])
10437 })
10438}
10439
10440#[inline(always)]
10451fn q8_round(t: f32) -> i8 {
10452 let t = t.clamp(-127.0, 127.0);
10453 let i = t as i32;
10454 let f = t - i as f32;
10455 let r = if f >= 0.5 {
10456 i + 1
10457 } else if f <= -0.5 {
10458 i - 1
10459 } else {
10460 i
10461 };
10462 r as i8
10463}
10464
10465#[inline]
10466fn with_krows<R>(n: usize, f: impl FnOnce(&mut [f32], &mut [f32]) -> R) -> R {
10467 KROWS.with(|s| {
10468 let b = unsafe { &mut *s.get() };
10471 for row in b.iter_mut() {
10472 if row.len() < n {
10473 row.resize(n, 0.0);
10474 }
10475 }
10476 let (a, b) = b.split_at_mut(1);
10477 f(&mut a[0][..n], &mut b[0][..n])
10478 })
10479}
10480
10481#[inline]
10482fn silu_mul_limited(mut gate: f32, mut up: f32, limit: f32) -> f32 {
10483 if limit > 0.0 {
10484 up = up.clamp(-limit, limit);
10485 gate = gate.min(limit);
10486 }
10487 gate / (1.0 + (-gate).exp()) * up
10488}
10489
10490fn split_act(x: &[f32]) -> SplitAct {
10491 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10492 let n = x.len();
10493 let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
10494 let thr = 8.0 * rms;
10495 let mut outliers: Vec<(usize, f32)> = Vec::new();
10499 let mut amax = 0f32;
10500 for (j, &v) in x.iter().enumerate() {
10501 let a = v.abs();
10502 if a > thr {
10503 outliers.push((j, v));
10504 } else if a > amax {
10505 amax = a;
10506 }
10507 }
10508 let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10509 let inv = 1.0 / sx;
10510 let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10511 xq.clear();
10512 xq.reserve(n);
10513 if outliers.is_empty() {
10514 xq.extend(
10515 x.iter()
10516 .map(|&v| q8_round(v * inv)),
10517 );
10518 } else {
10519 xq.extend(x.iter().map(|&v| {
10521 if v.abs() > thr {
10522 0
10523 } else {
10524 q8_round(v * inv)
10525 }
10526 }));
10527 }
10528 let xsum = xq.iter().map(|&v| v as i32).sum();
10529 SplitAct {
10530 xq,
10531 sx,
10532 outliers,
10533 xsum,
10534 }
10535}
10536
10537fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
10538 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10539 let n = x.len();
10540 let rms = (x
10541 .iter()
10542 .zip(col)
10543 .map(|(&a, &c)| {
10544 let v = a * c;
10545 (v * v) as f64
10546 })
10547 .sum::<f64>()
10548 / n.max(1) as f64)
10549 .sqrt() as f32;
10550 let thr = 8.0 * rms;
10551
10552 let mut outliers = Vec::new();
10553 let mut amax = 0f32;
10554 for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
10555 let v = a * c;
10556 let s = v.abs();
10557 if s > thr {
10558 outliers.push((j, v));
10559 } else if s > amax {
10560 amax = s;
10561 }
10562 }
10563
10564 let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10565 let inv = 1.0 / sx;
10566 let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10567 xq.clear();
10568 xq.reserve(n);
10569 if outliers.is_empty() {
10570 xq.extend(
10571 x.iter()
10572 .zip(col)
10573 .map(|(&a, &c)| q8_round((a * c) * inv)),
10574 );
10575 } else {
10576 xq.extend(x.iter().zip(col).map(|(&a, &c)| {
10577 let v = a * c;
10578 if v.abs() > thr {
10579 0
10580 } else {
10581 q8_round(v * inv)
10582 }
10583 }));
10584 }
10585 let xsum = xq.iter().map(|&v| v as i32).sum();
10586 SplitAct {
10587 xq,
10588 sx,
10589 outliers,
10590 xsum,
10591 }
10592}
10593
10594#[cfg(target_arch = "aarch64")]
10597#[target_feature(enable = "neon,dotprod")]
10598unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
10599 unsafe {
10601 use core::arch::aarch64::*;
10602 use core::arch::asm;
10603 let wp = w.as_ptr() as *const i8;
10604 let n = w.len();
10605 let (mut a0, mut a1, mut a2, mut a3) = (
10606 vdupq_n_s32(0),
10607 vdupq_n_s32(0),
10608 vdupq_n_s32(0),
10609 vdupq_n_s32(0),
10610 );
10611 let mut i = 0;
10612 while i + 64 <= n {
10613 let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10614 let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
10615 let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
10616 let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
10617 asm!(
10618 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
10619 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
10620 "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
10621 "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
10622 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10623 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
10624 w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
10625 options(pure, nomem, nostack),
10626 );
10627 i += 64;
10628 }
10629 while i + 16 <= n {
10630 let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10631 asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
10632 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
10633 i += 16;
10634 }
10635 let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
10636 while i < n {
10637 s += (*wp.add(i)) as i32 * xq[i] as i32;
10638 i += 1;
10639 }
10640 s
10641 }
10642}
10643
10644#[cfg(target_arch = "aarch64")]
10648#[target_feature(enable = "neon,dotprod")]
10649unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
10650 unsafe {
10652 use core::arch::aarch64::*;
10653 use core::arch::asm;
10654 let n = xq.len();
10655 let px = xq.as_ptr();
10656 let (p0, p1, p2, p3) = (
10657 w0.as_ptr() as *const i8,
10658 w1.as_ptr() as *const i8,
10659 w2.as_ptr() as *const i8,
10660 w3.as_ptr() as *const i8,
10661 );
10662 let (mut a0, mut a1, mut a2, mut a3) = (
10663 vdupq_n_s32(0),
10664 vdupq_n_s32(0),
10665 vdupq_n_s32(0),
10666 vdupq_n_s32(0),
10667 );
10668 let mut i = 0;
10669 while i + 16 <= n {
10670 let x = vld1q_s8(px.add(i));
10671 let v0 = vld1q_s8(p0.add(i));
10672 let v1 = vld1q_s8(p1.add(i));
10673 let v2 = vld1q_s8(p2.add(i));
10674 let v3 = vld1q_s8(p3.add(i));
10675 asm!(
10676 "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10677 "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10678 "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10679 "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10680 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10681 v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10682 options(pure, nomem, nostack),
10683 );
10684 i += 16;
10685 }
10686 let mut r = [
10687 vaddvq_s32(a0),
10688 vaddvq_s32(a1),
10689 vaddvq_s32(a2),
10690 vaddvq_s32(a3),
10691 ];
10692 while i < n {
10693 let xi = *px.add(i) as i32;
10694 r[0] += (*p0.add(i)) as i32 * xi;
10695 r[1] += (*p1.add(i)) as i32 * xi;
10696 r[2] += (*p2.add(i)) as i32 * xi;
10697 r[3] += (*p3.add(i)) as i32 * xi;
10698 i += 1;
10699 }
10700 r
10701 }
10702}
10703
10704#[cfg(target_arch = "aarch64")]
10711#[target_feature(enable = "neon,dotprod")]
10712unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
10713 unsafe {
10716 use core::arch::aarch64::*;
10717 use core::arch::asm;
10718 let n = xq.len();
10719 let px = xq.as_ptr();
10720 let pg = g.as_ptr() as *const i8;
10721 let (mut a0, mut a1, mut a2, mut a3) = (
10722 vdupq_n_s32(0),
10723 vdupq_n_s32(0),
10724 vdupq_n_s32(0),
10725 vdupq_n_s32(0),
10726 );
10727 let mut i = 0;
10728 while i + 16 <= n {
10729 let x = vld1q_s8(px.add(i));
10730 let base = pg.add(4 * i);
10731 let v0 = vld1q_s8(base);
10732 let v1 = vld1q_s8(base.add(16));
10733 let v2 = vld1q_s8(base.add(32));
10734 let v3 = vld1q_s8(base.add(48));
10735 asm!(
10736 "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10737 "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10738 "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10739 "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10740 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10741 v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10742 options(pure, nomem, nostack),
10743 );
10744 i += 16;
10745 }
10746 [
10747 vaddvq_s32(a0),
10748 vaddvq_s32(a1),
10749 vaddvq_s32(a2),
10750 vaddvq_s32(a3),
10751 ]
10752 }
10753}
10754
10755#[cfg(target_arch = "aarch64")]
10761fn q8_range_sdot(
10762 q: &[u8],
10763 rep: &[u8],
10764 row_scale: &[f32],
10765 act: &SplitAct,
10766 cols: usize,
10767 out_addr: SendMut,
10768 start: usize,
10769 end: usize,
10770) {
10771 let mut o = start;
10772 if !rep.is_empty() {
10775 while o < end && o % 4 != 0 {
10776 let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10777 unsafe { *out_addr.at(o) = v };
10778 o += 1;
10779 }
10780 }
10781 while o + 4 <= end {
10782 let r = if rep.is_empty() {
10783 unsafe {
10784 dot_i8_sdot_4rows(
10785 &q[o * cols..(o + 1) * cols],
10786 &q[(o + 1) * cols..(o + 2) * cols],
10787 &q[(o + 2) * cols..(o + 3) * cols],
10788 &q[(o + 3) * cols..(o + 4) * cols],
10789 &act.xq,
10790 )
10791 }
10792 } else {
10793 unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
10794 };
10795 for k in 0..4 {
10796 let mut acc = r[k] as f32 * act.sx;
10797 for &(j, xv) in &act.outliers {
10798 acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
10799 }
10800 unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
10802 }
10803 o += 4;
10804 }
10805 while o < end {
10806 let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10807 unsafe { *out_addr.at(o) = v };
10808 o += 1;
10809 }
10810}
10811
10812#[cfg(target_arch = "aarch64")]
10815#[allow(clippy::too_many_arguments)]
10816fn q8_range2_sdot(
10817 q: &[u8],
10818 row_scale: &[f32],
10819 a1: &SplitAct,
10820 a2: &SplitAct,
10821 cols: usize,
10822 p1: SendMut,
10823 p2: SendMut,
10824 start: usize,
10825 end: usize,
10826) {
10827 for o in start..end {
10828 let row = &q[o * cols..(o + 1) * cols];
10829 unsafe {
10831 *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
10832 *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
10833 }
10834 }
10835}
10836
10837#[allow(clippy::too_many_arguments)]
10839fn q8_range2_f32(
10840 q: &[u8],
10841 row_scale: &[f32],
10842 x1: &[f32],
10843 x2: &[f32],
10844 cols: usize,
10845 p1: SendMut,
10846 p2: SendMut,
10847 start: usize,
10848 end: usize,
10849) {
10850 for o in start..end {
10851 let row = &q[o * cols..(o + 1) * cols];
10852 unsafe {
10854 *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
10855 *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
10856 }
10857 }
10858}
10859
10860fn q8_range_f32(
10862 q: &[u8],
10863 row_scale: &[f32],
10864 xs: &[f32],
10865 cols: usize,
10866 out_addr: SendMut,
10867 start: usize,
10868 end: usize,
10869) {
10870 for o in start..end {
10871 let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
10872 unsafe { *out_addr.at(o) = v };
10874 }
10875}
10876
10877#[inline]
10881fn q8_row_dot(row: &[u8], act: &SplitAct) -> f32 {
10882 #[cfg(target_arch = "aarch64")]
10883 return row_dot_sdot(row, act);
10884 #[cfg(target_arch = "x86_64")]
10885 return row_dot_avx2(row, act);
10886 #[allow(unreachable_code)]
10887 q8_row_dot_scalar(row, act)
10888}
10889
10890#[allow(dead_code)]
10891fn q8_row_dot_scalar(row: &[u8], act: &SplitAct) -> f32 {
10892 let mut acc = 0i32;
10893 for (k, &b) in row.iter().enumerate() {
10894 acc += (b as i8) as i32 * act.xq[k] as i32;
10895 }
10896 let mut acc = acc as f32 * act.sx;
10897 for &(j, xv) in &act.outliers {
10898 acc += (row[j] as i8) as f32 * xv;
10899 }
10900 acc
10901}
10902
10903#[cfg(target_arch = "aarch64")]
10906#[inline]
10907fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
10908 let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
10909 for &(j, xv) in &act.outliers {
10910 acc += (row[j] as i8) as f32 * xv;
10911 }
10912 acc
10913}
10914
10915#[cfg(target_arch = "aarch64")]
10923#[target_feature(enable = "neon,dotprod")]
10924unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10925 unsafe {
10928 use core::arch::aarch64::*;
10929 use core::arch::asm;
10930 let lomask = vdupq_n_u8(0x0F);
10931 let eight = vdupq_n_s8(8);
10932 let mut acc = 0f32;
10933 for gi in 0..gpr {
10934 let g = g0 + gi;
10935 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10936 let b = vld1q_u8(packed.as_ptr().add(g * 16));
10937 let lo = vandq_u8(b, lomask);
10938 let hi = vshrq_n_u8::<4>(b);
10939 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10940 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10941 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
10942 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
10943 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
10944 asm!(
10945 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
10946 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
10947 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10948 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
10949 options(pure, nomem, nostack),
10950 );
10951 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10952 }
10953 acc
10954 }
10955}
10956
10957#[cfg(target_arch = "aarch64")]
10962#[target_feature(enable = "neon,dotprod")]
10963unsafe fn dot_q4_row_sdot2(
10964 packed: &[u8],
10965 scales: &[u8],
10966 g0: usize,
10967 gpr: usize,
10968 xq1: &[i8],
10969 xq2: &[i8],
10970) -> (f32, f32) {
10971 unsafe {
10974 use core::arch::aarch64::*;
10975 use core::arch::asm;
10976 let lomask = vdupq_n_u8(0x0F);
10977 let eight = vdupq_n_s8(8);
10978 let (mut acc1, mut acc2) = (0f32, 0f32);
10979 for gi in 0..gpr {
10980 let g = g0 + gi;
10981 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10982 let b = vld1q_u8(packed.as_ptr().add(g * 16));
10983 let lo = vandq_u8(b, lomask);
10984 let hi = vshrq_n_u8::<4>(b);
10985 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10986 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10987 let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
10988 let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
10989 let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
10990 let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
10991 let (mut a0, mut a1, mut b0, mut b1) = (
10992 vdupq_n_s32(0),
10993 vdupq_n_s32(0),
10994 vdupq_n_s32(0),
10995 vdupq_n_s32(0),
10996 );
10997 asm!(
10998 "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
10999 "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
11000 "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
11001 "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
11002 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
11003 b0 = inout(vreg) b0, b1 = inout(vreg) b1,
11004 e0 = in(vreg) e0, e1 = in(vreg) e1,
11005 x10 = in(vreg) x10, x11 = in(vreg) x11,
11006 x20 = in(vreg) x20, x21 = in(vreg) x21,
11007 options(pure, nomem, nostack),
11008 );
11009 acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
11010 acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
11011 }
11012 (acc1, acc2)
11013 }
11014}
11015
11016#[inline]
11021pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
11022 #[cfg(target_arch = "aarch64")]
11023 unsafe {
11024 return axpy_i8_f32_neon(acc, row, w);
11025 }
11026 #[cfg(target_arch = "x86_64")]
11027 if avx2_enabled() {
11028 return unsafe { axpy_i8_f32_avx2(acc, row, w) };
11029 }
11030 #[allow(unreachable_code)]
11031 {
11032 for (a, &b) in acc.iter_mut().zip(row) {
11033 *a += w * b as f32;
11034 }
11035 }
11036}
11037
11038#[cfg(target_arch = "x86_64")]
11040#[target_feature(enable = "avx2,fma")]
11041unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
11042 unsafe {
11044 use core::arch::x86_64::*;
11045 let n = acc.len().min(row.len());
11046 let ap = acc.as_mut_ptr();
11047 let rp = row.as_ptr();
11048 let wv = _mm256_set1_ps(w);
11049 let mut j = 0usize;
11050 while j + 16 <= n {
11051 let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
11052 let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
11053 let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
11054 let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
11055 let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
11056 _mm256_storeu_ps(ap.add(j), v0);
11057 _mm256_storeu_ps(ap.add(j + 8), v1);
11058 j += 16;
11059 }
11060 while j < n {
11061 *ap.add(j) += w * (*rp.add(j)) as f32;
11062 j += 1;
11063 }
11064 }
11065}
11066
11067#[cfg(target_arch = "aarch64")]
11068#[target_feature(enable = "neon")]
11069unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
11070 unsafe {
11072 use core::arch::aarch64::*;
11073 let n = acc.len().min(row.len());
11074 let ap = acc.as_mut_ptr();
11075 let rp = row.as_ptr();
11076 let wv = vdupq_n_f32(w);
11077 let mut j = 0usize;
11078 while j + 16 <= n {
11079 let rb = vld1q_s8(rp.add(j));
11080 let lo = vmovl_s8(vget_low_s8(rb));
11081 let hi = vmovl_s8(vget_high_s8(rb));
11082 for (off, half) in [(0, lo), (8, hi)] {
11083 let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
11084 let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
11085 let o = j + off;
11086 vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
11087 vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
11088 }
11089 j += 16;
11090 }
11091 while j < n {
11092 *ap.add(j) += w * (*rp.add(j)) as f32;
11093 j += 1;
11094 }
11095 }
11096}
11097
11098#[inline]
11101pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
11102 #[cfg(target_arch = "aarch64")]
11103 unsafe {
11104 return dot_i8_f32_neon(w, x);
11105 }
11106 #[cfg(target_arch = "x86_64")]
11107 if avx2_enabled() {
11108 return unsafe { dot_i8_f32_avx2(w, x) };
11109 }
11110 #[allow(unreachable_code)]
11111 {
11112 let mut sum = 0.0f32;
11113 for (j, &b) in w.iter().enumerate() {
11114 sum += (b as i8) as f32 * x[j];
11115 }
11116 sum
11117 }
11118}
11119
11120#[inline]
11124fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11125 #[cfg(target_arch = "aarch64")]
11126 unsafe {
11127 return dot_i8_col_f32_neon(w, x, col);
11128 }
11129 #[allow(unreachable_code)]
11130 {
11131 let mut sum = 0.0f32;
11132 for (j, &b) in w.iter().enumerate() {
11133 sum += (b as i8) as f32 * x[j] * col[j];
11134 }
11135 sum
11136 }
11137}
11138
11139#[cfg(target_arch = "aarch64")]
11140#[target_feature(enable = "neon")]
11141unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11142 unsafe {
11144 use core::arch::aarch64::*;
11145 let n = x.len();
11146 let wp = w.as_ptr() as *const i8;
11147 let xp = x.as_ptr();
11148 let cp = col.as_ptr();
11149 let (mut a0, mut a1, mut a2, mut a3) = (
11150 vdupq_n_f32(0.0),
11151 vdupq_n_f32(0.0),
11152 vdupq_n_f32(0.0),
11153 vdupq_n_f32(0.0),
11154 );
11155 let mut j = 0usize;
11156 while j + 16 <= n {
11157 let wb = vld1q_s8(wp.add(j));
11158 let lo = vmovl_s8(vget_low_s8(wb));
11159 let hi = vmovl_s8(vget_high_s8(wb));
11160 let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11161 let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11162 let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11163 let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11164 a0 = vfmaq_f32(
11165 a0,
11166 w0,
11167 vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
11168 );
11169 a1 = vfmaq_f32(
11170 a1,
11171 w1,
11172 vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
11173 );
11174 a2 = vfmaq_f32(
11175 a2,
11176 w2,
11177 vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
11178 );
11179 a3 = vfmaq_f32(
11180 a3,
11181 w3,
11182 vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
11183 );
11184 j += 16;
11185 }
11186 let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11187 while j < n {
11188 sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
11189 j += 1;
11190 }
11191 sum
11192 }
11193}
11194
11195#[cfg(target_arch = "aarch64")]
11196#[target_feature(enable = "neon")]
11197unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
11198 unsafe {
11200 use core::arch::aarch64::*;
11201 let n = x.len();
11202 let wp = w.as_ptr() as *const i8;
11203 let xp = x.as_ptr();
11204 let (mut a0, mut a1, mut a2, mut a3) = (
11205 vdupq_n_f32(0.0),
11206 vdupq_n_f32(0.0),
11207 vdupq_n_f32(0.0),
11208 vdupq_n_f32(0.0),
11209 );
11210 let mut j = 0usize;
11211 while j + 16 <= n {
11212 let wb = vld1q_s8(wp.add(j));
11213 let lo = vmovl_s8(vget_low_s8(wb));
11214 let hi = vmovl_s8(vget_high_s8(wb));
11215 let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11216 let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11217 let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11218 let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11219 a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
11220 a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
11221 a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
11222 a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
11223 j += 16;
11224 }
11225 let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11226 while j < n {
11227 sum += (*wp.add(j)) as f32 * *xp.add(j);
11228 j += 1;
11229 }
11230 sum
11231 }
11232}
11233
11234#[allow(clippy::too_many_arguments)]
11235fn qmatvec(
11236 q: &[u8],
11237 rep: &[u8],
11238 row_scale: &[f32],
11239 x: &[f32],
11240 col_field: &[f32],
11241 dtype: TensorDtype,
11242 rows: usize,
11243 cols: usize,
11244 out: &mut [f32],
11245 pool: Option<&Pool>,
11246) {
11247 debug_assert_eq!(out.len(), rows);
11248 #[cfg(not(target_arch = "aarch64"))]
11249 let _ = rep;
11250
11251 #[cfg(target_arch = "aarch64")]
11252 if sdot_enabled() {
11253 let act = if dtype == TensorDtype::Q8_2f {
11254 split_act_q8_2f(x, col_field)
11255 } else {
11256 split_act(x)
11257 };
11258 let out_addr = SendMut(out.as_mut_ptr());
11259 let run_range = |start: usize, end: usize| {
11260 q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
11261 };
11262 match pool {
11263 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11264 _ => run_range(0, rows),
11265 }
11266 return;
11267 }
11268 #[cfg(target_arch = "x86_64")]
11271 if avx2_a8w8_enabled() {
11272 let act = if dtype == TensorDtype::Q8_2f {
11273 split_act_q8_2f(x, col_field)
11274 } else {
11275 split_act(x)
11276 };
11277 let out_addr = SendMut(out.as_mut_ptr());
11278 let run_range = |start: usize, end: usize| {
11279 q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
11280 };
11281 match pool {
11282 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11283 _ => run_range(0, rows),
11284 }
11285 return;
11286 }
11287
11288 prescale_with(x, col_field, dtype, 1, |xs| {
11289 let out_addr = SendMut(out.as_mut_ptr());
11290 let run_range = move |start: usize, end: usize| {
11291 for o in start..end {
11292 let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
11293 unsafe { *out_addr.at(o) = v };
11295 }
11296 };
11297 match pool {
11298 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11299 _ => run_range(0, rows),
11300 }
11301 });
11302}
11303
11304#[allow(clippy::too_many_arguments)]
11305fn qmatvec2(
11306 q: &[u8],
11307 row_scale: &[f32],
11308 x1: &[f32],
11309 x2: &[f32],
11310 col_field: &[f32],
11311 dtype: TensorDtype,
11312 rows: usize,
11313 cols: usize,
11314 o1: &mut [f32],
11315 o2: &mut [f32],
11316 pool: Option<&Pool>,
11317) {
11318 #[cfg(target_arch = "aarch64")]
11319 if sdot_enabled() {
11320 let a1s = if dtype == TensorDtype::Q8_2f {
11321 split_act_q8_2f(x1, col_field)
11322 } else {
11323 split_act(x1)
11324 };
11325 let a2s = if dtype == TensorDtype::Q8_2f {
11326 split_act_q8_2f(x2, col_field)
11327 } else {
11328 split_act(x2)
11329 };
11330 let p1 = SendMut(o1.as_mut_ptr());
11331 let p2 = SendMut(o2.as_mut_ptr());
11332 let run_range = |start: usize, end: usize| {
11333 q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11334 };
11335 match pool {
11336 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11337 _ => run_range(0, rows),
11338 }
11339 return;
11340 }
11341 #[cfg(target_arch = "x86_64")]
11342 if avx2_a8w8_enabled() {
11343 let a1s = if dtype == TensorDtype::Q8_2f {
11344 split_act_q8_2f(x1, col_field)
11345 } else {
11346 split_act(x1)
11347 };
11348 let a2s = if dtype == TensorDtype::Q8_2f {
11349 split_act_q8_2f(x2, col_field)
11350 } else {
11351 split_act(x2)
11352 };
11353 let p1 = SendMut(o1.as_mut_ptr());
11354 let p2 = SendMut(o2.as_mut_ptr());
11355 let run_range = |start: usize, end: usize| {
11356 q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11357 };
11358 match pool {
11359 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11360 _ => run_range(0, rows),
11361 }
11362 return;
11363 }
11364
11365 prescale_with(x1, col_field, dtype, 1, |x1s| {
11366 prescale_with(x2, col_field, dtype, 2, |x2s| {
11367 let p1 = SendMut(o1.as_mut_ptr());
11368 let p2 = SendMut(o2.as_mut_ptr());
11369 let run_range = move |start: usize, end: usize| {
11370 for o in start..end {
11371 let row = &q[o * cols..(o + 1) * cols];
11372 let s1 = dot_i8_f32(row, x1s) * row_scale[o];
11373 let s2 = dot_i8_f32(row, x2s) * row_scale[o];
11374 unsafe {
11376 *p1.at(o) = s1;
11377 *p2.at(o) = s2;
11378 }
11379 }
11380 };
11381 match pool {
11382 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11383 _ => run_range(0, rows),
11384 }
11385 });
11386 });
11387}
11388
11389#[derive(Clone, Copy)]
11390struct SendMut(*mut f32);
11391unsafe impl Send for SendMut {}
11392unsafe impl Sync for SendMut {}
11393
11394impl SendMut {
11395 #[inline]
11396 fn at(self, i: usize) -> *mut f32 {
11397 unsafe { self.0.add(i) }
11398 }
11399}
11400
11401#[cfg(test)]
11402mod tests {
11403 #[test]
11407 fn q8_round_is_round_clamp() {
11408 let reference = |t: f32| t.round().clamp(-127.0, 127.0) as i8;
11409 let mut probe = vec![
11410 0.0f32,
11411 -0.0,
11412 f32::NAN,
11413 f32::INFINITY,
11414 f32::NEG_INFINITY,
11415 f32::MAX,
11416 f32::MIN,
11417 1e30,
11418 -1e30,
11419 f32::MIN_POSITIVE,
11420 -f32::MIN_POSITIVE,
11421 ];
11422 for k in -300i32..=300 {
11423 let h = k as f32 * 0.5;
11424 let up = f32::from_bits(h.to_bits() + 1);
11425 let down = f32::from_bits(h.to_bits().wrapping_sub(1));
11426 for t in [h, up, down] {
11427 probe.push(t);
11428 probe.push(-t);
11429 }
11430 }
11431 let mut t = -140.0f32;
11432 while t < 140.0 {
11433 probe.push(t);
11434 t += 0.000_731;
11435 }
11436 for t in probe {
11437 assert_eq!(q8_round(t), reference(t), "t = {t:e} ({:#x})", t.to_bits());
11438 }
11439 }
11440
11441 use super::*;
11442
11443 #[test]
11447 fn f32_matmat_equals_matvec_bitwise() {
11448 let (rows, cols, b) = (5usize, 37usize, 7usize);
11449 let w: Vec<f32> = (0..rows * cols).map(|i| ((i * 7919 % 113) as f32 - 56.0) / 37.0).collect();
11450 let xs: Vec<f32> = (0..b * cols).map(|i| ((i * 104_729 % 97) as f32 - 48.0) / 29.0).collect();
11451 let t = QTensor::from_f32(w, rows, cols);
11452 let mut out = vec![0f32; b * rows];
11453 t.matmat(&xs, b, &mut out, None);
11454 for bi in 0..b {
11455 let mut one = vec![0f32; rows];
11456 t.matvec(&xs[bi * cols..(bi + 1) * cols], &mut one, None);
11457 for o in 0..rows {
11458 assert_eq!(out[bi * rows + o].to_bits(), one[o].to_bits(), "bi {bi} row {o}");
11459 }
11460 }
11461 }
11462
11463 #[test]
11464 fn q2tp_i8_dot_matches_exact_on_grid() {
11465 let (rows, cols) = (5, 64);
11469 let gpr = cols / GROUP_SIZE;
11470 let chunks: Vec<u8> = (0..rows * gpr * Q2TP_CHUNK)
11474 .map(|i| (i as u32).wrapping_mul(2654435761) as u8)
11475 .collect();
11476 let scales: Vec<f32> = (0..gpr).map(|g| 0.5 + g as f32 * 0.25).collect();
11477 let x: Vec<f32> = (0..cols)
11478 .map(|i| if i % 3 == 0 { -1.0 } else { 1.0 })
11479 .collect();
11480 let act = split_act(&x);
11481 assert!(
11482 act.outliers.is_empty(),
11483 "on-grid input must have no outliers"
11484 );
11485 let gsum = q1_group_sums(&act.xq, gpr);
11486 for r in 0..rows {
11487 let exact = q2tp_row_exact(&chunks, r, gpr, &x, &scales);
11488 let fast = dot_q2tp_row_i8(&chunks, r, gpr, &act.xq, &gsum, &scales) * act.sx;
11489 assert!(
11490 (exact - fast).abs() <= exact.abs() * 1e-5 + 1e-5,
11491 "row {r}: exact {exact} vs i8 {fast}"
11492 );
11493 }
11494 }
11495
11496 #[cfg(target_arch = "x86_64")]
11497 #[test]
11498 fn q2tp_avx2_dot_matches_scalar_for_random_patterns() {
11499 if !std::arch::is_x86_feature_detected!("avx2") {
11504 return;
11505 }
11506 let mut seed = 0x9e3779b9u32;
11507 let mut next = || {
11508 seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
11509 seed
11510 };
11511 for _ in 0..20_000 {
11512 let mut ch = [0u8; Q2TP_CHUNK];
11513 let mut x = [0i8; GROUP_SIZE];
11514 for b in &mut ch {
11515 *b = next() as u8;
11516 }
11517 for v in &mut x {
11518 *v = (next() >> 24) as i8;
11519 }
11520 let mut reference = 0i32;
11521 for (k, &b) in ch.iter().enumerate() {
11522 reference += (b & 3) as i32 * x[k * 4] as i32;
11523 reference += ((b >> 2) & 3) as i32 * x[k * 4 + 1] as i32;
11524 reference += ((b >> 4) & 3) as i32 * x[k * 4 + 2] as i32;
11525 reference += ((b >> 6) & 3) as i32 * x[k * 4 + 3] as i32;
11526 }
11527 let got = unsafe { q2tp_code_dot_avx2(&ch, &x) };
11530 assert_eq!(got, reference, "packed q2 lane mismatch");
11531 }
11532 }
11533
11534 #[test]
11535 fn q2tp_affine_fuses_half_scale_correction_without_changing_raw_decode() {
11536 let (rows, cols) = (1usize, GROUP_SIZE);
11537 let mut bytes = vec![0u8; Q2TP_CHUNK + 4 + 1];
11538 bytes[..Q2TP_CHUNK].fill(0x24); bytes[Q2TP_CHUNK..Q2TP_CHUNK + 2].copy_from_slice(&0u16.to_le_bytes());
11542 bytes[Q2TP_CHUNK + 2..Q2TP_CHUNK + 4].copy_from_slice(&0u16.to_le_bytes());
11543 bytes[Q2TP_CHUNK + 4] = 1; let x = vec![1.0f32; cols];
11545 let mut raw = vec![0.0f32; rows];
11546 let mut affine = vec![0.0f32; rows];
11547 q2tp_matvec_for_test(&bytes, &x, rows, cols, &mut raw);
11548 q2tp_affine_matvec_for_test(&bytes, &x, rows, cols, &mut affine);
11549 assert_eq!(raw, vec![-24.0]);
11550 assert_eq!(affine, vec![-8.0]);
11551 assert!((affine[0] - (raw[0] + 0.5 * cols as f32)).abs() < 1e-6);
11552 }
11553
11554 #[test]
11555 fn q8_row_dot_fast_matches_scalar() {
11556 let cols = 96;
11559 let row: Vec<u8> = (0..cols)
11560 .map(|i| ((i * 37 % 251) - 125) as i8 as u8)
11561 .collect();
11562 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.13).sin()).collect();
11563 let act = split_act(&x);
11564 let fast = q8_row_dot(&row, &act);
11565 let scalar = q8_row_dot_scalar(&row, &act);
11566 assert!(
11567 (fast - scalar).abs() <= scalar.abs() * 1e-5 + 1e-5,
11568 "fast {fast} vs scalar {scalar}"
11569 );
11570 }
11571
11572 #[test]
11573 fn f32_matvec_matches_matvec_rows_bitexact() {
11574 let (rows, cols) = (300, 40);
11575 let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
11576 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
11577 let qt = QTensor::from_f32(w.clone(), rows, cols);
11578
11579 let mut a = vec![0.0f32; rows];
11580 matvec_rows(None, &w, &x, &mut a);
11581 let mut b = vec![0.0f32; rows];
11582 qt.matvec(&x, &mut b, None);
11583 assert_eq!(a, b);
11584 }
11585
11586 #[test]
11587 fn sdot_kernel_exact_on_grid() {
11588 eprintln!("sdot_enabled = {}", sdot_enabled());
11593 let (rows, cols) = (9, 80); let w: Vec<u8> = (0..rows * cols)
11595 .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11596 .collect();
11597 let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
11598 let x: Vec<f32> = (0..cols)
11599 .map(|i| match i % 3 {
11600 0 => 1.0,
11601 1 => -1.0,
11602 _ => 0.0,
11603 })
11604 .collect();
11605 let mut a = vec![0.0f32; rows];
11606 qmatvec(
11607 &w,
11608 &[],
11609 &scales,
11610 &x,
11611 &[],
11612 TensorDtype::Q8Row,
11613 rows,
11614 cols,
11615 &mut a,
11616 None,
11617 );
11618 for o in 0..rows {
11619 let mut acc = 0.0f32;
11620 for j in 0..cols {
11621 acc += (w[o * cols + j] as i8) as f32 * x[j];
11622 }
11623 let expect = acc * scales[o];
11624 assert!(
11625 (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11626 "row {o}: {} vs {expect}",
11627 a[o]
11628 );
11629 }
11630 }
11631
11632 #[test]
11633 fn q1_tbl_fast_path_matches_reference() {
11634 let (rows, cols) = (5, 256);
11639 let gpr = cols / GROUP_SIZE;
11640 let mut bytes = Vec::new();
11641 for t in 0..rows * gpr {
11642 let s = 0.007 + (t % 11) as f32 * 0.004;
11643 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11644 for j in 0..4 {
11645 bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
11646 }
11647 }
11648 let x: Vec<f32> = (0..cols)
11649 .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
11650 .collect();
11651 let mut w = vec![0.0f32; rows * cols];
11652 cortiq_core::quant::dequant_q1(&bytes, &mut w);
11653 let mut got = vec![0.0f32; rows];
11654 q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11655 for o in 0..rows {
11656 let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11657 assert!(
11658 (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11659 "row {o}: {} vs {expect}",
11660 got[o]
11661 );
11662 }
11663 let b = 5usize;
11666 let mut xs_all = Vec::new();
11667 for bi in 0..b {
11668 xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
11669 }
11670 let mut mm = vec![0.0f32; b * rows];
11671 q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
11672 for bi in 0..b {
11673 let mut single = vec![0.0f32; rows];
11674 q1_matvec(
11675 &bytes,
11676 &xs_all[bi * cols..(bi + 1) * cols],
11677 rows,
11678 cols,
11679 &mut single,
11680 None,
11681 );
11682 assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
11683 }
11684 }
11685
11686 #[test]
11687 fn q1_kernels_match_exact_reference() {
11688 let (rows, cols) = (7, 96);
11690 let gpr = cols / GROUP_SIZE;
11691 let mut bytes = Vec::new();
11692 for t in 0..rows * gpr {
11693 let s = 0.01 + (t % 13) as f32 * 0.003;
11694 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11695 for j in 0..4 {
11696 bytes.push(((t * 31 + j * 97) % 251) as u8);
11697 }
11698 }
11699 let x: Vec<f32> = (0..cols)
11701 .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
11702 .collect();
11703 let mut w = vec![0.0f32; rows * cols];
11705 cortiq_core::quant::dequant_q1(&bytes, &mut w);
11706 let mut expect = vec![0.0f32; rows];
11707 for o in 0..rows {
11708 expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11709 }
11710 let mut got = vec![0.0f32; rows];
11711 q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11712 for o in 0..rows {
11713 assert!(
11714 (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
11715 "row {o}: {} vs {}",
11716 got[o],
11717 expect[o]
11718 );
11719 }
11720 let x2: Vec<f32> = x.iter().map(|v| -v).collect();
11722 let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
11723 q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
11724 assert_eq!(a1, got);
11725 let mut xs = x.clone();
11726 xs.extend_from_slice(&x2);
11727 let mut mm = vec![0.0f32; 2 * rows];
11728 q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
11729 assert_eq!(&mm[..rows], got.as_slice());
11730 assert_eq!(&mm[rows..], a2.as_slice());
11731 }
11732
11733 #[test]
11734 fn repack_is_bit_identical() {
11735 let (rows, cols) = (267, 96); let w: Vec<u8> = (0..rows * cols)
11741 .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
11742 .collect();
11743 let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
11744 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
11745 let rep = q8_repack_layout(&w, rows, cols);
11746 for g in 0..rows / 4 {
11748 for c in 0..cols / 16 {
11749 for lane in 0..4 {
11750 assert_eq!(
11751 &rep[g * 4 * cols + c * 64 + lane * 16
11752 ..g * 4 * cols + c * 64 + lane * 16 + 16],
11753 &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
11754 );
11755 }
11756 }
11757 }
11758 let mut a = vec![0.0f32; rows];
11759 qmatvec(
11760 &w,
11761 &[],
11762 &scales,
11763 &x,
11764 &[],
11765 TensorDtype::Q8Row,
11766 rows,
11767 cols,
11768 &mut a,
11769 None,
11770 );
11771 let mut b = vec![0.0f32; rows];
11772 qmatvec(
11773 &w,
11774 &rep,
11775 &scales,
11776 &x,
11777 &[],
11778 TensorDtype::Q8Row,
11779 rows,
11780 cols,
11781 &mut b,
11782 None,
11783 );
11784 assert_eq!(a, b, "full-range repack output diverged");
11785
11786 #[cfg(target_arch = "aarch64")]
11787 if sdot_enabled() {
11788 let act = split_act(&x);
11790 let mut c1 = vec![0.0f32; rows];
11791 let mut c2 = vec![0.0f32; rows];
11792 q8_range_sdot(
11793 &w,
11794 &[],
11795 &scales,
11796 &act,
11797 cols,
11798 SendMut(c1.as_mut_ptr()),
11799 3,
11800 rows - 2,
11801 );
11802 q8_range_sdot(
11803 &w,
11804 &rep,
11805 &scales,
11806 &act,
11807 cols,
11808 SendMut(c2.as_mut_ptr()),
11809 3,
11810 rows - 2,
11811 );
11812 assert_eq!(c1, c2, "unaligned-range repack output diverged");
11813 }
11814 }
11815
11816 #[test]
11817 fn sdot_a8w8_noise_is_bounded() {
11818 let (rows, cols) = (16, 512);
11822 let w: Vec<u8> = (0..rows * cols)
11823 .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11824 .collect();
11825 let scales = vec![0.01f32; rows];
11826 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
11827 let mut a = vec![0.0f32; rows];
11828 qmatvec(
11829 &w,
11830 &[],
11831 &scales,
11832 &x,
11833 &[],
11834 TensorDtype::Q8Row,
11835 rows,
11836 cols,
11837 &mut a,
11838 None,
11839 );
11840 let (mut num, mut den) = (0f64, 0f64);
11841 for o in 0..rows {
11842 let mut acc = 0.0f32;
11843 for j in 0..cols {
11844 acc += (w[o * cols + j] as i8) as f32 * x[j];
11845 }
11846 let expect = acc * scales[o];
11847 num += ((a[o] - expect) as f64).powi(2);
11848 den += (expect as f64).powi(2);
11849 }
11850 let rel = (num / den.max(1e-12)).sqrt();
11851 assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
11852 }
11853
11854 #[test]
11855 fn i8_dot_neon_matches_scalar() {
11856 let n = 100;
11857 let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
11858 let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
11859 let mut scalar = 0.0f32;
11860 for j in 0..n {
11861 scalar += (w[j] as i8) as f32 * x[j];
11862 }
11863 let fast = dot_i8_f32(&w, &x);
11864 assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
11865 }
11866
11867 #[test]
11869 fn vbitmatvec_matches_full_dequant() {
11870 let (rows, cols) = (6, 64);
11871 let ng = cols / GROUP_SIZE;
11872 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
11874 let mut bytes = bits.clone();
11875 for g in 0..rows * ng {
11876 let s = 0.02 + 0.001 * g as f32;
11877 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11878 }
11879 for r in 0..rows {
11880 let b = bits[r] as usize;
11881 let (mut acc, mut nb) = (0u64, 0usize);
11882 let mut rowbytes = Vec::new();
11883 for i in 0..cols {
11884 let v = ((i * 7 + r * 13) % (1 << b)) as u64;
11885 acc = (acc << b) | v;
11886 nb += b;
11887 while nb >= 8 {
11888 nb -= 8;
11889 rowbytes.push(((acc >> nb) & 0xFF) as u8);
11890 }
11891 }
11892 if nb > 0 {
11893 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11894 }
11895 bytes.extend_from_slice(&rowbytes);
11896 }
11897 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
11898
11899 let mut reference = vec![0f32; rows * cols];
11900 cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
11901 let mut expect = vec![0f32; rows];
11902 for r in 0..rows {
11903 expect[r] = reference[r * cols..(r + 1) * cols]
11904 .iter()
11905 .zip(&x)
11906 .map(|(w, xv)| w * xv)
11907 .sum();
11908 }
11909 let mut got = vec![0f32; rows];
11910 let offsets = vbit_row_offsets(&bytes, rows, cols);
11911 vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
11912 let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
11916 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
11917 for r in 0..rows {
11918 assert!(
11919 (got[r] - expect[r]).abs() < tol * scale,
11920 "row {r}: {} vs {}",
11921 got[r],
11922 expect[r]
11923 );
11924 }
11925 }
11926
11927 #[test]
11932 #[cfg(target_arch = "x86_64")]
11933 fn vbit_matmat_blocked_matches_per_row() {
11934 let (rows, cols, b) = (64usize, 128usize, 9usize);
11935 let ng = cols / GROUP_SIZE;
11936 let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
11937 let mut bytes = bits.clone();
11938 for g in 0..rows * ng {
11939 let sc = 0.02 + 0.0005 * g as f32;
11940 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11941 }
11942 for r in 0..rows {
11943 let bw = bits[r] as usize;
11944 let (mut acc, mut nb) = (0u64, 0usize);
11945 let mut rowbytes = Vec::new();
11946 for i in 0..cols {
11947 let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
11948 acc = (acc << bw) | v;
11949 nb += bw;
11950 while nb >= 8 {
11951 nb -= 8;
11952 rowbytes.push(((acc >> nb) & 0xFF) as u8);
11953 }
11954 }
11955 if nb > 0 {
11956 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11957 }
11958 bytes.extend_from_slice(&rowbytes);
11959 }
11960 let x: Vec<f32> = (0..b * cols)
11961 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11962 .collect();
11963 let offsets = vbit_row_offsets(&bytes, rows, cols);
11964 let mut y_a = vec![0f32; b * rows];
11965 let mut y_b = vec![0f32; b * rows];
11966 unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11967 vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
11968 unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11969 vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
11970 unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11971 let max_d = y_a
11972 .iter()
11973 .zip(&y_b)
11974 .map(|(p, q)| (p - q).abs())
11975 .fold(0.0f32, f32::max);
11976 assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
11977 }
11978
11979 #[test]
11987 fn q4t_matmat_blocked_matches_per_row() {
11988 let (rows, cols, b) = (16usize, 64usize, 9usize);
11989 let gpr = cols / GROUP_SIZE;
11990 let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
11991 for r in 0..rows {
11992 for g in 0..gpr {
11993 let t = (r * gpr + g) * Q4_TILE;
11994 let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
11995 bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11996 for k in 0..16 {
11997 bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11998 }
11999 }
12000 }
12001 let x: Vec<f32> = (0..b * cols)
12002 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12003 .collect();
12004 let mut y_blk = vec![0f32; b * rows];
12005 let mut y_row = vec![0f32; b * rows];
12006 unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
12007 q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
12008 unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
12009 q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
12010 unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
12011 assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
12012 }
12013
12014 fn synth_q4tp(rows: usize, cols: usize) -> Vec<u8> {
12021 use cortiq_core::quant::{f32_to_f16, q4tp_code_stride, q4tp_put_code};
12022 let gpr = cols / GROUP_SIZE;
12023 let stride = q4tp_code_stride(gpr);
12024 let (params_off, codes_off, _) = q4tp_sections(rows, cols);
12025 let mut b = vec![0u8; codes_off + rows * stride];
12026 for r in 0..rows {
12027 for g in 0..gpr {
12028 let t = (r * gpr + g) * Q4TP_NIB;
12029 for k in 0..16 {
12030 b[t + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12031 }
12032 }
12033 let lo = -6.0 - 0.03 * (r % 17) as f32;
12034 let step = 0.01 + 0.004 * (r % 11) as f32;
12035 let p = params_off + r * 4;
12036 b[p..p + 2].copy_from_slice(&f32_to_f16(lo).to_le_bytes());
12037 b[p + 2..p + 4].copy_from_slice(&f32_to_f16(step).to_le_bytes());
12038 let crow = &mut b[codes_off + r * stride..codes_off + (r + 1) * stride];
12039 for g in 0..gpr {
12040 q4tp_put_code(crow, g, (r * 5 + g * 3) % 32);
12041 }
12042 }
12043 b
12044 }
12045
12046 fn q4tp_as_q4t(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
12050 let gpr = cols / GROUP_SIZE;
12051 let v = Q4tpView::new(bytes, rows, cols);
12052 let mut out = vec![0u8; rows * gpr * Q4_TILE];
12053 let mut sc = vec![0f32; gpr];
12054 for r in 0..rows {
12055 v.scales_into(r, gpr, &mut sc);
12056 for g in 0..gpr {
12057 let t = (r * gpr + g) * Q4_TILE;
12058 let s = sc[g];
12059 out[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12060 let src = (r * gpr + g) * Q4TP_NIB;
12061 out[t + 2..t + Q4_TILE].copy_from_slice(&v.nib[src..src + Q4TP_NIB]);
12062 }
12063 }
12064 out
12065 }
12066
12067 #[test]
12073 fn q4tp_exact_path_matches_dequant_reference() {
12074 let (rows, cols) = (256usize, 512usize);
12075 let gpr = cols / GROUP_SIZE;
12076 let bytes = synth_q4tp(rows, cols);
12077 let mut w = vec![0f32; rows * cols];
12078 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12079
12080 let x: Vec<f32> = (0..cols)
12081 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12082 .collect();
12083 let v = Q4tpView::new(&bytes, rows, cols);
12084 let mut sc = vec![0f32; gpr];
12085 for r in 0..rows {
12086 v.scales_into(r, gpr, &mut sc);
12087 let got = q4tp_row_exact(v.nib, r, gpr, &x, &sc);
12088 let want: f32 = (0..cols).map(|c| w[r * cols + c] * x[c]).sum();
12089 let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12093 assert!(
12094 (got - want).abs() <= 1e-5 * mag,
12095 "row {r}: kernel {got} vs dequant {want}"
12096 );
12097 }
12098 }
12099
12100 #[test]
12106 fn q4tp_matvec_matches_the_q4t_kernel_it_was_ported_from() {
12107 let (rows, cols) = (256usize, 512usize);
12108 let bytes = synth_q4tp(rows, cols);
12109 let twin = q4tp_as_q4t(&bytes, rows, cols);
12110 let x: Vec<f32> = (0..cols)
12111 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12112 .collect();
12113
12114 let mut got = vec![0f32; rows];
12115 q4tp_matvec(&bytes, &x, rows, cols, &mut got, None);
12116 let mut want = vec![0f32; rows];
12117 q4t_matvec(&twin, &x, rows, cols, &mut want, None);
12118
12119 let mut w = vec![0f32; rows * cols];
12122 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12123 for r in 0..rows {
12124 let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12125 assert!(
12126 (got[r] - want[r]).abs() <= 1e-3 * mag,
12127 "row {r}: q4tp {} vs q4t {}",
12128 got[r],
12129 want[r]
12130 );
12131 }
12132 }
12133
12134 #[test]
12139 fn q4tp_matmat_matches_the_q4t_kernel_it_was_ported_from() {
12140 let (rows, cols, b) = (256usize, 512usize, 5usize);
12141 let bytes = synth_q4tp(rows, cols);
12142 let twin = q4tp_as_q4t(&bytes, rows, cols);
12143 let xs: Vec<f32> = (0..b * cols)
12144 .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12145 .collect();
12146
12147 let mut got = vec![0f32; b * rows];
12148 q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, None);
12149 let mut want = vec![0f32; b * rows];
12150 q4t_matmat(&twin, &xs, b, rows, cols, &mut want, None);
12151
12152 let mut w = vec![0f32; rows * cols];
12153 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12154 for t in 0..b {
12155 for r in 0..rows {
12156 let mag: f32 = (0..cols)
12157 .map(|c| (w[r * cols + c] * xs[t * cols + c]).abs())
12158 .sum();
12159 let (g, wa) = (got[t * rows + r], want[t * rows + r]);
12160 assert!(
12161 (g - wa).abs() <= 1e-3 * mag,
12162 "batch {t} row {r}: q4tp {g} vs q4t {wa}"
12163 );
12164 }
12165 }
12166 }
12167
12168 #[test]
12169 fn q4tp_matvec2_matches_the_single_stream_kernel() {
12170 let (rows, cols) = (128usize, 256usize);
12171 let gpr = cols / GROUP_SIZE;
12172 let bytes = synth_q4tp(rows, cols);
12173 let xs: Vec<f32> = (0..2 * cols)
12174 .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12175 .collect();
12176
12177 let (mut o1, mut o2) = (vec![0f32; rows], vec![0f32; rows]);
12178 q4tp_matvec2(
12179 &bytes,
12180 &xs[..cols],
12181 &xs[cols..],
12182 rows,
12183 cols,
12184 &mut o1,
12185 &mut o2,
12186 None,
12187 );
12188
12189 let v = Q4tpView::new(&bytes, rows, cols);
12192 let mut sc = vec![0f32; gpr];
12193 for r in 0..rows {
12194 v.scales_into(r, gpr, &mut sc);
12195 assert_eq!(o1[r], q4tp_row_exact(v.nib, r, gpr, &xs[..cols], &sc));
12196 assert_eq!(o2[r], q4tp_row_exact(v.nib, r, gpr, &xs[cols..], &sc));
12197 }
12198 }
12199
12200 #[test]
12207 fn q4tp_matvec_keeps_pace_with_q4t() {
12208 let (rows, cols) = (4096usize, 3072usize);
12209 let bytes = synth_q4tp(rows, cols);
12210 let twin = q4tp_as_q4t(&bytes, rows, cols);
12211 let x: Vec<f32> = (0..cols).map(|i| (i % 97) as f32 / 97.0 - 0.5).collect();
12212 let mut o = vec![0f32; rows];
12213 let n = 12;
12214 let mut best = (f64::MAX, f64::MAX);
12215 for _ in 0..3 {
12218 let t0 = std::time::Instant::now();
12219 for _ in 0..n {
12220 q4t_matvec(&twin, &x, rows, cols, &mut o, None);
12221 }
12222 best.0 = best.0.min(t0.elapsed().as_secs_f64());
12223 let t0 = std::time::Instant::now();
12224 for _ in 0..n {
12225 q4tp_matvec(&bytes, &x, rows, cols, &mut o, None);
12226 }
12227 best.1 = best.1.min(t0.elapsed().as_secs_f64());
12228 }
12229 let ratio = best.1 / best.0;
12230 println!(
12231 "q4t {:.3} ms | q4tp {:.3} ms | {ratio:.2}x",
12232 best.0 * 1e3 / n as f64,
12233 best.1 * 1e3 / n as f64
12234 );
12235 assert!(ratio < 2.0, "q4tp matvec {ratio:.2}x slower than q4t");
12236 }
12237
12238 #[cfg(target_os = "macos")]
12239 #[test]
12240 fn q4t_matmat_accel_matches_dequant_reference() {
12241 if !accel_gemm_enabled() {
12242 return; }
12244 let (rows, cols, b) = (512usize, 1024usize, 8usize); let gpr = cols / GROUP_SIZE;
12246 let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
12247 for r in 0..rows {
12248 for g in 0..gpr {
12249 let t = (r * gpr + g) * Q4_TILE;
12250 let sc = 0.02 + 0.0005 * ((r * gpr + g) % 64) as f32;
12251 bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
12252 for k in 0..16 {
12253 bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12254 }
12255 }
12256 }
12257 let x: Vec<f32> = (0..b * cols)
12258 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12259 .collect();
12260 let mut got = vec![0f32; b * rows];
12261 q4t_matmat(&bytes, &x, b, rows, cols, &mut got, None);
12262 let mut w = vec![0f32; rows * cols];
12264 for r in 0..rows {
12265 for g in 0..gpr {
12266 let t = (r * gpr + g) * Q4_TILE;
12267 let s = f16_to_f32(u16::from_le_bytes([bytes[t], bytes[t + 1]]));
12268 for (k, &bb) in bytes[t + 2..t + Q4_TILE].iter().enumerate() {
12269 w[r * cols + g * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
12270 w[r * cols + g * GROUP_SIZE + k * 2 + 1] =
12271 (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
12272 }
12273 }
12274 }
12275 for bi in 0..b {
12276 for r in 0..rows {
12277 let want: f32 = (0..cols).map(|j| x[bi * cols + j] * w[r * cols + j]).sum();
12278 let d = (got[bi * rows + r] - want).abs();
12279 assert!(
12280 d <= want.abs().max(1.0) * 1e-4,
12281 "accel q4t GEMM diverged at ({bi},{r}): {} vs {want}",
12282 got[bi * rows + r]
12283 );
12284 }
12285 }
12286 }
12287
12288 #[test]
12289 fn q4matvec_matches_full_dequant() {
12290 let (rows, cols) = (8, 64);
12291 let groups = rows * cols / GROUP_SIZE;
12292 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12294 for i in 0..groups * 16 {
12295 bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12296 }
12297 for g in 0..groups {
12298 let s = 0.01 + 0.003 * g as f32;
12299 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12300 }
12301 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12302
12303 let mut reference = vec![0.0f32; rows * cols];
12304 cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12305 let mut expect = vec![0.0f32; rows];
12306 for r in 0..rows {
12307 expect[r] = reference[r * cols..(r + 1) * cols]
12308 .iter()
12309 .zip(&x)
12310 .map(|(w, xv)| w * xv)
12311 .sum();
12312 }
12313
12314 let mut got = vec![0.0f32; rows];
12315 q4matvec(&bytes, &x, rows, cols, &mut got, None);
12316 let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
12320 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12321 for r in 0..rows {
12322 assert!(
12323 (got[r] - expect[r]).abs() < tol * scale,
12324 "row {r}: {} vs {}",
12325 got[r],
12326 expect[r]
12327 );
12328 }
12329 }
12330
12331 #[test]
12334 fn vbitmatvec2_equals_two_singles() {
12335 let (rows, cols) = (6, 64);
12336 let ng = cols / GROUP_SIZE;
12337 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
12338 let mut bytes = bits.clone();
12339 for g in 0..rows * ng {
12340 let s = 0.02 + 0.001 * g as f32;
12341 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12342 }
12343 for r in 0..rows {
12344 let b = bits[r] as usize;
12345 let (mut acc, mut nb) = (0u64, 0usize);
12346 let mut rowbytes = Vec::new();
12347 for i in 0..cols {
12348 let v = ((i * 7 + r * 13) % (1 << b)) as u64;
12349 acc = (acc << b) | v;
12350 nb += b;
12351 while nb >= 8 {
12352 nb -= 8;
12353 rowbytes.push(((acc >> nb) & 0xFF) as u8);
12354 }
12355 }
12356 if nb > 0 {
12357 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12358 }
12359 bytes.extend_from_slice(&rowbytes);
12360 }
12361 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
12362 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
12363 let offsets = vbit_row_offsets(&bytes, rows, cols);
12364
12365 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12366 vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
12367 vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
12368 let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12369 vbitmatvec2(
12370 &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
12371 );
12372 assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
12373 assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
12374 }
12375
12376 #[test]
12378 fn q4matvec2_equals_two_singles() {
12379 let (rows, cols) = (8, 128);
12380 let groups = rows * cols / GROUP_SIZE;
12381 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12382 for i in 0..groups * 16 {
12383 bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12384 }
12385 for g in 0..groups {
12386 let s = 0.01 + 0.003 * g as f32;
12387 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12388 }
12389 let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12392 x1[9] = 250.0;
12393 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12394
12395 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12396 q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
12397 q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
12398 let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12399 q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
12400 assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
12401 assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
12402 }
12403
12404 #[test]
12407 fn matvec_many_equals_separate_matvecs() {
12408 use crate::pool::Pool;
12409 let (r1, r2, cols) = (300, 200, 64);
12410 let mk = |salt: usize, rows: usize| {
12411 QTensor::from_f32(
12412 (0..rows * cols)
12413 .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
12414 .collect(),
12415 rows,
12416 cols,
12417 )
12418 };
12419 let (a, b) = (mk(1, r1), mk(5, r2));
12420 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
12421 let pool = Pool::new(3);
12422
12423 let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
12424 a.matvec(&x, &mut ea, Some(&pool));
12425 b.matvec(&x, &mut eb, Some(&pool));
12426 let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
12427 QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12428 assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
12429 assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
12430 }
12431
12432 #[test]
12437 fn q4tp_matvec_many_equals_separate_matvecs() {
12438 use crate::pool::Pool;
12439 use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12440
12441 let (r1, r2, cols) = (300usize, 200usize, 64usize);
12442 let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12443 "arch_name": "tiny-q4tp",
12444 "hidden_size": cols,
12445 "intermediate_size": cols * 2,
12446 "num_layers": 1,
12447 "num_attention_heads": 2,
12448 "num_kv_heads": 1,
12449 "head_dim": 32,
12450 "vocab_size": r1,
12451 "layer_types": ["FullAttention"],
12452 "rms_norm_eps": 1e-6,
12453 "max_position_embeddings": 8,
12454 "linear_conv_kernel_dim": 0,
12455 "linear_num_key_heads": 0,
12456 "linear_num_value_heads": 0
12457 }))
12458 .unwrap();
12459 let header = CmfHeader {
12460 format: "cmf".into(),
12461 version: CMF_VERSION,
12462 arch,
12463 quant_type: QuantType::Q4Block,
12464 provenance: None,
12465 tokenizer_config: None,
12466 section_hashes: None,
12467 skills: Vec::new(),
12468 shard: None,
12469 calibration: None,
12470 routing: None,
12471 genome: None,
12472 lineage: Vec::new(),
12473 router: None,
12474 segments: Vec::new(),
12475 };
12476 let specs = [
12477 TensorSpec {
12478 name: "q".into(),
12479 dtype: TensorDtype::Q4TiledP,
12480 shape: vec![r1, cols],
12481 data: synth_q4tp(r1, cols),
12482 },
12483 TensorSpec {
12484 name: "kv".into(),
12485 dtype: TensorDtype::Q4TiledP,
12486 shape: vec![r2, cols],
12487 data: synth_q4tp(r2, cols),
12488 },
12489 ];
12490 let dir = std::env::temp_dir().join(format!("cmf-q4tp-many-{}", std::process::id()));
12491 std::fs::create_dir_all(&dir).unwrap();
12492 let path = dir.join("m.cmf");
12493 CmfModel::write(&path, &header, &specs, None, None).unwrap();
12494 let model = Arc::new(CmfModel::open(&path).unwrap());
12495 let (a, b) = (
12496 QTensor::from_model(&model, "q").unwrap(),
12497 QTensor::from_model(&model, "kv").unwrap(),
12498 );
12499 assert_eq!(a.model_dtype(), Some(TensorDtype::Q4TiledP));
12500 assert_eq!(b.model_dtype(), Some(TensorDtype::Q4TiledP));
12501 let x: Vec<f32> = (0..cols)
12502 .map(|i| ((i * 17 + 3) % 97) as f32 / 97.0 - 0.5)
12503 .collect();
12504 let pool = Pool::new(3);
12505 let (mut ea, mut eb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12506 a.matvec(&x, &mut ea, Some(&pool));
12507 b.matvec(&x, &mut eb, Some(&pool));
12508 let (mut ga, mut gb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12509 QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12510 assert_eq!(ea, ga, "Q4TP fused lane 1 must be bit-identical");
12511 assert_eq!(eb, gb, "Q4TP fused lane 2 must be bit-identical");
12512 let _ = std::fs::remove_dir_all(&dir);
12513 }
12514
12515 #[test]
12521 fn multi_token_moe_rows_equal_single_token_decode() {
12522 use crate::pool::Pool;
12523 use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12524
12525 let (h, inter, ne) = (64usize, 128usize, 3usize);
12526 let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12527 "arch_name": "tiny-q4tp-moe",
12528 "hidden_size": h,
12529 "intermediate_size": inter,
12530 "num_layers": 1,
12531 "num_attention_heads": 2,
12532 "num_kv_heads": 1,
12533 "head_dim": 32,
12534 "vocab_size": 8,
12535 "layer_types": ["FullAttention"],
12536 "rms_norm_eps": 1e-6,
12537 "max_position_embeddings": 8,
12538 "linear_conv_kernel_dim": 0,
12539 "linear_num_key_heads": 0,
12540 "linear_num_value_heads": 0
12541 }))
12542 .unwrap();
12543 let header = CmfHeader {
12544 format: "cmf".into(),
12545 version: CMF_VERSION,
12546 arch,
12547 quant_type: QuantType::Q4Block,
12548 provenance: None,
12549 tokenizer_config: None,
12550 section_hashes: None,
12551 skills: Vec::new(),
12552 shard: None,
12553 calibration: None,
12554 routing: None,
12555 genome: None,
12556 lineage: Vec::new(),
12557 router: None,
12558 segments: Vec::new(),
12559 };
12560 let mut specs = Vec::new();
12561 for e in 0..ne {
12562 for (k, (n, r, c)) in [("g", inter, h), ("u", inter, h), ("d", h, inter)]
12563 .into_iter()
12564 .enumerate()
12565 {
12566 let mut data = synth_q4tp(r, c);
12569 for (i, byte) in data[..r * (c / GROUP_SIZE) * Q4TP_NIB]
12570 .iter_mut()
12571 .enumerate()
12572 {
12573 *byte ^= ((i * (e * 3 + k + 1)) % 251) as u8;
12574 }
12575 specs.push(TensorSpec {
12576 name: format!("{n}{e}"),
12577 dtype: TensorDtype::Q4TiledP,
12578 shape: vec![r, c],
12579 data,
12580 });
12581 }
12582 }
12583 let dir = std::env::temp_dir().join(format!(
12584 "cmf-moe-rows-{}-{}",
12585 std::process::id(),
12586 FLOAT_ACTIVATIONS.get()
12587 ));
12588 std::fs::create_dir_all(&dir).unwrap();
12589 let path = dir.join("m.cmf");
12590 CmfModel::write(&path, &header, &specs, None, None).unwrap();
12591 let model = Arc::new(CmfModel::open(&path).unwrap());
12592 let t = |n: String| QTensor::from_model(&model, &n).unwrap();
12593 let g: Vec<QTensor> = (0..ne).map(|e| t(format!("g{e}"))).collect();
12594 let u: Vec<QTensor> = (0..ne).map(|e| t(format!("u{e}"))).collect();
12595 let d: Vec<QTensor> = (0..ne).map(|e| t(format!("d{e}"))).collect();
12596 let b = 4usize;
12597 let mut xs: Vec<f32> = (0..b * h)
12598 .map(|i| ((i * 31 + 7) % 89) as f32 / 89.0 - 0.5)
12599 .collect();
12600 xs[5] = 9.0; let routes: Vec<(Vec<usize>, Vec<f32>)> = vec![
12603 (vec![2, 0], vec![0.6, 0.4]),
12604 (vec![0, 1, 2], vec![0.2, 0.5, 0.3]),
12605 (vec![1], vec![1.0]),
12606 (vec![2, 1, 0], vec![0.25, 0.25, 0.5]),
12607 ];
12608 let pool = Pool::new(3);
12609 let mut want = vec![0f32; b * h];
12611 for (tk, (idx, w)) in routes.iter().enumerate() {
12612 let x = &xs[tk * h..(tk + 1) * h];
12613 let pairs: Vec<(&QTensor, &QTensor)> = idx.iter().map(|&e| (&g[e], &u[e])).collect();
12614 let mut gs: Vec<Vec<f32>> = idx.iter().map(|_| vec![0f32; inter]).collect();
12615 assert!(QTensor::moe_gate_up_many(&pairs, x, &mut gs, Some(&pool)));
12616 if FLOAT_ACTIVATIONS.get() {
12617 for (slot, &e) in idx.iter().enumerate() {
12618 let (mut gate, mut up) = (vec![0.0; inter], vec![0.0; inter]);
12619 g[e].matvec(x, &mut gate, Some(&pool));
12620 u[e].matvec(x, &mut up, Some(&pool));
12621 for (v, u) in gate.iter_mut().zip(up) {
12622 *v = (*v / (1.0 + (-*v).exp())) * u;
12623 }
12624 assert_eq!(gs[slot], gate, "float gate/up must equal ordinary matvecs");
12625 }
12626 }
12627 let downs: Vec<&QTensor> = idx.iter().map(|&e| &d[e]).collect();
12628 assert!(QTensor::moe_down_many(
12629 &downs,
12630 &gs,
12631 w,
12632 &mut want[tk * h..(tk + 1) * h],
12633 Some(&pool)
12634 ));
12635 }
12636 if FLOAT_ACTIVATIONS.get() {
12637 for (tk, (idx, w)) in routes.iter().enumerate() {
12638 let mut scalar = vec![0.0; h];
12639 for (&e, &weight) in idx.iter().zip(w) {
12640 let (mut gate, mut up, mut down) =
12641 (vec![0.0; inter], vec![0.0; inter], vec![0.0; h]);
12642 g[e].matvec(&xs[tk * h..(tk + 1) * h], &mut gate, Some(&pool));
12643 u[e].matvec(&xs[tk * h..(tk + 1) * h], &mut up, Some(&pool));
12644 for (v, u) in gate.iter_mut().zip(up) {
12645 *v = (*v / (1.0 + (-*v).exp())) * u;
12646 }
12647 d[e].matvec(&gate, &mut down, Some(&pool));
12648 for (v, d) in scalar.iter_mut().zip(down) {
12649 *v += weight * d;
12650 }
12651 }
12652 assert_eq!(
12653 &want[tk * h..(tk + 1) * h],
12654 scalar,
12655 "float many equals scalar experts"
12656 );
12657 }
12658 }
12659 let mut experts: Vec<usize> = Vec::new();
12661 let mut groups: Vec<Vec<usize>> = Vec::new();
12662 for (tk, (idx, _)) in routes.iter().enumerate() {
12663 for &e in idx {
12664 match experts.iter().position(|&x| x == e) {
12665 Some(k) => groups[k].push(tk),
12666 None => {
12667 experts.push(e);
12668 groups.push(vec![tk]);
12669 }
12670 }
12671 }
12672 }
12673 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
12674 let pairs: Vec<(&QTensor, &QTensor)> = experts.iter().map(|&e| (&g[e], &u[e])).collect();
12675 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
12676 assert!(QTensor::moe_gate_up_rows(
12677 &pairs,
12678 &groups,
12679 &xs,
12680 &mut gs,
12681 Some(&pool)
12682 ));
12683 let downs: Vec<&QTensor> = experts.iter().map(|&e| &d[e]).collect();
12684 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
12685 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; h]).collect();
12686 assert!(QTensor::moe_down_rows(
12687 &downs,
12688 &lens,
12689 &gs,
12690 &mut ds,
12691 Some(&pool)
12692 ));
12693 let slot = |tk: usize, e: usize| {
12694 let k = experts.iter().position(|&x| x == e).unwrap();
12695 groups[..k].iter().map(|g| g.len()).sum::<usize>()
12696 + groups[k].iter().position(|&x| x == tk).unwrap()
12697 };
12698 let mut got = vec![0f32; b * h];
12699 for (tk, (idx, w)) in routes.iter().enumerate() {
12700 for i in 0..h {
12701 let mut acc = 0f32;
12702 for (&e, &we) in idx.iter().zip(w) {
12703 acc += we * ds[slot(tk, e)][i];
12704 }
12705 got[tk * h + i] = acc;
12706 }
12707 }
12708 assert!(want.iter().any(|v| *v != 0.0));
12709 assert_eq!(
12710 want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12711 got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12712 "multi-token MoE must equal decode bit for bit"
12713 );
12714
12715 let b5 = 5usize;
12718 let x5: Vec<f32> = (0..b5 * h)
12719 .map(|i| ((i * 13 + 5) % 71) as f32 / 71.0 - 0.5)
12720 .collect();
12721 let mut mm = vec![0f32; b5 * inter];
12722 row_exact_scope(|| g[1].matmat(&x5, b5, &mut mm, Some(&pool)));
12723 for tk in 0..b5 {
12724 let mut mv = vec![0f32; inter];
12725 g[1].matvec(&x5[tk * h..(tk + 1) * h], &mut mv, Some(&pool));
12726 assert_eq!(
12727 mv.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12728 mm[tk * inter..(tk + 1) * inter]
12729 .iter()
12730 .map(|v| v.to_bits())
12731 .collect::<Vec<_>>(),
12732 "row-exact matmat token {tk}"
12733 );
12734 }
12735 let _ = std::fs::remove_dir_all(&dir);
12738 }
12739
12740 #[test]
12748 fn q4tp_matmat_fast_path_unchanged_outside_row_exact() {
12749 use crate::pool::Pool;
12750 use std::sync::atomic::Ordering::Relaxed;
12751 let _alt = Q4TP_ALT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
12752 Q4TP_ALT.store(2, Relaxed);
12755 let pool = Pool::new(3);
12756 for &(rows, cols, b) in &[(64usize, 256usize, 7usize), (320, 1024, 9)] {
12760 let bytes = synth_q4tp(rows, cols);
12761 let mut xs: Vec<f32> = (0..b * cols)
12762 .map(|i| ((i * 29 + 11) % 83) as f32 / 83.0 - 0.5)
12763 .collect();
12764 xs[3] = 7.5; let run = |exact: bool| {
12766 let mut out = vec![0f32; b * rows];
12767 q4tp_matmat_with(&bytes, &xs, b, rows, cols, &mut out, Some(&pool), exact);
12768 out
12769 };
12770 let (fast, exact) = (run(false), run(true));
12771 #[cfg(not(target_arch = "aarch64"))]
12772 let _ = fast;
12773 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
12774 let mut matvecs = vec![0f32; b * rows];
12775 for (bi, o) in matvecs.chunks_mut(rows).enumerate() {
12776 q4tp_matvec(
12777 &bytes,
12778 &xs[bi * cols..(bi + 1) * cols],
12779 rows,
12780 cols,
12781 o,
12782 Some(&pool),
12783 );
12784 }
12785 assert!(matvecs.iter().any(|v| *v != 0.0));
12786 assert_eq!(
12787 bits(&exact),
12788 bits(&matvecs),
12789 "{rows}x{cols} b={b}: row-exact matmat must equal per-token matvecs"
12790 );
12791 #[cfg(target_arch = "aarch64")]
12792 {
12793 let gpr = cols / GROUP_SIZE;
12795 let v = Q4tpView::new(&bytes, rows, cols);
12796 let mut old = vec![0f32; b * rows];
12797 let mut sc = vec![0f32; gpr];
12798 let a8w8 = a8w8_enabled();
12799 let blocked = sdot_enabled() && blocked_enabled();
12800 let acts: Vec<SplitAct> = (0..b)
12801 .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
12802 .collect();
12803 for r in 0..rows {
12804 v.scales_into(r, gpr, &mut sc);
12805 if !a8w8 {
12806 for bi in 0..b {
12807 let x = &xs[bi * cols..(bi + 1) * cols];
12808 old[bi * rows + r] = q4tp_row_exact(v.nib, r, gpr, x, &sc);
12809 }
12810 continue;
12811 }
12812 let finish = |d: f32, act: &SplitAct| {
12813 let mut acc = d * act.sx;
12814 for &(j, xv) in &act.outliers {
12815 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
12816 acc += w * s * xv;
12817 }
12818 acc
12819 };
12820 let mut bi = 0usize;
12821 while blocked && bi + 4 <= b {
12822 let xs4 = [
12823 acts[bi].xq.as_slice(),
12824 acts[bi + 1].xq.as_slice(),
12825 acts[bi + 2].xq.as_slice(),
12826 acts[bi + 3].xq.as_slice(),
12827 ];
12828 let d = unsafe { dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs4, &sc) };
12829 for k in 0..4 {
12830 old[(bi + k) * rows + r] = finish(d[k], &acts[bi + k]);
12831 }
12832 bi += 4;
12833 }
12834 for (bi, act) in acts.iter().enumerate().skip(bi) {
12835 let d = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc);
12836 old[bi * rows + r] = finish(d, act);
12837 }
12838 }
12839 assert_eq!(
12840 bits(&fast),
12841 bits(&old),
12842 "{rows}x{cols} b={b}: the fast path changed outside row_exact"
12843 );
12844 if blocked {
12847 assert_ne!(
12848 bits(&fast),
12849 bits(&matvecs),
12850 "{rows}x{cols} b={b}: the tuned tile no longer runs outside row_exact"
12851 );
12852 }
12853 }
12854 }
12855 Q4TP_ALT.store(0, Relaxed);
12856 }
12857
12858 #[test]
12859 fn row_exact_scopes_survive_overlap_nesting_and_unwind() {
12860 use std::sync::{Barrier, atomic::{AtomicUsize, Ordering}};
12861 let active = AtomicUsize::new(0);
12864 counted_row_exact_scope(&active, || {
12865 assert_eq!(active.load(Ordering::Acquire), 1);
12866 counted_row_exact_scope(&active, || {
12867 assert_eq!(active.load(Ordering::Acquire), 2);
12868 });
12869 assert_eq!(active.load(Ordering::Acquire), 1);
12870 });
12871 assert_eq!(active.load(Ordering::Acquire), 0);
12872
12873 let both_entered = Barrier::new(2);
12874 let release_last = Barrier::new(2);
12875 std::thread::scope(|s| {
12876 let first = s.spawn(|| counted_row_exact_scope(&active, || {
12877 both_entered.wait();
12878 }));
12879 let last = s.spawn(|| counted_row_exact_scope(&active, || {
12880 both_entered.wait();
12881 release_last.wait();
12882 }));
12883 first.join().unwrap();
12884 let after_first = active.load(Ordering::Acquire);
12885 release_last.wait();
12886 last.join().unwrap();
12887 assert_eq!(after_first, 1, "second request must remain exact");
12888 });
12889 assert_eq!(active.load(Ordering::Acquire), 0);
12890 let panic = std::panic::catch_unwind(|| {
12891 counted_row_exact_scope(&active, || panic!("scope unwind"));
12892 });
12893 assert!(panic.is_err());
12894 assert_eq!(active.load(Ordering::Acquire), 0);
12895 }
12896
12897 #[test]
12898 #[cfg(target_arch = "x86_64")]
12899 fn q4tp_float_avx2_is_bitwise_scalar() {
12900 if !avx2_enabled() {
12901 return;
12902 }
12903 for cols in [32, 64, 96, 2048, 4096] {
12904 let rows = 9;
12905 let bytes = synth_q4tp(rows, cols);
12906 let v = Q4tpView::new(&bytes, rows, cols);
12907 let gpr = cols / GROUP_SIZE;
12908 let mut sc = vec![0.0; gpr];
12909 for seed in 1..=5 {
12910 let xs: Vec<f32> = (0..cols)
12911 .map(|i| (((i * 104729 + seed * 8191) % 100003) as f32 - 50001.0) / 7919.0)
12912 .collect();
12913 for r in 0..rows {
12914 v.scales_into(r, gpr, &mut sc);
12915 let scalar = q4tp_row_float_scalar(v.nib, r, gpr, &xs, &sc);
12916 let vector = unsafe { q4tp_row_float_avx2(v.nib, r, gpr, &xs, &sc) };
12917 assert_eq!(
12918 scalar.to_bits(),
12919 vector.to_bits(),
12920 "cols={cols} row={r} seed={seed}"
12921 );
12922 }
12923 }
12924 }
12925 }
12926
12927 #[test]
12928 fn multi_token_moe_rows_float_equal_single_token_decode() {
12929 float_activations_scope(multi_token_moe_rows_equal_single_token_decode);
12930 }
12931
12932 #[test]
12933 fn full_gpu_q8_scope_is_nested_and_thread_local() {
12934 assert!(!FULL_GPU_Q8.get());
12935 let before = gpu_split_frac();
12936 {
12937 let _guard = enter_full_gpu_q8_scope();
12938 assert_eq!(gpu_split_frac(), 1.0);
12939 {
12940 let _nested = enter_full_gpu_q8_scope();
12941 }
12942 assert_eq!(gpu_split_frac(), 1.0);
12943 std::thread::spawn(|| assert!(!FULL_GPU_Q8.get())).join().unwrap();
12944 }
12945 assert!(!FULL_GPU_Q8.get());
12946 assert_eq!(gpu_split_frac(), before);
12947 }
12948
12949 #[test]
12950 fn float_activation_scope_is_nested_thread_local_and_unwind_safe() {
12951 assert!(!FLOAT_ACTIVATIONS.get());
12952 let before = a8w8_enabled();
12953 float_activations_scope(|| {
12954 assert!(!a8w8_enabled());
12955 float_activations_scope(|| assert!(!a8w8_enabled()));
12956 assert!(FLOAT_ACTIVATIONS.get());
12957 std::thread::spawn(|| assert!(!FLOAT_ACTIVATIONS.get()))
12958 .join()
12959 .unwrap();
12960 });
12961 assert!(!FLOAT_ACTIVATIONS.get());
12962 assert_eq!(a8w8_enabled(), before);
12963 let _ = std::panic::catch_unwind(|| float_activations_scope(|| panic!("test unwind")));
12964 assert!(!FLOAT_ACTIVATIONS.get());
12965 }
12966
12967 #[test]
12970 fn batched_matmat_equals_per_position_matvec() {
12971 let (rows, cols, b) = (8, 64, 5);
12972 let groups = rows * cols / GROUP_SIZE;
12974 let mut q4 = Vec::new();
12975 for i in 0..groups * 16 {
12976 q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12977 }
12978 for g in 0..groups {
12979 q4.extend_from_slice(
12980 &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12981 );
12982 }
12983 let ng = cols / GROUP_SIZE;
12985 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
12986 let mut vb = bits.clone();
12987 for g in 0..rows * ng {
12988 vb.extend_from_slice(
12989 &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
12990 );
12991 }
12992 for r in 0..rows {
12993 let bw = bits[r] as usize;
12994 let (mut acc, mut nb) = (0u64, 0usize);
12995 let mut rowbytes = Vec::new();
12996 for i in 0..cols {
12997 let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
12998 acc = (acc << bw) | v;
12999 nb += bw;
13000 while nb >= 8 {
13001 nb -= 8;
13002 rowbytes.push(((acc >> nb) & 0xFF) as u8);
13003 }
13004 }
13005 if nb > 0 {
13006 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
13007 }
13008 vb.extend_from_slice(&rowbytes);
13009 }
13010 let offsets = vbit_row_offsets(&vb, rows, cols);
13011
13012 let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
13013
13014 let mut got = vec![0f32; b * rows];
13016 q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
13017 for bi in 0..b {
13018 let mut expect = vec![0f32; rows];
13019 q4matvec(
13020 &q4,
13021 &xs[bi * cols..(bi + 1) * cols],
13022 rows,
13023 cols,
13024 &mut expect,
13025 None,
13026 );
13027 assert_eq!(
13028 &got[bi * rows..(bi + 1) * rows],
13029 &expect[..],
13030 "q4 batch pos {bi}"
13031 );
13032 }
13033
13034 let mut got = vec![0f32; b * rows];
13036 vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
13037 for bi in 0..b {
13038 let mut expect = vec![0f32; rows];
13039 vbitmatvec(
13040 &vb,
13041 &offsets,
13042 &xs[bi * cols..(bi + 1) * cols],
13043 rows,
13044 cols,
13045 &mut expect,
13046 None,
13047 );
13048 assert_eq!(
13049 &got[bi * rows..(bi + 1) * rows],
13050 &expect[..],
13051 "vbit batch pos {bi}"
13052 );
13053 }
13054 }
13055
13056 #[test]
13060 fn q4_tiled_matches_q4_block_bitexact() {
13061 let (rows, cols, b) = (8usize, 128usize, 3usize);
13062 let groups = rows * cols / GROUP_SIZE;
13063 let mut split = Vec::with_capacity(groups * 18);
13064 for i in 0..groups * 16 {
13065 split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
13066 }
13067 for g in 0..groups {
13068 split.extend_from_slice(
13069 &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
13070 );
13071 }
13072 let (packed, scales) = split.split_at(groups * 16);
13074 let mut tiled = Vec::with_capacity(groups * Q4_TILE);
13075 for g in 0..groups {
13076 tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
13077 tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
13078 }
13079
13080 let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
13081 x1[9] = 250.0; let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
13083
13084 let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
13085 q4matvec(&split, &x1, rows, cols, &mut a, None);
13086 q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
13087 assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
13088
13089 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
13090 let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
13091 q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
13092 q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
13093 assert_eq!(a1, t1);
13094 assert_eq!(a2, t2);
13095
13096 let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
13097 let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
13098 q4matmat(&split, &xs, b, rows, cols, &mut am, None);
13099 q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
13100 assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
13101 }
13102
13103 #[test]
13110 fn q4matvec_sdot_outlier_exact() {
13111 let (rows, cols) = (4, 128);
13112 let groups = rows * cols / GROUP_SIZE;
13113 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
13114 for i in 0..groups * 16 {
13115 bytes.push(((i * 11 + 5) % 256) as u8);
13116 }
13117 for g in 0..groups {
13118 let s = 0.02 + 0.002 * g as f32;
13119 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
13120 }
13121 let mut x: Vec<f32> = (0..cols)
13122 .map(|i| match i % 3 {
13123 0 => 1.0,
13124 1 => -1.0,
13125 _ => 0.0,
13126 })
13127 .collect();
13128 x[17] = 300.0; let mut reference = vec![0.0f32; rows * cols];
13131 cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
13132 let mut expect = vec![0.0f32; rows];
13133 for r in 0..rows {
13134 expect[r] = reference[r * cols..(r + 1) * cols]
13135 .iter()
13136 .zip(&x)
13137 .map(|(w, xv)| w * xv)
13138 .sum();
13139 }
13140 let mut got = vec![0.0f32; rows];
13141 q4matvec(&bytes, &x, rows, cols, &mut got, None);
13142 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
13143 for r in 0..rows {
13144 assert!(
13145 (got[r] - expect[r]).abs() < 2e-3 * scale,
13146 "row {r}: {} vs {} (outlier term must be exact)",
13147 got[r],
13148 expect[r]
13149 );
13150 }
13151 }
13152
13153 #[test]
13157 fn q1t_matvec_matches_reference() {
13158 use cortiq_core::quant::{dequant_q1t, f32_to_f16};
13159 let (rows, cols) = (3usize, 64usize); let gpr = cols / GROUP_SIZE;
13161 let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
13162 let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
13164 let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
13165 let mut bytes = Vec::new();
13166 for r in 0..rows {
13167 for g in 0..gpr {
13168 bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
13169 let mut c = [0u8; 7];
13170 for k in 0..GROUP_SIZE {
13171 let code = if is_out(r * cols + g * GROUP_SIZE + k) {
13173 0
13174 } else {
13175 ((k + r * 3 + g) % 3) as u8 };
13177 cortiq_core::quant::q1t_pack(&mut c, k, code);
13178 }
13179 bytes.extend_from_slice(&c);
13180 }
13181 }
13182 let mut row_ptr = vec![0u32; rows + 1];
13185 for &(idx, _) in &outliers {
13186 row_ptr[idx as usize / cols + 1] += 1;
13187 }
13188 for r in 0..rows {
13189 row_ptr[r + 1] += row_ptr[r];
13190 }
13191 for &p in &row_ptr {
13192 bytes.extend_from_slice(&p.to_le_bytes());
13193 }
13194 for &(idx, v) in &outliers {
13195 bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
13196 bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
13197 }
13198
13199 let mut refw = vec![0f32; rows * cols];
13200 dequant_q1t(&bytes, rows, cols, &mut refw);
13201 let x: Vec<f32> = (0..cols)
13204 .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13205 .collect();
13206 let mut expect = vec![0f32; rows];
13207 for r in 0..rows {
13208 let mut a = 0.0f32;
13209 for j in 0..cols {
13210 a += refw[r * cols + j] * x[j];
13211 }
13212 expect[r] = a;
13213 }
13214 let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
13215 let mut got = vec![0f32; rows];
13216 q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
13217 for r in 0..rows {
13218 assert!(
13219 (got[r] - expect[r]).abs() < tol(expect[r]),
13220 "row {r}: {} vs {}",
13221 got[r],
13222 expect[r]
13223 );
13224 }
13225 let x2: Vec<f32> = x.iter().chain(x.iter()).copied().collect();
13227 let mut gm = vec![0f32; 2 * rows];
13228 q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
13229 for r in 0..rows {
13230 assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
13231 assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
13232 }
13233 let xb: Vec<f32> = (0..cols)
13237 .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
13238 .collect();
13239 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13240 q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
13241 q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
13242 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13243 q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
13244 assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
13245 assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
13246 }
13247
13248 #[test]
13251 fn q1t_matvec2_odd_gpr_matches_singles() {
13252 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13253 let (rows, cols) = (5usize, 96usize); let gpr = cols / GROUP_SIZE;
13255 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13256 for r in 0..rows {
13257 for g in 0..gpr {
13258 bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
13259 let mut c = [0u8; 7];
13260 for k in 0..GROUP_SIZE {
13261 q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
13262 }
13263 bytes.extend_from_slice(&c);
13264 }
13265 }
13266 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13267 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13268 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13269 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13270 q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13271 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13272 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13273 assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
13274 assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
13275 }
13276
13277 #[test]
13281 #[ignore]
13282 fn q1t_matvec2_speed() {
13283 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13284 use std::time::Instant;
13285 let (rows, cols) = (8192usize, 4096usize);
13286 let gpr = cols / GROUP_SIZE;
13287 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13288 for r in 0..rows {
13289 for g in 0..gpr {
13290 let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13291 bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13292 let mut c = [0u8; 7];
13293 for k in 0..GROUP_SIZE {
13294 q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13295 }
13296 bytes.extend_from_slice(&c);
13297 }
13298 }
13299 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13300 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13301 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13302 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13303 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13305 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13306 let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
13307 for _ in 0..8 {
13308 let t0 = Instant::now();
13309 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13310 t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
13311 let t1 = Instant::now();
13312 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13313 q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13314 t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
13315 }
13316 assert_eq!(p1, s1);
13317 assert_eq!(p2, s2);
13318 println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
13319 }
13320
13321 #[test]
13325 #[ignore]
13326 fn q1t_matvec_speed() {
13327 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
13328 use std::time::Instant;
13329 let (rows, cols) = (8192usize, 4096usize); let gpr = cols / GROUP_SIZE;
13331 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
13332 for r in 0..rows {
13333 for g in 0..gpr {
13334 let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13335 bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13336 let mut c = [0u8; 7];
13337 for k in 0..GROUP_SIZE {
13338 q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13339 }
13340 bytes.extend_from_slice(&c);
13341 }
13342 }
13343 let (n, stride) = (rows * cols, 40usize); let mut row_ptr = vec![0u32; rows + 1];
13345 let mut idx = 0usize;
13346 while idx < n {
13347 row_ptr[idx / cols + 1] += 1;
13348 idx += stride;
13349 }
13350 for r in 0..rows {
13351 row_ptr[r + 1] += row_ptr[r];
13352 }
13353 for &p in &row_ptr {
13354 bytes.extend_from_slice(&p.to_le_bytes());
13355 }
13356 let mut idx = 0usize;
13357 while idx < n {
13358 bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
13359 bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
13360 idx += stride;
13361 }
13362 let x: Vec<f32> = (0..cols)
13365 .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13366 .collect();
13367 let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
13368
13369 let slow = |out: &mut [f32]| {
13371 let mut buf = vec![0f32; cols];
13372 for r in 0..rows {
13373 for g in 0..gpr {
13374 let off = (r * gpr + g) * Q1T_TILE;
13375 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
13376 let codes = &bytes[off + 2..off + Q1T_TILE];
13377 for k in 0..GROUP_SIZE {
13378 buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
13379 1 => s,
13380 2 => -s,
13381 _ => 0.0,
13382 };
13383 }
13384 }
13385 out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
13386 + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
13387 }
13388 };
13389 let iters = 5;
13390 let mut a = vec![0f32; rows];
13391 slow(&mut a); let t = Instant::now();
13393 for _ in 0..iters {
13394 slow(&mut a);
13395 }
13396 let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13397
13398 let mut b = vec![0f32; rows];
13399 q1t_matvec(&bytes, &x, rows, cols, &mut b, None); let t = Instant::now();
13401 for _ in 0..iters {
13402 q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
13403 }
13404 let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13405
13406 for r in 0..rows {
13407 assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
13408 }
13409 println!(
13410 "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms fused-LUT {fast_ms:.2} ms => {:.2}x",
13411 slow_ms / fast_ms
13412 );
13413 }
13414}
13415
13416#[cfg(test)]
13417mod gemm_bench {
13418 #[test]
13429 #[ignore]
13430 fn q4tp_matmat_throughput() {
13431 let _alt = super::Q4TP_ALT_TEST_LOCK
13432 .lock()
13433 .unwrap_or_else(|e| e.into_inner());
13434 let b: usize = std::env::var("CMF_BENCH_B")
13438 .ok()
13439 .and_then(|v| v.parse().ok())
13440 .unwrap_or(296);
13441 let (rows, cols) = (9216usize, 2304usize);
13442 let (_, _, _) = (rows, cols, b);
13443 let total =
13444 cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4TiledP, &[rows, cols])
13445 .unwrap();
13446 let (params_off, codes_off, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13451 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13452 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13453 let step = cortiq_core::quant::f32_to_f16(0.1);
13454 for r in 0..rows {
13455 let o = params_off + r * 4;
13456 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13457 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13458 }
13459 let _ = codes_off;
13460 let xs: Vec<f32> = (0..b * cols)
13461 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13462 .collect();
13463 let mut out = vec![0f32; b * rows];
13464 let pool = crate::pool::Pool::from_env();
13465 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13471 let reps: usize = std::env::var("CMF_BENCH_REPS")
13472 .ok()
13473 .and_then(|v| v.parse().ok())
13474 .unwrap_or(10);
13475 let mut best = [f64::MAX; 2];
13476 let mut sums = [0f32; 2];
13477 for _ in 0..reps {
13478 for (k, w) in [(0usize, 1u8), (1usize, 2u8)] {
13479 super::Q4TP_ALT.store(w, std::sync::atomic::Ordering::Relaxed);
13480 let t = std::time::Instant::now();
13481 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13482 best[k] = best[k].min(t.elapsed().as_secs_f64());
13483 sums[k] = out.iter().take(64).sum::<f32>();
13484 }
13485 }
13486 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13487 for (k, name) in ["previous", "tuned "].iter().enumerate() {
13488 println!(
13489 "q4tp matmat {rows}x{cols} b={b} {name}: {:.1} ms {:.1} GFLOP/s (checksum {:.3})",
13490 best[k] * 1e3,
13491 flops / best[k] / 1e9,
13492 sums[k]
13493 );
13494 }
13495 assert!(
13496 (sums[0] - sums[1]).abs() < 1e-2,
13497 "the tuned kernel changed the result: {} vs {}",
13498 sums[0],
13499 sums[1]
13500 );
13501 }
13502
13503 #[test]
13509 fn q4tp_matmat_blocked_matches_scalar() {
13510 use std::sync::atomic::Ordering::Relaxed;
13511 let _alt = super::Q4TP_ALT_TEST_LOCK
13512 .lock()
13513 .unwrap_or_else(|e| e.into_inner());
13514 for &(rows, cols, b) in &[
13522 (64usize, 128usize, 7usize),
13523 (33, 96, 4),
13524 (16, 256, 9),
13525 (192, 2304, 37),
13526 ] {
13527 let total = cortiq_core::quant::expected_nbytes(
13528 cortiq_core::TensorDtype::Q4TiledP,
13529 &[rows, cols],
13530 )
13531 .unwrap();
13532 let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13533 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 61 % 251) as u8).collect();
13534 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13535 let step = cortiq_core::quant::f32_to_f16(0.1);
13536 for r in 0..rows {
13537 let o = params_off + r * 4;
13538 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13539 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13540 }
13541 let xs: Vec<f32> = (0..b * cols)
13542 .map(|i| ((i % 89) as f32 - 44.0) / 44.0)
13543 .collect();
13544 let mut got = vec![0f32; b * rows];
13545 let mut want = vec![0f32; b * rows];
13546 let gpr = cols / 32;
13547 let view = super::Q4tpView::new(&bytes, rows, cols);
13548 let pool = crate::pool::Pool::from_env();
13549 super::Q4TP_ALT.store(2, Relaxed);
13550 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, pool.as_deref());
13551 super::Q4TP_ALT.store(1, Relaxed);
13552 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut want, pool.as_deref());
13553 super::Q4TP_ALT.store(0, Relaxed);
13554 let scale = want.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
13561 let (mut worst, mut at) = (0f32, 0usize);
13562 for (i, (g, w)) in got.iter().zip(&want).enumerate() {
13563 if (g - w).abs() > worst {
13564 worst = (g - w).abs();
13565 at = i;
13566 }
13567 }
13568 assert!(
13569 worst <= 1e-4 * scale,
13570 "{rows}x{cols} b={b}: blocked and scalar disagree by {worst:.3e} \
13571 (scale {scale:.3e}) at cell {at}: {} vs {}",
13572 got[at],
13573 want[at]
13574 );
13575
13576 let (mut e_blocked, mut e_scalar) = (0f64, 0f64);
13583 for bi in 0..b {
13584 let act = super::split_act(&xs[bi * cols..(bi + 1) * cols]);
13585 for r in 0..rows {
13586 let mut sc = vec![0f32; gpr];
13587 view.scales_into(r, gpr, &mut sc);
13588 let mut exact = 0f64;
13589 for j in 0..cols {
13590 let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13591 exact += w as f64 * sq as f64 * act.xq[j] as f64;
13592 }
13593 exact *= act.sx as f64;
13594 for &(j, xv) in &act.outliers {
13595 let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13596 exact += w as f64 * sq as f64 * xv as f64;
13597 }
13598 let i = bi * rows + r;
13599 e_blocked = e_blocked.max((got[i] as f64 - exact).abs());
13600 e_scalar = e_scalar.max((want[i] as f64 - exact).abs());
13601 }
13602 }
13603 println!(
13604 "{rows}x{cols} b={b}: worst error vs f64 — blocked {e_blocked:.3e}, \
13605 per-column {e_scalar:.3e}"
13606 );
13607 assert!(
13612 e_blocked <= 1e-5 * scale as f64 && e_scalar <= 1e-5 * scale as f64,
13613 "{rows}x{cols} b={b}: error against f64 too large — blocked \
13614 {e_blocked:.3e}, per-column {e_scalar:.3e}, scale {scale:.3e}"
13615 );
13616 }
13617 }
13618
13619 #[test]
13625 #[ignore]
13626 fn q4tp_matmat_row_exact_cost() {
13627 let _alt = super::Q4TP_ALT_TEST_LOCK
13628 .lock()
13629 .unwrap_or_else(|e| e.into_inner());
13630 let pool = crate::pool::Pool::from_env();
13631 let reps: usize = std::env::var("CMF_BENCH_REPS")
13632 .ok()
13633 .and_then(|v| v.parse().ok())
13634 .unwrap_or(30);
13635 for &(rows, cols) in &[(2048usize, 4096usize), (4096, 2048), (4096, 4096)] {
13636 let total = cortiq_core::quant::expected_nbytes(
13637 cortiq_core::TensorDtype::Q4TiledP,
13638 &[rows, cols],
13639 )
13640 .unwrap();
13641 let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13642 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13643 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13644 let step = cortiq_core::quant::f32_to_f16(0.1);
13645 for r in 0..rows {
13646 let o = params_off + r * 4;
13647 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13648 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13649 }
13650 for &b in &[2usize, 4, 5, 8] {
13651 let xs: Vec<f32> = (0..b * cols)
13652 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13653 .collect();
13654 let mut out = vec![0f32; b * rows];
13655 let mut best = [f64::MAX; 3];
13656 for _ in 0..reps {
13657 for (k, best_k) in best.iter_mut().enumerate() {
13658 let t = std::time::Instant::now();
13659 match k {
13660 0 | 1 => super::q4tp_matmat_with(
13661 &bytes,
13662 &xs,
13663 b,
13664 rows,
13665 cols,
13666 &mut out,
13667 pool.as_deref(),
13668 k == 1,
13669 ),
13670 _ => {
13671 for (bi, o) in out.chunks_mut(rows).enumerate() {
13672 super::q4tp_matvec(
13673 &bytes,
13674 &xs[bi * cols..(bi + 1) * cols],
13675 rows,
13676 cols,
13677 o,
13678 pool.as_deref(),
13679 );
13680 }
13681 }
13682 }
13683 *best_k = best_k.min(t.elapsed().as_secs_f64());
13684 }
13685 }
13686 println!(
13687 "q4tp {rows}x{cols} b={b}: fast {:.3} ms, row-exact {:.3} ms, \
13688 {b} matvecs {:.3} ms",
13689 best[0] * 1e3,
13690 best[1] * 1e3,
13691 best[2] * 1e3
13692 );
13693 }
13694 }
13695 }
13696
13697 #[test]
13701 #[ignore]
13702 fn q4t_matmat_throughput() {
13703 let (rows, cols, b) = (9216usize, 2304usize, 296usize);
13704 let total =
13705 cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4Tiled, &[rows, cols])
13706 .unwrap();
13707 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13710 let sc = cortiq_core::quant::f32_to_f16(0.02);
13711 for t in bytes.chunks_mut(super::Q4_TILE) {
13712 t[..2].copy_from_slice(&sc.to_le_bytes());
13713 }
13714 let xs: Vec<f32> = (0..b * cols)
13715 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13716 .collect();
13717 let mut out = vec![0f32; b * rows];
13718 let pool = crate::pool::Pool::from_env();
13719 super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13720 let reps: usize = std::env::var("CMF_BENCH_REPS")
13721 .ok()
13722 .and_then(|v| v.parse().ok())
13723 .unwrap_or(10);
13724 let mut best = f64::MAX;
13725 for _ in 0..reps {
13726 let t = std::time::Instant::now();
13727 super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13728 best = best.min(t.elapsed().as_secs_f64());
13729 }
13730 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13731 println!(
13732 "q4t matmat {rows}x{cols} b={b}: {:.1} ms {:.1} GFLOP/s (checksum {:.3})",
13733 best * 1e3,
13734 flops / best / 1e9,
13735 out.iter().take(64).sum::<f32>()
13736 );
13737 }
13738}