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::sync::Arc;
21
22pub enum QTensor {
23 F32 {
24 data: Vec<f32>,
25 rows: usize,
26 cols: usize,
27 },
28 Mapped {
29 model: Arc<CmfModel>,
30 idx: usize,
32 dtype: TensorDtype,
33 rows: usize,
34 cols: usize,
35 row_scale: Vec<f32>,
37 col_field: Vec<f32>,
39 vbit_offsets: Vec<usize>,
43 repack: Vec<u8>,
51 },
52}
53
54fn repack_enabled() -> bool {
61 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
62 *ON.get_or_init(|| {
63 std::env::var("CMF_REPACK")
64 .map(|v| v == "1")
65 .unwrap_or(cfg!(target_os = "android"))
66 })
67}
68
69fn q8_repack(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
73 #[cfg(target_arch = "aarch64")]
74 let arch_ok = sdot_enabled();
75 #[cfg(not(target_arch = "aarch64"))]
76 let arch_ok = false;
77 if !arch_ok || !repack_enabled() || rows < 256 || cols % 16 != 0 {
78 return Vec::new();
79 }
80 q8_repack_layout(bytes, rows, cols)
81}
82
83fn q8_repack_layout(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
86 let groups = rows / 4;
87 let mut rep = vec![0u8; groups * 4 * cols];
88 for g in 0..groups {
89 let dst = &mut rep[g * 4 * cols..(g + 1) * 4 * cols];
90 for c in 0..cols / 16 {
91 for lane in 0..4 {
92 let src = (g * 4 + lane) * cols + c * 16;
93 dst[c * 64 + lane * 16..c * 64 + lane * 16 + 16]
94 .copy_from_slice(&bytes[src..src + 16]);
95 }
96 }
97 }
98 rep
99}
100
101fn vbit_row_offsets(bytes: &[u8], rows: usize, cols: usize) -> Vec<usize> {
104 let ng = cols / GROUP_SIZE;
105 let bits = &bytes[..rows];
106 let mut offsets = Vec::with_capacity(rows + 1);
107 let mut off = rows + rows * ng * 2;
108 for r in 0..rows {
109 offsets.push(off);
110 off += (cols * bits[r] as usize).div_ceil(8);
111 }
112 offsets.push(off);
113 offsets
114}
115
116fn blocked_enabled() -> bool {
122 use std::sync::atomic::Ordering::Relaxed;
123 match BLOCKED_OVERRIDE.load(Relaxed) {
124 1 => false,
125 2 => true,
126 _ => {
127 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
128 *ON.get_or_init(|| {
129 std::env::var("CMF_X86_BLOCKED")
130 .map(|v| v != "0")
131 .unwrap_or(true)
132 })
133 }
134 }
135}
136
137static BLOCKED_OVERRIDE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
138
139pub fn set_blocked_override(on: Option<bool>) {
146 let v = match on {
147 None => 0,
148 Some(false) => 1,
149 Some(true) => 2,
150 };
151 BLOCKED_OVERRIDE.store(v, std::sync::atomic::Ordering::Relaxed);
152}
153
154fn gpu_lmhead_enabled() -> bool {
155 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
156 *ON.get_or_init(|| {
157 std::env::var("CMF_GPU_LMHEAD")
158 .map(|v| v != "0")
159 .unwrap_or(true)
160 })
161}
162
163fn gpu_split_frac() -> f32 {
164 if FULL_GPU_Q8.get() {
167 return 1.0;
168 }
169 static F: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
170 *F.get_or_init(|| {
171 std::env::var("CMF_GPU_SPLIT")
172 .ok()
173 .and_then(|v| v.parse::<f32>().ok())
174 .unwrap_or(0.5)
175 .clamp(0.0, 1.0)
176 })
177}
178
179impl QTensor {
180 pub fn from_f32(data: Vec<f32>, rows: usize, cols: usize) -> Self {
181 debug_assert_eq!(data.len(), rows * cols);
182 Self::F32 { data, rows, cols }
183 }
184
185 pub fn from_model(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
188 let idx = model
191 .tensor_index(name)
192 .ok_or_else(|| format!("tensor '{name}' not found in CMF directory"))?;
193 let entry = &model.tensors[idx];
194 if entry.shape.len() != 2 {
195 return Err(format!("QTensor::from_model needs 2-D, got '{name}'"));
196 }
197 let (rows, cols) = (entry.shape[0], entry.shape[1]);
198 let bytes = model.entry_bytes(entry);
199
200 match entry.dtype {
201 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
202 let n = rows * cols;
203 let scales_off = n;
204 let row_scale: Vec<f32> = (0..rows)
205 .map(|o| {
206 f16_to_f32(u16::from_le_bytes([
207 bytes[scales_off + o * 2],
208 bytes[scales_off + o * 2 + 1],
209 ]))
210 })
211 .collect();
212 let col_field: Vec<f32> = if entry.dtype == TensorDtype::Q8_2f {
213 let col_off = n + rows * 2;
214 (0..cols)
215 .map(|i| {
216 f16_to_f32(u16::from_le_bytes([
217 bytes[col_off + i * 2],
218 bytes[col_off + i * 2 + 1],
219 ]))
220 })
221 .collect()
222 } else {
223 Vec::new()
224 };
225 Ok(Self::Mapped {
226 model: model.clone(),
227 idx,
228 dtype: entry.dtype,
229 rows,
230 cols,
231 row_scale,
232 col_field,
233 vbit_offsets: Vec::new(),
234 repack: q8_repack(bytes, rows, cols),
235 })
236 }
237 TensorDtype::Vbit if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
239 model: model.clone(),
240 idx,
241 dtype: entry.dtype,
242 rows,
243 cols,
244 row_scale: Vec::new(),
245 col_field: Vec::new(),
246 vbit_offsets: vbit_row_offsets(bytes, rows, cols),
247 repack: Vec::new(),
248 }),
249 TensorDtype::VbitRo if cols % GROUP_SIZE == 0 => {
253 let (_, off_off, packed_off) = cortiq_core::quant::vbit_ro_sections(rows, cols);
254 let offsets: Vec<usize> = (0..=rows)
255 .map(|r| packed_off + cortiq_core::quant::vbit_ro_offset(bytes, off_off, r))
256 .collect();
257 Ok(Self::Mapped {
258 model: model.clone(),
259 idx,
260 dtype: entry.dtype,
261 rows,
262 cols,
263 row_scale: Vec::new(),
264 col_field: Vec::new(),
265 vbit_offsets: offsets,
266 repack: Vec::new(),
267 })
268 }
269 TensorDtype::Q4Tiled if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
275 model: model.clone(),
276 idx,
277 dtype: entry.dtype,
278 rows,
279 cols,
280 row_scale: Vec::new(),
281 col_field: Vec::new(),
282 vbit_offsets: Vec::new(),
283 repack: Vec::new(),
284 }),
285 TensorDtype::Q4TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
288 model: model.clone(),
289 idx,
290 dtype: entry.dtype,
291 rows,
292 cols,
293 row_scale: Vec::new(),
294 col_field: Vec::new(),
295 vbit_offsets: Vec::new(),
296 repack: Vec::new(),
297 }),
298 TensorDtype::Q2TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
300 model: model.clone(),
301 idx,
302 dtype: entry.dtype,
303 rows,
304 cols,
305 row_scale: Vec::new(),
306 col_field: Vec::new(),
307 vbit_offsets: Vec::new(),
308 repack: Vec::new(),
309 }),
310 TensorDtype::Q4Block if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
311 model: model.clone(),
312 idx,
313 dtype: entry.dtype,
314 rows,
315 cols,
316 row_scale: Vec::new(),
317 col_field: Vec::new(),
318 vbit_offsets: Vec::new(),
319 repack: Vec::new(),
320 }),
321 TensorDtype::Q1 if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
323 model: model.clone(),
324 idx,
325 dtype: entry.dtype,
326 rows,
327 cols,
328 row_scale: Vec::new(),
329 col_field: Vec::new(),
330 vbit_offsets: Vec::new(),
331 repack: Vec::new(),
332 }),
333 TensorDtype::Q1T if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
337 model: model.clone(),
338 idx,
339 dtype: entry.dtype,
340 rows,
341 cols,
342 row_scale: Vec::new(),
343 col_field: Vec::new(),
344 vbit_offsets: Vec::new(),
345 repack: Vec::new(),
346 }),
347 _ => {
349 let mut data = vec![0.0f32; rows * cols];
350 cortiq_core::quant::dequant_tensor(entry, bytes, &mut data)?;
351 Ok(Self::from_f32(data, rows, cols))
352 }
353 }
354 }
355
356 pub(crate) fn is_q1(&self) -> bool {
359 matches!(
360 self,
361 Self::Mapped {
362 dtype: TensorDtype::Q1,
363 ..
364 }
365 )
366 }
367
368 pub(crate) fn f32_parts(&self) -> Option<(&[f32], usize, usize)> {
371 match self {
372 Self::F32 { data, rows, cols } => Some((data, *rows, *cols)),
373 _ => None,
374 }
375 }
376
377 pub(crate) fn q1_parts(&self) -> Option<(usize, usize, usize)> {
384 if self.has_prism_contract() {
385 return None;
386 }
387 match self {
388 #[cfg(target_os = "macos")]
389 Self::Mapped {
390 dtype: TensorDtype::Q1T,
391 ..
392 } if !crate::gpu::metal_q1t_enabled() => None,
393 Self::Mapped {
394 idx,
395 dtype:
396 TensorDtype::Q1
397 | TensorDtype::Q1T
398 | TensorDtype::Q4Block
399 | TensorDtype::Q4Tiled
400 | TensorDtype::Q4TiledP
404 | TensorDtype::Q8Row
405 | TensorDtype::Q8_2f,
406 rows,
407 cols,
408 ..
409 } => Some((*idx, *rows, *cols)),
410 _ => None,
411 }
412 }
413
414 #[cfg(target_os = "macos")]
421 pub(crate) fn metal_graph_parts(&self) -> Option<(usize, usize, usize)> {
422 if let Some((model, idx, kind, _)) = self.graph_weight_descriptor() {
423 let name = &model.tensors[idx].name;
424 let forward = kind == 9 && crate::prism::is_forward_weight(model, name);
425 let affine = kind == 9 && crate::prism::is_affine_target(model, name);
426 if forward && affine {
427 let e = model.tensors.get(idx)?;
428 return Some((idx, *e.shape.first()?, *e.shape.get(1)?));
429 }
430 }
431 self.q1_parts()
432 }
433
434 pub(crate) fn q4t_parts(&self) -> Option<(usize, usize, usize)> {
440 if self.has_prism_contract() {
441 return None;
442 }
443 match self {
444 Self::Mapped {
445 idx,
446 dtype: TensorDtype::Q4Tiled,
447 rows,
448 cols,
449 ..
450 } => Some((*idx, *rows, *cols)),
451 _ => None,
452 }
453 }
454
455 pub(crate) fn q4tp_parts(&self) -> Option<(usize, usize, usize)> {
459 if self.has_prism_contract() {
460 return None;
461 }
462 match self {
463 Self::Mapped {
464 idx,
465 dtype: TensorDtype::Q4TiledP,
466 rows,
467 cols,
468 ..
469 } => Some((*idx, *rows, *cols)),
470 _ => None,
471 }
472 }
473
474 pub(crate) fn q8_row_parts(&self) -> Option<(usize, usize, usize, &[f32])> {
479 if self.has_prism_contract() {
480 return None;
481 }
482 match self {
483 Self::Mapped {
484 idx,
485 dtype: TensorDtype::Q8Row,
486 rows,
487 cols,
488 row_scale,
489 col_field,
490 ..
491 } if col_field.is_empty() => Some((*idx, *rows, *cols, row_scale)),
492 _ => None,
493 }
494 }
495
496 pub fn model_dtype(&self) -> Option<cortiq_core::TensorDtype> {
500 match self {
501 Self::Mapped { dtype, .. } => Some(*dtype),
502 _ => None,
503 }
504 }
505
506 pub fn model_idx(&self) -> Option<usize> {
511 match self {
512 Self::Mapped { idx, .. } => Some(*idx),
513 _ => None,
514 }
515 }
516
517 pub fn model_arc(&self) -> Option<std::sync::Arc<cortiq_core::CmfModel>> {
522 match self {
523 Self::Mapped { model, .. } => Some(model.clone()),
524 _ => None,
525 }
526 }
527
528 pub(crate) fn has_prism_contract(&self) -> bool {
533 matches!(self, Self::Mapped { model, .. } if crate::prism::has_contract(model))
534 }
535
536 pub fn rows(&self) -> usize {
537 match self {
538 Self::F32 { rows, .. } | Self::Mapped { rows, .. } => *rows,
539 }
540 }
541
542 pub(crate) fn mapped_q4t(&self) -> Option<(&Arc<CmfModel>, usize)> {
545 if self.has_prism_contract() {
546 return None;
547 }
548 match self {
549 Self::Mapped {
550 model,
551 idx,
552 dtype: TensorDtype::Q4Tiled,
553 ..
554 } => Some((model, *idx)),
555 _ => None,
556 }
557 }
558
559 pub fn mapped_q4tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
562 if self.has_prism_contract() {
563 return None;
564 }
565 match self {
566 Self::Mapped {
567 model,
568 idx,
569 dtype: TensorDtype::Q4TiledP,
570 ..
571 } => Some((model, *idx)),
572 _ => None,
573 }
574 }
575
576 pub fn mapped_device_gemm(&self) -> Option<(&Arc<CmfModel>, usize)> {
584 if self.has_prism_contract() {
585 return None;
586 }
587 match self {
588 Self::Mapped {
589 model,
590 idx,
591 dtype: TensorDtype::Q4TiledP | TensorDtype::Q8Row | TensorDtype::Q8_2f,
592 ..
593 } => Some((model, *idx)),
594 _ => None,
595 }
596 }
597
598 pub fn mapped_q2tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
601 if self.has_prism_contract() {
602 return None;
603 }
604 match self {
605 Self::Mapped {
606 model,
607 idx,
608 dtype: TensorDtype::Q2TiledP,
609 ..
610 } => Some((model, *idx)),
611 _ => None,
612 }
613 }
614
615 pub fn cols(&self) -> usize {
616 match self {
617 Self::F32 { cols, .. } | Self::Mapped { cols, .. } => *cols,
618 }
619 }
620
621 pub fn mapped_q1(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
624 if self.has_prism_contract() {
625 return None;
626 }
627 match self {
628 Self::Mapped {
629 model,
630 idx,
631 dtype: TensorDtype::Q1,
632 ..
633 } => Some((model, *idx)),
634 _ => None,
635 }
636 }
637
638 pub fn graph_weight(&self) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
648 if self.has_prism_contract() {
649 return None;
650 }
651 self.graph_weight_descriptor()
652 }
653
654 pub(crate) fn graph_weight_descriptor(
659 &self,
660 ) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
661 match self {
662 Self::Mapped {
663 model,
664 idx,
665 dtype: TensorDtype::Q8Row,
666 row_scale,
667 ..
668 } => Some((model, *idx, 0, row_scale.as_slice())),
669 Self::Mapped {
670 model,
671 idx,
672 dtype: TensorDtype::Q1,
673 ..
674 } => Some((model, *idx, 1, &[])),
675 Self::Mapped {
680 model,
681 idx,
682 dtype: TensorDtype::Q4Tiled,
683 ..
684 } => Some((model, *idx, 5, &[])),
685 Self::Mapped {
689 model,
690 idx,
691 dtype: TensorDtype::Q4TiledP,
692 ..
693 } => Some((model, *idx, 6, &[])),
694 Self::Mapped {
695 model,
696 idx,
697 dtype: TensorDtype::Q4Block,
698 ..
699 } => Some((model, *idx, 2, &[])),
700 Self::Mapped {
705 model,
706 idx,
707 dtype: TensorDtype::Q8_2f,
708 ..
709 } => Some((model, *idx, 7, &[])),
710 Self::Mapped {
711 model,
712 idx,
713 dtype: TensorDtype::Q1T,
714 ..
715 } => Some((model, *idx, 3, &[])),
716 Self::Mapped {
720 model,
721 idx,
722 dtype: TensorDtype::Q2TiledP,
723 ..
724 } => Some((model, *idx, 9, &[])),
725 _ => None,
726 }
727 }
728
729 pub fn as_f32(&self) -> Option<&[f32]> {
732 match self {
733 Self::F32 { data, .. } => Some(data),
734 Self::Mapped { .. } => None,
735 }
736 }
737
738 fn quant_bytes(&self) -> &[u8] {
739 match self {
740 Self::Mapped { model, idx, .. } => model.entry_bytes(&model.tensors[*idx]),
741 Self::F32 { .. } => unreachable!("quant_bytes on F32"),
742 }
743 }
744
745 pub fn row_f32(&self, r: usize, dst: &mut [f32]) {
747 let cols = self.cols();
748 debug_assert_eq!(dst.len(), cols);
749 match self {
750 Self::F32 { data, .. } => dst.copy_from_slice(&data[r * cols..(r + 1) * cols]),
751 Self::Mapped {
752 model,
753 idx,
754 dtype,
755 row_scale,
756 col_field,
757 vbit_offsets,
758 ..
759 } => {
760 if *dtype == TensorDtype::Q4Tiled {
761 let bytes = self.quant_bytes();
762 let gpr = cols / GROUP_SIZE;
763 for gi in 0..gpr {
764 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
765 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
766 for (k, &b) in tile[2..].iter().enumerate() {
767 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
768 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
769 }
770 }
771 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
772 crate::prism::inverse_embedding(model, dst);
773 }
774 return;
775 }
776 if *dtype == TensorDtype::Q4TiledP {
777 let bytes = self.quant_bytes();
778 let gpr = cols / GROUP_SIZE;
779 let v = Q4tpView::new(bytes, self.rows(), cols);
780 let mut sc = vec![0f32; gpr];
781 v.scales_into(r, gpr, &mut sc);
782 for gi in 0..gpr {
783 let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
784 let s = sc[gi];
785 for (k, &b) in tile.iter().enumerate() {
786 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
787 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
788 }
789 }
790 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
791 crate::prism::inverse_embedding(model, dst);
792 }
793 return;
794 }
795 if *dtype == TensorDtype::Q2TiledP {
796 let bytes = self.quant_bytes();
797 let gpr = cols / GROUP_SIZE;
798 let v = Q4tpView::new_q2(bytes, self.rows(), cols);
799 let mut sc = vec![0f32; gpr];
800 v.scales_into(r, gpr, &mut sc);
801 for gi in 0..gpr {
802 let ch =
803 &v.nib[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
804 let s = sc[gi];
805 for (k, &b) in ch.iter().enumerate() {
806 for j in 0..4 {
807 let center = if crate::prism::is_affine_target(
808 model,
809 &model.tensors[*idx].name,
810 ) {
811 1.0
812 } else {
813 1.5
814 };
815 dst[gi * GROUP_SIZE + k * 4 + j] =
816 (((b >> (2 * j)) & 3) as f32 - center) * s;
817 }
818 }
819 }
820 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
821 crate::prism::inverse_embedding(model, dst);
822 }
823 return;
824 }
825 if *dtype == TensorDtype::Q4Block {
826 let (packed, scales) = q4_split(self.quant_bytes(), self.rows(), cols);
827 let gpr = cols / GROUP_SIZE;
828 for gi in 0..gpr {
829 let g = r * gpr + gi;
830 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
831 for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
832 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
833 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
834 }
835 }
836 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
837 crate::prism::inverse_embedding(model, dst);
838 }
839 return;
840 }
841 if *dtype == TensorDtype::Q1 {
842 let bytes = self.quant_bytes();
843 let gpr = cols / GROUP_SIZE;
844 for gi in 0..gpr {
845 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
846 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
847 for (j, &b) in tile[2..].iter().enumerate() {
848 for k in 0..8 {
849 dst[gi * GROUP_SIZE + j * 8 + k] =
850 (((b >> k) & 1) as f32 * 2.0 - 1.0) * s;
851 }
852 }
853 }
854 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
855 crate::prism::inverse_embedding(model, dst);
856 }
857 return;
858 }
859 if *dtype == TensorDtype::Q1T {
860 let bytes = self.quant_bytes();
861 let gpr = cols / GROUP_SIZE;
862 let base_len = self.rows() * gpr * cortiq_core::quant::Q1T_TILE;
863 for gi in 0..gpr {
864 let off = (r * gpr + gi) * cortiq_core::quant::Q1T_TILE;
865 let s = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
866 bytes[off],
867 bytes[off + 1],
868 ]));
869 let codes = &bytes[off + 2..off + cortiq_core::quant::Q1T_TILE];
870 for k in 0..GROUP_SIZE {
871 dst[gi * GROUP_SIZE + k] = match cortiq_core::quant::q1t_code(codes, k)
872 {
873 1 => s,
874 2 => -s,
875 _ => 0.0,
876 };
877 }
878 }
879 let rows = self.rows();
881 let entries = base_len + (rows + 1) * 4;
882 if entries <= bytes.len() {
883 let ptrs = &bytes[base_len..base_len + (rows + 1) * 4];
884 let r0 = u32::from_le_bytes([
885 ptrs[r * 4],
886 ptrs[r * 4 + 1],
887 ptrs[r * 4 + 2],
888 ptrs[r * 4 + 3],
889 ]) as usize;
890 let r1 = u32::from_le_bytes([
891 ptrs[(r + 1) * 4],
892 ptrs[(r + 1) * 4 + 1],
893 ptrs[(r + 1) * 4 + 2],
894 ptrs[(r + 1) * 4 + 3],
895 ]) as usize;
896 let off = entries + r0 * 4;
897 for i in 0..r1 - r0 {
898 let item = &bytes[off + i * 4..off + i * 4 + 4];
899 let c = u16::from_le_bytes([item[0], item[1]]) as usize;
900 let v = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
901 item[2], item[3],
902 ]));
903 if c < cols {
904 dst[c] = v;
905 }
906 }
907 }
908 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
909 crate::prism::inverse_embedding(model, dst);
910 }
911 return;
912 }
913 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
914 let bytes = self.quant_bytes();
915 let rows = self.rows();
916 let ng = cols / GROUP_SIZE;
917 let bits = &bytes[..rows];
918 let sc_off = rows;
919 let off = vbit_offsets[r];
922 let b = bits[r] as usize;
923 let l = ((1usize << (b - 1)) - 1) as f32;
924 let data = &bytes[off..];
925 let (mut acc, mut nbits, mut byte_idx) = (0u64, 0usize, 0usize);
926 for (i, d) in dst.iter_mut().enumerate() {
927 while nbits < b {
928 acc = (acc << 8) | data[byte_idx] as u64;
929 byte_idx += 1;
930 nbits += 8;
931 }
932 let u = ((acc >> (nbits - b)) & ((1u64 << b) - 1)) as f32;
933 nbits -= b;
934 let so = (r * ng + i / GROUP_SIZE) * 2;
935 let sv = f16_to_f32(u16::from_le_bytes([
936 bytes[sc_off + so],
937 bytes[sc_off + so + 1],
938 ]));
939 *d = (u - l) * sv;
940 }
941 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
942 crate::prism::inverse_embedding(model, dst);
943 }
944 return;
945 }
946 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
947 let s = row_scale[r];
948 match dtype {
949 TensorDtype::Q8Row => {
950 for (d, &b) in dst.iter_mut().zip(q) {
951 *d = (b as i8) as f32 * s;
952 }
953 }
954 TensorDtype::Q8_2f => {
955 for (i, (d, &b)) in dst.iter_mut().zip(q).enumerate() {
956 *d = (b as i8) as f32 * s * col_field[i];
957 }
958 }
959 _ => unreachable!(),
960 }
961 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
962 crate::prism::inverse_embedding(model, dst);
963 }
964 }
965 }
966 }
967
968 pub fn sparse_col_ok(&self) -> bool {
973 match self {
974 Self::F32 { .. } => true,
975 Self::Mapped { dtype, .. } => {
976 matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
977 }
978 }
979 }
980
981 pub fn add_col_scaled(&self, c: usize, w: f32, out: &mut [f32]) {
985 let inter = self.cols();
986 let hidden = self.rows();
987 debug_assert_eq!(out.len(), hidden);
988 match self {
989 Self::F32 { data, .. } => {
990 for (k, o) in out.iter_mut().enumerate() {
991 *o += w * data[k * inter + c];
992 }
993 }
994 Self::Mapped {
995 dtype,
996 row_scale,
997 col_field,
998 ..
999 } => {
1000 let q = self.quant_bytes();
1001 let colf = if *dtype == TensorDtype::Q8_2f {
1002 col_field[c]
1003 } else {
1004 1.0
1005 };
1006 let wc = w * colf;
1007 for (k, o) in out.iter_mut().enumerate() {
1008 let b = q[k * inter + c] as i8 as f32;
1009 *o += wc * b * row_scale[k];
1010 }
1011 }
1012 }
1013 }
1014
1015 #[inline]
1024 pub fn prefetch_row(&self, r: usize) {
1025 let Self::Mapped { dtype, .. } = self else {
1026 return;
1027 };
1028 if !matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f) {
1029 return;
1030 }
1031 let cols = self.cols();
1032 let q = self.quant_bytes();
1033 let (a, b) = (r * cols, (r + 1) * cols);
1034 if b > q.len() {
1035 return;
1036 }
1037 let mut j = a;
1038 while j < b {
1039 unsafe { std::ptr::read_volatile(q.as_ptr().add(j)) };
1040 j += 512;
1041 }
1042 }
1043
1044 pub fn add_row_scaled(&self, r: usize, w: f32, out: &mut [f32], scratch: &mut [f32]) {
1053 let cols = self.cols();
1054 debug_assert_eq!(out.len(), cols);
1055 match self {
1056 Self::F32 { data, .. } => {
1057 let row = &data[r * cols..(r + 1) * cols];
1058 for (o, v) in out.iter_mut().zip(row) {
1059 *o += w * v;
1060 }
1061 }
1062 Self::Mapped {
1063 dtype,
1064 row_scale,
1065 col_field,
1066 ..
1067 } => match dtype {
1068 TensorDtype::Q8Row => {
1069 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1070 let ws = w * row_scale[r];
1071 let row: &[i8] =
1072 unsafe { std::slice::from_raw_parts(q.as_ptr() as *const i8, q.len()) };
1073 axpy_i8_f32(out, row, ws);
1074 }
1075 TensorDtype::Q8_2f => {
1076 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1077 let ws = w * row_scale[r];
1078 for ((o, b), c) in out.iter_mut().zip(q).zip(col_field) {
1079 *o += ws * c * (*b as i8 as f32);
1080 }
1081 }
1082 _ => {
1083 self.row_f32(r, scratch);
1084 for (o, v) in out.iter_mut().zip(scratch.iter()) {
1085 *o += w * v;
1086 }
1087 }
1088 },
1089 }
1090 }
1091
1092 pub fn row_dot(&self, r: usize, x: &[f32], scratch: &mut [f32]) -> f32 {
1096 let cols = self.cols();
1097 match self {
1098 Self::F32 { data, .. } => {
1099 let row = &data[r * cols..(r + 1) * cols];
1100 row.iter().zip(x).map(|(w, v)| w * v).sum()
1101 }
1102 Self::Mapped {
1103 model,
1104 idx,
1105 dtype,
1106 row_scale,
1107 col_field,
1108 ..
1109 } => {
1110 let prism_forward =
1111 crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1112 if prism_forward {
1113 let transformed = crate::prism::forward(model, &x[..cols]);
1114 let gpr = cols / GROUP_SIZE;
1115 match dtype {
1116 TensorDtype::Q2TiledP => {
1117 let v = Q4tpView::new_q2(self.quant_bytes(), self.rows(), cols);
1118 let mut sc = vec![0f32; gpr];
1119 v.scales_into(r, gpr, &mut sc);
1120 if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1121 return q2tp_affine_row_exact(v.nib, r, gpr, &transformed, &sc);
1122 }
1123 return q2tp_row_exact(v.nib, r, gpr, &transformed, &sc);
1124 }
1125 _ => {
1126 self.row_f32(r, scratch);
1127 return scratch.iter().zip(&transformed).map(|(w, v)| w * v).sum();
1128 }
1129 }
1130 }
1131 match dtype {
1132 TensorDtype::Q8Row => {
1133 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1134 dot_i8_f32(q, x) * row_scale[r]
1135 }
1136 TensorDtype::Q8_2f => {
1137 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1138 dot_i8_col_f32(q, x, col_field) * row_scale[r]
1139 }
1140 _ => {
1141 self.row_f32(r, scratch);
1142 scratch.iter().zip(x).map(|(w, v)| w * v).sum()
1143 }
1144 }
1145 }
1146 }
1147 }
1148
1149 pub fn matvec(&self, x: &[f32], out: &mut [f32], pool: Option<&Pool>) {
1152 match self {
1153 Self::F32 { data, .. } => matvec_rows(pool, data, x, out),
1157 Self::Mapped {
1158 model,
1159 idx,
1160 dtype,
1161 rows,
1162 cols,
1163 row_scale,
1164 col_field,
1165 vbit_offsets,
1166 repack,
1167 } => {
1168 let _ = (model, idx);
1169 assert!(
1179 out.len() >= *rows && x.len() >= *cols,
1180 "matvec {rows}x{cols}: out {} (need {rows}), x {} (need {cols})",
1181 out.len(),
1182 x.len(),
1183 );
1184 let prism_forward =
1185 crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1186 if *dtype == TensorDtype::Q2TiledP
1187 && std::env::var("CMF_Q2TP_TRACE").as_deref() == Ok("1")
1188 {
1189 use std::sync::atomic::{AtomicUsize, Ordering};
1190 static N: AtomicUsize = AtomicUsize::new(0);
1191 let n = N.fetch_add(1, Ordering::Relaxed);
1192 if n < 128 {
1193 eprintln!(
1194 "q2tp-dispatch #{n} name={} prism={} rows={} cols={} gpu={} optin={} layer={}",
1195 model.tensors[*idx].name,
1196 prism_forward,
1197 rows,
1198 cols,
1199 crate::gpu::enabled_here(),
1200 crate::gpu::q2tp_gpu_opt_in(),
1201 crate::gpu::cur_layer(),
1202 );
1203 }
1204 }
1205 if prism_forward {
1210 let transformed = crate::prism::forward(model, &x[..*cols]);
1211 match dtype {
1212 TensorDtype::Q4Block => {
1213 q4matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1214 }
1215 TensorDtype::Q4Tiled => {
1216 q4t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1217 }
1218 TensorDtype::Q4TiledP => {
1219 q4tp_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1220 }
1221 TensorDtype::Q2TiledP => {
1222 let affine =
1223 crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1224 if *rows * *cols >= 8_388_608
1225 && crate::gpu::enabled_here()
1226 && crate::gpu::q2tp_gpu_opt_in()
1227 {
1228 let gpu_ok = if affine {
1229 crate::gpu::q2tp_affine_matvec(
1230 model,
1231 *idx,
1232 &transformed,
1233 *rows,
1234 *cols,
1235 out,
1236 )
1237 } else {
1238 crate::gpu::q2tp_matvec(
1239 model,
1240 *idx,
1241 &transformed,
1242 *rows,
1243 *cols,
1244 out,
1245 )
1246 };
1247 if gpu_ok {
1248 return;
1249 }
1250 }
1251 if affine {
1252 q2tp_affine_matvec(
1253 self.quant_bytes(),
1254 &transformed,
1255 *rows,
1256 *cols,
1257 out,
1258 pool,
1259 )
1260 } else {
1261 q2tp_matvec(
1262 self.quant_bytes(),
1263 &transformed,
1264 *rows,
1265 *cols,
1266 out,
1267 pool,
1268 )
1269 }
1270 }
1271 TensorDtype::Q1 => {
1272 q1_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1273 }
1274 TensorDtype::Q1T => {
1275 q1t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1276 }
1277 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec(
1278 self.quant_bytes(),
1279 vbit_offsets,
1280 &transformed,
1281 *rows,
1282 *cols,
1283 out,
1284 pool,
1285 ),
1286 TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec(
1287 self.quant_bytes(),
1288 repack,
1289 row_scale,
1290 &transformed,
1291 col_field,
1292 *dtype,
1293 *rows,
1294 *cols,
1295 out,
1296 pool,
1297 ),
1298 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1299 }
1300 return;
1301 }
1302 if *dtype == TensorDtype::Q4Block {
1303 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1307 let t0 = std::time::Instant::now();
1308 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1309 crate::gpu::ProbeArm::Gpu => {
1310 if crate::gpu::q4b_matvec(model, *idx, x, *rows, *cols, out) {
1311 crate::gpu::probe_record(
1312 crate::gpu::OpClass::Matvec,
1313 true,
1314 t0.elapsed(),
1315 );
1316 return;
1317 }
1318 }
1319 crate::gpu::ProbeArm::CpuTimed => {
1320 q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1321 crate::gpu::probe_record(
1322 crate::gpu::OpClass::Matvec,
1323 false,
1324 t0.elapsed(),
1325 );
1326 return;
1327 }
1328 crate::gpu::ProbeArm::Cpu => {}
1329 }
1330 }
1331 q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1332 return;
1333 }
1334 if *dtype == TensorDtype::Q4Tiled {
1335 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1340 let t0 = std::time::Instant::now();
1341 let cls = crate::gpu::matvec_class(*rows, *cols);
1342 match crate::gpu::probe_arm(cls) {
1343 crate::gpu::ProbeArm::Gpu => {
1344 if crate::gpu::q4t_matvec(model, *idx, x, *rows, *cols, out) {
1345 crate::gpu::probe_record(cls, true, t0.elapsed());
1346 return;
1347 }
1348 }
1349 crate::gpu::ProbeArm::CpuTimed => {
1350 q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1351 crate::gpu::probe_record(cls, false, t0.elapsed());
1352 return;
1353 }
1354 crate::gpu::ProbeArm::Cpu => {}
1355 }
1356 }
1357 q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1358 return;
1359 }
1360 if *dtype == TensorDtype::Q4TiledP {
1361 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1367 let t0 = std::time::Instant::now();
1368 let cls = crate::gpu::matvec_class(*rows, *cols);
1369 match crate::gpu::probe_arm(cls) {
1370 crate::gpu::ProbeArm::Gpu => {
1371 if crate::gpu::q4tp_matvec(model, *idx, x, *rows, *cols, out) {
1372 crate::gpu::probe_record(cls, true, t0.elapsed());
1373 return;
1374 }
1375 }
1376 crate::gpu::ProbeArm::CpuTimed => {
1377 q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1378 crate::gpu::probe_record(cls, false, t0.elapsed());
1379 return;
1380 }
1381 crate::gpu::ProbeArm::Cpu => {}
1382 }
1383 }
1384 q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1385 return;
1386 }
1387 if *dtype == TensorDtype::Q2TiledP {
1388 q2tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1389 return;
1390 }
1391 if *dtype == TensorDtype::Q1 {
1392 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1397 let t0 = std::time::Instant::now();
1398 let arm = if crate::gpu::q1_force() {
1399 crate::gpu::ProbeArm::Gpu
1400 } else {
1401 crate::gpu::probe_arm(crate::gpu::OpClass::Matvec)
1402 };
1403 match arm {
1404 crate::gpu::ProbeArm::Gpu => {
1405 if crate::gpu::q1_matvec(model, *idx, x, *rows, *cols, out) {
1406 crate::gpu::probe_record(
1407 crate::gpu::OpClass::Matvec,
1408 true,
1409 t0.elapsed(),
1410 );
1411 return;
1412 }
1413 }
1414 crate::gpu::ProbeArm::CpuTimed => {
1415 q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1416 crate::gpu::probe_record(
1417 crate::gpu::OpClass::Matvec,
1418 false,
1419 t0.elapsed(),
1420 );
1421 return;
1422 }
1423 crate::gpu::ProbeArm::Cpu => {}
1424 }
1425 }
1426 q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1427 return;
1428 }
1429 if *dtype == TensorDtype::Q1T {
1430 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1434 let t0 = std::time::Instant::now();
1435 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1436 crate::gpu::ProbeArm::Gpu => {
1437 if crate::gpu::q1t_matvec(model, *idx, x, *rows, *cols, out) {
1438 q1t_add_overlay(self.quant_bytes(), x, *rows, *cols, out, pool);
1439 crate::gpu::probe_record(
1440 crate::gpu::OpClass::Matvec,
1441 true,
1442 t0.elapsed(),
1443 );
1444 return;
1445 }
1446 }
1447 crate::gpu::ProbeArm::CpuTimed => {
1448 q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1449 crate::gpu::probe_record(
1450 crate::gpu::OpClass::Matvec,
1451 false,
1452 t0.elapsed(),
1453 );
1454 return;
1455 }
1456 crate::gpu::ProbeArm::Cpu => {}
1457 }
1458 }
1459 q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1460 return;
1461 }
1462 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1463 vbitmatvec(self.quant_bytes(), vbit_offsets, x, *rows, *cols, out, pool);
1464 return;
1465 }
1466 let xs = prescale(x, col_field, *dtype);
1467 if *rows >= crate::gpu::min_rows()
1472 && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
1473 && gpu_lmhead_enabled()
1474 && crate::gpu::enabled_here()
1475 {
1476 let t0 = std::time::Instant::now();
1479 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1480 crate::gpu::ProbeArm::Gpu => {}
1481 crate::gpu::ProbeArm::CpuTimed => {
1482 qmatvec(
1483 self.quant_bytes(),
1484 repack,
1485 row_scale,
1486 x,
1487 col_field,
1488 *dtype,
1489 *rows,
1490 *cols,
1491 out,
1492 pool,
1493 );
1494 crate::gpu::probe_record(
1495 crate::gpu::OpClass::Matvec,
1496 false,
1497 t0.elapsed(),
1498 );
1499 return;
1500 }
1501 crate::gpu::ProbeArm::Cpu => {
1502 qmatvec(
1503 self.quant_bytes(),
1504 repack,
1505 row_scale,
1506 x,
1507 col_field,
1508 *dtype,
1509 *rows,
1510 *cols,
1511 out,
1512 pool,
1513 );
1514 return;
1515 }
1516 }
1517 let frac = gpu_split_frac();
1518 let cpu_rows = ((*rows as f32) * (1.0 - frac)) as usize;
1519 let (out_cpu, out_gpu) = out.split_at_mut(cpu_rows);
1520 let bytes = self.quant_bytes();
1521 let ok = std::thread::scope(|sc| {
1522 let g = sc.spawn(|| {
1523 crate::gpu::q8_matvec_range(
1524 model,
1525 *idx,
1526 cpu_rows,
1527 &row_scale[cpu_rows..],
1528 &xs,
1529 *rows - cpu_rows,
1530 *cols,
1531 out_gpu,
1532 )
1533 });
1534 if cpu_rows > 0 {
1535 let rep_cpu = if repack.is_empty() {
1538 &[][..]
1539 } else {
1540 &repack[..(cpu_rows / 4) * 4 * *cols]
1541 };
1542 qmatvec(
1543 &bytes[..cpu_rows * *cols],
1544 rep_cpu,
1545 &row_scale[..cpu_rows],
1546 x,
1547 col_field,
1548 *dtype,
1549 cpu_rows,
1550 *cols,
1551 out_cpu,
1552 pool,
1553 );
1554 }
1555 g.join().unwrap_or(false)
1556 });
1557 if ok {
1558 crate::gpu::probe_record(crate::gpu::OpClass::Matvec, true, t0.elapsed());
1559 return;
1560 }
1561 qmatvec(
1564 &bytes[cpu_rows * *cols..(*rows) * *cols],
1565 &[],
1566 &row_scale[cpu_rows..],
1567 x,
1568 col_field,
1569 *dtype,
1570 *rows - cpu_rows,
1571 *cols,
1572 out_gpu,
1573 pool,
1574 );
1575 return;
1576 }
1577 qmatvec(
1578 self.quant_bytes(),
1579 repack,
1580 row_scale,
1581 x,
1582 col_field,
1583 *dtype,
1584 *rows,
1585 *cols,
1586 out,
1587 pool,
1588 );
1589 }
1590 }
1591 }
1592
1593 pub fn matvec2(
1595 &self,
1596 x1: &[f32],
1597 x2: &[f32],
1598 o1: &mut [f32],
1599 o2: &mut [f32],
1600 pool: Option<&Pool>,
1601 ) {
1602 match self {
1603 Self::F32 { data, .. } => matvec_rows2(pool, data, x1, x2, o1, o2),
1604 Self::Mapped {
1605 model,
1606 idx,
1607 dtype,
1608 rows,
1609 cols,
1610 row_scale,
1611 col_field,
1612 vbit_offsets,
1613 ..
1614 } => {
1615 if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1616 let tx1 = crate::prism::forward(model, &x1[..*cols]);
1617 let tx2 = crate::prism::forward(model, &x2[..*cols]);
1618 match dtype {
1619 TensorDtype::Q4Block => {
1620 q4matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1621 }
1622 TensorDtype::Q4Tiled => {
1623 q4t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1624 }
1625 TensorDtype::Q4TiledP => {
1626 q4tp_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1627 }
1628 TensorDtype::Q2TiledP => {
1629 if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1630 q2tp_affine_matvec2(
1631 self.quant_bytes(),
1632 &tx1,
1633 &tx2,
1634 *rows,
1635 *cols,
1636 o1,
1637 o2,
1638 pool,
1639 )
1640 } else {
1641 q2tp_matvec2(
1642 self.quant_bytes(),
1643 &tx1,
1644 &tx2,
1645 *rows,
1646 *cols,
1647 o1,
1648 o2,
1649 pool,
1650 )
1651 }
1652 }
1653 TensorDtype::Q1 => {
1654 q1_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1655 }
1656 TensorDtype::Q1T => {
1657 q1t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1658 }
1659 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec2(
1660 self.quant_bytes(),
1661 vbit_offsets,
1662 &tx1,
1663 &tx2,
1664 *rows,
1665 *cols,
1666 o1,
1667 o2,
1668 pool,
1669 ),
1670 TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec2(
1671 self.quant_bytes(),
1672 row_scale,
1673 &tx1,
1674 &tx2,
1675 col_field,
1676 *dtype,
1677 *rows,
1678 *cols,
1679 o1,
1680 o2,
1681 pool,
1682 ),
1683 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1684 }
1685 return;
1686 }
1687 if *dtype == TensorDtype::Q4Block {
1688 q4matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1689 return;
1690 }
1691 if *dtype == TensorDtype::Q4Tiled {
1692 q4t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1693 return;
1694 }
1695 if *dtype == TensorDtype::Q4TiledP {
1696 q4tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1697 return;
1698 }
1699 if *dtype == TensorDtype::Q2TiledP {
1700 q2tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1701 return;
1702 }
1703 if *dtype == TensorDtype::Q1 {
1704 q1_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1705 return;
1706 }
1707 if *dtype == TensorDtype::Q1T {
1708 q1t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1714 return;
1715 }
1716 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1717 vbitmatvec2(
1718 self.quant_bytes(),
1719 vbit_offsets,
1720 x1,
1721 x2,
1722 *rows,
1723 *cols,
1724 o1,
1725 o2,
1726 pool,
1727 );
1728 return;
1729 }
1730 qmatvec2(
1731 self.quant_bytes(),
1732 row_scale,
1733 x1,
1734 x2,
1735 col_field,
1736 *dtype,
1737 *rows,
1738 *cols,
1739 o1,
1740 o2,
1741 pool,
1742 );
1743 }
1744 }
1745 }
1746}
1747
1748impl QTensor {
1749 pub fn q4tp_mapped(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
1757 if self.has_prism_contract() {
1758 return None;
1759 }
1760 match self {
1761 Self::Mapped {
1762 model, idx, dtype, ..
1763 } if *dtype == TensorDtype::Q4TiledP => Some((model, *idx)),
1764 _ => None,
1765 }
1766 }
1767
1768 pub fn matmat(&self, xs_all: &[f32], b: usize, out: &mut [f32], pool: Option<&Pool>) {
1769 let cols = self.cols();
1770 let rows = self.rows();
1771 debug_assert_eq!(xs_all.len(), b * cols);
1772 debug_assert_eq!(out.len(), b * rows);
1773 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Matmat);
1774 if crate::gptq_capture::capturing() {
1778 if let Self::Mapped { model, idx, .. } = self {
1779 crate::gptq_capture::accumulate(&model.tensors[*idx].name, xs_all, b, cols);
1780 }
1781 }
1782 match self {
1783 Self::F32 { data, .. } => {
1784 let out_addr = SendMut(out.as_mut_ptr());
1785 let run = |start: usize, end: usize| {
1786 for o in start..end {
1787 let row = &data[o * cols..(o + 1) * cols];
1788 for bi in 0..b {
1789 let x = &xs_all[bi * cols..(bi + 1) * cols];
1790 let mut acc = 0f32;
1791 for j in 0..cols {
1792 acc += row[j] * x[j];
1793 }
1794 unsafe { *out_addr.at(bi * rows + o) = acc };
1795 }
1796 }
1797 };
1798 dispatch_rows(pool, rows, &run);
1799 }
1800 Self::Mapped {
1801 model,
1802 idx,
1803 dtype,
1804 row_scale,
1805 col_field,
1806 vbit_offsets,
1807 ..
1808 } => {
1809 if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1810 let mut transformed = Vec::with_capacity(xs_all.len());
1811 for bi in 0..b {
1812 transformed.extend_from_slice(&crate::prism::forward(
1813 model,
1814 &xs_all[bi * cols..(bi + 1) * cols],
1815 ));
1816 }
1817 match dtype {
1818 TensorDtype::Q4Block => {
1819 q4matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1820 }
1821 TensorDtype::Q4Tiled => {
1822 q4t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1823 }
1824 TensorDtype::Q4TiledP => {
1825 q4tp_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1826 }
1827 TensorDtype::Q2TiledP => {
1828 let affine =
1829 crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1830 let gpu_batch_ok = if affine {
1836 b >= 2
1837 } else {
1838 b >= 32 && b * rows * cols >= 128_000_000
1839 };
1840 if gpu_batch_ok
1841 && cols % 32 == 0
1842 && crate::gpu::enabled_here()
1843 && crate::gpu::q2tp_gpu_opt_in()
1844 {
1845 let gpu_ok = if affine {
1846 crate::gpu::q2tp_affine_matmat(
1847 model,
1848 *idx,
1849 &transformed,
1850 b,
1851 rows,
1852 cols,
1853 out,
1854 )
1855 } else {
1856 crate::gpu::q2tp_matmat(
1857 model,
1858 *idx,
1859 &transformed,
1860 b,
1861 rows,
1862 cols,
1863 out,
1864 )
1865 };
1866 if gpu_ok {
1867 return;
1868 }
1869 }
1870 if affine
1874 && b == 1
1875 && cols % 32 == 0
1876 && crate::gpu::enabled_here()
1877 && crate::gpu::q2tp_gpu_opt_in()
1878 && crate::gpu::q2tp_affine_matvec(
1879 model,
1880 *idx,
1881 &transformed[..cols],
1882 rows,
1883 cols,
1884 &mut out[..rows],
1885 )
1886 {
1887 return;
1888 }
1889 if affine {
1890 q2tp_affine_matmat(
1891 self.quant_bytes(),
1892 &transformed,
1893 b,
1894 rows,
1895 cols,
1896 out,
1897 pool,
1898 )
1899 } else {
1900 q2tp_matmat(
1901 self.quant_bytes(),
1902 &transformed,
1903 b,
1904 rows,
1905 cols,
1906 out,
1907 pool,
1908 )
1909 }
1910 }
1911 TensorDtype::Q1 => {
1912 q1_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1913 }
1914 TensorDtype::Q1T => {
1915 q1t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1916 }
1917 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatmat(
1918 self.quant_bytes(),
1919 vbit_offsets,
1920 &transformed,
1921 b,
1922 rows,
1923 cols,
1924 out,
1925 pool,
1926 ),
1927 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
1928 let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1929 .map(|bi| {
1930 prescale(
1931 &transformed[bi * cols..(bi + 1) * cols],
1932 col_field,
1933 *dtype,
1934 )
1935 })
1936 .collect();
1937 qmatmat(self.quant_bytes(), row_scale, &pre, rows, cols, out, pool)
1938 }
1939 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1940 }
1941 return;
1942 }
1943 if *dtype == TensorDtype::Q4Block {
1944 q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1945 return;
1946 }
1947 if *dtype == TensorDtype::Q4TiledP {
1948 if b >= 32
1962 && b * rows * cols >= 128_000_000
1963 && cols % 32 == 0
1964 && !row_exact()
1965 && !crate::gpu::mm_killed()
1966 && crate::gpu::enabled_here()
1967 {
1968 let class = if b >= 128 {
1969 crate::gpu::OpClass::MatmatWide
1970 } else {
1971 crate::gpu::OpClass::Matmat
1972 };
1973 if let Self::Mapped { model, idx, .. } = self {
1974 if crate::mm_ab::on() {
1986 let mut g = vec![0f32; b * rows];
1987 let t = std::time::Instant::now();
1988 let took = crate::gpu::q4tp_matmat(
1989 model, *idx, xs_all, b, rows, cols, &mut g,
1990 );
1991 let dg = t.elapsed();
1992 let t = std::time::Instant::now();
1993 q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1994 let dc = t.elapsed();
1995 crate::mm_ab::record(b, rows, cols, took, dg, dc, &g, out);
1996 return;
1997 }
1998 let t0 = std::time::Instant::now();
1999 let resident = crate::gpu::weight_is_resident(model, *idx);
2003 match crate::gpu::probe_arm_cold_prefers_gpu(class, resident) {
2004 crate::gpu::ProbeArm::Gpu => {
2005 if crate::gpu::q4tp_matmat(
2006 model, *idx, xs_all, b, rows, cols, out,
2007 ) {
2008 let el = t0.elapsed();
2009 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2019 let budget = std::time::Duration::from_secs_f64(
2020 flops / 1.5e12 * 8.0 + 0.020,
2021 );
2022 crate::gpu::mm_budget_check(
2023 "q4tp matmat",
2024 el,
2025 budget,
2026 crate::gpu::probe_was_cold() || !resident,
2027 );
2028 crate::gpu::probe_record(class, true, el);
2029 return;
2030 }
2031 }
2032 crate::gpu::ProbeArm::CpuTimed => {
2033 q4tp_matmat(
2034 self.quant_bytes(),
2035 xs_all,
2036 b,
2037 rows,
2038 cols,
2039 out,
2040 pool,
2041 );
2042 crate::gpu::probe_record(class, false, t0.elapsed());
2043 return;
2044 }
2045 crate::gpu::ProbeArm::Cpu => {}
2046 }
2047 }
2048 }
2049 q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2050 return;
2051 }
2052 if *dtype == TensorDtype::Q2TiledP {
2053 if b >= 32
2059 && b * rows * cols >= 128_000_000
2060 && cols % 32 == 0
2061 && !crate::gpu::mm_killed()
2062 && crate::gpu::enabled_here()
2063 {
2064 let class = if b >= 128 {
2065 crate::gpu::OpClass::MatmatWide
2066 } else {
2067 crate::gpu::OpClass::Matmat
2068 };
2069 if let Self::Mapped { model, idx, .. } = self {
2070 let t0 = std::time::Instant::now();
2071 match crate::gpu::probe_arm(class) {
2072 crate::gpu::ProbeArm::Gpu => {
2073 if crate::gpu::q2tp_matmat(
2074 model, *idx, xs_all, b, rows, cols, out,
2075 ) {
2076 crate::gpu::probe_record(class, true, t0.elapsed());
2077 return;
2078 }
2079 }
2080 crate::gpu::ProbeArm::CpuTimed => {
2081 q2tp_matmat(
2082 self.quant_bytes(),
2083 xs_all,
2084 b,
2085 rows,
2086 cols,
2087 out,
2088 pool,
2089 );
2090 crate::gpu::probe_record(class, false, t0.elapsed());
2091 return;
2092 }
2093 crate::gpu::ProbeArm::Cpu => {}
2094 }
2095 }
2096 }
2097 q2tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2102 return;
2103 }
2104 if *dtype == TensorDtype::Q4Tiled {
2105 if b >= 32
2117 && b * rows * cols >= 128_000_000
2118 && cols % 32 == 0
2119 && !crate::gpu::mm_killed()
2120 && crate::gpu::enabled_here()
2121 {
2122 let class = if b >= 128 {
2123 crate::gpu::OpClass::MatmatWide
2124 } else {
2125 crate::gpu::OpClass::Matmat
2126 };
2127 if let Self::Mapped { model, idx, .. } = self {
2128 let t0 = std::time::Instant::now();
2129 match crate::gpu::probe_arm(class) {
2130 crate::gpu::ProbeArm::Gpu => {
2131 if crate::gpu::q4t_matmat(
2132 model, *idx, xs_all, b, rows, cols, out,
2133 ) {
2134 let el = t0.elapsed();
2135 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2145 let budget = std::time::Duration::from_secs_f64(
2146 flops / 1.5e12 * 8.0 + 0.020,
2147 );
2148 crate::gpu::mm_budget_check(
2149 "q4t matmat",
2150 el,
2151 budget,
2152 crate::gpu::probe_was_cold(),
2153 );
2154 crate::gpu::probe_record(class, true, el);
2155 return;
2156 }
2157 }
2158 crate::gpu::ProbeArm::CpuTimed => {
2159 q4t_matmat(
2160 self.quant_bytes(),
2161 xs_all,
2162 b,
2163 rows,
2164 cols,
2165 out,
2166 pool,
2167 );
2168 crate::gpu::probe_record(class, false, t0.elapsed());
2169 return;
2170 }
2171 crate::gpu::ProbeArm::Cpu => {}
2172 }
2173 }
2174 }
2175 q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2176 return;
2177 }
2178 if *dtype == TensorDtype::Q1 {
2179 if b >= 32
2182 && b * rows * cols >= 128_000_000
2183 && cols % 64 == 0
2184 && crate::gpu::enabled_here()
2185 {
2186 if let Self::Mapped { model, idx, .. } = self {
2187 let t0 = std::time::Instant::now();
2188 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2189 crate::gpu::ProbeArm::Gpu => {
2190 if crate::gpu::q1_matmat(
2191 model, *idx, xs_all, b, rows, cols, out,
2192 ) {
2193 crate::gpu::probe_record(
2194 crate::gpu::OpClass::Matmat,
2195 true,
2196 t0.elapsed(),
2197 );
2198 return;
2199 }
2200 }
2201 crate::gpu::ProbeArm::CpuTimed => {
2202 q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2203 crate::gpu::probe_record(
2204 crate::gpu::OpClass::Matmat,
2205 false,
2206 t0.elapsed(),
2207 );
2208 return;
2209 }
2210 crate::gpu::ProbeArm::Cpu => {}
2211 }
2212 }
2213 }
2214 q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2215 return;
2216 }
2217 if *dtype == TensorDtype::Q1T {
2218 if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
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::q1t_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 q1t_matmat(
2238 self.quant_bytes(),
2239 xs_all,
2240 b,
2241 rows,
2242 cols,
2243 out,
2244 pool,
2245 );
2246 crate::gpu::probe_record(
2247 crate::gpu::OpClass::Matmat,
2248 false,
2249 t0.elapsed(),
2250 );
2251 return;
2252 }
2253 crate::gpu::ProbeArm::Cpu => {}
2254 }
2255 }
2256 }
2257 q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2258 return;
2259 }
2260 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
2261 vbitmatmat(
2262 self.quant_bytes(),
2263 vbit_offsets,
2264 xs_all,
2265 b,
2266 rows,
2267 cols,
2268 out,
2269 pool,
2270 );
2271 return;
2272 }
2273 let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
2274 .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
2275 .collect();
2276 if row_exact()
2282 && (1..=4).contains(&b)
2283 && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
2284 && crate::gpu::enabled_here()
2285 && crate::gpu::wgpu_active()
2286 {
2287 let flat: Vec<f32> = pre.iter().flat_map(|v| v.iter().copied()).collect();
2288 if crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out) {
2289 return;
2290 }
2291 }
2292 if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2298 if let Self::Mapped { model, idx, .. } = self {
2299 let t0 = std::time::Instant::now();
2300 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2301 crate::gpu::ProbeArm::Gpu
2302 if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
2303 && !crate::gpu::q8_resident_or_upload(model, *idx) =>
2304 {
2305 let q = self.quant_bytes();
2309 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2310 return;
2311 }
2312 crate::gpu::ProbeArm::Gpu => {
2313 let flat: Vec<f32> =
2314 pre.iter().flat_map(|v| v.iter().copied()).collect();
2315 if crate::gpu::q8_matmat(
2316 model, *idx, row_scale, &flat, b, rows, cols, out,
2317 ) {
2318 crate::gpu::probe_record(
2319 crate::gpu::OpClass::Matmat,
2320 true,
2321 t0.elapsed(),
2322 );
2323 return;
2324 }
2325 }
2326 crate::gpu::ProbeArm::CpuTimed => {
2327 let q = self.quant_bytes();
2328 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2329 crate::gpu::probe_record(
2330 crate::gpu::OpClass::Matmat,
2331 false,
2332 t0.elapsed(),
2333 );
2334 return;
2335 }
2336 crate::gpu::ProbeArm::Cpu => {}
2337 }
2338 }
2339 }
2340 let q = self.quant_bytes();
2341 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2342 }
2343 }
2344 }
2345}
2346
2347impl QTensor {
2348 pub fn device_matmat(&self, xs: &[f32], b: usize, out: &mut [f32]) -> bool {
2358 let (rows, cols) = (self.rows(), self.cols());
2359 let Self::Mapped {
2360 model,
2361 idx,
2362 dtype,
2363 row_scale,
2364 col_field,
2365 ..
2366 } = self
2367 else {
2368 return false;
2369 };
2370 if crate::prism::has_contract(model) {
2371 return false;
2372 }
2373 match *dtype {
2374 TensorDtype::Q4TiledP => crate::gpu::q4tp_matmat(model, *idx, xs, b, rows, cols, out),
2375 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
2379 if *dtype == TensorDtype::Q8_2f
2382 && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2383 && crate::gpu::q8_matmat_2f(
2384 model, *idx, row_scale, col_field, xs, b, rows, cols, out,
2385 )
2386 {
2387 return true;
2388 }
2389 let flat: Vec<f32> = (0..b)
2390 .flat_map(|bi| {
2391 prescale(&xs[bi * cols..(bi + 1) * cols], col_field, *dtype).into_owned()
2392 })
2393 .collect();
2394 crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out)
2395 }
2396 _ => false,
2397 }
2398 }
2399
2400 pub fn matvec_many<const N: usize>(
2407 ts: [&QTensor; N],
2408 x: &[f32],
2409 mut outs: [&mut [f32]; N],
2410 pool: Option<&Pool>,
2411 ) {
2412 let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2413 if ts.iter().any(|t| t.has_prism_contract()) {
2414 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2418 t.matvec(x, o, pool);
2419 }
2420 return;
2421 }
2422 let uniform_q8 = ts.iter().all(|t| {
2423 matches!(
2424 t,
2425 Self::Mapped {
2426 dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2427 ..
2428 }
2429 )
2430 });
2431 let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2432 let uniform_q4 = ts.iter().all(|t| {
2433 matches!(
2434 t,
2435 Self::Mapped {
2436 dtype: TensorDtype::Q4Block,
2437 ..
2438 }
2439 )
2440 });
2441 let uniform_vbit = ts.iter().all(|t| {
2442 matches!(
2443 t,
2444 Self::Mapped {
2445 dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2446 ..
2447 }
2448 )
2449 });
2450 let uniform_q1 = ts.iter().all(|t| {
2451 matches!(
2452 t,
2453 Self::Mapped {
2454 dtype: TensorDtype::Q1,
2455 ..
2456 }
2457 )
2458 });
2459 let uniform_q1t = ts.iter().all(|t| {
2460 matches!(
2461 t,
2462 Self::Mapped {
2463 dtype: TensorDtype::Q1T,
2464 ..
2465 }
2466 )
2467 });
2468 let uniform_q4tp = ts.iter().all(|t| {
2473 matches!(
2474 t,
2475 Self::Mapped {
2476 dtype: TensorDtype::Q4TiledP,
2477 ..
2478 }
2479 )
2480 }) && ts
2481 .iter()
2482 .all(|t| t.cols() == ts[0].cols() && t.cols() % GROUP_SIZE == 0);
2483 let Some(pool) = pool else {
2484 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2485 t.matvec(x, o, None);
2486 }
2487 return;
2488 };
2489 if total_rows < 256
2490 || !(uniform_q8
2491 || uniform_f32
2492 || uniform_q4
2493 || uniform_vbit
2494 || uniform_q1
2495 || uniform_q1t
2496 || uniform_q4tp)
2497 {
2498 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2499 t.matvec(x, o, Some(pool));
2500 }
2501 return;
2502 }
2503
2504 if uniform_q4tp {
2505 let cols = ts[0].cols();
2511 let gpr = cols / GROUP_SIZE;
2512 let views: Vec<Q4tpView> = ts
2513 .iter()
2514 .map(|t| Q4tpView::new(t.quant_bytes(), t.rows(), cols))
2515 .collect();
2516 let rows_of: Vec<usize> = ts.iter().map(|t| t.rows()).collect();
2517 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2518 let locate = |flat: usize| -> (usize, usize) {
2520 let mut acc = 0;
2521 for (i, &r) in rows_of.iter().enumerate() {
2522 if flat < acc + r {
2523 return (i, flat - acc);
2524 }
2525 acc += r;
2526 }
2527 (rows_of.len() - 1, 0)
2528 };
2529 let (views, outs_addr) = (&views, &outs_addr);
2530 if a8w8_enabled() {
2531 let act = split_act(x);
2532 let act = &act;
2533 let run = |start: usize, end: usize| {
2534 let mut sc = vec![0f32; gpr];
2535 for flat in start..end {
2536 let (t, r) = locate(flat);
2537 let v = &views[t];
2538 v.scales_into(r, gpr, &mut sc);
2539 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
2540 for &(j, xv) in &act.outliers {
2541 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
2542 acc += w * s * xv;
2543 }
2544 unsafe { *outs_addr[t].at(r) = acc };
2546 }
2547 };
2548 pool.run_rows(total_rows, &run);
2549 } else {
2550 let run = |start: usize, end: usize| {
2551 let mut sc = vec![0f32; gpr];
2552 for flat in start..end {
2553 let (t, r) = locate(flat);
2554 let v = &views[t];
2555 v.scales_into(r, gpr, &mut sc);
2556 unsafe { *outs_addr[t].at(r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
2558 }
2559 };
2560 pool.run_rows(total_rows, &run);
2561 }
2562 return;
2563 }
2564
2565 if uniform_q1 {
2566 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2569 if a8w8_enabled() {
2570 let act = split_act(x);
2571 let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
2572 let (act, gsum) = (&act, &gsum);
2573 let closures: [_; N] = std::array::from_fn(|i| {
2574 let (bytes, gpr, out) =
2575 (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2576 move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
2577 });
2578 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2579 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2580 pool.run_many(&parts);
2581 } else {
2582 let closures: [_; N] = std::array::from_fn(|i| {
2583 let (bytes, gpr, out) =
2584 (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2585 move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
2586 });
2587 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2588 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2589 pool.run_many(&parts);
2590 }
2591 return;
2592 }
2593
2594 if uniform_q1t {
2595 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2599 const TILE: usize = cortiq_core::quant::Q1T_TILE;
2600 if a8w8_enabled() {
2601 let act = split_act(x);
2602 let act = &act;
2603 let x_ref = x;
2604 let closures: [_; N] = std::array::from_fn(|i| {
2605 let bytes = ts[i].quant_bytes();
2606 let (rows, cols) = (ts[i].rows(), ts[i].cols());
2607 let gpr = cols / GROUP_SIZE;
2608 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2609 let out = outs_addr[i];
2610 move |s: usize, e: usize| {
2611 q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
2612 }
2613 });
2614 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2615 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2616 pool.run_many(&parts);
2617 } else {
2618 let x_ref = x;
2619 let closures: [_; N] = std::array::from_fn(|i| {
2620 let bytes = ts[i].quant_bytes();
2621 let (rows, cols) = (ts[i].rows(), ts[i].cols());
2622 let gpr = cols / GROUP_SIZE;
2623 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2624 let out = outs_addr[i];
2625 move |s: usize, e: usize| {
2626 q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
2627 }
2628 });
2629 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2630 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2631 pool.run_many(&parts);
2632 }
2633 return;
2634 }
2635
2636 if uniform_q4 || uniform_vbit {
2637 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2638 if a8w8_enabled() {
2640 let act = split_act(x);
2641 let act = &act;
2642 if uniform_q4 {
2643 let closures: [_; N] = std::array::from_fn(|i| {
2644 let (packed, scales) =
2645 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2646 let (gpr, cols, out) =
2647 (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
2648 move |s: usize, e: usize| {
2649 q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
2650 }
2651 });
2652 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2653 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2654 pool.run_many(&parts);
2655 } else {
2656 let closures: [_; N] = std::array::from_fn(|i| {
2657 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2658 unreachable!()
2659 };
2660 let (bytes, rows, cols, out) = (
2661 ts[i].quant_bytes(),
2662 ts[i].rows(),
2663 ts[i].cols(),
2664 outs_addr[i],
2665 );
2666 move |s: usize, e: usize| {
2667 vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
2668 }
2669 });
2670 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2671 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2672 pool.run_many(&parts);
2673 }
2674 return;
2675 }
2676 if uniform_q4 {
2677 let closures: [_; N] = std::array::from_fn(|i| {
2678 let (packed, scales) =
2679 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2680 let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2681 move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
2682 });
2683 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2684 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2685 pool.run_many(&parts);
2686 } else {
2687 let closures: [_; N] = std::array::from_fn(|i| {
2688 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2689 unreachable!()
2690 };
2691 let (bytes, rows, cols, out) = (
2692 ts[i].quant_bytes(),
2693 ts[i].rows(),
2694 ts[i].cols(),
2695 outs_addr[i],
2696 );
2697 move |s: usize, e: usize| {
2698 vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
2699 }
2700 });
2701 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2702 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2703 pool.run_many(&parts);
2704 }
2705 return;
2706 }
2707
2708 if uniform_f32 {
2709 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2710 let closures: [_; N] = std::array::from_fn(|i| {
2711 let Self::F32 { data, cols, .. } = ts[i] else {
2712 unreachable!()
2713 };
2714 let out = outs_addr[i];
2715 move |start: usize, end: usize| {
2716 for o in start..end {
2717 let row = &data[o * cols..(o + 1) * cols];
2718 let mut sum = 0.0f32;
2719 for j in 0..*cols {
2720 sum += row[j] * x[j];
2721 }
2722 unsafe { *out.at(o) = sum };
2724 }
2725 }
2726 });
2727 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2728 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2729 pool.run_many(&parts);
2730 return;
2731 }
2732
2733 struct Ctx<'a> {
2736 bytes: &'a [u8],
2737 #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
2738 rep: &'a [u8],
2739 row_scale: &'a [f32],
2740 cols: usize,
2741 xs: std::borrow::Cow<'a, [f32]>,
2742 }
2743 let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2744 let Self::Mapped {
2745 dtype,
2746 cols,
2747 row_scale,
2748 col_field,
2749 repack,
2750 ..
2751 } = ts[i]
2752 else {
2753 unreachable!()
2754 };
2755 Ctx {
2756 bytes: ts[i].quant_bytes(),
2757 rep: repack,
2758 row_scale,
2759 cols: *cols,
2760 xs: prescale(x, col_field, *dtype),
2761 }
2762 });
2763 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2764 #[cfg(target_arch = "aarch64")]
2765 if sdot_enabled() {
2766 let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2767 let closures: [_; N] = std::array::from_fn(|i| {
2768 let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2769 move |start: usize, end: usize| {
2770 q8_range_sdot(c.bytes, c.rep, c.row_scale, act, c.cols, out, start, end)
2771 }
2772 });
2773 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2774 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2775 pool.run_many(&parts);
2776 return;
2777 }
2778 #[cfg(target_arch = "x86_64")]
2779 if avx2_a8w8_enabled() {
2780 let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2781 let closures: [_; N] = std::array::from_fn(|i| {
2782 let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2783 move |start: usize, end: usize| {
2784 q8_range_avx2(c.bytes, c.row_scale, act, c.cols, out, start, end)
2785 }
2786 });
2787 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2788 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2789 pool.run_many(&parts);
2790 return;
2791 }
2792 let closures: [_; N] = std::array::from_fn(|i| {
2793 let (c, out) = (&ctxs[i], outs_addr[i]);
2794 move |start: usize, end: usize| {
2795 q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
2796 }
2797 });
2798 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2799 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2800 pool.run_many(&parts);
2801 }
2802}
2803
2804impl QTensor {
2805 #[allow(clippy::needless_range_loop)]
2810 pub fn matvec2_many<const N: usize>(
2811 ts: [&QTensor; N],
2812 x1: &[f32],
2813 x2: &[f32],
2814 mut o1s: [&mut [f32]; N],
2815 mut o2s: [&mut [f32]; N],
2816 pool: Option<&Pool>,
2817 ) {
2818 let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2819 if ts.iter().any(|t| t.has_prism_contract()) {
2820 for i in 0..N {
2821 ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2822 }
2823 return;
2824 }
2825 let uniform_q8 = ts.iter().all(|t| {
2826 matches!(
2827 t,
2828 Self::Mapped {
2829 dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2830 ..
2831 }
2832 )
2833 });
2834 let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2835 let uniform_q4 = ts.iter().all(|t| {
2836 matches!(
2837 t,
2838 Self::Mapped {
2839 dtype: TensorDtype::Q4Block,
2840 ..
2841 }
2842 )
2843 });
2844 let uniform_vbit = ts.iter().all(|t| {
2845 matches!(
2846 t,
2847 Self::Mapped {
2848 dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2849 ..
2850 }
2851 )
2852 });
2853 let fusable = pool.is_some()
2854 && total_rows >= 256
2855 && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
2856 if !fusable {
2857 for i in 0..N {
2858 ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2859 }
2860 return;
2861 }
2862 let pool = pool.unwrap();
2863
2864 if uniform_q4 || uniform_vbit {
2865 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2866 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2867 if a8w8_enabled() {
2869 let a1 = split_act(x1);
2870 let a2 = split_act(x2);
2871 let (a1, a2) = (&a1, &a2);
2872 if uniform_q4 {
2873 let closures: [_; N] = std::array::from_fn(|i| {
2874 let (packed, scales) =
2875 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2876 let (gpr, cols, o1, o2) =
2877 (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
2878 move |s: usize, e: usize| {
2879 q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
2880 }
2881 });
2882 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2883 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2884 pool.run_many(&parts);
2885 } else {
2886 let closures: [_; N] = std::array::from_fn(|i| {
2887 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2888 unreachable!()
2889 };
2890 let (bytes, rows, cols, o1, o2) = (
2891 ts[i].quant_bytes(),
2892 ts[i].rows(),
2893 ts[i].cols(),
2894 p1[i],
2895 p2[i],
2896 );
2897 move |s: usize, e: usize| {
2898 vbit_range2_a8w8(
2899 bytes,
2900 vbit_offsets,
2901 x1,
2902 x2,
2903 a1,
2904 a2,
2905 rows,
2906 cols,
2907 o1,
2908 o2,
2909 s,
2910 e,
2911 )
2912 }
2913 });
2914 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2915 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2916 pool.run_many(&parts);
2917 }
2918 return;
2919 }
2920 if uniform_q4 {
2921 let closures: [_; N] = std::array::from_fn(|i| {
2922 let (packed, scales) =
2923 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2924 let (gpr, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
2925 move |s: usize, e: usize| {
2926 q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
2927 }
2928 });
2929 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2930 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2931 pool.run_many(&parts);
2932 } else {
2933 let closures: [_; N] = std::array::from_fn(|i| {
2934 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2935 unreachable!()
2936 };
2937 let (bytes, rows, cols, o1, o2) = (
2938 ts[i].quant_bytes(),
2939 ts[i].rows(),
2940 ts[i].cols(),
2941 p1[i],
2942 p2[i],
2943 );
2944 move |s: usize, e: usize| {
2945 vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
2946 }
2947 });
2948 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2949 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2950 pool.run_many(&parts);
2951 }
2952 return;
2953 }
2954
2955 if uniform_f32 {
2956 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2957 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2958 let closures: [_; N] = std::array::from_fn(|i| {
2959 let Self::F32 { data, cols, .. } = ts[i] else {
2960 unreachable!()
2961 };
2962 let (o1, o2) = (p1[i], p2[i]);
2963 move |start: usize, end: usize| {
2964 for o in start..end {
2965 let row = &data[o * cols..(o + 1) * cols];
2966 let (mut s1, mut s2) = (0.0f32, 0.0f32);
2967 for j in 0..*cols {
2968 s1 += row[j] * x1[j];
2969 s2 += row[j] * x2[j];
2970 }
2971 unsafe {
2973 *o1.at(o) = s1;
2974 *o2.at(o) = s2;
2975 }
2976 }
2977 }
2978 });
2979 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2980 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2981 pool.run_many(&parts);
2982 return;
2983 }
2984
2985 struct Ctx<'a> {
2986 bytes: &'a [u8],
2987 row_scale: &'a [f32],
2988 cols: usize,
2989 xs1: std::borrow::Cow<'a, [f32]>,
2990 xs2: std::borrow::Cow<'a, [f32]>,
2991 }
2992 let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2993 let Self::Mapped {
2994 dtype,
2995 cols,
2996 row_scale,
2997 col_field,
2998 ..
2999 } = ts[i]
3000 else {
3001 unreachable!()
3002 };
3003 Ctx {
3004 bytes: ts[i].quant_bytes(),
3005 row_scale,
3006 cols: *cols,
3007 xs1: prescale(x1, col_field, *dtype),
3008 xs2: prescale(x2, col_field, *dtype),
3009 }
3010 });
3011 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3012 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3013 #[cfg(target_arch = "aarch64")]
3014 if sdot_enabled() {
3015 let acts: [(SplitAct, SplitAct); N] =
3016 std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3017 let closures: [_; N] = std::array::from_fn(|i| {
3018 let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3019 move |start: usize, end: usize| {
3020 q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3021 }
3022 });
3023 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3024 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3025 pool.run_many(&parts);
3026 return;
3027 }
3028 #[cfg(target_arch = "x86_64")]
3029 if avx2_a8w8_enabled() {
3030 let acts: [(SplitAct, SplitAct); N] =
3031 std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3032 let closures: [_; N] = std::array::from_fn(|i| {
3033 let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3034 move |start: usize, end: usize| {
3035 q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3036 }
3037 });
3038 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3039 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3040 pool.run_many(&parts);
3041 return;
3042 }
3043 let closures: [_; N] = std::array::from_fn(|i| {
3044 let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
3045 move |start: usize, end: usize| {
3046 q8_range2_f32(
3047 c.bytes,
3048 c.row_scale,
3049 &c.xs1,
3050 &c.xs2,
3051 c.cols,
3052 o1,
3053 o2,
3054 start,
3055 end,
3056 )
3057 }
3058 });
3059 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3060 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3061 pool.run_many(&parts);
3062 }
3063
3064 pub fn matvec_silu_mul(
3069 gate: &QTensor,
3070 up: &QTensor,
3071 x: &[f32],
3072 out: &mut [f32],
3073 pool: Option<&Pool>,
3074 ) -> bool {
3075 Self::matvec_silu_mul_limited(gate, up, x, out, 0.0, pool)
3076 }
3077
3078 pub fn matvec_silu_mul_limited(
3084 gate: &QTensor,
3085 up: &QTensor,
3086 x: &[f32],
3087 out: &mut [f32],
3088 limit: f32,
3089 pool: Option<&Pool>,
3090 ) -> bool {
3091 if gate.has_prism_contract() || up.has_prism_contract() {
3092 return false;
3096 }
3097 let inter = gate.rows();
3098 debug_assert_eq!(up.rows(), inter);
3099 debug_assert_eq!(out.len(), inter);
3100 debug_assert_eq!(gate.cols(), up.cols());
3101 if !a8w8_enabled() {
3102 return false;
3103 }
3104 let act = split_act(x);
3105 let act = &act;
3106 let x_ref = x;
3107 let out_addr = SendMut(out.as_mut_ptr());
3108
3109 match (gate, up) {
3110 (
3112 Self::Mapped {
3113 dtype: TensorDtype::Q4Block,
3114 ..
3115 },
3116 Self::Mapped {
3117 dtype: TensorDtype::Q4Block,
3118 ..
3119 },
3120 ) => {
3121 let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
3122 let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
3123 let gpr = gate.cols() / GROUP_SIZE;
3124 let cols = gate.cols();
3125 let run = move |start: usize, end: usize| {
3126 for r in start..end {
3127 let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
3128 let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
3129 for &(j, xv) in &act.outliers {
3130 let flat = r * cols + j;
3131 let gb = gp[flat / 2];
3132 let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
3133 let gsc = f16_to_f32(u16::from_le_bytes([
3134 gs[(flat / GROUP_SIZE) * 2],
3135 gs[(flat / GROUP_SIZE) * 2 + 1],
3136 ]));
3137 gv += ((gn as i32 - 8) as f32) * gsc * xv;
3138 let ub = up_p[flat / 2];
3139 let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
3140 let usc = f16_to_f32(u16::from_le_bytes([
3141 up_s[(flat / GROUP_SIZE) * 2],
3142 up_s[(flat / GROUP_SIZE) * 2 + 1],
3143 ]));
3144 uv += ((un as i32 - 8) as f32) * usc * xv;
3145 }
3146 let silu_g = gv / (1.0 + (-gv).exp());
3147 unsafe { *out_addr.at(r) = silu_g * uv };
3149 }
3150 };
3151 dispatch_rows(pool, inter, &run);
3152 true
3153 }
3154 (
3158 Self::Mapped {
3159 dtype: TensorDtype::Q4Tiled,
3160 ..
3161 },
3162 Self::Mapped {
3163 dtype: TensorDtype::Q4Tiled,
3164 ..
3165 },
3166 ) => {
3167 let g_bytes = gate.quant_bytes();
3168 let u_bytes = up.quant_bytes();
3169 let gpr = gate.cols() / GROUP_SIZE;
3170 let run = move |start: usize, end: usize| {
3171 for r in start..end {
3172 let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3173 let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3174 for &(j, xv) in &act.outliers {
3175 let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
3176 gv += w * s * xv;
3177 let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
3178 uv += w * s * xv;
3179 }
3180 let silu_g = gv / (1.0 + (-gv).exp());
3181 unsafe { *out_addr.at(r) = silu_g * uv };
3183 }
3184 };
3185 dispatch_rows(pool, inter, &run);
3186 true
3187 }
3188 (
3191 Self::Mapped {
3192 dtype: TensorDtype::Q4TiledP,
3193 ..
3194 },
3195 Self::Mapped {
3196 dtype: TensorDtype::Q4TiledP,
3197 ..
3198 },
3199 ) => {
3200 let cols = gate.cols();
3201 let gpr = cols / GROUP_SIZE;
3202 let gv_view = Q4tpView::new(gate.quant_bytes(), inter, cols);
3203 let uv_view = Q4tpView::new(up.quant_bytes(), inter, cols);
3204 let run = |start: usize, end: usize| {
3205 let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3206 for r in start..end {
3207 gv_view.scales_into(r, gpr, &mut gsc);
3208 uv_view.scales_into(r, gpr, &mut usc);
3209 let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3210 let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3211 for &(j, xv) in &act.outliers {
3212 let (w, s) = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3213 gv += w * s * xv;
3214 let (w, s) = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3215 uv += w * s * xv;
3216 }
3217 let silu_g = gv / (1.0 + (-gv).exp());
3218 unsafe { *out_addr.at(r) = silu_g * uv };
3220 }
3221 };
3222 dispatch_rows(pool, inter, &run);
3223 true
3224 }
3225 (
3231 Self::Mapped {
3232 dtype: TensorDtype::Q1,
3233 ..
3234 },
3235 Self::Mapped {
3236 dtype: TensorDtype::Q1,
3237 ..
3238 },
3239 ) => {
3240 let g_bytes = gate.quant_bytes();
3241 let u_bytes = up.quant_bytes();
3242 let gpr = gate.cols() / GROUP_SIZE;
3243 let gsum = q1_group_sums(&act.xq, gpr);
3244 let gsum = &gsum;
3245 let run = move |start: usize, end: usize| {
3246 for r in start..end {
3247 let mut gv = dot_q1_row_i8(g_bytes, r, gpr, &act.xq, gsum) * act.sx;
3248 let mut uv = dot_q1_row_i8(u_bytes, r, gpr, &act.xq, gsum) * act.sx;
3249 for &(j, xv) in &act.outliers {
3250 let (w, s) = q1_outlier(g_bytes, r, gpr, j);
3251 gv += w * s * xv;
3252 let (w, s) = q1_outlier(u_bytes, r, gpr, j);
3253 uv += w * s * xv;
3254 }
3255 let silu_g = gv / (1.0 + (-gv).exp());
3256 unsafe { *out_addr.at(r) = silu_g * uv };
3258 }
3259 };
3260 dispatch_rows(pool, inter, &run);
3261 true
3262 }
3263 (
3267 Self::Mapped {
3268 dtype: TensorDtype::Q2TiledP,
3269 ..
3270 },
3271 Self::Mapped {
3272 dtype: TensorDtype::Q2TiledP,
3273 ..
3274 },
3275 ) => {
3276 let cols = gate.cols();
3277 let gpr = cols / GROUP_SIZE;
3278 let gv_view = Q4tpView::new_q2(gate.quant_bytes(), inter, cols);
3279 let uv_view = Q4tpView::new_q2(up.quant_bytes(), inter, cols);
3280 let gsum = q1_group_sums(&act.xq, gpr);
3281 let gsum = &gsum;
3282 let run = move |start: usize, end: usize| {
3283 let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3284 for r in start..end {
3285 gv_view.scales_into(r, gpr, &mut gsc);
3286 uv_view.scales_into(r, gpr, &mut usc);
3287 let mut gv =
3288 dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, &gsc) * act.sx;
3289 let mut uv =
3290 dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, &usc) * act.sx;
3291 for &(j, xv) in &act.outliers {
3292 let (w, s) = q2tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3293 gv += w * s * xv;
3294 let (w, s) = q2tp_outlier(uv_view.nib, r, gpr, j, &usc);
3295 uv += w * s * xv;
3296 }
3297 let silu_g = gv / (1.0 + (-gv).exp());
3298 unsafe { *out_addr.at(r) = silu_g * uv };
3300 }
3301 };
3302 dispatch_rows(pool, inter, &run);
3303 true
3304 }
3305 (
3310 Self::Mapped {
3311 dtype: TensorDtype::Q8Row,
3312 row_scale: g_rs,
3313 ..
3314 },
3315 Self::Mapped {
3316 dtype: TensorDtype::Q8Row,
3317 row_scale: u_rs,
3318 ..
3319 },
3320 ) => {
3321 let g_bytes = gate.quant_bytes();
3322 let u_bytes = up.quant_bytes();
3323 let cols = gate.cols();
3324 let run = move |start: usize, end: usize| {
3325 for r in start..end {
3326 let gv = q8_row_dot(&g_bytes[r * cols..(r + 1) * cols], act) * g_rs[r];
3327 let uv = q8_row_dot(&u_bytes[r * cols..(r + 1) * cols], act) * u_rs[r];
3328 let silu_g = gv / (1.0 + (-gv).exp());
3329 unsafe { *out_addr.at(r) = silu_g * uv };
3331 }
3332 };
3333 dispatch_rows(pool, inter, &run);
3334 true
3335 }
3336 (
3338 Self::Mapped {
3339 dtype: TensorDtype::Q1T,
3340 ..
3341 },
3342 Self::Mapped {
3343 dtype: TensorDtype::Q1T,
3344 ..
3345 },
3346 ) => {
3347 const TILE: usize = cortiq_core::quant::Q1T_TILE;
3348 let g_bytes = gate.quant_bytes();
3349 let u_bytes = up.quant_bytes();
3350 let gpr = gate.cols() / GROUP_SIZE;
3351 let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
3352 let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
3353 let run = move |start: usize, end: usize| {
3354 for r in start..end {
3355 let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3356 let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3357 for &(j, xv) in &act.outliers {
3358 gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
3359 uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
3360 }
3361 gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
3362 uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
3363 let silu_g = gv / (1.0 + (-gv).exp());
3364 unsafe { *out_addr.at(r) = silu_g * uv };
3366 }
3367 };
3368 dispatch_rows(pool, inter, &run);
3369 true
3370 }
3371 _ => false,
3372 }
3373 }
3374
3375 pub fn moe_gate_up_many(
3390 pairs: &[(&QTensor, &QTensor)],
3391 x: &[f32],
3392 outs: &mut [Vec<f32>],
3393 pool: Option<&Pool>,
3394 ) -> bool {
3395 if pairs.is_empty() || pairs.len() != outs.len() {
3396 return false;
3397 }
3398 if !a8w8_enabled() {
3399 let groups = vec![vec![0]; pairs.len()];
3400 return Self::moe_gate_up_rows(pairs, &groups, x, outs, pool);
3401 }
3402 let inter = pairs[0].0.rows();
3403 let cols = pairs[0].0.cols();
3404 if cols % GROUP_SIZE != 0 {
3405 return false;
3406 }
3407 let gpr = cols / GROUP_SIZE;
3408 let q2 = matches!(
3411 pairs[0].0,
3412 Self::Mapped {
3413 dtype: TensorDtype::Q2TiledP,
3414 ..
3415 }
3416 );
3417 let want = if q2 {
3418 TensorDtype::Q2TiledP
3419 } else {
3420 TensorDtype::Q4TiledP
3421 };
3422 let mut views = Vec::with_capacity(pairs.len() * 2);
3423 for ((g, u), o) in pairs.iter().zip(outs.iter()) {
3424 let both = matches!(g, Self::Mapped { dtype, .. } if *dtype == want)
3425 && matches!(u, Self::Mapped { dtype, .. } if *dtype == want);
3426 if !both
3427 || g.rows() != inter
3428 || u.rows() != inter
3429 || g.cols() != cols
3430 || u.cols() != cols
3431 || o.len() != inter
3432 {
3433 return false;
3434 }
3435 let mk = if q2 { Q4tpView::new_q2 } else { Q4tpView::new };
3436 views.push(mk(g.quant_bytes(), inter, cols));
3437 views.push(mk(u.quant_bytes(), inter, cols));
3438 }
3439 let act = split_act(x);
3440 let gsum = if q2 {
3441 q1_group_sums(&act.xq, gpr)
3442 } else {
3443 Vec::new()
3444 };
3445 let (act, gsum) = (&act, &gsum);
3446 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3447 let (views, ptrs) = (&views, &ptrs);
3448 let run = |start: usize, end: usize| {
3449 let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3450 for flat in start..end {
3451 let (e, r) = (flat / inter, flat % inter);
3452 let gv_view = &views[e * 2];
3453 let uv_view = &views[e * 2 + 1];
3454 gv_view.scales_into(r, gpr, &mut gsc);
3455 uv_view.scales_into(r, gpr, &mut usc);
3456 let (mut gv, mut uv) = if q2 {
3457 (
3458 dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, &gsc) * act.sx,
3459 dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, &usc) * act.sx,
3460 )
3461 } else {
3462 (
3463 dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx,
3464 dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx,
3465 )
3466 };
3467 for &(j, xv) in &act.outliers {
3468 let (og, ou) = if q2 {
3469 (
3470 q2tp_outlier(gv_view.nib, r, gpr, j, &gsc),
3471 q2tp_outlier(uv_view.nib, r, gpr, j, &usc),
3472 )
3473 } else {
3474 (
3475 q4tp_outlier(gv_view.nib, r, gpr, j, &gsc),
3476 q4tp_outlier(uv_view.nib, r, gpr, j, &usc),
3477 )
3478 };
3479 gv += og.0 * og.1 * xv;
3480 uv += ou.0 * ou.1 * xv;
3481 }
3482 let silu_g = gv / (1.0 + (-gv).exp());
3483 unsafe { *ptrs[e].at(r) = silu_g * uv };
3485 }
3486 };
3487 dispatch_rows(pool, pairs.len() * inter, &run);
3488 true
3489 }
3490
3491 pub fn moe_down_many(
3500 downs: &[&QTensor],
3501 gs: &[Vec<f32>],
3502 weights: &[f32],
3503 out: &mut [f32],
3504 pool: Option<&Pool>,
3505 ) -> bool {
3506 if downs.is_empty() || downs.len() != gs.len() || downs.len() != weights.len() {
3507 return false;
3508 }
3509 if !a8w8_enabled() {
3510 let mut terms = vec![vec![0.0; out.len()]; downs.len()];
3511 if !Self::moe_down_rows(downs, &vec![1; downs.len()], gs, &mut terms, pool) {
3512 return false;
3513 }
3514 out.fill(0.0);
3515 for (row, &w) in terms.iter().zip(weights) {
3516 for (o, &v) in out.iter_mut().zip(row) {
3517 *o += w * v;
3518 }
3519 }
3520 return true;
3521 }
3522 let rows = out.len();
3523 let cols = downs[0].cols();
3524 if cols % GROUP_SIZE != 0 {
3525 return false;
3526 }
3527 let gpr = cols / GROUP_SIZE;
3528 let mut views = Vec::with_capacity(downs.len());
3529 for (d, g) in downs.iter().zip(gs.iter()) {
3530 if !matches!(
3531 d,
3532 Self::Mapped {
3533 dtype: TensorDtype::Q4TiledP,
3534 ..
3535 }
3536 ) || d.rows() != rows
3537 || d.cols() != cols
3538 || g.len() != cols
3539 {
3540 return false;
3541 }
3542 views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3543 }
3544 let acts: Vec<SplitAct> = gs.iter().map(|g| split_act(g)).collect();
3546 let out_addr = SendMut(out.as_mut_ptr());
3554 let (views, acts, weights) = (&views, &acts, &weights);
3555 let run = |start: usize, end: usize| {
3556 let mut sc = vec![0f32; gpr];
3557 for r in start..end {
3558 let mut acc = 0f32;
3559 for (e, v) in views.iter().enumerate() {
3560 v.scales_into(r, gpr, &mut sc);
3561 let a = &acts[e];
3562 let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3563 for &(j, xv) in &a.outliers {
3564 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3565 d += w * s * xv;
3566 }
3567 acc += weights[e] * d;
3568 }
3569 unsafe { *out_addr.at(r) = acc };
3571 }
3572 };
3573 dispatch_rows(pool, rows, &run);
3574 true
3575 }
3576
3577 pub fn moe_gate_up_rows(
3589 pairs: &[(&QTensor, &QTensor)],
3590 groups: &[Vec<usize>],
3591 xs: &[f32],
3592 outs: &mut [Vec<f32>],
3593 pool: Option<&Pool>,
3594 ) -> bool {
3595 if pairs.is_empty() || pairs.len() != groups.len() {
3596 return false;
3597 }
3598 let inter = pairs[0].0.rows();
3599 let cols = pairs[0].0.cols();
3600 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
3601 if cols == 0 || cols % GROUP_SIZE != 0 || outs.len() != n_pairs || xs.len() % cols != 0 {
3602 return false;
3603 }
3604 let b = xs.len() / cols;
3605 let gpr = cols / GROUP_SIZE;
3606 let mut views = Vec::with_capacity(pairs.len() * 2);
3607 for (g, u) in pairs {
3608 let q4tp = |t: &QTensor| {
3609 matches!(
3610 t,
3611 Self::Mapped {
3612 dtype: TensorDtype::Q4TiledP,
3613 ..
3614 }
3615 )
3616 };
3617 if g.has_prism_contract()
3618 || u.has_prism_contract()
3619 || !q4tp(g)
3620 || !q4tp(u)
3621 || g.rows() != inter
3622 || u.rows() != inter
3623 || g.cols() != cols
3624 || u.cols() != cols
3625 {
3626 return false;
3627 }
3628 views.push(Q4tpView::new(g.quant_bytes(), inter, cols));
3629 views.push(Q4tpView::new(u.quant_bytes(), inter, cols));
3630 }
3631 if outs.iter().any(|o| o.len() != inter) || groups.iter().flatten().any(|&t| t >= b) {
3632 return false;
3633 }
3634 let quantized = a8w8_enabled();
3635 let acts: Vec<SplitAct> = if quantized {
3636 (0..b)
3637 .map(|t| split_act(&xs[t * cols..(t + 1) * cols]))
3638 .collect()
3639 } else {
3640 Vec::new()
3641 };
3642 let mut offs = Vec::with_capacity(groups.len());
3643 let mut o = 0usize;
3644 for g in groups {
3645 offs.push(o);
3646 o += g.len();
3647 }
3648 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3649 let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3650 let run = |start: usize, end: usize| {
3651 let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3652 for flat in start..end {
3653 let (e, r) = (flat / inter, flat % inter);
3654 let (gv_view, uv_view) = (&views[e * 2], &views[e * 2 + 1]);
3655 gv_view.scales_into(r, gpr, &mut gsc);
3656 uv_view.scales_into(r, gpr, &mut usc);
3657 for (k, &t) in groups[e].iter().enumerate() {
3658 if !quantized {
3659 let x = &xs[t * cols..(t + 1) * cols];
3660 let gv = q4tp_row_exact(gv_view.nib, r, gpr, x, &gsc);
3661 let uv = q4tp_row_exact(uv_view.nib, r, gpr, x, &usc);
3662 unsafe { *ptrs[offs[e] + k].at(r) = (gv / (1.0 + (-gv).exp())) * uv };
3663 continue;
3664 }
3665 let act = &acts[t];
3666 let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3667 let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3668 for &(j, xv) in &act.outliers {
3669 let og = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3670 let ou = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3671 gv += og.0 * og.1 * xv;
3672 uv += ou.0 * ou.1 * xv;
3673 }
3674 let silu_g = gv / (1.0 + (-gv).exp());
3675 unsafe { *ptrs[offs[e] + k].at(r) = silu_g * uv };
3678 }
3679 }
3680 };
3681 dispatch_rows(pool, pairs.len() * inter, &run);
3682 true
3683 }
3684
3685 pub fn moe_down_rows(
3692 downs: &[&QTensor],
3693 group_lens: &[usize],
3694 gs: &[Vec<f32>],
3695 outs: &mut [Vec<f32>],
3696 pool: Option<&Pool>,
3697 ) -> bool {
3698 if downs.is_empty() || downs.len() != group_lens.len() {
3699 return false;
3700 }
3701 let rows = downs[0].rows();
3702 let cols = downs[0].cols();
3703 let n_pairs: usize = group_lens.iter().sum();
3704 if cols == 0 || cols % GROUP_SIZE != 0 || gs.len() != n_pairs || outs.len() != n_pairs {
3705 return false;
3706 }
3707 let gpr = cols / GROUP_SIZE;
3708 let mut views = Vec::with_capacity(downs.len());
3709 for d in downs {
3710 if d.has_prism_contract()
3711 || !matches!(
3712 d,
3713 Self::Mapped {
3714 dtype: TensorDtype::Q4TiledP,
3715 ..
3716 }
3717 )
3718 || d.rows() != rows
3719 || d.cols() != cols
3720 {
3721 return false;
3722 }
3723 views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3724 }
3725 if gs.iter().any(|g| g.len() != cols) || outs.iter().any(|o| o.len() != rows) {
3726 return false;
3727 }
3728 let quantized = a8w8_enabled();
3729 let acts: Vec<SplitAct> = if quantized {
3730 gs.iter().map(|g| split_act(g)).collect()
3731 } else {
3732 Vec::new()
3733 };
3734 let mut offs = Vec::with_capacity(group_lens.len());
3735 let mut o = 0usize;
3736 for &l in group_lens {
3737 offs.push(o);
3738 o += l;
3739 }
3740 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3741 let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3742 let run = |start: usize, end: usize| {
3743 let mut sc = vec![0f32; gpr];
3744 for flat in start..end {
3745 let (e, r) = (flat / rows, flat % rows);
3746 let v = &views[e];
3747 v.scales_into(r, gpr, &mut sc);
3748 for k in 0..group_lens[e] {
3749 if !quantized {
3750 let d = q4tp_row_exact(v.nib, r, gpr, &gs[offs[e] + k], &sc);
3751 unsafe { *ptrs[offs[e] + k].at(r) = d };
3752 continue;
3753 }
3754 let a = &acts[offs[e] + k];
3755 let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3756 for &(j, xv) in &a.outliers {
3757 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3758 d += w * s * xv;
3759 }
3760 unsafe { *ptrs[offs[e] + k].at(r) = d };
3762 }
3763 }
3764 };
3765 dispatch_rows(pool, downs.len() * rows, &run);
3766 true
3767 }
3768}
3769
3770#[cfg(target_os = "macos")]
3775mod accel_blas {
3776 #[link(name = "Accelerate", kind = "framework")]
3777 unsafe extern "C" {
3778 pub fn cblas_sgemm(
3779 order: i32,
3780 trans_a: i32,
3781 trans_b: i32,
3782 m: i32,
3783 n: i32,
3784 k: i32,
3785 alpha: f32,
3786 a: *const f32,
3787 lda: i32,
3788 b: *const f32,
3789 ldb: i32,
3790 beta: f32,
3791 c: *mut f32,
3792 ldc: i32,
3793 );
3794 }
3795}
3796
3797#[cfg(target_os = "macos")]
3798pub(crate) fn accel_gemm_enabled() -> bool {
3799 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3800 *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
3801}
3802
3803#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3806pub(crate) fn accel_gemm_enabled() -> bool {
3807 true
3808}
3809
3810#[cfg(target_arch = "aarch64")]
3817#[allow(clippy::too_many_arguments)]
3818pub(crate) fn neon_gemm_rm(
3819 m: usize,
3820 n: usize,
3821 k: usize,
3822 alpha: f32,
3823 a: &[f32],
3824 lda: usize,
3825 b_mat: &[f32],
3826 ldb: usize,
3827 b_rows_are_n: bool,
3828 c: &mut [f32],
3829 ldc: usize,
3830) {
3831 debug_assert!(a.len() >= (m - 1) * lda + k);
3832 debug_assert!(c.len() >= (m - 1) * ldc + n);
3833 unsafe {
3835 use core::arch::aarch64::*;
3836 let mut i = 0usize;
3837 while i < m {
3838 let mi = (m - i).min(4);
3839 let mut j = 0usize;
3840 while j < n {
3841 let nj = (n - j).min(8);
3842 if mi == 4 && nj == 8 {
3843 let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3844 let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3845 let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3846 let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3847 for p in 0..k {
3848 let (b0, b1) = if b_rows_are_n {
3849 let base = b_mat.as_ptr().add(j * ldb + p);
3852 let g = |o: usize| *base.add(o * ldb);
3853 ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
3854 } else {
3855 let base = b_mat.as_ptr().add(p * ldb + j);
3856 (
3857 [*base, *base.add(1), *base.add(2), *base.add(3)],
3858 [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
3859 )
3860 };
3861 let bv0 = vld1q_f32(b0.as_ptr());
3862 let bv1 = vld1q_f32(b1.as_ptr());
3863 let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
3864 let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
3865 let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
3866 let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
3867 c0a = vfmaq_f32(c0a, a0, bv0);
3868 c0b = vfmaq_f32(c0b, a0, bv1);
3869 c1a = vfmaq_f32(c1a, a1, bv0);
3870 c1b = vfmaq_f32(c1b, a1, bv1);
3871 c2a = vfmaq_f32(c2a, a2, bv0);
3872 c2b = vfmaq_f32(c2b, a2, bv1);
3873 c3a = vfmaq_f32(c3a, a3, bv0);
3874 c3b = vfmaq_f32(c3b, a3, bv1);
3875 }
3876 let al = vdupq_n_f32(alpha);
3877 for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
3878 .iter()
3879 .enumerate()
3880 {
3881 let dst = c.as_mut_ptr().add((i + r) * ldc + j);
3882 vst1q_f32(dst, vmulq_f32(*ca, al));
3883 vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
3884 }
3885 } else {
3886 for r in 0..mi {
3887 for q in 0..nj {
3888 let mut acc = 0f32;
3889 for p in 0..k {
3890 let bv = if b_rows_are_n {
3891 b_mat[(j + q) * ldb + p]
3892 } else {
3893 b_mat[p * ldb + j + q]
3894 };
3895 acc += a[(i + r) * lda + p] * bv;
3896 }
3897 c[(i + r) * ldc + j + q] = acc * alpha;
3898 }
3899 }
3900 }
3901 j += nj;
3902 }
3903 i += mi;
3904 }
3905 }
3906}
3907
3908#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3910#[allow(clippy::too_many_arguments)]
3911pub(crate) fn sgemm_rm(
3912 m: usize,
3913 n: usize,
3914 k: usize,
3915 alpha: f32,
3916 a: &[f32],
3917 lda: usize,
3918 b_mat: &[f32],
3919 ldb: usize,
3920 b_rows_are_n: bool,
3921 c: &mut [f32],
3922 ldc: usize,
3923) {
3924 neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3925}
3926
3927#[allow(clippy::too_many_arguments)]
3931pub fn sgemm_public(
3932 m: usize,
3933 n: usize,
3934 k: usize,
3935 alpha: f32,
3936 a: &[f32],
3937 lda: usize,
3938 b_mat: &[f32],
3939 ldb: usize,
3940 b_rows_are_n: bool,
3941 c: &mut [f32],
3942 ldc: usize,
3943) {
3944 #[cfg(any(target_os = "macos", target_arch = "aarch64"))]
3945 {
3946 sgemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3947 }
3948 #[cfg(not(any(target_os = "macos", target_arch = "aarch64")))]
3953 {
3954 for i in 0..m {
3955 for j in 0..n {
3956 let mut acc = 0f32;
3957 for p in 0..k {
3958 let bv = if b_rows_are_n {
3959 b_mat[j * ldb + p]
3960 } else {
3961 b_mat[p * ldb + j]
3962 };
3963 acc += a[i * lda + p] * bv;
3964 }
3965 c[i * ldc + j] = alpha * acc;
3966 }
3967 }
3968 }
3969}
3970
3971#[cfg(target_os = "macos")]
3974#[allow(clippy::too_many_arguments)]
3975pub(crate) fn sgemm_rm(
3976 m: usize,
3977 n: usize,
3978 k: usize,
3979 alpha: f32,
3980 a: &[f32],
3981 lda: usize,
3982 b_mat: &[f32],
3983 ldb: usize,
3984 b_rows_are_n: bool,
3985 c: &mut [f32],
3986 ldc: usize,
3987) {
3988 debug_assert!(a.len() >= (m - 1) * lda + k);
3989 debug_assert!(c.len() >= (m - 1) * ldc + n);
3990 #[cfg(target_arch = "aarch64")]
3995 if std::env::var("CMF_FORCE_NEON_GEMM")
3996 .map(|v| v == "1")
3997 .unwrap_or(false)
3998 {
3999 return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4000 }
4001 unsafe {
4002 accel_blas::cblas_sgemm(
4003 101, 111, if b_rows_are_n { 112 } else { 111 },
4006 m as i32,
4007 n as i32,
4008 k as i32,
4009 alpha,
4010 a.as_ptr(),
4011 lda as i32,
4012 b_mat.as_ptr(),
4013 ldb as i32,
4014 0.0,
4015 c.as_mut_ptr(),
4016 ldc as i32,
4017 );
4018 }
4019}
4020
4021#[cfg(target_os = "macos")]
4028fn qmatmat_accel(
4029 q: &[u8],
4030 row_scale: &[f32],
4031 pre: &[std::borrow::Cow<'_, [f32]>],
4032 rows: usize,
4033 cols: usize,
4034 out: &mut [f32],
4035 pool: Option<&Pool>,
4036) {
4037 const TR: usize = 2048;
4042 let b = pre.len();
4043 thread_local! {
4044 static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4045 static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4046 }
4047 XPANEL.with(|xp| {
4048 WTILE.with(|wt| {
4049 let mut xpanel = xp.borrow_mut();
4050 xpanel.clear();
4051 for x in pre {
4052 xpanel.extend_from_slice(x);
4053 }
4054 let mut wtile = wt.borrow_mut();
4055 wtile.resize(TR * cols, 0.0);
4056 let mut r0 = 0usize;
4057 while r0 < rows {
4058 let tr = TR.min(rows - r0);
4059 let wt_addr = SendMut(wtile.as_mut_ptr());
4061 let run = |start: usize, end: usize| {
4062 for r in start..end {
4063 let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
4064 let s = row_scale[r0 + r];
4065 let dst =
4067 unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
4068 for (d, &v) in dst.iter_mut().zip(row) {
4069 *d = (v as i8) as f32 * s;
4070 }
4071 }
4072 };
4073 dispatch_rows(pool, tr, &run);
4074 unsafe {
4076 accel_blas::cblas_sgemm(
4077 101, 111, 112, b as i32,
4081 tr as i32,
4082 cols as i32,
4083 1.0,
4084 xpanel.as_ptr(),
4085 cols as i32,
4086 wtile.as_ptr(),
4087 cols as i32,
4088 0.0,
4089 out.as_mut_ptr().add(r0),
4090 rows as i32,
4091 );
4092 }
4093 r0 += tr;
4094 }
4095 })
4096 });
4097}
4098
4099fn qmatmat(
4100 q: &[u8],
4101 row_scale: &[f32],
4102 pre: &[std::borrow::Cow<'_, [f32]>],
4103 rows: usize,
4104 cols: usize,
4105 out: &mut [f32],
4106 pool: Option<&Pool>,
4107) {
4108 let b = pre.len();
4109 debug_assert_eq!(out.len(), b * rows);
4110 #[cfg(target_os = "macos")]
4115 if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
4116 qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
4117 return;
4118 }
4119 #[cfg(target_arch = "aarch64")]
4120 if sdot_enabled() {
4121 let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4122 let out_addr = SendMut(out.as_mut_ptr());
4123 let blocked_ok = blocked_enabled();
4126 let use_i8mm = i8mm_enabled();
4127 if blocked_ok {
4128 let run = |start: usize, end: usize| {
4129 let mut o = start;
4130 while o < end {
4131 if o + 2 <= end {
4132 let r0 = &q[o * cols..(o + 1) * cols];
4133 let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4134 let mut bi = 0usize;
4135 while bi + 4 <= acts.len() {
4136 let xs = [
4137 acts[bi].xq.as_slice(),
4138 acts[bi + 1].xq.as_slice(),
4139 acts[bi + 2].xq.as_slice(),
4140 acts[bi + 3].xq.as_slice(),
4141 ];
4142 let d = if use_i8mm {
4143 unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
4144 } else {
4145 unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
4146 };
4147 for (r, row) in [r0, r1].into_iter().enumerate() {
4148 for k in 0..4 {
4149 let act = &acts[bi + k];
4150 let mut v = d[r][k] as f32 * act.sx;
4151 for &(j, xv) in &act.outliers {
4152 v += (row[j] as i8) as f32 * xv;
4153 }
4154 unsafe {
4155 *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4156 };
4157 }
4158 }
4159 bi += 4;
4160 }
4161 while bi < acts.len() {
4162 for (r, row) in [r0, r1].into_iter().enumerate() {
4163 let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
4164 unsafe { *out_addr.at(bi * rows + o + r) = v };
4165 }
4166 bi += 1;
4167 }
4168 o += 2;
4169 } else {
4170 let row = &q[o * cols..(o + 1) * cols];
4171 for (bi, act) in acts.iter().enumerate() {
4172 let v = row_dot_sdot(row, act) * row_scale[o];
4173 unsafe { *out_addr.at(bi * rows + o) = v };
4174 }
4175 o += 1;
4176 }
4177 }
4178 };
4179 dispatch_rows(pool, rows, &run);
4180 return;
4181 }
4182 let run = |start: usize, end: usize| {
4183 for o in start..end {
4184 let row = &q[o * cols..(o + 1) * cols];
4185 for (bi, act) in acts.iter().enumerate() {
4186 let v = row_dot_sdot(row, act) * row_scale[o];
4187 unsafe { *out_addr.at(bi * rows + o) = v };
4188 }
4189 }
4190 };
4191 dispatch_rows(pool, rows, &run);
4192 return;
4193 }
4194 #[cfg(target_arch = "x86_64")]
4199 if avx2_a8w8_enabled() {
4200 let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4201 let out_addr = SendMut(out.as_mut_ptr());
4202 let blocked_ok = blocked_enabled();
4205 if !avx512vnni_enabled() && blocked_ok && !row_exact() {
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 = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
4221 for (r, row) in [r0, r1].into_iter().enumerate() {
4222 for k in 0..4 {
4223 let act = &acts[bi + k];
4224 let mut v = d[r][k] as f32 * act.sx;
4225 for &(j, xv) in &act.outliers {
4226 v += (row[j] as i8) as f32 * xv;
4227 }
4228 unsafe {
4229 *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4230 };
4231 }
4232 }
4233 bi += 4;
4234 }
4235 while bi < acts.len() {
4236 for (r, row) in [r0, r1].into_iter().enumerate() {
4237 let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
4238 unsafe { *out_addr.at(bi * rows + o + r) = v };
4239 }
4240 bi += 1;
4241 }
4242 o += 2;
4243 } else {
4244 let row = &q[o * cols..(o + 1) * cols];
4245 for (bi, act) in acts.iter().enumerate() {
4246 let v = row_dot_avx2(row, act) * row_scale[o];
4247 unsafe { *out_addr.at(bi * rows + o) = v };
4248 }
4249 o += 1;
4250 }
4251 }
4252 };
4253 dispatch_rows(pool, rows, &run);
4254 return;
4255 }
4256 let run = |start: usize, end: usize| {
4257 for o in start..end {
4258 let row = &q[o * cols..(o + 1) * cols];
4259 for (bi, act) in acts.iter().enumerate() {
4260 let v = row_dot_avx2(row, act) * row_scale[o];
4261 unsafe { *out_addr.at(bi * rows + o) = v };
4262 }
4263 }
4264 };
4265 dispatch_rows(pool, rows, &run);
4266 return;
4267 }
4268 let out_addr = SendMut(out.as_mut_ptr());
4269 let run = |start: usize, end: usize| {
4270 for o in start..end {
4271 let row = &q[o * cols..(o + 1) * cols];
4272 for (bi, x) in pre.iter().enumerate() {
4273 let mut acc = 0f32;
4274 for j in 0..cols {
4275 acc += (row[j] as i8) as f32 * x[j];
4276 }
4277 unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
4278 }
4279 }
4280 };
4281 dispatch_rows(pool, rows, &run);
4282}
4283
4284fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
4287 match pool {
4288 Some(pool) if rows >= 256 => pool.run_rows(rows, run),
4289 _ => run(0, rows),
4290 }
4291}
4292
4293fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
4295 let groups = rows * cols / GROUP_SIZE;
4296 bytes.split_at(groups * 16)
4297}
4298
4299#[inline]
4304fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
4305 #[cfg(target_arch = "aarch64")]
4306 unsafe {
4307 return vbit_fill4_neon(data, buf);
4308 }
4309 #[cfg(target_arch = "x86_64")]
4310 if avx2_enabled() {
4311 return unsafe { vbit_fill4_avx2(data, buf) };
4312 }
4313 #[allow(unreachable_code)]
4314 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4315 let u = unpack8::<4>(&data[blk * 4..]);
4316 for k in 0..8 {
4317 chunk[k] = (u[k] - 7) as i8 as u8;
4318 }
4319 }
4320}
4321
4322#[cfg(target_arch = "aarch64")]
4323#[target_feature(enable = "neon")]
4324unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
4325 unsafe {
4328 use core::arch::aarch64::*;
4329 let n = buf.len();
4330 let mask = vdupq_n_u8(0x0F);
4331 let seven = vdupq_n_s8(7);
4332 let mut g = 0usize;
4333 while g * 32 + 32 <= n {
4334 let b = vld1q_u8(data.as_ptr().add(g * 16));
4335 let hi = vshrq_n_u8::<4>(b);
4336 let lo = vandq_u8(b, mask);
4337 let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
4338 let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
4339 vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
4340 vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
4341 g += 1;
4342 }
4343 }
4344}
4345
4346#[cfg(target_arch = "x86_64")]
4347#[target_feature(enable = "avx2")]
4348unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
4349 unsafe {
4351 use core::arch::x86_64::*;
4352 let n = buf.len();
4353 let mask = _mm_set1_epi8(0x0F);
4354 let seven = _mm256_set1_epi8(7);
4355 let mut g = 0usize;
4356 while g * 32 + 32 <= n {
4357 let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
4358 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
4359 let lo = _mm_and_si128(b, mask);
4360 let z = _mm256_sub_epi8(
4361 _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
4362 seven,
4363 );
4364 _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
4365 g += 1;
4366 }
4367 }
4368}
4369
4370#[inline(always)]
4375fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
4376 let mut acc = 0u64;
4377 for i in 0..B {
4378 acc = (acc << 8) | data[i] as u64;
4379 }
4380 let mask = (1u64 << B) - 1;
4381 let mut out = [0i32; 8];
4382 for (k, o) in out.iter_mut().enumerate() {
4383 *o = ((acc >> ((7 - k) * B)) & mask) as i32;
4384 }
4385 out
4386}
4387
4388#[allow(clippy::too_many_arguments)]
4394fn vbitmatvec(
4395 bytes: &[u8],
4396 offsets: &[usize],
4397 x: &[f32],
4398 rows: usize,
4399 cols: usize,
4400 out: &mut [f32],
4401 pool: Option<&Pool>,
4402) {
4403 debug_assert_eq!(out.len(), rows);
4404 debug_assert_eq!(offsets.len(), rows + 1);
4405
4406 if a8w8_enabled() {
4410 let act = split_act(x);
4411 let out_addr = SendMut(out.as_mut_ptr());
4412 let run = move |start: usize, end: usize| {
4413 vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
4414 };
4415 dispatch_rows(pool, rows, &run);
4416 return;
4417 }
4418
4419 let out_addr = SendMut(out.as_mut_ptr());
4420 let run = move |start: usize, end: usize| {
4421 vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
4422 };
4423 dispatch_rows(pool, rows, &run);
4424}
4425
4426#[allow(clippy::too_many_arguments)]
4430fn vbit_range_a8w8(
4431 bytes: &[u8],
4432 offsets: &[usize],
4433 x: &[f32],
4434 act: &SplitAct,
4435 rows: usize,
4436 cols: usize,
4437 out: SendMut,
4438 start: usize,
4439 end: usize,
4440) {
4441 let ng = cols / GROUP_SIZE;
4442 let bits = &bytes[..rows];
4443 let sc_off = rows;
4444 let row_dot = |r: usize| -> f32 {
4445 let b = bits[r] as usize;
4446 let l = (1i32 << (b - 1)) - 1;
4447 let mask = (1u64 << b) - 1;
4448 let data = &bytes[offsets[r]..offsets[r + 1]];
4449 if b == 8 {
4450 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4452 let mut dot = 0f32;
4453 for g in 0..ng {
4454 let so = (r * ng + g) * 2;
4455 let sgf = f16_to_f32(u16::from_le_bytes([
4456 bytes[sc_off + so],
4457 bytes[sc_off + so + 1],
4458 ]));
4459 let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4460 let mut gd = 0f32;
4461 for &xv in xg.iter() {
4462 if nbits < 8 {
4463 acc = (acc << 8) | data[idx] as u64;
4464 idx += 1;
4465 nbits += 8;
4466 }
4467 let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4468 nbits -= 8;
4469 gd += (u - l) as f32 * xv;
4470 }
4471 dot += gd * sgf;
4472 }
4473 return dot;
4474 }
4475 thread_local! {
4479 static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
4480 const { std::cell::RefCell::new(Vec::new()) };
4481 }
4482 #[inline(always)]
4483 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4484 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4485 let u = unpack8::<B>(&data[blk * B..]);
4486 for k in 0..8 {
4487 chunk[k] = (u[k] - l) as i8 as u8;
4488 }
4489 }
4490 }
4491 let _ = mask;
4492 VBIT_SCRATCH.with(|scratch| {
4493 let mut buf = scratch.borrow_mut();
4494 buf.resize(cols, 0);
4495 match b {
4496 3 => fill::<3>(data, l, &mut buf),
4497 4 => vbit_fill4(data, &mut buf),
4498 5 => fill::<5>(data, l, &mut buf),
4499 6 => fill::<6>(data, l, &mut buf),
4500 _ => unreachable!(),
4501 }
4502 let mut dot = 0f32;
4503 for g in 0..ng {
4504 let so = (r * ng + g) * 2;
4505 let s = f16_to_f32(u16::from_le_bytes([
4506 bytes[sc_off + so],
4507 bytes[sc_off + so + 1],
4508 ]));
4509 let d = dot_i8_i8(
4510 &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4511 &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4512 ) as f32
4513 * act.sx;
4514 dot += d * s;
4515 }
4516 for &(j, xv) in &act.outliers {
4517 let so = (r * ng + j / GROUP_SIZE) * 2;
4518 let s = f16_to_f32(u16::from_le_bytes([
4519 bytes[sc_off + so],
4520 bytes[sc_off + so + 1],
4521 ]));
4522 dot += (buf[j] as i8) as f32 * s * xv;
4524 }
4525 dot
4526 })
4527 };
4528 for r in start..end {
4529 unsafe { *out.at(r) = row_dot(r) };
4531 }
4532}
4533
4534#[allow(clippy::too_many_arguments)]
4536fn vbit_range_f32(
4537 bytes: &[u8],
4538 offsets: &[usize],
4539 x: &[f32],
4540 rows: usize,
4541 cols: usize,
4542 out: SendMut,
4543 start: usize,
4544 end: usize,
4545) {
4546 let ng = cols / GROUP_SIZE;
4547 let bits = &bytes[..rows];
4548 let sc_off = rows;
4549 #[inline(always)]
4553 fn dot_row<const B: usize>(
4554 data: &[u8],
4555 bytes: &[u8],
4556 sc_off: usize,
4557 r: usize,
4558 ng: usize,
4559 x: &[f32],
4560 ) -> f32 {
4561 let l = ((1i32 << (B - 1)) - 1) as f32;
4562 let gbytes = GROUP_SIZE * B / 8;
4563 let mut dot = 0f32;
4564 for g in 0..ng {
4565 let so = (r * ng + g) * 2;
4566 let s = f16_to_f32(u16::from_le_bytes([
4567 bytes[sc_off + so],
4568 bytes[sc_off + so + 1],
4569 ]));
4570 let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4571 let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4572 let mut gd = 0f32;
4573 for blk in 0..GROUP_SIZE / 8 {
4574 let u = unpack8::<B>(&gd0[blk * B..]);
4575 let xb = &xg[blk * 8..blk * 8 + 8];
4576 for k in 0..8 {
4577 gd += (u[k] as f32 - l) * xb[k];
4578 }
4579 }
4580 dot += gd * s;
4581 }
4582 dot
4583 }
4584 for r in start..end {
4585 let data = &bytes[offsets[r]..offsets[r + 1]];
4586 let v = match bits[r] {
4587 3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
4588 4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
4589 5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
4590 6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
4591 8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
4592 b => unreachable!("vbit bit-width {b} (validated at load)"),
4593 };
4594 unsafe { *out.at(r) = v };
4596 }
4597}
4598
4599#[allow(clippy::too_many_arguments)]
4604fn vbitmatvec2(
4605 bytes: &[u8],
4606 offsets: &[usize],
4607 x1: &[f32],
4608 x2: &[f32],
4609 rows: usize,
4610 cols: usize,
4611 o1: &mut [f32],
4612 o2: &mut [f32],
4613 pool: Option<&Pool>,
4614) {
4615 debug_assert_eq!(o1.len(), rows);
4616 debug_assert_eq!(o2.len(), rows);
4617
4618 if a8w8_enabled() {
4619 let a1 = split_act(x1);
4620 let a2 = split_act(x2);
4621 let p1 = SendMut(o1.as_mut_ptr());
4622 let p2 = SendMut(o2.as_mut_ptr());
4623 let run = move |start: usize, end: usize| {
4624 vbit_range2_a8w8(
4625 bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
4626 )
4627 };
4628 dispatch_rows(pool, rows, &run);
4629 return;
4630 }
4631
4632 let p1 = SendMut(o1.as_mut_ptr());
4633 let p2 = SendMut(o2.as_mut_ptr());
4634 let run = move |start: usize, end: usize| {
4635 vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
4636 };
4637 dispatch_rows(pool, rows, &run);
4638}
4639
4640#[allow(clippy::too_many_arguments)]
4644fn vbit_range2_a8w8(
4645 bytes: &[u8],
4646 offsets: &[usize],
4647 x1: &[f32],
4648 x2: &[f32],
4649 a1: &SplitAct,
4650 a2: &SplitAct,
4651 rows: usize,
4652 cols: usize,
4653 p1: SendMut,
4654 p2: SendMut,
4655 start: usize,
4656 end: usize,
4657) {
4658 let ng = cols / GROUP_SIZE;
4659 let bits = &bytes[..rows];
4660 let sc_off = rows;
4661 let row_dots = |r: usize| -> (f32, f32) {
4662 let b = bits[r] as usize;
4663 let l = (1i32 << (b - 1)) - 1;
4664 let data = &bytes[offsets[r]..offsets[r + 1]];
4665 if b == 8 {
4666 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4669 let (mut d1, mut d2) = (0f32, 0f32);
4670 for g in 0..ng {
4671 let so = (r * ng + g) * 2;
4672 let sgf = f16_to_f32(u16::from_le_bytes([
4673 bytes[sc_off + so],
4674 bytes[sc_off + so + 1],
4675 ]));
4676 let (mut g1, mut g2) = (0f32, 0f32);
4677 for k in 0..GROUP_SIZE {
4678 if nbits < 8 {
4679 acc = (acc << 8) | data[idx] as u64;
4680 idx += 1;
4681 nbits += 8;
4682 }
4683 let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4684 nbits -= 8;
4685 let w = (u - l) as f32;
4686 g1 += w * x1[g * GROUP_SIZE + k];
4687 g2 += w * x2[g * GROUP_SIZE + k];
4688 }
4689 d1 += g1 * sgf;
4690 d2 += g2 * sgf;
4691 }
4692 return (d1, d2);
4693 }
4694 thread_local! {
4695 static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
4696 const { std::cell::RefCell::new(Vec::new()) };
4697 }
4698 #[inline(always)]
4699 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4700 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4701 let u = unpack8::<B>(&data[blk * B..]);
4702 for k in 0..8 {
4703 chunk[k] = (u[k] - l) as i8 as u8;
4704 }
4705 }
4706 }
4707 VBIT_SCRATCH2.with(|scratch| {
4708 let mut buf = scratch.borrow_mut();
4709 buf.resize(cols, 0);
4710 match b {
4711 3 => fill::<3>(data, l, &mut buf),
4712 4 => vbit_fill4(data, &mut buf),
4713 5 => fill::<5>(data, l, &mut buf),
4714 6 => fill::<6>(data, l, &mut buf),
4715 _ => unreachable!(),
4716 }
4717 let (mut d1, mut d2) = (0f32, 0f32);
4718 for g in 0..ng {
4719 let so = (r * ng + g) * 2;
4720 let s = f16_to_f32(u16::from_le_bytes([
4721 bytes[sc_off + so],
4722 bytes[sc_off + so + 1],
4723 ]));
4724 let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4725 let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
4726 let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
4727 d1 += v1 * s;
4728 d2 += v2 * s;
4729 }
4730 for &(j, xv) in &a1.outliers {
4731 let so = (r * ng + j / GROUP_SIZE) * 2;
4732 let s = f16_to_f32(u16::from_le_bytes([
4733 bytes[sc_off + so],
4734 bytes[sc_off + so + 1],
4735 ]));
4736 d1 += (buf[j] as i8) as f32 * s * xv;
4737 }
4738 for &(j, xv) in &a2.outliers {
4739 let so = (r * ng + j / GROUP_SIZE) * 2;
4740 let s = f16_to_f32(u16::from_le_bytes([
4741 bytes[sc_off + so],
4742 bytes[sc_off + so + 1],
4743 ]));
4744 d2 += (buf[j] as i8) as f32 * s * xv;
4745 }
4746 (d1, d2)
4747 })
4748 };
4749 for r in start..end {
4750 let (v1, v2) = row_dots(r);
4751 unsafe {
4753 *p1.at(r) = v1;
4754 *p2.at(r) = v2;
4755 }
4756 }
4757}
4758
4759#[allow(clippy::too_many_arguments)]
4763fn vbit_range2_f32(
4764 bytes: &[u8],
4765 offsets: &[usize],
4766 x1: &[f32],
4767 x2: &[f32],
4768 rows: usize,
4769 cols: usize,
4770 p1: SendMut,
4771 p2: SendMut,
4772 start: usize,
4773 end: usize,
4774) {
4775 let ng = cols / GROUP_SIZE;
4776 let bits = &bytes[..rows];
4777 let sc_off = rows;
4778 #[inline(always)]
4779 #[allow(clippy::too_many_arguments)]
4780 fn dot_row2<const B: usize>(
4781 data: &[u8],
4782 bytes: &[u8],
4783 sc_off: usize,
4784 r: usize,
4785 ng: usize,
4786 x1: &[f32],
4787 x2: &[f32],
4788 ) -> (f32, f32) {
4789 let l = ((1i32 << (B - 1)) - 1) as f32;
4790 let gbytes = GROUP_SIZE * B / 8;
4791 let (mut d1, mut d2) = (0f32, 0f32);
4792 for g in 0..ng {
4793 let so = (r * ng + g) * 2;
4794 let s = f16_to_f32(u16::from_le_bytes([
4795 bytes[sc_off + so],
4796 bytes[sc_off + so + 1],
4797 ]));
4798 let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4799 let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4800 let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4801 let (mut g1, mut g2) = (0f32, 0f32);
4802 for blk in 0..GROUP_SIZE / 8 {
4803 let u = unpack8::<B>(&gd0[blk * B..]);
4804 for k in 0..8 {
4805 let w = u[k] as f32 - l;
4806 g1 += w * x1g[blk * 8 + k];
4807 g2 += w * x2g[blk * 8 + k];
4808 }
4809 }
4810 d1 += g1 * s;
4811 d2 += g2 * s;
4812 }
4813 (d1, d2)
4814 }
4815 for r in start..end {
4816 let data = &bytes[offsets[r]..offsets[r + 1]];
4817 let (v1, v2) = match bits[r] {
4818 3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
4819 4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
4820 5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
4821 6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
4822 8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
4823 b => unreachable!("vbit bit-width {b} (validated at load)"),
4824 };
4825 unsafe {
4827 *p1.at(r) = v1;
4828 *p2.at(r) = v2;
4829 }
4830 }
4831}
4832
4833#[inline]
4840#[allow(unreachable_code)]
4841fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4842 #[cfg(target_arch = "aarch64")]
4843 unsafe {
4844 return dot_q4t_row_sdot(bytes, r, gpr, xq);
4845 }
4846 #[cfg(target_arch = "x86_64")]
4847 unsafe {
4848 if vnni_tiles_enabled() {
4849 return dot_q4t_row_vnni(bytes, r, gpr, xq);
4850 }
4851 return dot_q4t_row_avx2(bytes, r, gpr, xq);
4852 }
4853 let mut acc = 0f32;
4854 for gi in 0..gpr {
4855 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
4856 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
4857 let mut d = 0i32;
4858 for (k, &b) in tile[2..].iter().enumerate() {
4859 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
4860 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
4861 }
4862 acc += d as f32 * s;
4863 }
4864 acc
4865}
4866
4867#[cfg(target_arch = "aarch64")]
4868#[target_feature(enable = "neon,dotprod")]
4869unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4870 unsafe {
4873 use core::arch::aarch64::*;
4874 use core::arch::asm;
4875 let lomask = vdupq_n_u8(0x0F);
4876 let eight = vdupq_n_s8(8);
4877 let mut acc = 0f32;
4878 for gi in 0..gpr {
4879 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4880 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4881 let b = vld1q_u8(t.add(2));
4882 let lo = vandq_u8(b, lomask);
4883 let hi = vshrq_n_u8::<4>(b);
4884 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
4885 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
4886 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
4887 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
4888 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4889 asm!(
4890 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
4891 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
4892 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4893 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
4894 options(pure, nomem, nostack),
4895 );
4896 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
4897 }
4898 acc
4899 }
4900}
4901
4902#[cfg(target_arch = "x86_64")]
4903#[target_feature(enable = "avx2")]
4904unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4905 unsafe {
4907 use core::arch::x86_64::*;
4908 let lomask = _mm_set1_epi8(0x0F);
4909 let eight = _mm256_set1_epi8(8);
4910 let ones = _mm256_set1_epi16(1);
4911 let mut acc = 0f32;
4912 for gi in 0..gpr {
4913 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4914 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4915 let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4916 let lo = _mm_and_si128(b, lomask);
4917 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4918 let w = _mm256_sub_epi8(
4919 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4920 eight,
4921 );
4922 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4923 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4924 let d = _mm256_madd_epi16(p16, ones);
4925 let hi128 = _mm256_extracti128_si256::<1>(d);
4926 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
4927 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
4928 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
4929 acc += _mm_cvtsi128_si32(s32) as f32 * s;
4930 }
4931 acc
4932 }
4933}
4934
4935#[cfg(target_arch = "x86_64")]
4939#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
4940unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4941 unsafe {
4943 use core::arch::x86_64::*;
4944 let lomask = _mm_set1_epi8(0x0F);
4945 let eight = _mm256_set1_epi8(8);
4946 let mut acc = 0f32;
4947 for gi in 0..gpr {
4948 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4949 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4950 let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4951 let lo = _mm_and_si128(b, lomask);
4952 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4953 let w = _mm256_sub_epi8(
4954 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4955 eight,
4956 );
4957 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4958 let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4959 acc += d as f32 * s;
4960 }
4961 acc
4962 }
4963}
4964
4965#[cfg(target_arch = "x86_64")]
4970#[target_feature(enable = "avx2,fma")]
4975unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
4976 unsafe {
4978 use core::arch::x86_64::*;
4979 let lomask = _mm_set1_epi8(0x0F);
4980 let eight = _mm256_set1_epi8(8);
4981 let ones = _mm256_set1_epi16(1);
4982 let mut f0 = _mm256_setzero_ps();
4995 let mut f1 = _mm256_setzero_ps();
4996 let mut f2 = _mm256_setzero_ps();
4997 let mut f3 = _mm256_setzero_ps();
4998 for gi in 0..gpr {
4999 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5000 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5001 let sv = _mm256_set1_ps(s);
5002 let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5003 let lo = _mm_and_si128(bb, lomask);
5004 let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5005 let w = _mm256_sub_epi8(
5006 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5007 eight,
5008 );
5009 let aw = _mm256_abs_epi8(w);
5010 let off = gi * GROUP_SIZE;
5011 let dot = |xq: &[i8]| {
5012 let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5013 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5014 _mm256_cvtepi32_ps(_mm256_madd_epi16(p16, ones))
5015 };
5016 f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5017 f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5018 f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5019 f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5020 }
5021 [
5022 hsum256_ps(f0),
5023 hsum256_ps(f1),
5024 hsum256_ps(f2),
5025 hsum256_ps(f3),
5026 ]
5027 }
5028}
5029
5030#[cfg(target_arch = "x86_64")]
5033#[target_feature(enable = "avx2")]
5034#[inline]
5035unsafe fn hsum256_ps(v: core::arch::x86_64::__m256) -> f32 {
5036 unsafe {
5038 use core::arch::x86_64::*;
5039 let hi = _mm256_extractf128_ps::<1>(v);
5040 let s = _mm_add_ps(_mm256_castps256_ps128(v), hi);
5041 let s = _mm_add_ps(s, _mm_movehl_ps(s, s));
5042 let s = _mm_add_ss(s, _mm_shuffle_ps::<0x55>(s, s));
5043 _mm_cvtss_f32(s)
5044 }
5045}
5046
5047#[cfg(target_arch = "x86_64")]
5049#[target_feature(enable = "avx2,fma,avx512f,avx512bw,avx512vl,avx512vnni")]
5050unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5051 unsafe {
5053 use core::arch::x86_64::*;
5054 let lomask = _mm_set1_epi8(0x0F);
5055 let eight = _mm256_set1_epi8(8);
5056 let mut f0 = _mm256_setzero_ps();
5059 let mut f1 = _mm256_setzero_ps();
5060 let mut f2 = _mm256_setzero_ps();
5061 let mut f3 = _mm256_setzero_ps();
5062 for gi in 0..gpr {
5063 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5064 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5065 let sv = _mm256_set1_ps(s);
5066 let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5067 let lo = _mm_and_si128(bb, lomask);
5068 let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5069 let w = _mm256_sub_epi8(
5070 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5071 eight,
5072 );
5073 let aw = _mm256_abs_epi8(w);
5074 let off = gi * GROUP_SIZE;
5075 let dot = |xq: &[i8]| {
5076 let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5077 _mm256_cvtepi32_ps(_mm256_dpbusd_epi32(
5078 _mm256_setzero_si256(),
5079 aw,
5080 _mm256_sign_epi8(x, w),
5081 ))
5082 };
5083 f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5084 f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5085 f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5086 f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5087 }
5088 let acc = [
5089 hsum256_ps(f0),
5090 hsum256_ps(f1),
5091 hsum256_ps(f2),
5092 hsum256_ps(f3),
5093 ];
5094 acc
5095 }
5096}
5097
5098#[cfg(target_arch = "aarch64")]
5103#[target_feature(enable = "neon,dotprod")]
5104unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5105 unsafe {
5107 use core::arch::aarch64::*;
5108 use core::arch::asm;
5109 let lomask = vdupq_n_u8(0x0F);
5110 let eight = vdupq_n_s8(8);
5111 let mut acc = [0f32; 4];
5112 for gi in 0..gpr {
5113 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5114 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5115 let b = vld1q_u8(t.add(2));
5116 let lo = vandq_u8(b, lomask);
5117 let hi = vshrq_n_u8::<4>(b);
5118 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5119 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5120 for (k, xq) in xs.iter().enumerate() {
5121 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5122 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5123 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5124 asm!(
5125 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5126 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5127 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5128 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5129 options(pure, nomem, nostack),
5130 );
5131 acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5132 }
5133 }
5134 acc
5135 }
5136}
5137
5138#[inline]
5140fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
5141 let gi = j / GROUP_SIZE;
5142 let k = j % GROUP_SIZE;
5143 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5144 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5145 let byte = tile[2 + k / 2];
5146 let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5147 ((nib as i32 - 8) as f32, s)
5148}
5149
5150#[inline]
5153fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
5154 let mut acc = 0f32;
5155 for gi in 0..gpr {
5156 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5157 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5158 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5159 let mut ga = 0f32;
5160 for (k, &b) in tile[2..].iter().enumerate() {
5161 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5162 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5163 }
5164 acc += ga * s;
5165 }
5166 acc
5167}
5168
5169struct Q4tpView<'a> {
5173 nib: &'a [u8],
5174 params: &'a [u8],
5175 codes: &'a [u8],
5176 stride: usize,
5177 zero_rung: bool,
5179}
5180
5181impl<'a> Q4tpView<'a> {
5182 fn new(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5183 let (params_off, codes_off, stride) = q4tp_sections(rows, cols);
5184 Self {
5185 nib: &bytes[..params_off],
5186 params: &bytes[params_off..codes_off],
5187 codes: &bytes[codes_off..],
5188 stride,
5189 zero_rung: false,
5190 }
5191 }
5192
5193 fn new_q2(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5195 let (params_off, codes_off, stride) = q2tp_sections(rows, cols);
5196 Self {
5197 nib: &bytes[..params_off],
5198 params: &bytes[params_off..codes_off],
5199 codes: &bytes[codes_off..],
5200 stride,
5201 zero_rung: true,
5202 }
5203 }
5204
5205 #[inline]
5221 fn scales_into(&self, r: usize, gpr: usize, out: &mut [f32]) {
5222 let tab = if self.zero_rung {
5223 q2tp_ladder(self.params, r)
5224 } else {
5225 q4tp_ladder(self.params, r)
5226 };
5227 let codes = &self.codes[r * self.stride..(r + 1) * self.stride];
5228 let out = &mut out[..gpr];
5229 let mut chunks = out.chunks_exact_mut(8);
5230 let mut ci = 0usize;
5231 for c in &mut chunks {
5232 let w = u64::from(codes[ci])
5233 | u64::from(codes[ci + 1]) << 8
5234 | u64::from(codes[ci + 2]) << 16
5235 | u64::from(codes[ci + 3]) << 24
5236 | u64::from(codes[ci + 4]) << 32;
5237 for (k, o) in c.iter_mut().enumerate() {
5238 *o = tab[((w >> (5 * k)) & 31) as usize];
5239 }
5240 ci += 5;
5241 }
5242 let tail = &codes[ci..];
5245 for (k, o) in chunks.into_remainder().iter_mut().enumerate() {
5246 *o = tab[q4tp_code(tail, k)];
5247 }
5248 }
5249}
5250
5251#[inline]
5252fn dot_q4tp_row_i8(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5253 #[cfg(target_arch = "aarch64")]
5254 unsafe {
5255 return dot_q4tp_row_sdot(nib, r, gpr, xq, scales);
5256 }
5257 #[cfg(target_arch = "x86_64")]
5258 unsafe {
5259 if vnni_tiles_enabled() {
5260 return dot_q4tp_row_vnni(nib, r, gpr, xq, scales);
5261 }
5262 return dot_q4tp_row_avx2(nib, r, gpr, xq, scales);
5263 }
5264 #[allow(unreachable_code)]
5265 {
5266 let mut acc = 0f32;
5267 for gi in 0..gpr {
5268 let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5269 let s = scales[gi];
5270 let mut d = 0i32;
5271 for (k, &b) in tile.iter().enumerate() {
5272 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5273 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5274 }
5275 acc += d as f32 * s;
5276 }
5277 acc
5278 }
5279}
5280
5281#[cfg(target_arch = "aarch64")]
5284#[target_feature(enable = "neon,dotprod")]
5285unsafe fn dot_q4tp_row_sdot(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5286 unsafe {
5289 use core::arch::aarch64::*;
5290 use core::arch::asm;
5291 let lomask = vdupq_n_u8(0x0F);
5292 let eight = vdupq_n_s8(8);
5293 let mut acc = 0f32;
5294 for gi in 0..gpr {
5295 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5296 let s = *scales.get_unchecked(gi);
5297 let b = vld1q_u8(t);
5298 let lo = vandq_u8(b, lomask);
5299 let hi = vshrq_n_u8::<4>(b);
5300 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5301 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5302 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5303 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5304 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5305 asm!(
5306 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5307 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5308 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5309 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5310 options(pure, nomem, nostack),
5311 );
5312 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5313 }
5314 acc
5315 }
5316}
5317
5318#[cfg(target_arch = "x86_64")]
5319#[target_feature(enable = "avx2")]
5320unsafe fn dot_q4tp_row_avx2(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5321 unsafe {
5323 use core::arch::x86_64::*;
5324 let lomask = _mm_set1_epi8(0x0F);
5325 let eight = _mm256_set1_epi8(8);
5326 let ones = _mm256_set1_epi16(1);
5327 let mut acc = 0f32;
5328 for gi in 0..gpr {
5329 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5330 let s = *scales.get_unchecked(gi);
5331 let b = _mm_loadu_si128(t as *const __m128i);
5332 let lo = _mm_and_si128(b, lomask);
5333 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5334 let w = _mm256_sub_epi8(
5335 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5336 eight,
5337 );
5338 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5339 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5340 let d = _mm256_madd_epi16(p16, ones);
5341 let hi128 = _mm256_extracti128_si256::<1>(d);
5342 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5343 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5344 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5345 acc += _mm_cvtsi128_si32(s32) as f32 * s;
5346 }
5347 acc
5348 }
5349}
5350
5351#[cfg(target_arch = "x86_64")]
5354#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5355unsafe fn dot_q4tp_row_vnni(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5356 unsafe {
5358 use core::arch::x86_64::*;
5359 let lomask = _mm_set1_epi8(0x0F);
5360 let eight = _mm256_set1_epi8(8);
5361 let mut acc = 0f32;
5362 for gi in 0..gpr {
5363 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5364 let s = *scales.get_unchecked(gi);
5365 let b = _mm_loadu_si128(t as *const __m128i);
5366 let lo = _mm_and_si128(b, lomask);
5367 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5368 let w = _mm256_sub_epi8(
5369 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5370 eight,
5371 );
5372 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5373 acc += dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w)) as f32 * s;
5374 }
5375 acc
5376 }
5377}
5378
5379#[inline]
5382fn q4tp_row_exact(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5383 #[cfg(target_arch = "x86_64")]
5384 if avx2_enabled() {
5385 return unsafe { q4tp_row_float_avx2(nib, r, gpr, x, scales) };
5387 }
5388 q4tp_row_float_scalar(nib, r, gpr, x, scales)
5389}
5390
5391#[inline]
5392fn q4tp_row_float_scalar(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5393 let mut acc = 0f32;
5394 for gi in 0..gpr {
5395 let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5396 let s = scales[gi];
5397 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5398 let mut ga = 0f32;
5399 for (k, &b) in tile.iter().enumerate() {
5400 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5401 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5402 }
5403 acc += ga * s;
5404 }
5405 acc
5406}
5407
5408#[cfg(target_arch = "x86_64")]
5412#[target_feature(enable = "avx2")]
5413unsafe fn q4tp_row_float_avx2(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5414 unsafe {
5417 use core::arch::x86_64::*;
5418 let mask = _mm_set1_epi8(15);
5419 let eight = _mm_set1_epi8(8);
5420 let order = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
5421 let mut acc = 0.0f32;
5422 for gi in 0..gpr {
5423 let packed = _mm_loadu_si128(nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB).cast());
5424 let lo = _mm_and_si128(packed, mask);
5425 let hi = _mm_and_si128(_mm_srli_epi16::<4>(packed), mask);
5426 let w0 = _mm_sub_epi8(_mm_unpacklo_epi8(lo, hi), eight);
5427 let w1 = _mm_sub_epi8(_mm_unpackhi_epi8(lo, hi), eight);
5428 let xp = x.as_ptr().add(gi * GROUP_SIZE);
5429 let a = _mm256_mul_ps(
5430 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w0)),
5431 _mm256_loadu_ps(xp),
5432 );
5433 let b = _mm256_mul_ps(
5434 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w0))),
5435 _mm256_loadu_ps(xp.add(8)),
5436 );
5437 let c = _mm256_mul_ps(
5438 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w1)),
5439 _mm256_loadu_ps(xp.add(16)),
5440 );
5441 let d = _mm256_mul_ps(
5442 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w1))),
5443 _mm256_loadu_ps(xp.add(24)),
5444 );
5445 let mut pairs = [0.0f32; 16];
5446 _mm256_storeu_ps(
5447 pairs.as_mut_ptr(),
5448 _mm256_permutevar8x32_ps(_mm256_hadd_ps(a, b), order),
5449 );
5450 _mm256_storeu_ps(
5451 pairs.as_mut_ptr().add(8),
5452 _mm256_permutevar8x32_ps(_mm256_hadd_ps(c, d), order),
5453 );
5454 let mut ga = 0.0f32;
5455 for v in pairs {
5456 ga += v;
5457 }
5458 acc += ga * scales[gi];
5459 }
5460 acc
5461 }
5462}
5463
5464#[inline]
5467fn q4tp_outlier(nib: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5468 let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5469 let byte = nib[(r * gpr + gi) * Q4TP_NIB + k / 2];
5470 let n = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5471 ((n as i32 - 8) as f32, scales[gi])
5472}
5473
5474fn q4tp_matvec(
5476 bytes: &[u8],
5477 x: &[f32],
5478 rows: usize,
5479 cols: usize,
5480 out: &mut [f32],
5481 pool: Option<&Pool>,
5482) {
5483 debug_assert_eq!(out.len(), rows);
5484 let gpr = cols / GROUP_SIZE;
5485 let v = Q4tpView::new(bytes, rows, cols);
5486 let out_addr = SendMut(out.as_mut_ptr());
5487 if a8w8_enabled() {
5488 let act = split_act(x);
5489 let run = |start: usize, end: usize| {
5490 with_krow(gpr, |sc| {
5492 for r in start..end {
5493 v.scales_into(r, gpr, sc);
5494 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
5495 for &(j, xv) in &act.outliers {
5496 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
5497 acc += w * s * xv;
5498 }
5499 unsafe { *out_addr.at(r) = acc };
5501 }
5502 })
5503 };
5504 dispatch_rows(pool, rows, &run);
5505 return;
5506 }
5507 let run = |start: usize, end: usize| {
5508 with_krow(gpr, |sc| {
5509 for r in start..end {
5510 v.scales_into(r, gpr, sc);
5511 unsafe { *out_addr.at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
5513 }
5514 })
5515 };
5516 dispatch_rows(pool, rows, &run);
5517}
5518
5519#[allow(clippy::too_many_arguments)]
5522fn q4tp_matvec2(
5523 bytes: &[u8],
5524 x1: &[f32],
5525 x2: &[f32],
5526 rows: usize,
5527 cols: usize,
5528 o1: &mut [f32],
5529 o2: &mut [f32],
5530 pool: Option<&Pool>,
5531) {
5532 let gpr = cols / GROUP_SIZE;
5533 let v = Q4tpView::new(bytes, rows, cols);
5534 let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5535 let run = |start: usize, end: usize| {
5536 let mut sc = vec![0f32; gpr];
5537 for r in start..end {
5538 v.scales_into(r, gpr, &mut sc);
5539 unsafe {
5541 *p1.at(r) = q4tp_row_exact(v.nib, r, gpr, x1, &sc);
5542 *p2.at(r) = q4tp_row_exact(v.nib, r, gpr, x2, &sc);
5543 }
5544 }
5545 };
5546 dispatch_rows(pool, rows, &run);
5547}
5548
5549#[inline]
5552fn q2tp_outlier(chunks: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5553 let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5554 let byte = chunks[(r * gpr + gi) * Q2TP_CHUNK + k / 4];
5555 let c = (byte >> (2 * (k % 4))) & 3;
5556 (c as f32 - 1.5, scales[gi])
5557}
5558
5559#[cfg(target_arch = "x86_64")]
5560const Q2TP_DECODE_U32: [u32; 256] = {
5561 let mut tab = [0u32; 256];
5562 let mut b = 0usize;
5563 while b < 256 {
5564 tab[b] = ((b as u32) & 3)
5565 | ((((b as u32) >> 2) & 3) << 8)
5566 | ((((b as u32) >> 4) & 3) << 16)
5567 | ((((b as u32) >> 6) & 3) << 24);
5568 b += 1;
5569 }
5570 tab
5571};
5572
5573#[cfg(target_arch = "x86_64")]
5577#[target_feature(enable = "avx2")]
5578unsafe fn q2tp_code_dot_avx2(ch: &[u8], x: &[i8]) -> i32 {
5579 use core::arch::x86_64::*;
5580 debug_assert!(ch.len() >= Q2TP_CHUNK && x.len() >= GROUP_SIZE);
5581 let codes = _mm256_setr_epi32(
5582 Q2TP_DECODE_U32[ch[0] as usize] as i32,
5583 Q2TP_DECODE_U32[ch[1] as usize] as i32,
5584 Q2TP_DECODE_U32[ch[2] as usize] as i32,
5585 Q2TP_DECODE_U32[ch[3] as usize] as i32,
5586 Q2TP_DECODE_U32[ch[4] as usize] as i32,
5587 Q2TP_DECODE_U32[ch[5] as usize] as i32,
5588 Q2TP_DECODE_U32[ch[6] as usize] as i32,
5589 Q2TP_DECODE_U32[ch[7] as usize] as i32,
5590 );
5591 let xv = unsafe { _mm256_loadu_si256(x.as_ptr().cast()) };
5592 let pair = _mm256_maddubs_epi16(codes, xv);
5593 let quad = _mm256_madd_epi16(pair, _mm256_set1_epi16(1));
5594 let sum128 = _mm_add_epi32(
5595 _mm256_castsi256_si128(quad),
5596 _mm256_extracti128_si256(quad, 1),
5597 );
5598 let sum64 = _mm_hadd_epi32(sum128, sum128);
5599 _mm_cvtsi128_si32(_mm_hadd_epi32(sum64, sum64))
5600}
5601
5602#[inline]
5609fn dot_q2tp_row_i8(
5610 chunks: &[u8],
5611 r: usize,
5612 gpr: usize,
5613 xq: &[i8],
5614 gsum: &[i32],
5615 scales: &[f32],
5616) -> f32 {
5617 let mut acc = 0f32;
5618 let base = r * gpr * Q2TP_CHUNK;
5619 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5620 let mut codes = [0i8; GROUP_SIZE];
5621 #[cfg(target_arch = "x86_64")]
5622 let avx2 = std::arch::is_x86_feature_detected!("avx2");
5623 for gi in 0..gpr {
5624 let ch = &chunks[base + gi * Q2TP_CHUNK..base + (gi + 1) * Q2TP_CHUNK];
5625 let xg = &xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5626 #[cfg(target_arch = "aarch64")]
5627 let dot = unsafe {
5633 use core::arch::aarch64::*;
5634 let b = vld1_u8(ch.as_ptr());
5635 let three = vdup_n_u8(3);
5636 let c0 = vreinterpret_s8_u8(vand_u8(b, three));
5637 let c1 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 2), three));
5638 let c2 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 4), three));
5639 let c3 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 6), three));
5640 let x4 = vld4_s8(xg.as_ptr());
5641 let mut acc4 = vdupq_n_s32(0);
5642 acc4 = vpadalq_s16(acc4, vmull_s8(c0, x4.0));
5643 acc4 = vpadalq_s16(acc4, vmull_s8(c1, x4.1));
5644 acc4 = vpadalq_s16(acc4, vmull_s8(c2, x4.2));
5645 acc4 = vpadalq_s16(acc4, vmull_s8(c3, x4.3));
5646 vaddvq_s32(acc4)
5647 };
5648 #[cfg(target_arch = "x86_64")]
5649 let dot: i32 = if avx2 {
5650 unsafe { q2tp_code_dot_avx2(ch, xg) }
5653 } else {
5654 ch.iter()
5655 .enumerate()
5656 .map(|(k, &b)| {
5657 ((b & 3) as i32) * xg[k * 4] as i32
5658 + (((b >> 2) & 3) as i32) * xg[k * 4 + 1] as i32
5659 + (((b >> 4) & 3) as i32) * xg[k * 4 + 2] as i32
5660 + (((b >> 6) & 3) as i32) * xg[k * 4 + 3] as i32
5661 })
5662 .sum()
5663 };
5664 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5665 let dot: i32 = {
5666 for (k, &b) in ch.iter().enumerate() {
5667 codes[k * 4] = (b & 3) as i8;
5668 codes[k * 4 + 1] = ((b >> 2) & 3) as i8;
5669 codes[k * 4 + 2] = ((b >> 4) & 3) as i8;
5670 codes[k * 4 + 3] = ((b >> 6) & 3) as i8;
5671 }
5672 codes
5673 .iter()
5674 .zip(xg)
5675 .map(|(&c, &x)| c as i32 * x as i32)
5676 .sum()
5677 };
5678 acc += scales[gi] * (dot as f32 - 1.5 * gsum[gi] as f32);
5679 }
5680 acc
5681}
5682
5683fn q2tp_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5687 q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.5)
5688}
5689
5690#[inline]
5694fn q2tp_affine_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5695 q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.0)
5696}
5697
5698#[inline]
5699fn q2tp_row_exact_center(
5700 chunks: &[u8],
5701 r: usize,
5702 gpr: usize,
5703 x: &[f32],
5704 scales: &[f32],
5705 center: f32,
5706) -> f32 {
5707 let mut acc = 0f32;
5708 for gi in 0..gpr {
5709 let ch = &chunks[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
5710 let s = scales[gi];
5711 let xb = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5712 let mut g = 0f32;
5713 for (k, &b) in ch.iter().enumerate() {
5714 g += ((b & 3) as f32 - center) * xb[k * 4]
5715 + (((b >> 2) & 3) as f32 - center) * xb[k * 4 + 1]
5716 + (((b >> 4) & 3) as f32 - center) * xb[k * 4 + 2]
5717 + (((b >> 6) & 3) as f32 - center) * xb[k * 4 + 3];
5718 }
5719 acc += s * g;
5720 }
5721 acc
5722}
5723
5724fn q2tp_matvec(
5725 bytes: &[u8],
5726 x: &[f32],
5727 rows: usize,
5728 cols: usize,
5729 out: &mut [f32],
5730 pool: Option<&Pool>,
5731) {
5732 q2tp_matvec_mode(bytes, x, rows, cols, out, pool, false);
5733}
5734
5735fn q2tp_affine_matvec(
5736 bytes: &[u8],
5737 x: &[f32],
5738 rows: usize,
5739 cols: usize,
5740 out: &mut [f32],
5741 pool: Option<&Pool>,
5742) {
5743 q2tp_matvec_mode(bytes, x, rows, cols, out, pool, true);
5744}
5745
5746fn q2tp_matvec_mode(
5747 bytes: &[u8],
5748 x: &[f32],
5749 rows: usize,
5750 cols: usize,
5751 out: &mut [f32],
5752 pool: Option<&Pool>,
5753 affine: bool,
5754) {
5755 debug_assert_eq!(out.len(), rows);
5756 let gpr = cols / GROUP_SIZE;
5757 let v = Q4tpView::new_q2(bytes, rows, cols);
5758 let out_addr = SendMut(out.as_mut_ptr());
5759 if !affine && a8w8_enabled() {
5764 let act = split_act(x);
5765 let gsum = q1_group_sums(&act.xq, gpr);
5766 let (act, gsum) = (&act, &gsum);
5767 let run = move |start: usize, end: usize| {
5768 with_krow(gpr, |sc| {
5769 for r in start..end {
5770 v.scales_into(r, gpr, sc);
5771 let mut acc = dot_q2tp_row_i8(v.nib, r, gpr, &act.xq, gsum, sc) * act.sx;
5772 for &(j, xv) in &act.outliers {
5773 let (w, s) = q2tp_outlier(v.nib, r, gpr, j, sc);
5774 acc += w * s * xv;
5775 }
5776 unsafe { *out_addr.at(r) = acc };
5778 }
5779 })
5780 };
5781 dispatch_rows(pool, rows, &run);
5782 return;
5783 }
5784 let run = |start: usize, end: usize| {
5785 with_krow(gpr, |sc| {
5786 for r in start..end {
5787 v.scales_into(r, gpr, sc);
5788 unsafe {
5790 *out_addr.at(r) = if affine {
5791 q2tp_affine_row_exact(v.nib, r, gpr, x, sc)
5792 } else {
5793 q2tp_row_exact(v.nib, r, gpr, x, sc)
5794 }
5795 };
5796 }
5797 })
5798 };
5799 dispatch_rows(pool, rows, &run);
5800}
5801
5802#[allow(clippy::too_many_arguments)]
5804fn q2tp_matvec2(
5805 bytes: &[u8],
5806 x1: &[f32],
5807 x2: &[f32],
5808 rows: usize,
5809 cols: usize,
5810 o1: &mut [f32],
5811 o2: &mut [f32],
5812 pool: Option<&Pool>,
5813) {
5814 q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, false);
5815}
5816
5817#[allow(clippy::too_many_arguments)]
5818fn q2tp_affine_matvec2(
5819 bytes: &[u8],
5820 x1: &[f32],
5821 x2: &[f32],
5822 rows: usize,
5823 cols: usize,
5824 o1: &mut [f32],
5825 o2: &mut [f32],
5826 pool: Option<&Pool>,
5827) {
5828 q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, true);
5829}
5830
5831#[allow(clippy::too_many_arguments)]
5832fn q2tp_matvec2_mode(
5833 bytes: &[u8],
5834 x1: &[f32],
5835 x2: &[f32],
5836 rows: usize,
5837 cols: usize,
5838 o1: &mut [f32],
5839 o2: &mut [f32],
5840 pool: Option<&Pool>,
5841 affine: bool,
5842) {
5843 let gpr = cols / GROUP_SIZE;
5844 let v = Q4tpView::new_q2(bytes, rows, cols);
5845 let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5846 let run = |start: usize, end: usize| {
5847 let mut sc = vec![0f32; gpr];
5848 for r in start..end {
5849 v.scales_into(r, gpr, &mut sc);
5850 unsafe {
5852 *p1.at(r) = if affine {
5853 q2tp_affine_row_exact(v.nib, r, gpr, x1, &sc)
5854 } else {
5855 q2tp_row_exact(v.nib, r, gpr, x1, &sc)
5856 };
5857 *p2.at(r) = if affine {
5858 q2tp_affine_row_exact(v.nib, r, gpr, x2, &sc)
5859 } else {
5860 q2tp_row_exact(v.nib, r, gpr, x2, &sc)
5861 };
5862 }
5863 }
5864 };
5865 dispatch_rows(pool, rows, &run);
5866}
5867
5868pub fn q2tp_matvec_for_test(bytes: &[u8], x: &[f32], rows: usize, cols: usize, out: &mut [f32]) {
5875 let gpr = cols / GROUP_SIZE;
5880 let v = Q4tpView::new_q2(bytes, rows, cols);
5881 with_krow(gpr, |sc| {
5882 for r in 0..rows {
5883 v.scales_into(r, gpr, sc);
5884 out[r] = q2tp_row_exact(v.nib, r, gpr, x, sc);
5885 }
5886 });
5887}
5888
5889pub fn q2tp_affine_matvec_for_test(
5892 bytes: &[u8],
5893 x: &[f32],
5894 rows: usize,
5895 cols: usize,
5896 out: &mut [f32],
5897) {
5898 q2tp_affine_matvec(bytes, x, rows, cols, out, None);
5899}
5900
5901pub fn q2tp_matmat_for_test(
5902 bytes: &[u8],
5903 xs_all: &[f32],
5904 b: usize,
5905 rows: usize,
5906 cols: usize,
5907 out: &mut [f32],
5908) {
5909 q2tp_matmat(bytes, xs_all, b, rows, cols, out, None);
5910}
5911
5912fn q2tp_matmat(
5913 bytes: &[u8],
5914 xs_all: &[f32],
5915 b: usize,
5916 rows: usize,
5917 cols: usize,
5918 out: &mut [f32],
5919 pool: Option<&Pool>,
5920) {
5921 q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, false);
5922}
5923
5924fn q2tp_affine_matmat(
5925 bytes: &[u8],
5926 xs_all: &[f32],
5927 b: usize,
5928 rows: usize,
5929 cols: usize,
5930 out: &mut [f32],
5931 pool: Option<&Pool>,
5932) {
5933 q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, true);
5934}
5935
5936fn q2tp_matmat_mode(
5937 bytes: &[u8],
5938 xs_all: &[f32],
5939 b: usize,
5940 rows: usize,
5941 cols: usize,
5942 out: &mut [f32],
5943 pool: Option<&Pool>,
5944 affine: bool,
5945) {
5946 debug_assert_eq!(out.len(), b * rows);
5947 let gpr = cols / GROUP_SIZE;
5948 let v = Q4tpView::new_q2(bytes, rows, cols);
5949 let out_addr = SendMut(out.as_mut_ptr());
5950 let run = |start: usize, end: usize| {
5951 let mut sc = vec![0f32; gpr];
5952 for r in start..end {
5953 v.scales_into(r, gpr, &mut sc);
5954 for bi in 0..b {
5955 let x = &xs_all[bi * cols..(bi + 1) * cols];
5956 unsafe {
5958 *out_addr.at(bi * rows + r) = if affine {
5959 q2tp_affine_row_exact(v.nib, r, gpr, x, &sc)
5960 } else {
5961 q2tp_row_exact(v.nib, r, gpr, x, &sc)
5962 }
5963 };
5964 }
5965 }
5966 };
5967 dispatch_rows(pool, rows, &run);
5968}
5969
5970#[cfg(target_arch = "aarch64")]
5978#[target_feature(enable = "neon,dotprod")]
5979unsafe fn dot_q4tp_row_1x4_sdot_v1(
5980 nib: &[u8],
5981 r: usize,
5982 gpr: usize,
5983 xs: [&[i8]; 4],
5984 scales: &[f32],
5985) -> [f32; 4] {
5986 unsafe {
5987 use core::arch::aarch64::*;
5988 use core::arch::asm;
5989 let lomask = vdupq_n_u8(0x0F);
5990 let eight = vdupq_n_s8(8);
5991 let (mut f0, mut f1, mut f2, mut f3) = (0f32, 0f32, 0f32, 0f32);
5992 for gi in 0..gpr {
5993 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5994 let s = *scales.get_unchecked(gi);
5995 let bb = vld1q_u8(t);
5996 let lo = vandq_u8(bb, lomask);
5997 let hi = vshrq_n_u8::<4>(bb);
5998 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5999 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6000 let mut d = [0f32; 4];
6001 for (k, dk) in d.iter_mut().enumerate() {
6002 let x0 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE));
6003 let x1 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE + 16));
6004 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6005 asm!(
6006 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6007 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6008 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6009 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6010 options(pure, nomem, nostack),
6011 );
6012 *dk = vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
6013 }
6014 f0 += d[0];
6015 f1 += d[1];
6016 f2 += d[2];
6017 f3 += d[3];
6018 }
6019 [f0, f1, f2, f3]
6020 }
6021}
6022
6023#[allow(dead_code)]
6030static Q4TP_ALT: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
6031
6032#[cfg(test)]
6035static Q4TP_ALT_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
6036
6037#[cfg(target_arch = "x86_64")]
6043fn q4tp_blocked_x86() -> bool {
6044 match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6045 1 => false,
6046 2 => avx512vnni_enabled(),
6051 _ => blocked_enabled() && avx512vnni_enabled(),
6055 }
6056}
6057
6058#[cfg(target_arch = "aarch64")]
6060#[allow(dead_code)]
6061fn q4tp_v1() -> bool {
6062 match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6063 1 => true,
6064 2 => false,
6065 _ => {
6066 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6067 *ON.get_or_init(|| std::env::var("CMF_Q4TP_V1").is_ok_and(|v| v != "0"))
6068 }
6069 }
6070}
6071
6072#[cfg(target_arch = "x86_64")]
6084#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6085unsafe fn dot_q4tp_2x8_avx512(
6086 nib: &[u8],
6087 r0: usize,
6088 gpr: usize,
6089 xs: [&[i8]; 8],
6090 sc0: &[f32],
6091 sc1: &[f32],
6092) -> [[f32; 8]; 2] {
6093 unsafe {
6096 use core::arch::x86_64::*;
6097 let lomask = _mm256_set1_epi8(0x0F);
6098 let eight = _mm256_set1_epi8(8);
6099 let zero = _mm512_setzero_si512();
6100 let mut v0 = [_mm512_setzero_ps(); 8];
6101 let mut v1 = [_mm512_setzero_ps(); 8];
6102 let pairs = gpr / 2;
6103 let unpack = |r: usize, gi: usize| -> (__m512i, __mmask64) {
6104 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6105 let bb = _mm256_loadu_si256(t as *const __m256i);
6106 let lo = _mm256_and_si256(bb, lomask);
6107 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6108 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6109 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6110 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6111 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6112 (_mm512_abs_epi8(w), _mm512_movepi8_mask(w))
6113 };
6114 for gp in 0..pairs {
6115 let gi = gp * 2;
6116 let (wa0, neg0) = unpack(r0, gi);
6117 let (wa1, neg1) = unpack(r0 + 1, gi);
6118 let off = gi * GROUP_SIZE;
6119 let sv = |sc: &[f32]| {
6120 _mm512_insertf32x8::<1>(
6121 _mm512_castps256_ps512(_mm256_set1_ps(*sc.get_unchecked(gi))),
6122 _mm256_set1_ps(*sc.get_unchecked(gi + 1)),
6123 )
6124 };
6125 let s0 = sv(sc0);
6126 let s1 = sv(sc1);
6127 for k in 0..8 {
6128 let xv = _mm512_loadu_si512(xs[k].as_ptr().add(off) as *const __m512i);
6129 let d0 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6130 zero,
6131 wa0,
6132 _mm512_mask_sub_epi8(xv, neg0, zero, xv),
6133 ));
6134 let d1 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6135 zero,
6136 wa1,
6137 _mm512_mask_sub_epi8(xv, neg1, zero, xv),
6138 ));
6139 v0[k] = _mm512_fmadd_ps(d0, s0, v0[k]);
6140 v1[k] = _mm512_fmadd_ps(d1, s1, v1[k]);
6141 }
6142 }
6143 let mut acc = [[0f32; 8]; 2];
6144 for k in 0..8 {
6145 acc[0][k] = _mm512_reduce_add_ps(v0[k]);
6146 acc[1][k] = _mm512_reduce_add_ps(v1[k]);
6147 }
6148 if gpr % 2 == 1 {
6149 let off = (gpr - 1) * GROUP_SIZE;
6150 for j in off..off + GROUP_SIZE {
6151 let (w0, sa) = q4tp_outlier(nib, r0, gpr, j, sc0);
6152 let (w1, sb) = q4tp_outlier(nib, r0 + 1, gpr, j, sc1);
6153 for k in 0..8 {
6154 let x = *xs[k].get_unchecked(j) as f32;
6155 acc[0][k] += w0 * sa * x;
6156 acc[1][k] += w1 * sb * x;
6157 }
6158 }
6159 }
6160 acc
6161 }
6162}
6163
6164#[cfg(target_arch = "x86_64")]
6169#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6170unsafe fn dot_q4tp_row_1x8_avx512(
6171 nib: &[u8],
6172 r: usize,
6173 gpr: usize,
6174 xs: [&[i8]; 8],
6175 scales: &[f32],
6176) -> [f32; 8] {
6177 unsafe {
6179 use core::arch::x86_64::*;
6180 let lomask = _mm256_set1_epi8(0x0F);
6181 let eight = _mm256_set1_epi8(8);
6182 let zero = _mm512_setzero_si512();
6183 let (mut v0, mut v1, mut v2, mut v3) = (
6184 _mm512_setzero_ps(),
6185 _mm512_setzero_ps(),
6186 _mm512_setzero_ps(),
6187 _mm512_setzero_ps(),
6188 );
6189 let (mut v4, mut v5, mut v6, mut v7) = (
6190 _mm512_setzero_ps(),
6191 _mm512_setzero_ps(),
6192 _mm512_setzero_ps(),
6193 _mm512_setzero_ps(),
6194 );
6195 let pairs = gpr / 2;
6196 for gp in 0..pairs {
6197 let gi = gp * 2;
6198 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6199 let bb = _mm256_loadu_si256(t as *const __m256i);
6200 let lo = _mm256_and_si256(bb, lomask);
6201 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6202 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6207 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6208 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6209 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6210 let wabs = _mm512_abs_epi8(w);
6211 let neg = _mm512_movepi8_mask(w);
6212 let off = gi * GROUP_SIZE;
6213 let sv = _mm512_insertf32x8::<1>(
6214 _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6215 _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6216 );
6217 let dot = |x: &[i8]| -> __m512 {
6218 let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6219 let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6220 _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6221 };
6222 v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6223 v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6224 v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6225 v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6226 v4 = _mm512_fmadd_ps(dot(xs[4]), sv, v4);
6227 v5 = _mm512_fmadd_ps(dot(xs[5]), sv, v5);
6228 v6 = _mm512_fmadd_ps(dot(xs[6]), sv, v6);
6229 v7 = _mm512_fmadd_ps(dot(xs[7]), sv, v7);
6230 }
6231 let mut acc = [
6232 _mm512_reduce_add_ps(v0),
6233 _mm512_reduce_add_ps(v1),
6234 _mm512_reduce_add_ps(v2),
6235 _mm512_reduce_add_ps(v3),
6236 _mm512_reduce_add_ps(v4),
6237 _mm512_reduce_add_ps(v5),
6238 _mm512_reduce_add_ps(v6),
6239 _mm512_reduce_add_ps(v7),
6240 ];
6241 if gpr % 2 == 1 {
6244 let off = (gpr - 1) * GROUP_SIZE;
6245 for j in off..off + GROUP_SIZE {
6246 let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6247 let ws = w * s;
6248 for k in 0..8 {
6249 acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6250 }
6251 }
6252 }
6253 acc
6254 }
6255}
6256
6257#[cfg(target_arch = "x86_64")]
6270#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6271unsafe fn dot_q4tp_row_1x4_avx512(
6272 nib: &[u8],
6273 r: usize,
6274 gpr: usize,
6275 xs: [&[i8]; 4],
6276 scales: &[f32],
6277) -> [f32; 4] {
6278 unsafe {
6280 use core::arch::x86_64::*;
6281 let lomask = _mm256_set1_epi8(0x0F);
6282 let eight = _mm256_set1_epi8(8);
6283 let zero = _mm512_setzero_si512();
6284 let (mut v0, mut v1, mut v2, mut v3) = (
6285 _mm512_setzero_ps(),
6286 _mm512_setzero_ps(),
6287 _mm512_setzero_ps(),
6288 _mm512_setzero_ps(),
6289 );
6290 let pairs = gpr / 2;
6291 for gp in 0..pairs {
6292 let gi = gp * 2;
6293 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6294 let bb = _mm256_loadu_si256(t as *const __m256i);
6295 let lo = _mm256_and_si256(bb, lomask);
6296 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6297 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6302 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6303 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6304 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6305 let wabs = _mm512_abs_epi8(w);
6306 let neg = _mm512_movepi8_mask(w);
6307 let off = gi * GROUP_SIZE;
6308 let sv = _mm512_insertf32x8::<1>(
6309 _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6310 _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6311 );
6312 let dot = |x: &[i8]| -> __m512 {
6313 let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6314 let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6315 _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6316 };
6317 v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6318 v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6319 v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6320 v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6321 }
6322 let mut acc = [
6323 _mm512_reduce_add_ps(v0),
6324 _mm512_reduce_add_ps(v1),
6325 _mm512_reduce_add_ps(v2),
6326 _mm512_reduce_add_ps(v3),
6327 ];
6328 if gpr % 2 == 1 {
6331 let off = (gpr - 1) * GROUP_SIZE;
6332 for j in off..off + GROUP_SIZE {
6333 let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6334 let ws = w * s;
6335 for k in 0..4 {
6336 acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6337 }
6338 }
6339 }
6340 acc
6341 }
6342}
6343
6344#[cfg(target_arch = "aarch64")]
6348#[target_feature(enable = "neon,dotprod")]
6349unsafe fn dot_q4tp_row_1x4_sdot(
6350 nib: &[u8],
6351 r: usize,
6352 gpr: usize,
6353 xs: [&[i8]; 4],
6354 scales: &[f32],
6355) -> [f32; 4] {
6356 unsafe {
6358 use core::arch::aarch64::*;
6359 use core::arch::asm;
6360 let lomask = vdupq_n_u8(0x0F);
6361 let eight = vdupq_n_s8(8);
6362 let (mut v0, mut v1, mut v2, mut v3) = (
6378 vdupq_n_f32(0.0),
6379 vdupq_n_f32(0.0),
6380 vdupq_n_f32(0.0),
6381 vdupq_n_f32(0.0),
6382 );
6383 for gi in 0..gpr {
6384 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6385 let s = *scales.get_unchecked(gi);
6386 let bb = vld1q_u8(t);
6387 let lo = vandq_u8(bb, lomask);
6388 let hi = vshrq_n_u8::<4>(bb);
6389 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6390 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6391 let off = gi * GROUP_SIZE;
6392 let dot4 = |x: &[i8]| -> int32x4_t {
6393 let x0 = vld1q_s8(x.as_ptr().add(off));
6394 let x1 = vld1q_s8(x.as_ptr().add(off + 16));
6395 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6396 asm!(
6397 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6398 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6399 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6400 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6401 options(pure, nomem, nostack),
6402 );
6403 vaddq_s32(a0, a1)
6404 };
6405 v0 = vfmaq_n_f32(v0, vcvtq_f32_s32(dot4(xs[0])), s);
6406 v1 = vfmaq_n_f32(v1, vcvtq_f32_s32(dot4(xs[1])), s);
6407 v2 = vfmaq_n_f32(v2, vcvtq_f32_s32(dot4(xs[2])), s);
6408 v3 = vfmaq_n_f32(v3, vcvtq_f32_s32(dot4(xs[3])), s);
6409 }
6410 [
6411 vaddvq_f32(v0),
6412 vaddvq_f32(v1),
6413 vaddvq_f32(v2),
6414 vaddvq_f32(v3),
6415 ]
6416 }
6417}
6418
6419fn q4tp_matmat(
6427 bytes: &[u8],
6428 xs_all: &[f32],
6429 b: usize,
6430 rows: usize,
6431 cols: usize,
6432 out: &mut [f32],
6433 pool: Option<&Pool>,
6434) {
6435 q4tp_matmat_with(bytes, xs_all, b, rows, cols, out, pool, row_exact())
6436}
6437
6438#[allow(clippy::too_many_arguments)]
6442fn q4tp_matmat_with(
6443 bytes: &[u8],
6444 xs_all: &[f32],
6445 b: usize,
6446 rows: usize,
6447 cols: usize,
6448 out: &mut [f32],
6449 pool: Option<&Pool>,
6450 exact: bool,
6451) {
6452 debug_assert_eq!(out.len(), b * rows);
6453 let gpr = cols / GROUP_SIZE;
6454 let v = Q4tpView::new(bytes, rows, cols);
6455
6456 #[cfg(target_os = "macos")]
6460 if !exact && b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6461 dequant_matmat_accel(
6462 &|r, dst| {
6463 let mut sc = [0f32; 32];
6464 let mut scv;
6465 let s: &[f32] = if gpr <= 32 {
6466 v.scales_into(r, gpr, &mut sc);
6467 &sc[..gpr]
6468 } else {
6469 scv = vec![0f32; gpr];
6470 v.scales_into(r, gpr, &mut scv);
6471 &scv
6472 };
6473 for gi in 0..gpr {
6474 let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
6475 for (k, &bb) in tile.iter().enumerate() {
6476 dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s[gi];
6477 dst[gi * GROUP_SIZE + k * 2 + 1] =
6478 (((bb >> 4) & 0x0F) as f32 - 8.0) * s[gi];
6479 }
6480 }
6481 },
6482 xs_all,
6483 b,
6484 rows,
6485 cols,
6486 out,
6487 pool,
6488 );
6489 return;
6490 }
6491
6492 let out_addr = SendMut(out.as_mut_ptr());
6493 if a8w8_enabled() {
6494 let acts: Vec<SplitAct> = (0..b)
6495 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6496 .collect();
6497 let acts = &acts;
6498 #[cfg(target_arch = "aarch64")]
6505 let blocked_ok = sdot_enabled() && blocked_enabled();
6506 #[cfg(target_arch = "x86_64")]
6513 let blocked_ok = q4tp_blocked_x86() && !exact;
6514 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
6515 let blocked_ok = {
6516 let _ = exact;
6517 false
6518 };
6519 let panel_cols: usize = std::env::var("CMF_Q4TP_PANEL")
6529 .ok()
6530 .and_then(|v| v.parse().ok())
6531 .filter(|v| *v > 0)
6532 .unwrap_or(256);
6533 let run = |start: usize, end: usize| {
6534 for abase in (0..acts.len()).step_by(panel_cols) {
6535 let alen = (acts.len() - abase).min(panel_cols);
6536 let mut sc = vec![0f32; gpr];
6537 #[cfg(target_arch = "x86_64")]
6538 let mut r_lo = start;
6539 #[cfg(target_arch = "x86_64")]
6540 if blocked_ok && alen >= 8 {
6541 let mut sc1 = vec![0f32; gpr];
6542 while r_lo + 2 <= end {
6543 v.scales_into(r_lo, gpr, &mut sc);
6544 v.scales_into(r_lo + 1, gpr, &mut sc1);
6545 let mut bi = 0usize;
6546 while bi + 8 <= alen {
6547 let xs = [
6548 acts[abase + bi].xq.as_slice(),
6549 acts[abase + bi + 1].xq.as_slice(),
6550 acts[abase + bi + 2].xq.as_slice(),
6551 acts[abase + bi + 3].xq.as_slice(),
6552 acts[abase + bi + 4].xq.as_slice(),
6553 acts[abase + bi + 5].xq.as_slice(),
6554 acts[abase + bi + 6].xq.as_slice(),
6555 acts[abase + bi + 7].xq.as_slice(),
6556 ];
6557 let d = unsafe { dot_q4tp_2x8_avx512(v.nib, r_lo, gpr, xs, &sc, &sc1) };
6558 for (row, dr, scr) in [(r_lo, &d[0], &sc), (r_lo + 1, &d[1], &sc1)] {
6559 for k in 0..8 {
6560 let act = &acts[abase + bi + k];
6561 let mut acc = dr[k] * act.sx;
6562 for &(j, xv) in &act.outliers {
6563 let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6564 acc += w * s * xv;
6565 }
6566 unsafe { *out_addr.at((abase + bi + k) * rows + row) = acc };
6568 }
6569 }
6570 bi += 8;
6571 }
6572 for row in [r_lo, r_lo + 1] {
6575 let scr: &[f32] = if row == r_lo { &sc } else { &sc1 };
6576 for b2 in bi..alen {
6577 let act = &acts[abase + b2];
6578 let xs4 = [
6579 act.xq.as_slice(),
6580 act.xq.as_slice(),
6581 act.xq.as_slice(),
6582 act.xq.as_slice(),
6583 ];
6584 let d =
6585 unsafe { dot_q4tp_row_1x4_avx512(v.nib, row, gpr, xs4, scr) };
6586 let mut acc = d[0] * act.sx;
6587 for &(j, xv) in &act.outliers {
6588 let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6589 acc += w * s * xv;
6590 }
6591 unsafe { *out_addr.at((abase + b2) * rows + row) = acc };
6593 }
6594 }
6595 r_lo += 2;
6596 }
6597 }
6598 #[cfg(target_arch = "x86_64")]
6599 let row_start = r_lo;
6600 #[cfg(not(target_arch = "x86_64"))]
6601 let row_start = start;
6602 for r in row_start..end {
6603 v.scales_into(r, gpr, &mut sc);
6604 let mut bi = 0usize;
6605 #[cfg(target_arch = "x86_64")]
6606 if blocked_ok {
6607 while bi + 8 <= alen {
6608 let xs = [
6609 acts[abase + bi].xq.as_slice(),
6610 acts[abase + bi + 1].xq.as_slice(),
6611 acts[abase + bi + 2].xq.as_slice(),
6612 acts[abase + bi + 3].xq.as_slice(),
6613 acts[abase + bi + 4].xq.as_slice(),
6614 acts[abase + bi + 5].xq.as_slice(),
6615 acts[abase + bi + 6].xq.as_slice(),
6616 acts[abase + bi + 7].xq.as_slice(),
6617 ];
6618 let d = unsafe { dot_q4tp_row_1x8_avx512(v.nib, r, gpr, xs, &sc) };
6619 for k in 0..8 {
6620 let act = &acts[abase + bi + k];
6621 let mut acc = d[k] * act.sx;
6622 for &(j, xv) in &act.outliers {
6623 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6624 acc += w * s * xv;
6625 }
6626 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6628 }
6629 bi += 8;
6630 }
6631 while bi + 4 <= alen {
6632 let xs = [
6633 acts[abase + bi].xq.as_slice(),
6634 acts[abase + bi + 1].xq.as_slice(),
6635 acts[abase + bi + 2].xq.as_slice(),
6636 acts[abase + bi + 3].xq.as_slice(),
6637 ];
6638 let d = unsafe { dot_q4tp_row_1x4_avx512(v.nib, r, gpr, xs, &sc) };
6639 for k in 0..4 {
6640 let act = &acts[abase + bi + k];
6641 let mut acc = d[k] * act.sx;
6642 for &(j, xv) in &act.outliers {
6643 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6644 acc += w * s * xv;
6645 }
6646 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6648 }
6649 bi += 4;
6650 }
6651 }
6652 #[cfg(target_arch = "aarch64")]
6653 if blocked_ok {
6654 while bi + 4 <= alen {
6655 let xs = [
6656 acts[abase + bi].xq.as_slice(),
6657 acts[abase + bi + 1].xq.as_slice(),
6658 acts[abase + bi + 2].xq.as_slice(),
6659 acts[abase + bi + 3].xq.as_slice(),
6660 ];
6661 let d = unsafe {
6662 if exact || q4tp_v1() {
6663 dot_q4tp_row_1x4_sdot_v1(v.nib, r, gpr, xs, &sc)
6664 } else {
6665 dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs, &sc)
6666 }
6667 };
6668 for k in 0..4 {
6669 let act = &acts[abase + bi + k];
6670 let mut acc = d[k] * act.sx;
6671 for &(j, xv) in &act.outliers {
6672 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6673 acc += w * s * xv;
6674 }
6675 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6677 }
6678 bi += 4;
6679 }
6680 }
6681 let _ = blocked_ok;
6682 while bi < alen {
6683 let act = &acts[abase + bi];
6684 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
6685 for &(j, xv) in &act.outliers {
6686 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6687 acc += w * s * xv;
6688 }
6689 unsafe { *out_addr.at((abase + bi) * rows + r) = acc };
6691 bi += 1;
6692 }
6693 }
6694 }
6695 };
6696 dispatch_rows(pool, rows, &run);
6697 return;
6698 }
6699
6700 let run = |start: usize, end: usize| {
6701 let mut sc = vec![0f32; gpr];
6702 for r in start..end {
6703 v.scales_into(r, gpr, &mut sc);
6704 for bi in 0..b {
6705 let x = &xs_all[bi * cols..(bi + 1) * cols];
6706 unsafe { *out_addr.at(bi * rows + r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
6708 }
6709 }
6710 };
6711 dispatch_rows(pool, rows, &run);
6712}
6713
6714fn q4t_matvec(
6716 bytes: &[u8],
6717 x: &[f32],
6718 rows: usize,
6719 cols: usize,
6720 out: &mut [f32],
6721 pool: Option<&Pool>,
6722) {
6723 debug_assert_eq!(out.len(), rows);
6724 let gpr = cols / GROUP_SIZE;
6725 let out_addr = SendMut(out.as_mut_ptr());
6726 if a8w8_enabled() {
6727 let act = split_act(x);
6728 let run = move |start: usize, end: usize| {
6729 for r in start..end {
6730 let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6731 for &(j, xv) in &act.outliers {
6732 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6733 acc += w * s * xv;
6734 }
6735 unsafe { *out_addr.at(r) = acc };
6737 }
6738 };
6739 dispatch_rows(pool, rows, &run);
6740 return;
6741 }
6742 let run = move |start: usize, end: usize| {
6743 for r in start..end {
6744 unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
6746 }
6747 };
6748 dispatch_rows(pool, rows, &run);
6749}
6750
6751#[allow(clippy::too_many_arguments)]
6753fn q4t_matvec2(
6754 bytes: &[u8],
6755 x1: &[f32],
6756 x2: &[f32],
6757 rows: usize,
6758 cols: usize,
6759 o1: &mut [f32],
6760 o2: &mut [f32],
6761 pool: Option<&Pool>,
6762) {
6763 let gpr = cols / GROUP_SIZE;
6764 let p1 = SendMut(o1.as_mut_ptr());
6765 let p2 = SendMut(o2.as_mut_ptr());
6766 if a8w8_enabled() {
6767 let a1 = split_act(x1);
6768 let a2 = split_act(x2);
6769 let run = move |start: usize, end: usize| {
6770 for r in start..end {
6771 let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
6772 let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
6773 for &(j, xv) in &a1.outliers {
6774 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6775 v1 += w * s * xv;
6776 }
6777 for &(j, xv) in &a2.outliers {
6778 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6779 v2 += w * s * xv;
6780 }
6781 unsafe {
6783 *p1.at(r) = v1;
6784 *p2.at(r) = v2;
6785 }
6786 }
6787 };
6788 dispatch_rows(pool, rows, &run);
6789 return;
6790 }
6791 let run = move |start: usize, end: usize| {
6792 for r in start..end {
6793 unsafe {
6795 *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
6796 *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
6797 }
6798 }
6799 };
6800 dispatch_rows(pool, rows, &run);
6801}
6802
6803#[allow(clippy::too_many_arguments)]
6805#[cfg(target_os = "macos")]
6811fn dequant_matmat_accel(
6812 dequant_row: &(dyn Fn(usize, &mut [f32]) + Sync),
6813 xs_all: &[f32],
6814 b: usize,
6815 rows: usize,
6816 cols: usize,
6817 out: &mut [f32],
6818 pool: Option<&Pool>,
6819) {
6820 const TR: usize = 2048;
6821 thread_local! {
6822 static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
6823 }
6824 WTILE.with(|wt| {
6825 let mut wtile = wt.borrow_mut();
6826 wtile.resize(TR * cols, 0.0);
6827 let mut r0 = 0usize;
6828 while r0 < rows {
6829 let tr = TR.min(rows - r0);
6830 let wt_addr = SendMut(wtile.as_mut_ptr());
6831 let run = |start: usize, end: usize| {
6832 for r in start..end {
6833 let dst = unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
6835 dequant_row(r0 + r, dst);
6836 }
6837 };
6838 dispatch_rows(pool, tr, &run);
6839 unsafe {
6840 accel_blas::cblas_sgemm(
6841 101, 111, 112, b as i32,
6845 tr as i32,
6846 cols as i32,
6847 1.0,
6848 xs_all.as_ptr(),
6849 cols as i32,
6850 wtile.as_ptr(),
6851 cols as i32,
6852 0.0,
6853 out.as_mut_ptr().add(r0),
6854 rows as i32,
6855 );
6856 }
6857 r0 += tr;
6858 }
6859 });
6860}
6861
6862fn q4t_matmat(
6863 bytes: &[u8],
6864 xs_all: &[f32],
6865 b: usize,
6866 rows: usize,
6867 cols: usize,
6868 out: &mut [f32],
6869 pool: Option<&Pool>,
6870) {
6871 debug_assert_eq!(out.len(), b * rows);
6872 let gpr = cols / GROUP_SIZE;
6873 #[cfg(target_os = "macos")]
6877 if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6878 dequant_matmat_accel(
6879 &|r, dst| {
6880 for gi in 0..gpr {
6881 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
6882 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
6883 for (k, &bb) in tile[2..].iter().enumerate() {
6884 dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
6885 dst[gi * GROUP_SIZE + k * 2 + 1] = (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
6886 }
6887 }
6888 },
6889 xs_all,
6890 b,
6891 rows,
6892 cols,
6893 out,
6894 pool,
6895 );
6896 return;
6897 }
6898 let out_addr = SendMut(out.as_mut_ptr());
6899 if a8w8_enabled() {
6900 let acts: Vec<SplitAct> = (0..b)
6901 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6902 .collect();
6903 let acts = &acts;
6904 #[cfg(target_arch = "x86_64")]
6905 let blocked_ok = avx2_enabled() && blocked_enabled();
6906 #[cfg(target_arch = "aarch64")]
6907 let blocked_ok = sdot_enabled() && blocked_enabled();
6908 #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
6909 let blocked_ok = false;
6910 let run = move |start: usize, end: usize| {
6911 for r in start..end {
6912 let mut bi = 0usize;
6913 #[cfg(target_arch = "aarch64")]
6914 if blocked_ok {
6915 while bi + 4 <= acts.len() {
6916 let xs = [
6917 acts[bi].xq.as_slice(),
6918 acts[bi + 1].xq.as_slice(),
6919 acts[bi + 2].xq.as_slice(),
6920 acts[bi + 3].xq.as_slice(),
6921 ];
6922 let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
6923 for k in 0..4 {
6924 let act = &acts[bi + k];
6925 let mut acc = d[k] * act.sx;
6926 for &(j, xv) in &act.outliers {
6927 let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6928 acc += w * sc * xv;
6929 }
6930 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
6932 }
6933 bi += 4;
6934 }
6935 }
6936 #[cfg(target_arch = "x86_64")]
6937 if blocked_ok {
6938 while bi + 4 <= acts.len() {
6939 let xs = [
6940 acts[bi].xq.as_slice(),
6941 acts[bi + 1].xq.as_slice(),
6942 acts[bi + 2].xq.as_slice(),
6943 acts[bi + 3].xq.as_slice(),
6944 ];
6945 let d = unsafe {
6946 if vnni_tiles_enabled() {
6947 dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
6948 } else {
6949 dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
6950 }
6951 };
6952 for k in 0..4 {
6953 let act = &acts[bi + k];
6954 let mut acc = d[k] * act.sx;
6955 for &(j, xv) in &act.outliers {
6956 let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6957 acc += w * sc * xv;
6958 }
6959 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
6961 }
6962 bi += 4;
6963 }
6964 }
6965 let _ = blocked_ok;
6966 while bi < acts.len() {
6967 let act = &acts[bi];
6968 let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6969 for &(j, xv) in &act.outliers {
6970 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6971 acc += w * s * xv;
6972 }
6973 unsafe { *out_addr.at(bi * rows + r) = acc };
6975 bi += 1;
6976 }
6977 }
6978 };
6979 dispatch_rows(pool, rows, &run);
6980 return;
6981 }
6982 let run = move |start: usize, end: usize| {
6983 for r in start..end {
6984 for bi in 0..b {
6985 let x = &xs_all[bi * cols..(bi + 1) * cols];
6986 unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
6988 }
6989 }
6990 };
6991 dispatch_rows(pool, rows, &run);
6992}
6993
6994fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
7003 (0..gpr)
7004 .map(|gi| {
7005 xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
7006 .iter()
7007 .map(|&v| v as i32)
7008 .sum()
7009 })
7010 .collect()
7011}
7012
7013#[inline]
7017#[allow(unreachable_code)]
7018#[cfg(target_arch = "x86_64")]
7023#[target_feature(enable = "avx2")]
7024unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7025 unsafe {
7027 use core::arch::x86_64::*;
7028 let expand = _mm256_setr_epi8(
7030 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,
7031 3, 3, 3,
7032 );
7033 let bitsel = _mm256_setr_epi8(
7034 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7035 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7036 );
7037 let ones8 = _mm256_set1_epi8(1);
7038 let ones16 = _mm256_set1_epi16(1);
7039 let mut acc = 0f32;
7040 for gi in 0..gpr {
7041 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7042 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7043 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7044 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7045 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7046 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7047 let sel = _mm256_and_si256(x, mask);
7048 let p16 = _mm256_maddubs_epi16(ones8, sel);
7050 let d32 = _mm256_madd_epi16(p16, ones16);
7051 let hi128 = _mm256_extracti128_si256::<1>(d32);
7052 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7053 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7054 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7055 let msum = _mm_cvtsi128_si32(s32);
7056 let d = 2 * msum - gsum[gi];
7059 acc += d as f32 * s;
7060 }
7061 acc
7062 }
7063}
7064
7065#[cfg(target_arch = "x86_64")]
7068#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7069unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7070 unsafe {
7072 use core::arch::x86_64::*;
7073 let expand = _mm256_setr_epi8(
7074 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,
7075 3, 3, 3,
7076 );
7077 let bitsel = _mm256_setr_epi8(
7078 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7079 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7080 );
7081 let ones8 = _mm256_set1_epi8(1);
7082 let mut acc = 0f32;
7083 for gi in 0..gpr {
7084 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7085 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7086 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7087 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7088 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7089 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7090 let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7091 let d = 2 * msum - gsum[gi];
7092 acc += d as f32 * s;
7093 }
7094 acc
7095 }
7096}
7097
7098#[cfg(target_arch = "x86_64")]
7100#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7101unsafe fn dot_q1_row_1x4_vnni(
7102 bytes: &[u8],
7103 r: usize,
7104 gpr: usize,
7105 xs: [&[i8]; 4],
7106 gsums: [&[i32]; 4],
7107) -> [f32; 4] {
7108 unsafe {
7110 use core::arch::x86_64::*;
7111 let expand = _mm256_setr_epi8(
7112 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,
7113 3, 3, 3,
7114 );
7115 let bitsel = _mm256_setr_epi8(
7116 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7117 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7118 );
7119 let ones8 = _mm256_set1_epi8(1);
7120 let mut acc = [0f32; 4];
7121 for gi in 0..gpr {
7122 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7123 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7124 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7125 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7126 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7127 for (k, xq) in xs.iter().enumerate() {
7128 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7129 let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7130 let d = 2 * msum - gsums[k][gi];
7131 acc[k] += d as f32 * s;
7132 }
7133 }
7134 acc
7135 }
7136}
7137
7138#[cfg(target_arch = "x86_64")]
7141#[target_feature(enable = "avx2")]
7142unsafe fn dot_q1_row_1x4_avx2(
7143 bytes: &[u8],
7144 r: usize,
7145 gpr: usize,
7146 xs: [&[i8]; 4],
7147 gsums: [&[i32]; 4],
7148) -> [f32; 4] {
7149 unsafe {
7151 use core::arch::x86_64::*;
7152 let expand = _mm256_setr_epi8(
7153 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,
7154 3, 3, 3,
7155 );
7156 let bitsel = _mm256_setr_epi8(
7157 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7158 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7159 );
7160 let ones8 = _mm256_set1_epi8(1);
7161 let ones16 = _mm256_set1_epi16(1);
7162 let mut acc = [0f32; 4];
7163 for gi in 0..gpr {
7164 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7165 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7166 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7167 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7168 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7169 for (k, xq) in xs.iter().enumerate() {
7170 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7171 let sel = _mm256_and_si256(x, mask);
7172 let p16 = _mm256_maddubs_epi16(ones8, sel);
7173 let d32 = _mm256_madd_epi16(p16, ones16);
7174 let hi128 = _mm256_extracti128_si256::<1>(d32);
7175 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7176 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7177 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7178 let msum = _mm_cvtsi128_si32(s32);
7179 let d = 2 * msum - gsums[k][gi];
7180 acc[k] += d as f32 * s;
7181 }
7182 }
7183 acc
7184 }
7185}
7186
7187#[allow(unreachable_code)]
7188fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7189 #[cfg(target_arch = "aarch64")]
7190 unsafe {
7191 return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
7192 }
7193 #[cfg(target_arch = "x86_64")]
7194 if avx2_enabled() {
7195 unsafe {
7196 if vnni_tiles_enabled() {
7197 return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
7198 }
7199 return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
7200 }
7201 }
7202 let _ = gsum;
7203 let mut acc = 0f32;
7204 for gi in 0..gpr {
7205 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7206 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7207 let mut d = 0i32;
7208 for (j, &b) in tile[2..].iter().enumerate() {
7209 for k in 0..8 {
7210 let w = ((b >> k) & 1) as i32 * 2 - 1;
7211 d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
7212 }
7213 }
7214 acc += d as f32 * s;
7215 }
7216 acc
7217}
7218
7219#[cfg(target_arch = "aarch64")]
7228#[target_feature(enable = "neon,dotprod")]
7229unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7230 unsafe {
7233 use core::arch::aarch64::*;
7234 use core::arch::asm;
7235 const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7236 let m = vld1q_u8(MASKS.as_ptr());
7237 macro_rules! tile_dot {
7239 ($t:expr, $x:expr) => {{
7240 let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
7241 let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
7242 let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7243 let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7244 let x0 = vld1q_s8($x);
7245 let x1 = vld1q_s8($x.add(16));
7246 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7247 asm!(
7248 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7249 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7250 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7251 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7252 options(pure, nomem, nostack),
7253 );
7254 vaddq_s32(a0, a1)
7255 }};
7256 }
7257 const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7266 const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7267 const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7268 const IW11: [u8; 16] = [
7269 10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7270 ];
7271 const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7272 let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7273 let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7274 let isc = vld1_u8(ISC.as_ptr());
7275 macro_rules! tile_dot_tbl {
7277 ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
7278 let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
7279 let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
7280 let x0 = vld1q_s8($x);
7281 let x1 = vld1q_s8($x.add(16));
7282 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7283 asm!(
7284 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7285 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7286 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7287 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7288 options(pure, nomem, nostack),
7289 );
7290 vaddq_s32(a0, a1)
7291 }};
7292 }
7293 let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7294 let row_base = r * gpr * Q1_TILE;
7295 let abs_end = bytes.len();
7296 let xp = xq.as_ptr();
7297 let gp = gsum.as_ptr();
7298 let mut accv = vdupq_n_f32(0.0);
7299 let mut gi = 0;
7300 while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7303 let t0 = base.add(gi * Q1_TILE);
7304 let ld_a = vld1q_u8(t0);
7305 let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7306 let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
7307 let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
7308 let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
7309 let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
7310 let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7312 let g = vld1q_s32(gp.add(gi));
7313 let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7314 let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7315 let scf: float32x4_t;
7316 asm!(
7317 "fcvtl {o:v}.4s, {i:v}.4h",
7318 o = out(vreg) scf, i = in(vreg) sc16,
7319 options(pure, nomem, nostack),
7320 );
7321 accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
7322 gi += 4;
7323 }
7324 let mut acc = vaddvq_f32(accv);
7325 while gi < gpr {
7326 let t = base.add(gi * Q1_TILE);
7327 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7328 let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
7329 acc += (-(2 * d + *gp.add(gi))) as f32 * s;
7330 gi += 1;
7331 }
7332 acc
7333 }
7334}
7335
7336#[cfg(target_arch = "aarch64")]
7341#[target_feature(enable = "neon,dotprod")]
7342unsafe fn dot_q1_row_1x4_sdot(
7343 bytes: &[u8],
7344 r: usize,
7345 gpr: usize,
7346 xs: [&[i8]; 4],
7347 gs: [&[i32]; 4],
7348) -> [f32; 4] {
7349 unsafe {
7351 use core::arch::aarch64::*;
7352 use core::arch::asm;
7353 const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7354 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 m = vld1q_u8(MASKS.as_ptr());
7362 let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7363 let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7364 let isc = vld1_u8(ISC.as_ptr());
7365 macro_rules! sdot2 {
7366 ($w0:expr, $w1:expr, $x:expr) => {{
7367 let x0 = vld1q_s8($x);
7368 let x1 = vld1q_s8($x.add(16));
7369 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7370 asm!(
7371 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7372 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7373 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7374 w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
7375 options(pure, nomem, nostack),
7376 );
7377 vaddq_s32(a0, a1)
7378 }};
7379 }
7380 let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7381 let row_base = r * gpr * Q1_TILE;
7382 let abs_end = bytes.len();
7383 let mut accv = [vdupq_n_f32(0.0); 4];
7384 let mut gi = 0;
7385 while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7386 let t0 = base.add(gi * Q1_TILE);
7387 let ld_a = vld1q_u8(t0);
7388 let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7389 let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
7391 let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
7392 let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
7393 let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
7394 let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
7395 let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
7396 let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
7397 let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
7398 let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7399 let scf: float32x4_t;
7400 asm!(
7401 "fcvtl {o:v}.4s, {i:v}.4h",
7402 o = out(vreg) scf, i = in(vreg) sc16,
7403 options(pure, nomem, nostack),
7404 );
7405 for k in 0..4 {
7406 let xp = xs[k].as_ptr();
7407 let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
7408 let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
7409 let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
7410 let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
7411 let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7412 let g = vld1q_s32(gs[k].as_ptr().add(gi));
7413 let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7414 accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
7415 }
7416 gi += 4;
7417 }
7418 let mut acc = [
7419 vaddvq_f32(accv[0]),
7420 vaddvq_f32(accv[1]),
7421 vaddvq_f32(accv[2]),
7422 vaddvq_f32(accv[3]),
7423 ];
7424 while gi < gpr {
7425 let t = base.add(gi * Q1_TILE);
7426 let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7427 let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
7428 let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
7429 let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7430 let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7431 for k in 0..4 {
7432 let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
7433 acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
7434 }
7435 gi += 1;
7436 }
7437 acc
7438 }
7439}
7440
7441#[inline]
7443fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
7444 let gi = j / GROUP_SIZE;
7445 let k = j % GROUP_SIZE;
7446 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7447 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7448 let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
7449 ((bit as i32 * 2 - 1) as f32, s)
7450}
7451
7452#[inline]
7454fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
7455 let mut acc = 0f32;
7456 for gi in 0..gpr {
7457 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7458 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7459 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
7460 let mut ga = 0f32;
7461 for (j, &b) in tile[2..].iter().enumerate() {
7462 for k in 0..8 {
7463 ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
7464 }
7465 }
7466 acc += ga * s;
7467 }
7468 acc
7469}
7470
7471#[allow(clippy::too_many_arguments)]
7474fn q1_range_a8w8(
7475 bytes: &[u8],
7476 gpr: usize,
7477 act: &SplitAct,
7478 gsum: &[i32],
7479 out: SendMut,
7480 start: usize,
7481 end: usize,
7482) {
7483 for r in start..end {
7484 let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
7485 for &(j, xv) in &act.outliers {
7486 let (w, s) = q1_outlier(bytes, r, gpr, j);
7487 acc += w * s * xv;
7488 }
7489 unsafe { *out.at(r) = acc };
7491 }
7492}
7493
7494fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
7496 for r in start..end {
7497 unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
7499 }
7500}
7501
7502fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
7507 let entries = base_len + (rows + 1) * 4;
7508 (base_len, entries, entries <= bytes.len())
7509}
7510
7511#[inline]
7513fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
7514 let o = rp_off + r * 4;
7515 u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
7516}
7517
7518const SIGN5: [[f32; 5]; 256] = {
7522 let mut lut = [[0.0f32; 5]; 256];
7523 let pow3 = [1u16, 3, 9, 27, 81];
7524 let mut byte = 0usize;
7525 while byte < 256 {
7526 let mut i = 0usize;
7527 while i < 5 {
7528 let code = (byte as u16 / pow3[i]) % 3;
7529 lut[byte][i] = if code == 1 {
7530 1.0
7531 } else if code == 2 {
7532 -1.0
7533 } else {
7534 0.0
7535 };
7536 i += 1;
7537 }
7538 byte += 1;
7539 }
7540 lut
7541};
7542
7543const SIGN5_I8: [[i8; 5]; 256] = {
7545 let mut lut = [[0i8; 5]; 256];
7546 let pow3 = [1u16, 3, 9, 27, 81];
7547 let mut byte = 0usize;
7548 while byte < 256 {
7549 let mut i = 0usize;
7550 while i < 5 {
7551 let code = (byte as u16 / pow3[i]) % 3;
7552 lut[byte][i] = if code == 1 {
7553 1
7554 } else if code == 2 {
7555 -1
7556 } else {
7557 0
7558 };
7559 i += 1;
7560 }
7561 byte += 1;
7562 }
7563 lut
7564};
7565
7566const SIGN5_U64: [u64; 256] = {
7572 let mut lut = [0u64; 256];
7573 let pow3 = [1u16, 3, 9, 27, 81];
7574 let mut byte = 0usize;
7575 while byte < 256 {
7576 let mut v = 0u64;
7577 let mut i = 0usize;
7578 while i < 5 {
7579 let code = (byte as u16 / pow3[i]) % 3;
7580 let s: u8 = if code == 1 {
7581 1
7582 } else if code == 2 {
7583 0xFF
7584 } else {
7585 0
7586 };
7587 v |= (s as u64) << (i * 8);
7588 i += 1;
7589 }
7590 lut[byte] = v;
7591 byte += 1;
7592 }
7593 lut
7594};
7595
7596#[inline]
7601fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
7602 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7603 let off = (r * gpr + j / GROUP_SIZE) * TILE;
7604 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7605 let within = j % GROUP_SIZE;
7606 SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
7607}
7608
7609#[cfg(target_arch = "aarch64")]
7612#[target_feature(enable = "neon,dotprod")]
7613#[inline]
7614unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
7615 unsafe {
7617 use core::arch::aarch64::*;
7618 use core::arch::asm;
7619 let w0 = vld1q_s8(w);
7620 let w1 = vld1q_s8(w.add(16));
7621 let x0 = vld1q_s8(x);
7622 let x1 = vld1q_s8(x.add(16));
7623 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7624 asm!(
7625 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7626 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7627 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7628 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7629 options(pure, nomem, nostack),
7630 );
7631 vaddvq_s32(vaddq_s32(a0, a1))
7632 }
7633}
7634
7635#[cfg(target_arch = "x86_64")]
7638#[target_feature(enable = "avx2")]
7639#[inline]
7640unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
7641 unsafe {
7643 use core::arch::x86_64::*;
7644 let wv = _mm256_loadu_si256(w as *const __m256i);
7645 let xv = _mm256_loadu_si256(x as *const __m256i);
7646 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7647 let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
7648 let hi128 = _mm256_extracti128_si256::<1>(d);
7649 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
7650 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7651 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7652 _mm_cvtsi128_si32(s32)
7653 }
7654}
7655
7656#[inline]
7661fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
7662 debug_assert!(dst.len() >= 40);
7663 unsafe {
7666 let p = dst.as_mut_ptr();
7667 for bi in 0..7 {
7668 core::ptr::write_unaligned(
7669 p.add(bi * 5) as *mut u64,
7670 SIGN5_U64[*codes.add(bi) as usize],
7671 );
7672 }
7673 }
7674}
7675
7676#[inline]
7681fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
7682 #[cfg(target_arch = "aarch64")]
7683 unsafe {
7684 return sdot32_i8(w, x);
7685 }
7686 #[cfg(target_arch = "x86_64")]
7687 unsafe {
7688 return i8dot32_avx2(w, x);
7689 }
7690 #[allow(unreachable_code)]
7691 unsafe {
7692 let mut s = 0i32;
7693 for k in 0..GROUP_SIZE {
7694 s += *w.add(k) as i32 * *x.add(k) as i32;
7695 }
7696 s
7697 }
7698}
7699
7700#[inline]
7701unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
7702 let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
7703 (
7704 SIGN5_U64[*codes as usize],
7705 SIGN5_U64[*codes.add(1) as usize],
7706 SIGN5_U64[*codes.add(2) as usize],
7707 SIGN5_U64[*codes.add(3) as usize],
7708 SIGN5_U64[*codes.add(4) as usize],
7709 SIGN5_U64[*codes.add(5) as usize],
7710 SIGN5_U64[*codes.add(6) as usize],
7711 )
7712 };
7713
7714 let u0 = s0 | (s1 << 40);
7715 let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
7716 let u2 = (s3 >> 8) | (s4 << 32);
7717 let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
7718
7719 (u0, u1, u2, u3)
7720}
7721
7722#[cfg(target_arch = "aarch64")]
7726#[target_feature(enable = "neon,dotprod")]
7727unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7728 use core::arch::aarch64::*;
7729 use core::arch::asm;
7730 unsafe {
7731 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7732 let mut acc = 0f32;
7733 let bytes_ptr = bytes.as_ptr();
7734 let xq_ptr = xq.as_ptr();
7735 let row_off = r * gpr * TILE;
7736
7737 let gpr2 = gpr & !1;
7738 let mut gi = 0;
7739 while gi < gpr2 {
7740 let off0 = row_off + gi * TILE;
7741 let off1 = off0 + TILE;
7742 let s0 = f16_to_f32(u16::from_le_bytes([
7743 *bytes_ptr.add(off0),
7744 *bytes_ptr.add(off0 + 1),
7745 ]));
7746 let s1 = f16_to_f32(u16::from_le_bytes([
7747 *bytes_ptr.add(off1),
7748 *bytes_ptr.add(off1 + 1),
7749 ]));
7750
7751 let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
7752 let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
7753
7754 let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
7755 let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
7756 let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
7757 let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
7758
7759 let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7760 let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7761 let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
7762 let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
7763
7764 let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
7765 let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7766 asm!(
7767 "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
7768 "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
7769 "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
7770 "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
7771 a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
7772 a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
7773 w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
7774 w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
7775 options(pure, nomem, nostack),
7776 );
7777 let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
7778 let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
7779 acc += d0 as f32 * s0 + d1 as f32 * s1;
7780 gi += 2;
7781 }
7782
7783 if gi < gpr {
7784 let off = row_off + gi * TILE;
7785 let s = f16_to_f32(u16::from_le_bytes([
7786 *bytes_ptr.add(off),
7787 *bytes_ptr.add(off + 1),
7788 ]));
7789 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7790 let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
7791 let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
7792 let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7793 let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7794 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7795 asm!(
7796 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7797 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7798 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7799 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7800 options(pure, nomem, nostack),
7801 );
7802 let d = vaddvq_s32(vaddq_s32(a0, a1));
7803 acc += d as f32 * s;
7804 }
7805 acc
7806 }
7807}
7808
7809#[cfg(target_arch = "x86_64")]
7811#[target_feature(enable = "avx2")]
7812unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7813 use core::arch::x86_64::*;
7814 unsafe {
7815 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7816 let mut acc = 0f32;
7817 let bytes_ptr = bytes.as_ptr();
7818 let xq_ptr = xq.as_ptr();
7819 let row_off = r * gpr * TILE;
7820
7821 let ones = _mm256_set1_epi16(1);
7822 for gi in 0..gpr {
7823 let off = row_off + gi * TILE;
7824 let s = f16_to_f32(u16::from_le_bytes([
7825 *bytes_ptr.add(off),
7826 *bytes_ptr.add(off + 1),
7827 ]));
7828 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7829 let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7830 let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7831 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7832 let d256 = _mm256_madd_epi16(p16, ones);
7833 let d128 = _mm_add_epi32(
7834 _mm256_castsi256_si128(d256),
7835 _mm256_extracti128_si256(d256, 1),
7836 );
7837 let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
7838 let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
7839 acc += d32 as f32 * s;
7840 }
7841 acc
7842 }
7843}
7844
7845#[cfg(target_arch = "x86_64")]
7847#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7848unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7849 use core::arch::x86_64::*;
7850 unsafe {
7852 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7853 let mut acc = 0f32;
7854 let bytes_ptr = bytes.as_ptr();
7855 let xq_ptr = xq.as_ptr();
7856 let row_off = r * gpr * TILE;
7857 for gi in 0..gpr {
7858 let off = row_off + gi * TILE;
7859 let s = f16_to_f32(u16::from_le_bytes([
7860 *bytes_ptr.add(off),
7861 *bytes_ptr.add(off + 1),
7862 ]));
7863 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7864 let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7865 let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7866 let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7867 acc += d as f32 * s;
7868 }
7869 acc
7870 }
7871}
7872
7873#[inline]
7877fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7878 #[cfg(target_arch = "aarch64")]
7879 unsafe {
7880 return q1t_dot_row_sdot(bytes, r, gpr, xq);
7881 }
7882 #[cfg(target_arch = "x86_64")]
7883 unsafe {
7884 if vnni_tiles_enabled() {
7885 return q1t_dot_row_vnni(bytes, r, gpr, xq);
7886 }
7887 return q1t_dot_row_avx2(bytes, r, gpr, xq);
7888 }
7889 #[allow(unreachable_code)]
7890 {
7891 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7892 let mut acc = 0f32;
7893 let mut sg = [0i8; GROUP_SIZE + 8]; for gi in 0..gpr {
7895 let off = (r * gpr + gi) * TILE;
7896 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7897 q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
7898 let mut d = 0i32;
7899 for k in 0..GROUP_SIZE {
7900 d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
7901 }
7902 acc += d as f32 * s;
7903 }
7904 acc
7905 }
7906}
7907
7908fn q1t_row_outlier_correction(
7915 bytes: &[u8],
7916 r: usize,
7917 rp_off: usize,
7918 entries_off: usize,
7919 has_ov: bool,
7920 x: &[f32],
7921) -> f32 {
7922 if !has_ov {
7923 return 0.0;
7924 }
7925 let (c0, c1) = (
7926 q1t_rowptr(bytes, rp_off, r),
7927 q1t_rowptr(bytes, rp_off, r + 1),
7928 );
7929 let mut corr = 0f32;
7930 for p in c0..c1 {
7931 let e = entries_off + p * 4;
7932 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
7933 let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
7934 corr += val * x[col];
7935 }
7936 corr
7937}
7938
7939fn q1t_dequant_row(
7943 bytes: &[u8],
7944 r: usize,
7945 gpr: usize,
7946 rp_off: usize,
7947 entries_off: usize,
7948 has_ov: bool,
7949 buf: &mut [f32],
7950) {
7951 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7952 for g in 0..gpr {
7953 let off = (r * gpr + g) * TILE;
7954 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7955 let codes = &bytes[off + 2..off + TILE];
7956 let bc = g * GROUP_SIZE;
7957 for bi in 0..6 {
7959 let lut = &SIGN5[codes[bi] as usize];
7960 let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
7961 for i in 0..5 {
7962 d[i] = lut[i] * s;
7963 }
7964 }
7965 let lut = &SIGN5[codes[6] as usize];
7966 buf[bc + 30] = lut[0] * s;
7967 buf[bc + 31] = lut[1] * s;
7968 }
7969 if !has_ov {
7970 return;
7971 }
7972 let (c0, c1) = (
7973 q1t_rowptr(bytes, rp_off, r),
7974 q1t_rowptr(bytes, rp_off, r + 1),
7975 );
7976 for p in c0..c1 {
7977 let e = entries_off + p * 4;
7978 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
7979 buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
7980 }
7981}
7982
7983fn q1t_add_overlay(
7987 bytes: &[u8],
7988 x: &[f32],
7989 rows: usize,
7990 cols: usize,
7991 out: &mut [f32],
7992 pool: Option<&Pool>,
7993) {
7994 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7995 let gpr = cols / GROUP_SIZE;
7996 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
7997 if !has_ov {
7998 return;
7999 }
8000 let out_addr = SendMut(out.as_mut_ptr());
8001 let run = move |start: usize, end: usize| {
8002 for r in start..end {
8003 let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8004 unsafe { *out_addr.at(r) += corr };
8006 }
8007 };
8008 dispatch_rows(pool, rows, &run);
8009}
8010
8011#[allow(clippy::too_many_arguments)]
8014fn q1t_range_a8w8(
8015 bytes: &[u8],
8016 gpr: usize,
8017 rp_off: usize,
8018 ent_off: usize,
8019 has_ov: bool,
8020 act: &SplitAct,
8021 x: &[f32],
8022 out: SendMut,
8023 start: usize,
8024 end: usize,
8025) {
8026 for r in start..end {
8027 let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8028 for &(j, xv) in &act.outliers {
8029 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8030 }
8031 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8032 unsafe { *out.at(r) = acc };
8034 }
8035}
8036
8037#[allow(clippy::too_many_arguments)]
8040fn q1t_range_f32_batch(
8041 bytes: &[u8],
8042 gpr: usize,
8043 rp_off: usize,
8044 ent_off: usize,
8045 has_ov: bool,
8046 x: &[f32],
8047 out: SendMut,
8048 start: usize,
8049 end: usize,
8050) {
8051 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8052 let mut sg = [0f32; GROUP_SIZE];
8053 for r in start..end {
8054 let mut acc = 0f32;
8055 for g in 0..gpr {
8056 let off = (r * gpr + g) * TILE;
8057 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8058 let codes = &bytes[off + 2..off + TILE];
8059 let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8060 for bi in 0..6 {
8061 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8062 }
8063 let lut = &SIGN5[codes[6] as usize];
8064 sg[30] = lut[0];
8065 sg[31] = lut[1];
8066 let mut gsum = 0f32;
8067 for k in 0..GROUP_SIZE {
8068 gsum += sg[k] * xg[k];
8069 }
8070 acc += s * gsum;
8071 }
8072 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8073 unsafe { *out.at(r) = acc };
8075 }
8076}
8077
8078fn q1t_matvec(
8082 bytes: &[u8],
8083 x: &[f32],
8084 rows: usize,
8085 cols: usize,
8086 out: &mut [f32],
8087 pool: Option<&Pool>,
8088) {
8089 debug_assert_eq!(out.len(), rows);
8090 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8091 let gpr = cols / GROUP_SIZE;
8092 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8093 let out_addr = SendMut(out.as_mut_ptr());
8094 if a8w8_enabled() {
8098 let act = split_act(x);
8099 let act = &act;
8100 let run = move |start: usize, end: usize| {
8101 for r in start..end {
8102 let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8103 for &(j, xv) in &act.outliers {
8104 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8105 }
8106 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8107 unsafe { *out_addr.at(r) = acc };
8109 }
8110 };
8111 dispatch_rows(pool, rows, &run);
8112 return;
8113 }
8114 let run = move |start: usize, end: usize| {
8115 let mut sg = [0f32; GROUP_SIZE];
8119 for r in start..end {
8120 let mut acc = 0f32;
8121 for g in 0..gpr {
8122 let off = (r * gpr + g) * TILE;
8123 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8124 let codes = &bytes[off + 2..off + TILE];
8125 let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8126 for bi in 0..6 {
8127 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8128 }
8129 let lut = &SIGN5[codes[6] as usize];
8130 sg[30] = lut[0];
8131 sg[31] = lut[1];
8132 let mut gsum = 0f32;
8133 for k in 0..GROUP_SIZE {
8134 gsum += sg[k] * xg[k];
8135 }
8136 acc += s * gsum;
8137 }
8138 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8139 unsafe { *out_addr.at(r) = acc };
8140 }
8141 };
8142 dispatch_rows(pool, rows, &run);
8143}
8144
8145#[cfg(target_arch = "aarch64")]
8151#[target_feature(enable = "neon,dotprod")]
8152unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
8153 use core::arch::aarch64::*;
8154 use core::arch::asm;
8155 unsafe {
8157 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8158 let bytes_ptr = bytes.as_ptr();
8159 let row_off = r * gpr * TILE;
8160 let xp = [xa.as_ptr(), xb.as_ptr()];
8161 let mut acc = [0f32; 2];
8162 macro_rules! sdot2 {
8163 ($w0:expr, $w1:expr, $x:expr) => {{
8164 let x0 = vld1q_s8($x);
8165 let x1 = vld1q_s8($x.add(16));
8166 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
8167 asm!(
8168 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
8169 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
8170 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
8171 w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
8172 options(pure, nomem, nostack),
8173 );
8174 vaddvq_s32(vaddq_s32(a0, a1))
8175 }};
8176 }
8177 let gpr2 = gpr & !1;
8178 let mut gi = 0;
8179 while gi < gpr2 {
8180 let off0 = row_off + gi * TILE;
8181 let off1 = off0 + TILE;
8182 let s0 = f16_to_f32(u16::from_le_bytes([
8183 *bytes_ptr.add(off0),
8184 *bytes_ptr.add(off0 + 1),
8185 ]));
8186 let s1 = f16_to_f32(u16::from_le_bytes([
8187 *bytes_ptr.add(off1),
8188 *bytes_ptr.add(off1 + 1),
8189 ]));
8190 let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
8191 let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
8192 let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
8193 let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
8194 let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
8195 let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
8196 for k in 0..2 {
8197 let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
8198 let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
8199 acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
8200 }
8201 gi += 2;
8202 }
8203 if gi < gpr {
8204 let off = row_off + gi * TILE;
8205 let s = f16_to_f32(u16::from_le_bytes([
8206 *bytes_ptr.add(off),
8207 *bytes_ptr.add(off + 1),
8208 ]));
8209 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
8210 let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
8211 let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
8212 for k in 0..2 {
8213 let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
8214 acc[k] += d as f32 * s;
8215 }
8216 }
8217 acc
8218 }
8219}
8220
8221fn q1t_matvec2(
8227 bytes: &[u8],
8228 x1: &[f32],
8229 x2: &[f32],
8230 rows: usize,
8231 cols: usize,
8232 o1: &mut [f32],
8233 o2: &mut [f32],
8234 pool: Option<&Pool>,
8235) {
8236 debug_assert_eq!(o1.len(), rows);
8237 debug_assert_eq!(o2.len(), rows);
8238 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8239 let gpr = cols / GROUP_SIZE;
8240 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8241 let out1 = SendMut(o1.as_mut_ptr());
8242 let out2 = SendMut(o2.as_mut_ptr());
8243 if a8w8_enabled() {
8244 let a1 = split_act(x1);
8245 let a2 = split_act(x2);
8246 let (a1, a2) = (&a1, &a2);
8247 let run = move |start: usize, end: usize| {
8248 for r in start..end {
8249 #[cfg(target_arch = "aarch64")]
8250 let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
8253 #[cfg(not(target_arch = "aarch64"))]
8254 let ds = [
8255 q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
8256 q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
8257 ];
8258 let mut acc1 = ds[0] * a1.sx;
8259 for &(j, xv) in &a1.outliers {
8260 acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
8261 }
8262 acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8263 let mut acc2 = ds[1] * a2.sx;
8264 for &(j, xv) in &a2.outliers {
8265 acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
8266 }
8267 acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8268 unsafe {
8270 *out1.at(r) = acc1;
8271 *out2.at(r) = acc2;
8272 }
8273 }
8274 };
8275 dispatch_rows(pool, rows, &run);
8276 return;
8277 }
8278 let run = move |start: usize, end: usize| {
8279 let mut sg = [0f32; GROUP_SIZE];
8282 for r in start..end {
8283 let mut acc1 = 0f32;
8284 let mut acc2 = 0f32;
8285 for g in 0..gpr {
8286 let off = (r * gpr + g) * TILE;
8287 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8288 let codes = &bytes[off + 2..off + TILE];
8289 for bi in 0..6 {
8290 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8291 }
8292 let lut = &SIGN5[codes[6] as usize];
8293 sg[30] = lut[0];
8294 sg[31] = lut[1];
8295 let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8296 let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8297 let mut gsum1 = 0f32;
8298 for k in 0..GROUP_SIZE {
8299 gsum1 += sg[k] * xg1[k];
8300 }
8301 acc1 += s * gsum1;
8302 let mut gsum2 = 0f32;
8303 for k in 0..GROUP_SIZE {
8304 gsum2 += sg[k] * xg2[k];
8305 }
8306 acc2 += s * gsum2;
8307 }
8308 acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8309 acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8310 unsafe {
8312 *out1.at(r) = acc1;
8313 *out2.at(r) = acc2;
8314 }
8315 }
8316 };
8317 dispatch_rows(pool, rows, &run);
8318}
8319
8320fn q1t_matmat(
8323 bytes: &[u8],
8324 xs: &[f32],
8325 b: usize,
8326 rows: usize,
8327 cols: usize,
8328 out: &mut [f32],
8329 pool: Option<&Pool>,
8330) {
8331 debug_assert_eq!(out.len(), b * rows);
8332 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8333 let gpr = cols / GROUP_SIZE;
8334 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8335 let out_addr = SendMut(out.as_mut_ptr());
8336 if a8w8_enabled() {
8340 let acts: Vec<SplitAct> = (0..b)
8341 .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
8342 .collect();
8343 let acts = &acts;
8344 let run = move |start: usize, end: usize| {
8345 let mut sg = vec![0i8; cols + 8]; let mut sc = vec![0f32; gpr]; let mut accs = vec![0f32; b]; for r in start..end {
8349 for g in 0..gpr {
8350 let off = (r * gpr + g) * TILE;
8351 sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8352 q1t_unpack_group_i8(
8353 bytes.as_ptr().wrapping_add(off + 2),
8354 &mut sg[g * GROUP_SIZE..],
8355 );
8356 }
8357 for bi in 0..b {
8358 let act = &acts[bi];
8359 let mut isum = 0f32;
8360 for g in 0..gpr {
8361 let d = q1t_i8dot32(
8362 sg.as_ptr().wrapping_add(g * GROUP_SIZE),
8363 act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
8364 );
8365 isum += d as f32 * sc[g];
8366 }
8367 let mut acc = isum * act.sx;
8368 for &(j, xv) in &act.outliers {
8369 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8370 }
8371 accs[bi] = acc;
8372 }
8373 if has_ov {
8377 let (c0, c1) = (
8378 q1t_rowptr(bytes, rp_off, r),
8379 q1t_rowptr(bytes, rp_off, r + 1),
8380 );
8381 for p in c0..c1 {
8382 let e = ent_off + p * 4;
8383 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8384 let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8385 for bi in 0..b {
8386 accs[bi] += val * xs[bi * cols + col];
8387 }
8388 }
8389 }
8390 for bi in 0..b {
8391 unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
8392 }
8393 }
8394 };
8395 dispatch_rows(pool, rows, &run);
8396 return;
8397 }
8398 let run = move |start: usize, end: usize| {
8399 let mut buf = vec![0f32; cols];
8400 for r in start..end {
8401 q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
8402 for bi in 0..b {
8403 let xr = &xs[bi * cols..(bi + 1) * cols];
8404 let mut acc = 0f32;
8405 for j in 0..cols {
8406 acc += buf[j] * xr[j];
8407 }
8408 unsafe { *out_addr.at(bi * rows + r) = acc };
8409 }
8410 }
8411 };
8412 dispatch_rows(pool, rows, &run);
8413}
8414
8415fn q1_matvec(
8416 bytes: &[u8],
8417 x: &[f32],
8418 rows: usize,
8419 cols: usize,
8420 out: &mut [f32],
8421 pool: Option<&Pool>,
8422) {
8423 debug_assert_eq!(out.len(), rows);
8424 let gpr = cols / GROUP_SIZE;
8425 let out_addr = SendMut(out.as_mut_ptr());
8426 if a8w8_enabled() {
8427 let act = split_act(x);
8428 let gsum = q1_group_sums(&act.xq, gpr);
8429 let (act, gsum) = (&act, &gsum);
8430 let run = move |start: usize, end: usize| {
8431 q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
8432 };
8433 dispatch_rows(pool, rows, &run);
8434 return;
8435 }
8436 let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
8437 dispatch_rows(pool, rows, &run);
8438}
8439
8440#[allow(clippy::too_many_arguments)]
8442fn q1_matvec2(
8443 bytes: &[u8],
8444 x1: &[f32],
8445 x2: &[f32],
8446 rows: usize,
8447 cols: usize,
8448 o1: &mut [f32],
8449 o2: &mut [f32],
8450 pool: Option<&Pool>,
8451) {
8452 let gpr = cols / GROUP_SIZE;
8453 let p1 = SendMut(o1.as_mut_ptr());
8454 let p2 = SendMut(o2.as_mut_ptr());
8455 if a8w8_enabled() {
8456 let a1 = split_act(x1);
8457 let a2 = split_act(x2);
8458 let g1 = q1_group_sums(&a1.xq, gpr);
8459 let g2 = q1_group_sums(&a2.xq, gpr);
8460 let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
8461 let run = move |start: usize, end: usize| {
8462 for r in start..end {
8463 let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
8464 let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
8465 for &(j, xv) in &a1.outliers {
8466 let (w, s) = q1_outlier(bytes, r, gpr, j);
8467 v1 += w * s * xv;
8468 }
8469 for &(j, xv) in &a2.outliers {
8470 let (w, s) = q1_outlier(bytes, r, gpr, j);
8471 v2 += w * s * xv;
8472 }
8473 unsafe {
8475 *p1.at(r) = v1;
8476 *p2.at(r) = v2;
8477 }
8478 }
8479 };
8480 dispatch_rows(pool, rows, &run);
8481 return;
8482 }
8483 let run = move |start: usize, end: usize| {
8484 for r in start..end {
8485 unsafe {
8487 *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
8488 *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
8489 }
8490 }
8491 };
8492 dispatch_rows(pool, rows, &run);
8493}
8494
8495#[allow(clippy::too_many_arguments)]
8497fn q1_matmat(
8498 bytes: &[u8],
8499 xs_all: &[f32],
8500 b: usize,
8501 rows: usize,
8502 cols: usize,
8503 out: &mut [f32],
8504 pool: Option<&Pool>,
8505) {
8506 debug_assert_eq!(out.len(), b * rows);
8507 let gpr = cols / GROUP_SIZE;
8508 let out_addr = SendMut(out.as_mut_ptr());
8509 if a8w8_enabled() {
8510 let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
8511 .map(|bi| {
8512 let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
8513 let gsum = q1_group_sums(&act.xq, gpr);
8514 (act, gsum)
8515 })
8516 .collect();
8517 let acts = &acts;
8518 #[cfg(target_arch = "x86_64")]
8519 let blocked_ok = avx2_enabled() && blocked_enabled();
8520 #[cfg(target_arch = "aarch64")]
8521 let blocked_ok = sdot_enabled() && blocked_enabled();
8522 let run = move |start: usize, end: usize| {
8523 for r in start..end {
8524 let mut bi = 0usize;
8525 #[cfg(target_arch = "aarch64")]
8528 if blocked_ok {
8529 while bi + 4 <= acts.len() {
8530 let xs = [
8531 acts[bi].0.xq.as_slice(),
8532 acts[bi + 1].0.xq.as_slice(),
8533 acts[bi + 2].0.xq.as_slice(),
8534 acts[bi + 3].0.xq.as_slice(),
8535 ];
8536 let gs = [
8537 acts[bi].1.as_slice(),
8538 acts[bi + 1].1.as_slice(),
8539 acts[bi + 2].1.as_slice(),
8540 acts[bi + 3].1.as_slice(),
8541 ];
8542 let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
8543 for k in 0..4 {
8544 let (act, _) = &acts[bi + k];
8545 let mut acc = d[k] * act.sx;
8546 for &(j, xv) in &act.outliers {
8547 let (w, sc) = q1_outlier(bytes, r, gpr, j);
8548 acc += w * sc * xv;
8549 }
8550 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8552 }
8553 bi += 4;
8554 }
8555 }
8556 #[cfg(target_arch = "x86_64")]
8557 if blocked_ok {
8558 while bi + 4 <= acts.len() {
8559 let xs = [
8560 acts[bi].0.xq.as_slice(),
8561 acts[bi + 1].0.xq.as_slice(),
8562 acts[bi + 2].0.xq.as_slice(),
8563 acts[bi + 3].0.xq.as_slice(),
8564 ];
8565 let gs = [
8566 acts[bi].1.as_slice(),
8567 acts[bi + 1].1.as_slice(),
8568 acts[bi + 2].1.as_slice(),
8569 acts[bi + 3].1.as_slice(),
8570 ];
8571 let d = unsafe {
8572 if vnni_tiles_enabled() {
8573 dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
8574 } else {
8575 dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
8576 }
8577 };
8578 for k in 0..4 {
8579 let (act, _) = &acts[bi + k];
8580 let mut acc = d[k] * act.sx;
8581 for &(j, xv) in &act.outliers {
8582 let (w, sc) = q1_outlier(bytes, r, gpr, j);
8583 acc += w * sc * xv;
8584 }
8585 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8587 }
8588 bi += 4;
8589 }
8590 }
8591 while bi < acts.len() {
8592 let (act, gsum) = &acts[bi];
8593 let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
8594 for &(j, xv) in &act.outliers {
8595 let (w, s) = q1_outlier(bytes, r, gpr, j);
8596 acc += w * s * xv;
8597 }
8598 unsafe { *out_addr.at(bi * rows + r) = acc };
8600 bi += 1;
8601 }
8602 }
8603 };
8604 dispatch_rows(pool, rows, &run);
8605 return;
8606 }
8607 let run = move |start: usize, end: usize| {
8608 for r in start..end {
8609 for bi in 0..b {
8610 let x = &xs_all[bi * cols..(bi + 1) * cols];
8611 unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
8613 }
8614 }
8615 };
8616 dispatch_rows(pool, rows, &run);
8617}
8618
8619fn q4matvec(
8625 bytes: &[u8],
8626 x: &[f32],
8627 rows: usize,
8628 cols: usize,
8629 out: &mut [f32],
8630 pool: Option<&Pool>,
8631) {
8632 debug_assert_eq!(out.len(), rows);
8633 let (packed, scales) = q4_split(bytes, rows, cols);
8634 let gpr = cols / GROUP_SIZE;
8635 let out_addr = SendMut(out.as_mut_ptr());
8636
8637 if a8w8_enabled() {
8638 let act = split_act(x);
8639 let run = move |start: usize, end: usize| {
8640 q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
8641 };
8642 dispatch_rows(pool, rows, &run);
8643 return;
8644 }
8645
8646 let run =
8647 move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
8648 dispatch_rows(pool, rows, &run);
8649}
8650
8651#[inline]
8654#[allow(unreachable_code)]
8655#[cfg(target_arch = "x86_64")]
8660#[target_feature(enable = "avx2")]
8661unsafe fn dot_q4b_row_1x4_avx2(
8662 buf: &[u8],
8663 scales: &[u8],
8664 g0: usize,
8665 gpr: usize,
8666 xs: [&[i8]; 4],
8667) -> [f32; 4] {
8668 unsafe {
8670 use core::arch::x86_64::*;
8671 let ones = _mm256_set1_epi16(1);
8672 let mut acc = [0f32; 4];
8673 for gi in 0..gpr {
8674 let s = f16_to_f32(u16::from_le_bytes([
8675 scales[(g0 + gi) * 2],
8676 scales[(g0 + gi) * 2 + 1],
8677 ]));
8678 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8679 let aw = _mm256_abs_epi8(w);
8680 for (k, xq) in xs.iter().enumerate() {
8681 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8682 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8683 let d = _mm256_madd_epi16(p16, ones);
8684 let hi128 = _mm256_extracti128_si256::<1>(d);
8685 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8686 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8687 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8688 acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
8689 }
8690 }
8691 acc
8692 }
8693}
8694
8695#[cfg(target_arch = "x86_64")]
8697#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8698unsafe fn dot_q4b_row_1x4_vnni(
8699 buf: &[u8],
8700 scales: &[u8],
8701 g0: usize,
8702 gpr: usize,
8703 xs: [&[i8]; 4],
8704) -> [f32; 4] {
8705 unsafe {
8707 use core::arch::x86_64::*;
8708 let mut acc = [0f32; 4];
8709 for gi in 0..gpr {
8710 let s = f16_to_f32(u16::from_le_bytes([
8711 scales[(g0 + gi) * 2],
8712 scales[(g0 + gi) * 2 + 1],
8713 ]));
8714 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8715 let aw = _mm256_abs_epi8(w);
8716 for (k, xq) in xs.iter().enumerate() {
8717 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8718 let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8719 acc[k] += d as f32 * s;
8720 }
8721 }
8722 acc
8723 }
8724}
8725
8726#[cfg(target_arch = "x86_64")]
8732#[target_feature(enable = "avx2")]
8733unsafe fn dot_q4b_row_1x4_sx_avx2(
8734 buf: &[u8],
8735 scales: &[u8],
8736 g0: usize,
8737 gpr: usize,
8738 xs: [&[i8]; 4],
8739 sxs: [f32; 4],
8740) -> [f32; 4] {
8741 unsafe {
8743 use core::arch::x86_64::*;
8744 let ones = _mm256_set1_epi16(1);
8745 let mut acc = [0f32; 4];
8746 for gi in 0..gpr {
8747 let s = f16_to_f32(u16::from_le_bytes([
8748 scales[(g0 + gi) * 2],
8749 scales[(g0 + gi) * 2 + 1],
8750 ]));
8751 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8752 let aw = _mm256_abs_epi8(w);
8753 for (k, xq) in xs.iter().enumerate() {
8754 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8755 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8756 let d = _mm256_madd_epi16(p16, ones);
8757 let hi128 = _mm256_extracti128_si256::<1>(d);
8758 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8759 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8760 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8761 acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
8762 }
8763 }
8764 acc
8765 }
8766}
8767
8768#[cfg(target_arch = "x86_64")]
8771#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8772unsafe fn dot_q4b_row_1x4_sx_vnni(
8773 buf: &[u8],
8774 scales: &[u8],
8775 g0: usize,
8776 gpr: usize,
8777 xs: [&[i8]; 4],
8778 sxs: [f32; 4],
8779) -> [f32; 4] {
8780 unsafe {
8782 use core::arch::x86_64::*;
8783 let mut acc = [0f32; 4];
8784 for gi in 0..gpr {
8785 let s = f16_to_f32(u16::from_le_bytes([
8786 scales[(g0 + gi) * 2],
8787 scales[(g0 + gi) * 2 + 1],
8788 ]));
8789 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8790 let aw = _mm256_abs_epi8(w);
8791 for (k, xq) in xs.iter().enumerate() {
8792 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8793 let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8794 acc[k] += (d as f32 * sxs[k]) * s;
8795 }
8796 }
8797 acc
8798 }
8799}
8800
8801#[allow(unreachable_code)]
8802fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
8803 #[cfg(target_arch = "aarch64")]
8804 unsafe {
8805 return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
8806 }
8807 #[cfg(target_arch = "x86_64")]
8808 unsafe {
8809 return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
8810 }
8811 let mut acc = 0f32;
8812 for gi in 0..gpr {
8813 let g = g0 + gi;
8814 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8815 let mut d = 0i32;
8816 for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
8817 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
8818 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
8819 }
8820 acc += d as f32 * s;
8821 }
8822 acc
8823}
8824
8825#[inline]
8827#[allow(unreachable_code)]
8828fn dot_q4_row_i8_2(
8829 packed: &[u8],
8830 scales: &[u8],
8831 g0: usize,
8832 gpr: usize,
8833 xq1: &[i8],
8834 xq2: &[i8],
8835) -> (f32, f32) {
8836 #[cfg(target_arch = "aarch64")]
8837 unsafe {
8838 return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
8839 }
8840 #[cfg(target_arch = "x86_64")]
8841 unsafe {
8842 return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
8843 }
8844 (
8845 dot_q4_row_i8(packed, scales, g0, gpr, xq1),
8846 dot_q4_row_i8(packed, scales, g0, gpr, xq2),
8847 )
8848}
8849
8850#[allow(clippy::too_many_arguments)]
8853fn q4_range_a8w8(
8854 packed: &[u8],
8855 scales: &[u8],
8856 gpr: usize,
8857 cols: usize,
8858 act: &SplitAct,
8859 out: SendMut,
8860 start: usize,
8861 end: usize,
8862) {
8863 for r in start..end {
8864 let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
8865 for &(j, xv) in &act.outliers {
8867 let flat = r * cols + j;
8868 let byte = packed[flat / 2];
8869 let nib = if flat & 1 == 0 {
8870 byte & 0x0F
8871 } else {
8872 byte >> 4
8873 };
8874 let s = f16_to_f32(u16::from_le_bytes([
8875 scales[(flat / GROUP_SIZE) * 2],
8876 scales[(flat / GROUP_SIZE) * 2 + 1],
8877 ]));
8878 acc += ((nib as i32 - 8) as f32) * s * xv;
8879 }
8880 unsafe { *out.at(r) = acc };
8882 }
8883}
8884
8885#[allow(clippy::too_many_arguments)]
8888fn q4_range2_a8w8(
8889 packed: &[u8],
8890 scales: &[u8],
8891 gpr: usize,
8892 cols: usize,
8893 a1: &SplitAct,
8894 a2: &SplitAct,
8895 p1: SendMut,
8896 p2: SendMut,
8897 start: usize,
8898 end: usize,
8899) {
8900 for r in start..end {
8901 let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
8902 let mut acc1 = s1 * a1.sx;
8903 let mut acc2 = s2 * a2.sx;
8904 let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
8906 for &(j, xv) in outliers {
8907 let flat = r * cols + j;
8908 let byte = packed[flat / 2];
8909 let nib = if flat & 1 == 0 {
8910 byte & 0x0F
8911 } else {
8912 byte >> 4
8913 };
8914 let s = f16_to_f32(u16::from_le_bytes([
8915 scales[(flat / GROUP_SIZE) * 2],
8916 scales[(flat / GROUP_SIZE) * 2 + 1],
8917 ]));
8918 *acc += ((nib as i32 - 8) as f32) * s * xv;
8919 }
8920 };
8921 fix(&a1.outliers, &mut acc1);
8922 fix(&a2.outliers, &mut acc2);
8923 unsafe {
8925 *p1.at(r) = acc1;
8926 *p2.at(r) = acc2;
8927 }
8928 }
8929}
8930
8931fn q4_range_f32(
8933 packed: &[u8],
8934 scales: &[u8],
8935 gpr: usize,
8936 x: &[f32],
8937 out: SendMut,
8938 start: usize,
8939 end: usize,
8940) {
8941 for r in start..end {
8942 let mut acc = 0f32;
8943 for gi in 0..gpr {
8944 let g = r * gpr + gi;
8945 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8946 let pk = &packed[g * 16..(g + 1) * 16];
8947 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
8948 let mut ga = 0f32;
8949 for (k, &b) in pk.iter().enumerate() {
8950 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
8951 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
8952 }
8953 acc += ga * s;
8954 }
8955 unsafe { *out.at(r) = acc };
8957 }
8958}
8959
8960#[allow(clippy::too_many_arguments)]
8964fn q4matvec2(
8965 bytes: &[u8],
8966 x1: &[f32],
8967 x2: &[f32],
8968 rows: usize,
8969 cols: usize,
8970 o1: &mut [f32],
8971 o2: &mut [f32],
8972 pool: Option<&Pool>,
8973) {
8974 debug_assert_eq!(o1.len(), rows);
8975 debug_assert_eq!(o2.len(), rows);
8976 let (packed, scales) = q4_split(bytes, rows, cols);
8977 let gpr = cols / GROUP_SIZE;
8978
8979 if a8w8_enabled() {
8980 let a1 = split_act(x1);
8981 let a2 = split_act(x2);
8982 let p1 = SendMut(o1.as_mut_ptr());
8983 let p2 = SendMut(o2.as_mut_ptr());
8984 let run = move |start: usize, end: usize| {
8985 q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
8986 };
8987 dispatch_rows(pool, rows, &run);
8988 return;
8989 }
8990
8991 let p1 = SendMut(o1.as_mut_ptr());
8992 let p2 = SendMut(o2.as_mut_ptr());
8993 let run = move |start: usize, end: usize| {
8994 q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
8995 };
8996 dispatch_rows(pool, rows, &run);
8997}
8998
8999#[allow(clippy::too_many_arguments)]
9001fn q4_range2_f32(
9002 packed: &[u8],
9003 scales: &[u8],
9004 gpr: usize,
9005 x1: &[f32],
9006 x2: &[f32],
9007 p1: SendMut,
9008 p2: SendMut,
9009 start: usize,
9010 end: usize,
9011) {
9012 for r in start..end {
9013 let (mut acc1, mut acc2) = (0f32, 0f32);
9014 for gi in 0..gpr {
9015 let g = r * gpr + gi;
9016 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9017 let pk = &packed[g * 16..(g + 1) * 16];
9018 let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9019 let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9020 let (mut g1, mut g2) = (0f32, 0f32);
9021 for (k, &b) in pk.iter().enumerate() {
9022 let wl = (b & 0x0F) as f32 - 8.0;
9023 let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
9024 g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
9025 g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
9026 }
9027 acc1 += g1 * s;
9028 acc2 += g2 * s;
9029 }
9030 unsafe {
9032 *p1.at(r) = acc1;
9033 *p2.at(r) = acc2;
9034 }
9035 }
9036}
9037
9038thread_local! {
9039 static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
9042 static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9043}
9044
9045#[allow(clippy::too_many_arguments)]
9051fn q4matmat(
9052 bytes: &[u8],
9053 xs_all: &[f32],
9054 b: usize,
9055 rows: usize,
9056 cols: usize,
9057 out: &mut [f32],
9058 pool: Option<&Pool>,
9059) {
9060 debug_assert_eq!(xs_all.len(), b * cols);
9061 debug_assert_eq!(out.len(), b * rows);
9062 let (packed, scales) = q4_split(bytes, rows, cols);
9063 let gpr = cols / GROUP_SIZE;
9064 let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9065
9066 if a8w8_enabled() {
9067 let acts: Vec<SplitAct> = (0..b)
9068 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9069 .collect();
9070 let acts = &acts;
9071 let out_addr = SendMut(out.as_mut_ptr());
9072 let run = move |start: usize, end: usize| {
9073 ROW_I8.with(|rb| {
9074 let mut buf = rb.borrow_mut();
9075 buf.resize(cols, 0);
9076 for r in start..end {
9077 for gi in 0..gpr {
9081 let g = r * gpr + gi;
9082 for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9083 buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
9084 buf[gi * GROUP_SIZE + k * 2 + 1] =
9085 (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
9086 }
9087 }
9088 let mut bi = 0usize;
9089 #[cfg(target_arch = "x86_64")]
9090 if avx2_enabled() && blocked_enabled() {
9091 while bi + 4 <= acts.len() {
9092 let xs = [
9093 acts[bi].xq.as_slice(),
9094 acts[bi + 1].xq.as_slice(),
9095 acts[bi + 2].xq.as_slice(),
9096 acts[bi + 3].xq.as_slice(),
9097 ];
9098 let d = unsafe {
9099 if vnni_tiles_enabled() {
9100 dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
9101 } else {
9102 dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
9103 }
9104 };
9105 for k in 0..4 {
9106 let act = &acts[bi + k];
9107 let mut acc = d[k] * act.sx;
9108 for &(j, xv) in &act.outliers {
9109 acc += (buf[j] as i8) as f32
9110 * gscale((r * cols + j) / GROUP_SIZE)
9111 * xv;
9112 }
9113 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
9115 }
9116 bi += 4;
9117 }
9118 }
9119 while bi < acts.len() {
9120 let act = &acts[bi];
9121 let mut acc = 0f32;
9122 for gi in 0..gpr {
9123 let d = dot_i8_i8(
9124 &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9125 &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9126 );
9127 acc += d as f32 * gscale(r * gpr + gi);
9128 }
9129 acc *= act.sx;
9130 for &(j, xv) in &act.outliers {
9132 acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
9133 }
9134 unsafe { *out_addr.at(bi * rows + r) = acc };
9136 bi += 1;
9137 }
9138 }
9139 })
9140 };
9141 dispatch_rows(pool, rows, &run);
9142 return;
9143 }
9144
9145 let out_addr = SendMut(out.as_mut_ptr());
9146 let run = move |start: usize, end: usize| {
9147 ROW_F32.with(|rb| {
9148 let mut buf = rb.borrow_mut();
9149 buf.resize(cols, 0.0);
9150 for r in start..end {
9151 for gi in 0..gpr {
9154 let g = r * gpr + gi;
9155 for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9156 buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
9157 buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
9158 }
9159 }
9160 for bi in 0..b {
9161 let x = &xs_all[bi * cols..(bi + 1) * cols];
9162 let mut acc = 0f32;
9163 for gi in 0..gpr {
9164 let mut ga = 0f32;
9165 for k in 0..GROUP_SIZE / 2 {
9170 let e = gi * GROUP_SIZE + k * 2;
9171 ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
9172 }
9173 acc += ga * gscale(r * gpr + gi);
9174 }
9175 unsafe { *out_addr.at(bi * rows + r) = acc };
9177 }
9178 }
9179 })
9180 };
9181 dispatch_rows(pool, rows, &run);
9182}
9183
9184#[allow(clippy::too_many_arguments)]
9189fn vbitmatmat(
9190 bytes: &[u8],
9191 offsets: &[usize],
9192 xs_all: &[f32],
9193 b: usize,
9194 rows: usize,
9195 cols: usize,
9196 out: &mut [f32],
9197 pool: Option<&Pool>,
9198) {
9199 debug_assert_eq!(xs_all.len(), b * cols);
9200 debug_assert_eq!(out.len(), b * rows);
9201 debug_assert_eq!(offsets.len(), rows + 1);
9202 let ng = cols / GROUP_SIZE;
9203 let bits = &bytes[..rows];
9204 let sc_off = rows;
9205 let gscale = |r: usize, g: usize| {
9206 let so = (r * ng + g) * 2;
9207 f16_to_f32(u16::from_le_bytes([
9208 bytes[sc_off + so],
9209 bytes[sc_off + so + 1],
9210 ]))
9211 };
9212
9213 let decode_f32 = |r: usize, dst: &mut [f32]| {
9215 let bw = bits[r] as usize;
9216 let l = ((1i32 << (bw - 1)) - 1) as f32;
9217 let data = &bytes[offsets[r]..offsets[r + 1]];
9218 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
9219 for d in dst.iter_mut() {
9220 while nbits < bw {
9221 acc = (acc << 8) | data[idx] as u64;
9222 idx += 1;
9223 nbits += 8;
9224 }
9225 let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
9226 nbits -= bw;
9227 *d = u - l;
9228 }
9229 };
9230
9231 if a8w8_enabled() {
9232 let acts: Vec<SplitAct> = (0..b)
9233 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9234 .collect();
9235 let acts = &acts;
9236 let out_addr = SendMut(out.as_mut_ptr());
9237 let run = move |start: usize, end: usize| {
9238 for r in start..end {
9239 let bw = bits[r] as usize;
9240 if bw == 8 {
9241 ROW_F32.with(|rb| {
9244 let mut buf = rb.borrow_mut();
9245 buf.resize(cols, 0.0);
9246 decode_f32(r, &mut buf);
9247 for bi in 0..b {
9248 let x = &xs_all[bi * cols..(bi + 1) * cols];
9249 let mut dot = 0f32;
9250 for g in 0..ng {
9251 let mut gd = 0f32;
9252 for k in 0..GROUP_SIZE {
9253 gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9254 }
9255 dot += gd * gscale(r, g);
9256 }
9257 unsafe { *out_addr.at(bi * rows + r) = dot };
9259 }
9260 });
9261 continue;
9262 }
9263 let l = (1i32 << (bw - 1)) - 1;
9264 let data = &bytes[offsets[r]..offsets[r + 1]];
9265 ROW_I8.with(|rb| {
9266 let mut buf = rb.borrow_mut();
9267 buf.resize(cols, 0);
9268 #[inline(always)]
9269 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
9270 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
9271 let u = unpack8::<B>(&data[blk * B..]);
9272 for k in 0..8 {
9273 chunk[k] = (u[k] - l) as i8 as u8;
9274 }
9275 }
9276 }
9277 match bw {
9278 3 => fill::<3>(data, l, &mut buf),
9279 4 => vbit_fill4(data, &mut buf),
9280 5 => fill::<5>(data, l, &mut buf),
9281 6 => fill::<6>(data, l, &mut buf),
9282 _ => unreachable!("vbit bit-width {bw} (validated at load)"),
9283 }
9284 let mut bi = 0usize;
9285 #[cfg(target_arch = "x86_64")]
9289 if avx2_enabled() && blocked_enabled() {
9290 while bi + 4 <= acts.len() {
9291 let xs = [
9292 acts[bi].xq.as_slice(),
9293 acts[bi + 1].xq.as_slice(),
9294 acts[bi + 2].xq.as_slice(),
9295 acts[bi + 3].xq.as_slice(),
9296 ];
9297 let sxs = [
9298 acts[bi].sx,
9299 acts[bi + 1].sx,
9300 acts[bi + 2].sx,
9301 acts[bi + 3].sx,
9302 ];
9303 let d = unsafe {
9304 if vnni_tiles_enabled() {
9305 dot_q4b_row_1x4_sx_vnni(
9306 &buf,
9307 &bytes[sc_off..],
9308 r * ng,
9309 ng,
9310 xs,
9311 sxs,
9312 )
9313 } else {
9314 dot_q4b_row_1x4_sx_avx2(
9315 &buf,
9316 &bytes[sc_off..],
9317 r * ng,
9318 ng,
9319 xs,
9320 sxs,
9321 )
9322 }
9323 };
9324 for k in 0..4 {
9325 let act = &acts[bi + k];
9326 let mut dot = d[k];
9327 for &(j, xv) in &act.outliers {
9328 dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9329 }
9330 unsafe { *out_addr.at((bi + k) * rows + r) = dot };
9332 }
9333 bi += 4;
9334 }
9335 }
9336 while bi < acts.len() {
9337 let act = &acts[bi];
9338 let mut dot = 0f32;
9339 for g in 0..ng {
9340 let d = dot_i8_i8(
9341 &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9342 &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9343 ) as f32
9344 * act.sx;
9345 dot += d * gscale(r, g);
9346 }
9347 for &(j, xv) in &act.outliers {
9348 dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9349 }
9350 unsafe { *out_addr.at(bi * rows + r) = dot };
9352 bi += 1;
9353 }
9354 });
9355 }
9356 };
9357 dispatch_rows(pool, rows, &run);
9358 return;
9359 }
9360
9361 let out_addr = SendMut(out.as_mut_ptr());
9362 let run = move |start: usize, end: usize| {
9363 ROW_F32.with(|rb| {
9364 let mut buf = rb.borrow_mut();
9365 buf.resize(cols, 0.0);
9366 for r in start..end {
9367 decode_f32(r, &mut buf);
9368 for bi in 0..b {
9369 let x = &xs_all[bi * cols..(bi + 1) * cols];
9370 let mut dot = 0f32;
9371 for g in 0..ng {
9372 let mut gd = 0f32;
9373 for k in 0..GROUP_SIZE {
9374 gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9375 }
9376 dot += gd * gscale(r, g);
9377 }
9378 unsafe { *out_addr.at(bi * rows + r) = dot };
9380 }
9381 }
9382 })
9383 };
9384 dispatch_rows(pool, rows, &run);
9385}
9386
9387pub(crate) fn gpu_batch_job<'a>(
9391 t: &'a QTensor,
9392 x: &[f32],
9393) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
9394 match t {
9395 QTensor::Mapped {
9396 model,
9397 idx,
9398 dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
9399 rows,
9400 cols,
9401 row_scale,
9402 col_field,
9403 ..
9404 } => Some((
9405 model.clone(),
9406 crate::gpu::BatchJob {
9407 idx: *idx,
9408 rows: *rows,
9409 cols: *cols,
9410 row_scale,
9411 xs: prescale(x, col_field, *dt).into_owned(),
9412 layout: crate::gpu::BatchLayout::Q8,
9413 },
9414 )),
9415 QTensor::Mapped {
9417 model,
9418 idx,
9419 dtype: TensorDtype::Q1,
9420 rows,
9421 cols,
9422 ..
9423 } => Some((
9424 model.clone(),
9425 crate::gpu::BatchJob {
9426 idx: *idx,
9427 rows: *rows,
9428 cols: *cols,
9429 row_scale: &[],
9430 xs: x.to_vec(),
9431 layout: crate::gpu::BatchLayout::Q1,
9432 },
9433 )),
9434 QTensor::Mapped {
9439 model,
9440 idx,
9441 dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
9442 rows,
9443 cols,
9444 ..
9445 } => Some((
9446 model.clone(),
9447 crate::gpu::BatchJob {
9448 idx: *idx,
9449 rows: *rows,
9450 cols: *cols,
9451 row_scale: &[],
9452 xs: x.to_vec(),
9453 layout: if *dt == TensorDtype::Q4Tiled {
9454 crate::gpu::BatchLayout::Q4t
9455 } else {
9456 crate::gpu::BatchLayout::Q4tp
9457 },
9458 },
9459 )),
9460 _ => None,
9461 }
9462}
9463
9464thread_local! {
9465 static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9466 static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9467}
9468
9469pub(crate) fn prescale<'a>(
9470 x: &'a [f32],
9471 col_field: &[f32],
9472 dtype: TensorDtype,
9473) -> std::borrow::Cow<'a, [f32]> {
9474 if dtype == TensorDtype::Q8_2f {
9475 x.iter().zip(col_field).map(|(a, c)| a * c).collect()
9476 } else {
9477 std::borrow::Cow::Borrowed(x)
9478 }
9479}
9480
9481pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
9484 x: &[f32],
9485 col_field: &[f32],
9486 dtype: TensorDtype,
9487 buf_id: u8,
9488 f: F,
9489) -> R {
9490 if dtype == TensorDtype::Q8_2f {
9491 if buf_id == 1 {
9492 PRESCALE_BUF1.with(|b| {
9493 let mut buf = b.borrow_mut();
9494 buf.clear();
9495 buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9496 f(&buf)
9497 })
9498 } else {
9499 PRESCALE_BUF2.with(|b| {
9500 let mut buf = b.borrow_mut();
9501 buf.clear();
9502 buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9503 f(&buf)
9504 })
9505 }
9506 } else {
9507 f(x)
9508 }
9509}
9510
9511#[cfg(target_arch = "x86_64")]
9516pub(crate) fn avx2_enabled() -> bool {
9517 use std::sync::OnceLock;
9518 static ON: OnceLock<bool> = OnceLock::new();
9519 *ON.get_or_init(|| {
9520 std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
9521 && std::arch::is_x86_feature_detected!("avx2")
9522 && std::arch::is_x86_feature_detected!("fma")
9523 })
9524}
9525
9526#[cfg(target_arch = "x86_64")]
9531fn avx2_a8w8_enabled() -> bool {
9532 if FLOAT_ACTIVATIONS.get() {
9533 return false;
9534 }
9535 use std::sync::OnceLock;
9536 static ON: OnceLock<bool> = OnceLock::new();
9537 *ON.get_or_init(|| {
9538 avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
9539 })
9540}
9541
9542thread_local! {
9543 static FULL_GPU_Q8: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9544}
9545
9546pub(crate) fn enter_full_gpu_q8_scope() -> impl Drop {
9549 struct Restore(bool, std::marker::PhantomData<std::rc::Rc<()>>);
9550 impl Drop for Restore {
9551 fn drop(&mut self) {
9552 FULL_GPU_Q8.set(self.0);
9553 }
9554 }
9555 Restore(FULL_GPU_Q8.replace(true), std::marker::PhantomData)
9556}
9557
9558thread_local! {
9562 static FLOAT_ACTIVATIONS: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9563}
9564
9565pub(crate) fn float_activations_scope<R>(f: impl FnOnce() -> R) -> R {
9566 struct Restore(bool);
9567 impl Drop for Restore {
9568 fn drop(&mut self) {
9569 FLOAT_ACTIVATIONS.set(self.0);
9570 }
9571 }
9572 let _restore = Restore(FLOAT_ACTIVATIONS.replace(true));
9573 f()
9574}
9575
9576static ROW_EXACT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
9588
9589pub(crate) fn row_exact() -> bool {
9590 ROW_EXACT.load(std::sync::atomic::Ordering::Acquire) != 0
9591}
9592
9593fn counted_row_exact_scope<R>(active: &std::sync::atomic::AtomicUsize, f: impl FnOnce() -> R) -> R {
9594 struct Restore<'a>(&'a std::sync::atomic::AtomicUsize);
9595 impl Drop for Restore<'_> {
9596 fn drop(&mut self) {
9597 self.0.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
9598 }
9599 }
9600 active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
9601 let _restore = Restore(active);
9602 f()
9603}
9604
9605pub(crate) fn row_exact_scope<R>(f: impl FnOnce() -> R) -> R {
9607 counted_row_exact_scope(&ROW_EXACT, f)
9608}
9609
9610#[inline]
9614pub(crate) fn a8w8_enabled() -> bool {
9615 #[cfg(target_arch = "aarch64")]
9616 {
9617 sdot_enabled()
9618 }
9619 #[cfg(target_arch = "x86_64")]
9620 {
9621 avx2_a8w8_enabled()
9622 }
9623 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
9624 {
9625 false
9626 }
9627}
9628
9629#[inline]
9632#[allow(unreachable_code)]
9633fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
9634 #[cfg(target_arch = "aarch64")]
9635 unsafe {
9636 return dot_i8_sdot(w, xq);
9637 }
9638 #[cfg(target_arch = "x86_64")]
9639 unsafe {
9640 if avx512vnni_enabled() {
9641 return dot_i8_i8_vnni(w, xq);
9642 }
9643 return dot_i8_i8_avx2(w, xq);
9644 }
9645 w.iter()
9646 .zip(xq)
9647 .map(|(&a, &b)| (a as i8) as i32 * b as i32)
9648 .sum()
9649}
9650
9651#[cfg(target_arch = "x86_64")]
9655fn avx512vnni_enabled() -> bool {
9656 use std::sync::OnceLock;
9657 static ON: OnceLock<bool> = OnceLock::new();
9658 *ON.get_or_init(|| {
9659 std::env::var("CMF_AVX512")
9660 .map(|v| v != "0")
9661 .unwrap_or(true)
9662 && std::arch::is_x86_feature_detected!("avx512f")
9663 && std::arch::is_x86_feature_detected!("avx512bw")
9664 && std::arch::is_x86_feature_detected!("avx512vl")
9665 && std::arch::is_x86_feature_detected!("avx512vnni")
9666 })
9667}
9668
9669#[cfg(target_arch = "x86_64")]
9677fn vnni_tiles_enabled() -> bool {
9678 use std::sync::OnceLock;
9679 static ON: OnceLock<bool> = OnceLock::new();
9680 *ON.get_or_init(|| {
9681 std::env::var("CMF_VNNI_TILES")
9682 .map(|v| v != "0")
9683 .unwrap_or(true)
9684 && avx512vnni_enabled()
9685 })
9686}
9687
9688#[cfg(target_arch = "x86_64")]
9693#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9694#[inline]
9695unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
9696 unsafe {
9698 use core::arch::x86_64::*;
9699 let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
9700 let hi128 = _mm256_extracti128_si256::<1>(d);
9701 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9702 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9703 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9704 _mm_cvtsi128_si32(s32)
9705 }
9706}
9707
9708#[cfg(target_arch = "x86_64")]
9713#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9714unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
9715 unsafe {
9717 use core::arch::x86_64::*;
9718 let n = w.len();
9719 let mut j = 0usize;
9720 let mut total: i32;
9721 {
9726 #[inline(always)]
9727 unsafe fn step(
9728 w: *const u8,
9729 x: *const i8,
9730 acc: core::arch::x86_64::__m512i,
9731 ) -> core::arch::x86_64::__m512i {
9732 unsafe {
9733 use core::arch::x86_64::*;
9734 let wv = _mm512_loadu_si512(w as *const _);
9735 let xv = _mm512_loadu_si512(x as *const _);
9736 let aw = _mm512_abs_epi8(wv);
9737 let neg = _mm512_movepi8_mask(wv);
9738 let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
9739 _mm512_dpbusd_epi32(acc, aw, sx)
9740 }
9741 }
9742 let (mut a0, mut a1, mut a2, mut a3) = (
9743 _mm512_setzero_si512(),
9744 _mm512_setzero_si512(),
9745 _mm512_setzero_si512(),
9746 _mm512_setzero_si512(),
9747 );
9748 while j + 256 <= n {
9749 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9750 a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
9751 a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
9752 a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
9753 j += 256;
9754 }
9755 while j + 64 <= n {
9756 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9757 j += 64;
9758 }
9759 let s01 = _mm512_add_epi32(a0, a1);
9760 let s23 = _mm512_add_epi32(a2, a3);
9761 total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
9762 }
9763 if j + 32 <= n {
9765 let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9766 let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9767 let d = _mm256_dpbusd_epi32(
9768 _mm256_setzero_si256(),
9769 _mm256_abs_epi8(wv),
9770 _mm256_sign_epi8(xv, wv),
9771 );
9772 let hi128 = _mm256_extracti128_si256::<1>(d);
9773 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9774 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9775 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9776 total += _mm_cvtsi128_si32(s32);
9777 j += 32;
9778 }
9779 while j < n {
9780 total += (w[j] as i8) as i32 * xq[j] as i32;
9781 j += 1;
9782 }
9783 total
9784 }
9785}
9786
9787#[cfg(target_arch = "x86_64")]
9789#[target_feature(enable = "avx2,fma")]
9790unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
9791 unsafe {
9793 use core::arch::x86_64::*;
9794 let n = x.len();
9795 let wp = w.as_ptr();
9796 let xp = x.as_ptr();
9797 let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
9798 let mut j = 0usize;
9799 while j + 16 <= n {
9800 let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
9801 let lo = _mm256_cvtepi8_epi32(wb);
9802 let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
9803 a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
9804 a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
9805 j += 16;
9806 }
9807 let acc = _mm256_add_ps(a0, a1);
9808 let hi128 = _mm256_extractf128_ps::<1>(acc);
9809 let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
9810 let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
9811 let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
9812 let mut sum = _mm_cvtss_f32(s32);
9813 while j < n {
9814 sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
9815 j += 1;
9816 }
9817 sum
9818 }
9819}
9820
9821#[cfg(target_arch = "x86_64")]
9826#[target_feature(enable = "avx2")]
9827unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
9828 unsafe {
9830 use core::arch::x86_64::*;
9831 let n = w.len();
9832 let ones = _mm256_set1_epi16(1);
9833 let mut acc = _mm256_setzero_si256();
9834 let mut j = 0usize;
9835 while j + 32 <= n {
9836 let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9837 let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9838 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
9839 acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
9840 j += 32;
9841 }
9842 let hi128 = _mm256_extracti128_si256::<1>(acc);
9843 let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
9844 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9845 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9846 let mut s = _mm_cvtsi128_si32(s32);
9847 while j < n {
9848 s += (w[j] as i8) as i32 * xq[j] as i32;
9849 j += 1;
9850 }
9851 s
9852 }
9853}
9854
9855#[cfg(target_arch = "aarch64")]
9859#[target_feature(enable = "neon,i8mm")]
9860unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9861 unsafe {
9863 use core::arch::aarch64::*;
9864 use core::arch::asm;
9865 let n = w0.len();
9866 let w0p = w0.as_ptr() as *const i8;
9867 let w1p = w1.as_ptr() as *const i8;
9868 let mut acc01 = vdupq_n_s32(0);
9871 let mut acc23 = vdupq_n_s32(0);
9872 let mut i = 0usize;
9873 while i + 8 <= n {
9874 let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
9875 let xb01 = vcombine_s8(
9876 vld1_s8(xs[0].as_ptr().add(i)),
9877 vld1_s8(xs[1].as_ptr().add(i)),
9878 );
9879 let xb23 = vcombine_s8(
9880 vld1_s8(xs[2].as_ptr().add(i)),
9881 vld1_s8(xs[3].as_ptr().add(i)),
9882 );
9883 asm!(
9884 "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
9885 "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
9886 a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
9887 w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
9888 options(pure, nomem, nostack),
9889 );
9890 i += 8;
9891 }
9892 let mut out = [[0i32; 4]; 2];
9893 let a01: [i32; 4] = core::mem::transmute(acc01);
9894 let a23: [i32; 4] = core::mem::transmute(acc23);
9895 out[0][0] = a01[0];
9896 out[0][1] = a01[1];
9897 out[1][0] = a01[2];
9898 out[1][1] = a01[3];
9899 out[0][2] = a23[0];
9900 out[0][3] = a23[1];
9901 out[1][2] = a23[2];
9902 out[1][3] = a23[3];
9903 if i < n {
9904 for (k, x) in xs.iter().enumerate() {
9905 for j in i..n {
9906 out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9907 out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9908 }
9909 }
9910 }
9911 out
9912 }
9913}
9914
9915#[cfg(target_arch = "aarch64")]
9919#[target_feature(enable = "neon,dotprod")]
9920unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9921 unsafe {
9923 use core::arch::aarch64::*;
9924 use core::arch::asm;
9925 let n = w0.len();
9926 let w0p = w0.as_ptr() as *const i8;
9927 let w1p = w1.as_ptr() as *const i8;
9928 let mut acc = [[vdupq_n_s32(0); 4]; 2];
9929 let mut i = 0usize;
9930 while i + 16 <= n {
9931 let wv0 = vld1q_s8(w0p.add(i));
9932 let wv1 = vld1q_s8(w1p.add(i));
9933 for (k, x) in xs.iter().enumerate() {
9934 let xv = vld1q_s8(x.as_ptr().add(i));
9935 let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
9936 asm!(
9937 "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
9938 "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
9939 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
9940 w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
9941 options(pure, nomem, nostack),
9942 );
9943 acc[0][k] = a0;
9944 acc[1][k] = a1;
9945 }
9946 i += 16;
9947 }
9948 let mut out = [[0i32; 4]; 2];
9949 for r in 0..2 {
9950 for k in 0..4 {
9951 out[r][k] = vaddvq_s32(acc[r][k]);
9952 }
9953 }
9954 if i < n {
9955 for (k, x) in xs.iter().enumerate() {
9956 for j in i..n {
9957 out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9958 out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9959 }
9960 }
9961 }
9962 out
9963 }
9964}
9965
9966#[cfg(target_arch = "x86_64")]
9972#[target_feature(enable = "avx2")]
9973unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9974 unsafe {
9976 use core::arch::x86_64::*;
9977 let n = w0.len();
9978 let ones = _mm256_set1_epi16(1);
9979 let mut acc = [[_mm256_setzero_si256(); 4]; 2];
9980 let mut j = 0usize;
9981 while j + 32 <= n {
9982 let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
9983 let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
9984 let aw0 = _mm256_abs_epi8(wv0);
9985 let aw1 = _mm256_abs_epi8(wv1);
9986 for (k, x) in xs.iter().enumerate() {
9987 let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
9988 let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
9989 acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
9990 let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
9991 acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
9992 }
9993 j += 32;
9994 }
9995 let mut out = [[0i32; 4]; 2];
9996 for r in 0..2 {
9997 for k in 0..4 {
9998 let a = acc[r][k];
9999 let hi128 = _mm256_extracti128_si256::<1>(a);
10000 let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
10001 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10002 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10003 out[r][k] = _mm_cvtsi128_si32(s32);
10004 }
10005 }
10006 if j < n {
10007 for (k, x) in xs.iter().enumerate() {
10008 for i in j..n {
10009 out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
10010 out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
10011 }
10012 }
10013 }
10014 out
10015 }
10016}
10017
10018#[cfg(target_arch = "x86_64")]
10023#[inline]
10024fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
10025 let dot = if avx512vnni_enabled() && row.len() >= 64 {
10026 (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
10027 } else {
10028 unsafe { dot_i8_i8_avx2(row, &act.xq) }
10029 };
10030 let mut acc = dot as f32 * act.sx;
10031 for &(j, xv) in &act.outliers {
10032 acc += (row[j] as i8) as f32 * xv;
10033 }
10034 acc
10035}
10036
10037#[cfg(target_arch = "x86_64")]
10041#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
10042unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
10043 unsafe {
10045 use core::arch::x86_64::*;
10046 let n = w.len();
10047 let flip = _mm512_set1_epi8(-128); #[inline(always)]
10049 unsafe fn step(
10050 w: *const u8,
10051 x: *const i8,
10052 flip: core::arch::x86_64::__m512i,
10053 acc: core::arch::x86_64::__m512i,
10054 ) -> core::arch::x86_64::__m512i {
10055 unsafe {
10056 use core::arch::x86_64::*;
10057 let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
10058 _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
10059 }
10060 }
10061 let (mut a0, mut a1, mut a2, mut a3) = (
10062 _mm512_setzero_si512(),
10063 _mm512_setzero_si512(),
10064 _mm512_setzero_si512(),
10065 _mm512_setzero_si512(),
10066 );
10067 let mut j = 0usize;
10068 while j + 256 <= n {
10069 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10070 a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
10071 a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
10072 a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
10073 j += 256;
10074 }
10075 while j + 64 <= n {
10076 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10077 j += 64;
10078 }
10079 let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
10080 _mm512_add_epi32(a0, a1),
10081 _mm512_add_epi32(a2, a3),
10082 ));
10083 while j < n {
10085 total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
10086 j += 1;
10087 }
10088 total
10089 }
10090}
10091
10092#[cfg(target_arch = "x86_64")]
10098#[target_feature(enable = "avx2")]
10099unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10100 unsafe {
10103 use core::arch::x86_64::*;
10104 let lomask = _mm_set1_epi8(0x0F);
10105 let eight = _mm256_set1_epi8(8);
10106 let ones = _mm256_set1_epi16(1);
10107 let mut acc = 0f32;
10108 for gi in 0..gpr {
10109 let g = g0 + gi;
10110 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10111 let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10112 let lo = _mm_and_si128(b, lomask);
10113 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10114 let w = _mm256_sub_epi8(
10115 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10116 eight,
10117 );
10118 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10119 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
10120 let d = _mm256_madd_epi16(p16, ones);
10121 let hi128 = _mm256_extracti128_si256::<1>(d);
10122 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10123 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10124 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10125 acc += _mm_cvtsi128_si32(s32) as f32 * s;
10126 }
10127 acc
10128 }
10129}
10130
10131#[cfg(target_arch = "x86_64")]
10134#[target_feature(enable = "avx2")]
10135unsafe fn dot_q4_row_avx2_2(
10136 packed: &[u8],
10137 scales: &[u8],
10138 g0: usize,
10139 gpr: usize,
10140 xq1: &[i8],
10141 xq2: &[i8],
10142) -> (f32, f32) {
10143 unsafe {
10145 use core::arch::x86_64::*;
10146 let lomask = _mm_set1_epi8(0x0F);
10147 let eight = _mm256_set1_epi8(8);
10148 let ones = _mm256_set1_epi16(1);
10149 let (mut acc1, mut acc2) = (0f32, 0f32);
10150 #[inline(always)]
10151 unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
10152 unsafe {
10153 use core::arch::x86_64::*;
10154 let hi128 = _mm256_extracti128_si256::<1>(d);
10155 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10156 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10157 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10158 _mm_cvtsi128_si32(s32)
10159 }
10160 }
10161 for gi in 0..gpr {
10162 let g = g0 + gi;
10163 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10164 let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10165 let lo = _mm_and_si128(b, lomask);
10166 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10167 let w = _mm256_sub_epi8(
10168 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10169 eight,
10170 );
10171 let aw = _mm256_abs_epi8(w);
10172 let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10173 let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10174 let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
10175 let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
10176 acc1 += hsum(d1) as f32 * s;
10177 acc2 += hsum(d2) as f32 * s;
10178 }
10179 (acc1, acc2)
10180 }
10181}
10182
10183#[cfg(target_arch = "x86_64")]
10185fn q8_range_avx2(
10186 q: &[u8],
10187 row_scale: &[f32],
10188 act: &SplitAct,
10189 cols: usize,
10190 out_addr: SendMut,
10191 start: usize,
10192 end: usize,
10193) {
10194 for o in start..end {
10195 let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10196 unsafe { *out_addr.at(o) = v };
10198 }
10199}
10200
10201#[cfg(target_arch = "x86_64")]
10203#[allow(clippy::too_many_arguments)]
10204fn q8_range2_avx2(
10205 q: &[u8],
10206 row_scale: &[f32],
10207 a1: &SplitAct,
10208 a2: &SplitAct,
10209 cols: usize,
10210 p1: SendMut,
10211 p2: SendMut,
10212 start: usize,
10213 end: usize,
10214) {
10215 for o in start..end {
10216 let row = &q[o * cols..(o + 1) * cols];
10217 unsafe {
10219 *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
10220 *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
10221 }
10222 }
10223}
10224
10225#[cfg(target_arch = "aarch64")]
10236fn i8mm_enabled() -> bool {
10237 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10238 *ON.get_or_init(|| {
10239 std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
10240 && std::arch::is_aarch64_feature_detected!("i8mm")
10241 })
10242}
10243
10244#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
10248fn sdot_enabled() -> bool {
10249 if FLOAT_ACTIVATIONS.get() {
10250 return false;
10251 }
10252 use std::sync::OnceLock;
10253 static ON: OnceLock<bool> = OnceLock::new();
10254 *ON.get_or_init(|| {
10255 let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
10256 if !want {
10257 return false;
10258 }
10259
10260 #[cfg(target_arch = "aarch64")]
10261 {
10262 if std::arch::is_aarch64_feature_detected!("dotprod") {
10263 return true;
10264 }
10265 #[cfg(target_os = "android")]
10266 {
10267 if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
10268 if cpuinfo.lines().any(|l| {
10269 (l.starts_with("Features") || l.starts_with("features"))
10270 && l.contains("asimddp")
10271 }) {
10272 return true;
10273 }
10274 }
10275 }
10276 false
10277 }
10278 #[cfg(not(target_arch = "aarch64"))]
10279 {
10280 false
10281 }
10282 })
10283}
10284
10285struct SplitAct {
10290 xq: Vec<i8>,
10291 sx: f32,
10292 outliers: Vec<(usize, f32)>,
10293 #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
10296 xsum: i32,
10297}
10298
10299thread_local! {
10300 static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
10303 const { std::cell::RefCell::new(Vec::new()) };
10304}
10305
10306impl Drop for SplitAct {
10307 fn drop(&mut self) {
10308 let buf = std::mem::take(&mut self.xq);
10309 if buf.capacity() > 0 {
10310 XQ_FREE.with(|f| {
10311 let mut f = f.borrow_mut();
10312 if f.len() < 16 {
10313 f.push(buf);
10314 }
10315 });
10316 }
10317 }
10318}
10319
10320thread_local! {
10321 static KROW: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
10328}
10329
10330#[inline]
10333fn with_krow<R>(n: usize, f: impl FnOnce(&mut [f32]) -> R) -> R {
10334 KROW.with(|s| {
10335 let mut b = s.borrow_mut();
10336 if b.len() < n {
10337 b.resize(n, 0.0);
10338 }
10339 f(&mut b[..n])
10340 })
10341}
10342
10343#[inline(always)]
10354fn q8_round(t: f32) -> i8 {
10355 let t = t.clamp(-127.0, 127.0);
10356 let i = t as i32;
10357 let f = t - i as f32;
10358 let r = if f >= 0.5 {
10359 i + 1
10360 } else if f <= -0.5 {
10361 i - 1
10362 } else {
10363 i
10364 };
10365 r as i8
10366}
10367
10368fn split_act(x: &[f32]) -> SplitAct {
10369 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10370 let n = x.len();
10371 let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
10372 let thr = 8.0 * rms;
10373 let mut outliers: Vec<(usize, f32)> = Vec::new();
10377 let mut amax = 0f32;
10378 for (j, &v) in x.iter().enumerate() {
10379 let a = v.abs();
10380 if a > thr {
10381 outliers.push((j, v));
10382 } else if a > amax {
10383 amax = a;
10384 }
10385 }
10386 let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10387 let inv = 1.0 / sx;
10388 let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10389 xq.clear();
10390 xq.reserve(n);
10391 if outliers.is_empty() {
10392 xq.extend(
10393 x.iter()
10394 .map(|&v| q8_round(v * inv)),
10395 );
10396 } else {
10397 xq.extend(x.iter().map(|&v| {
10399 if v.abs() > thr {
10400 0
10401 } else {
10402 q8_round(v * inv)
10403 }
10404 }));
10405 }
10406 let xsum = xq.iter().map(|&v| v as i32).sum();
10407 SplitAct {
10408 xq,
10409 sx,
10410 outliers,
10411 xsum,
10412 }
10413}
10414
10415fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
10416 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10417 let n = x.len();
10418 let rms = (x
10419 .iter()
10420 .zip(col)
10421 .map(|(&a, &c)| {
10422 let v = a * c;
10423 (v * v) as f64
10424 })
10425 .sum::<f64>()
10426 / n.max(1) as f64)
10427 .sqrt() as f32;
10428 let thr = 8.0 * rms;
10429
10430 let mut outliers = Vec::new();
10431 let mut amax = 0f32;
10432 for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
10433 let v = a * c;
10434 let s = v.abs();
10435 if s > thr {
10436 outliers.push((j, v));
10437 } else if s > amax {
10438 amax = s;
10439 }
10440 }
10441
10442 let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10443 let inv = 1.0 / sx;
10444 let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10445 xq.clear();
10446 xq.reserve(n);
10447 if outliers.is_empty() {
10448 xq.extend(
10449 x.iter()
10450 .zip(col)
10451 .map(|(&a, &c)| q8_round((a * c) * inv)),
10452 );
10453 } else {
10454 xq.extend(x.iter().zip(col).map(|(&a, &c)| {
10455 let v = a * c;
10456 if v.abs() > thr {
10457 0
10458 } else {
10459 q8_round(v * inv)
10460 }
10461 }));
10462 }
10463 let xsum = xq.iter().map(|&v| v as i32).sum();
10464 SplitAct {
10465 xq,
10466 sx,
10467 outliers,
10468 xsum,
10469 }
10470}
10471
10472#[cfg(target_arch = "aarch64")]
10475#[target_feature(enable = "neon,dotprod")]
10476unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
10477 unsafe {
10479 use core::arch::aarch64::*;
10480 use core::arch::asm;
10481 let wp = w.as_ptr() as *const i8;
10482 let n = w.len();
10483 let (mut a0, mut a1, mut a2, mut a3) = (
10484 vdupq_n_s32(0),
10485 vdupq_n_s32(0),
10486 vdupq_n_s32(0),
10487 vdupq_n_s32(0),
10488 );
10489 let mut i = 0;
10490 while i + 64 <= n {
10491 let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10492 let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
10493 let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
10494 let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
10495 asm!(
10496 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
10497 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
10498 "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
10499 "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
10500 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10501 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
10502 w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
10503 options(pure, nomem, nostack),
10504 );
10505 i += 64;
10506 }
10507 while i + 16 <= n {
10508 let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10509 asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
10510 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
10511 i += 16;
10512 }
10513 let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
10514 while i < n {
10515 s += (*wp.add(i)) as i32 * xq[i] as i32;
10516 i += 1;
10517 }
10518 s
10519 }
10520}
10521
10522#[cfg(target_arch = "aarch64")]
10526#[target_feature(enable = "neon,dotprod")]
10527unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
10528 unsafe {
10530 use core::arch::aarch64::*;
10531 use core::arch::asm;
10532 let n = xq.len();
10533 let px = xq.as_ptr();
10534 let (p0, p1, p2, p3) = (
10535 w0.as_ptr() as *const i8,
10536 w1.as_ptr() as *const i8,
10537 w2.as_ptr() as *const i8,
10538 w3.as_ptr() as *const i8,
10539 );
10540 let (mut a0, mut a1, mut a2, mut a3) = (
10541 vdupq_n_s32(0),
10542 vdupq_n_s32(0),
10543 vdupq_n_s32(0),
10544 vdupq_n_s32(0),
10545 );
10546 let mut i = 0;
10547 while i + 16 <= n {
10548 let x = vld1q_s8(px.add(i));
10549 let v0 = vld1q_s8(p0.add(i));
10550 let v1 = vld1q_s8(p1.add(i));
10551 let v2 = vld1q_s8(p2.add(i));
10552 let v3 = vld1q_s8(p3.add(i));
10553 asm!(
10554 "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10555 "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10556 "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10557 "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10558 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10559 v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10560 options(pure, nomem, nostack),
10561 );
10562 i += 16;
10563 }
10564 let mut r = [
10565 vaddvq_s32(a0),
10566 vaddvq_s32(a1),
10567 vaddvq_s32(a2),
10568 vaddvq_s32(a3),
10569 ];
10570 while i < n {
10571 let xi = *px.add(i) as i32;
10572 r[0] += (*p0.add(i)) as i32 * xi;
10573 r[1] += (*p1.add(i)) as i32 * xi;
10574 r[2] += (*p2.add(i)) as i32 * xi;
10575 r[3] += (*p3.add(i)) as i32 * xi;
10576 i += 1;
10577 }
10578 r
10579 }
10580}
10581
10582#[cfg(target_arch = "aarch64")]
10589#[target_feature(enable = "neon,dotprod")]
10590unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
10591 unsafe {
10594 use core::arch::aarch64::*;
10595 use core::arch::asm;
10596 let n = xq.len();
10597 let px = xq.as_ptr();
10598 let pg = g.as_ptr() as *const i8;
10599 let (mut a0, mut a1, mut a2, mut a3) = (
10600 vdupq_n_s32(0),
10601 vdupq_n_s32(0),
10602 vdupq_n_s32(0),
10603 vdupq_n_s32(0),
10604 );
10605 let mut i = 0;
10606 while i + 16 <= n {
10607 let x = vld1q_s8(px.add(i));
10608 let base = pg.add(4 * i);
10609 let v0 = vld1q_s8(base);
10610 let v1 = vld1q_s8(base.add(16));
10611 let v2 = vld1q_s8(base.add(32));
10612 let v3 = vld1q_s8(base.add(48));
10613 asm!(
10614 "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10615 "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10616 "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10617 "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10618 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10619 v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10620 options(pure, nomem, nostack),
10621 );
10622 i += 16;
10623 }
10624 [
10625 vaddvq_s32(a0),
10626 vaddvq_s32(a1),
10627 vaddvq_s32(a2),
10628 vaddvq_s32(a3),
10629 ]
10630 }
10631}
10632
10633#[cfg(target_arch = "aarch64")]
10639fn q8_range_sdot(
10640 q: &[u8],
10641 rep: &[u8],
10642 row_scale: &[f32],
10643 act: &SplitAct,
10644 cols: usize,
10645 out_addr: SendMut,
10646 start: usize,
10647 end: usize,
10648) {
10649 let mut o = start;
10650 if !rep.is_empty() {
10653 while o < end && o % 4 != 0 {
10654 let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10655 unsafe { *out_addr.at(o) = v };
10656 o += 1;
10657 }
10658 }
10659 while o + 4 <= end {
10660 let r = if rep.is_empty() {
10661 unsafe {
10662 dot_i8_sdot_4rows(
10663 &q[o * cols..(o + 1) * cols],
10664 &q[(o + 1) * cols..(o + 2) * cols],
10665 &q[(o + 2) * cols..(o + 3) * cols],
10666 &q[(o + 3) * cols..(o + 4) * cols],
10667 &act.xq,
10668 )
10669 }
10670 } else {
10671 unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
10672 };
10673 for k in 0..4 {
10674 let mut acc = r[k] as f32 * act.sx;
10675 for &(j, xv) in &act.outliers {
10676 acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
10677 }
10678 unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
10680 }
10681 o += 4;
10682 }
10683 while o < end {
10684 let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10685 unsafe { *out_addr.at(o) = v };
10686 o += 1;
10687 }
10688}
10689
10690#[cfg(target_arch = "aarch64")]
10693#[allow(clippy::too_many_arguments)]
10694fn q8_range2_sdot(
10695 q: &[u8],
10696 row_scale: &[f32],
10697 a1: &SplitAct,
10698 a2: &SplitAct,
10699 cols: usize,
10700 p1: SendMut,
10701 p2: SendMut,
10702 start: usize,
10703 end: usize,
10704) {
10705 for o in start..end {
10706 let row = &q[o * cols..(o + 1) * cols];
10707 unsafe {
10709 *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
10710 *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
10711 }
10712 }
10713}
10714
10715#[allow(clippy::too_many_arguments)]
10717fn q8_range2_f32(
10718 q: &[u8],
10719 row_scale: &[f32],
10720 x1: &[f32],
10721 x2: &[f32],
10722 cols: usize,
10723 p1: SendMut,
10724 p2: SendMut,
10725 start: usize,
10726 end: usize,
10727) {
10728 for o in start..end {
10729 let row = &q[o * cols..(o + 1) * cols];
10730 unsafe {
10732 *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
10733 *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
10734 }
10735 }
10736}
10737
10738fn q8_range_f32(
10740 q: &[u8],
10741 row_scale: &[f32],
10742 xs: &[f32],
10743 cols: usize,
10744 out_addr: SendMut,
10745 start: usize,
10746 end: usize,
10747) {
10748 for o in start..end {
10749 let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
10750 unsafe { *out_addr.at(o) = v };
10752 }
10753}
10754
10755#[inline]
10759fn q8_row_dot(row: &[u8], act: &SplitAct) -> f32 {
10760 #[cfg(target_arch = "aarch64")]
10761 return row_dot_sdot(row, act);
10762 #[cfg(target_arch = "x86_64")]
10763 return row_dot_avx2(row, act);
10764 #[allow(unreachable_code)]
10765 q8_row_dot_scalar(row, act)
10766}
10767
10768#[allow(dead_code)]
10769fn q8_row_dot_scalar(row: &[u8], act: &SplitAct) -> f32 {
10770 let mut acc = 0i32;
10771 for (k, &b) in row.iter().enumerate() {
10772 acc += (b as i8) as i32 * act.xq[k] as i32;
10773 }
10774 let mut acc = acc as f32 * act.sx;
10775 for &(j, xv) in &act.outliers {
10776 acc += (row[j] as i8) as f32 * xv;
10777 }
10778 acc
10779}
10780
10781#[cfg(target_arch = "aarch64")]
10784#[inline]
10785fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
10786 let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
10787 for &(j, xv) in &act.outliers {
10788 acc += (row[j] as i8) as f32 * xv;
10789 }
10790 acc
10791}
10792
10793#[cfg(target_arch = "aarch64")]
10801#[target_feature(enable = "neon,dotprod")]
10802unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10803 unsafe {
10806 use core::arch::aarch64::*;
10807 use core::arch::asm;
10808 let lomask = vdupq_n_u8(0x0F);
10809 let eight = vdupq_n_s8(8);
10810 let mut acc = 0f32;
10811 for gi in 0..gpr {
10812 let g = g0 + gi;
10813 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10814 let b = vld1q_u8(packed.as_ptr().add(g * 16));
10815 let lo = vandq_u8(b, lomask);
10816 let hi = vshrq_n_u8::<4>(b);
10817 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10818 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10819 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
10820 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
10821 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
10822 asm!(
10823 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
10824 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
10825 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10826 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
10827 options(pure, nomem, nostack),
10828 );
10829 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10830 }
10831 acc
10832 }
10833}
10834
10835#[cfg(target_arch = "aarch64")]
10840#[target_feature(enable = "neon,dotprod")]
10841unsafe fn dot_q4_row_sdot2(
10842 packed: &[u8],
10843 scales: &[u8],
10844 g0: usize,
10845 gpr: usize,
10846 xq1: &[i8],
10847 xq2: &[i8],
10848) -> (f32, f32) {
10849 unsafe {
10852 use core::arch::aarch64::*;
10853 use core::arch::asm;
10854 let lomask = vdupq_n_u8(0x0F);
10855 let eight = vdupq_n_s8(8);
10856 let (mut acc1, mut acc2) = (0f32, 0f32);
10857 for gi in 0..gpr {
10858 let g = g0 + gi;
10859 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10860 let b = vld1q_u8(packed.as_ptr().add(g * 16));
10861 let lo = vandq_u8(b, lomask);
10862 let hi = vshrq_n_u8::<4>(b);
10863 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10864 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10865 let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
10866 let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
10867 let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
10868 let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
10869 let (mut a0, mut a1, mut b0, mut b1) = (
10870 vdupq_n_s32(0),
10871 vdupq_n_s32(0),
10872 vdupq_n_s32(0),
10873 vdupq_n_s32(0),
10874 );
10875 asm!(
10876 "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
10877 "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
10878 "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
10879 "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
10880 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10881 b0 = inout(vreg) b0, b1 = inout(vreg) b1,
10882 e0 = in(vreg) e0, e1 = in(vreg) e1,
10883 x10 = in(vreg) x10, x11 = in(vreg) x11,
10884 x20 = in(vreg) x20, x21 = in(vreg) x21,
10885 options(pure, nomem, nostack),
10886 );
10887 acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10888 acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
10889 }
10890 (acc1, acc2)
10891 }
10892}
10893
10894#[inline]
10899pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
10900 #[cfg(target_arch = "aarch64")]
10901 unsafe {
10902 return axpy_i8_f32_neon(acc, row, w);
10903 }
10904 #[cfg(target_arch = "x86_64")]
10905 if avx2_enabled() {
10906 return unsafe { axpy_i8_f32_avx2(acc, row, w) };
10907 }
10908 #[allow(unreachable_code)]
10909 {
10910 for (a, &b) in acc.iter_mut().zip(row) {
10911 *a += w * b as f32;
10912 }
10913 }
10914}
10915
10916#[cfg(target_arch = "x86_64")]
10918#[target_feature(enable = "avx2,fma")]
10919unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
10920 unsafe {
10922 use core::arch::x86_64::*;
10923 let n = acc.len().min(row.len());
10924 let ap = acc.as_mut_ptr();
10925 let rp = row.as_ptr();
10926 let wv = _mm256_set1_ps(w);
10927 let mut j = 0usize;
10928 while j + 16 <= n {
10929 let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
10930 let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
10931 let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
10932 let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
10933 let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
10934 _mm256_storeu_ps(ap.add(j), v0);
10935 _mm256_storeu_ps(ap.add(j + 8), v1);
10936 j += 16;
10937 }
10938 while j < n {
10939 *ap.add(j) += w * (*rp.add(j)) as f32;
10940 j += 1;
10941 }
10942 }
10943}
10944
10945#[cfg(target_arch = "aarch64")]
10946#[target_feature(enable = "neon")]
10947unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
10948 unsafe {
10950 use core::arch::aarch64::*;
10951 let n = acc.len().min(row.len());
10952 let ap = acc.as_mut_ptr();
10953 let rp = row.as_ptr();
10954 let wv = vdupq_n_f32(w);
10955 let mut j = 0usize;
10956 while j + 16 <= n {
10957 let rb = vld1q_s8(rp.add(j));
10958 let lo = vmovl_s8(vget_low_s8(rb));
10959 let hi = vmovl_s8(vget_high_s8(rb));
10960 for (off, half) in [(0, lo), (8, hi)] {
10961 let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
10962 let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
10963 let o = j + off;
10964 vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
10965 vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
10966 }
10967 j += 16;
10968 }
10969 while j < n {
10970 *ap.add(j) += w * (*rp.add(j)) as f32;
10971 j += 1;
10972 }
10973 }
10974}
10975
10976#[inline]
10979pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
10980 #[cfg(target_arch = "aarch64")]
10981 unsafe {
10982 return dot_i8_f32_neon(w, x);
10983 }
10984 #[cfg(target_arch = "x86_64")]
10985 if avx2_enabled() {
10986 return unsafe { dot_i8_f32_avx2(w, x) };
10987 }
10988 #[allow(unreachable_code)]
10989 {
10990 let mut sum = 0.0f32;
10991 for (j, &b) in w.iter().enumerate() {
10992 sum += (b as i8) as f32 * x[j];
10993 }
10994 sum
10995 }
10996}
10997
10998#[inline]
11002fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11003 #[cfg(target_arch = "aarch64")]
11004 unsafe {
11005 return dot_i8_col_f32_neon(w, x, col);
11006 }
11007 #[allow(unreachable_code)]
11008 {
11009 let mut sum = 0.0f32;
11010 for (j, &b) in w.iter().enumerate() {
11011 sum += (b as i8) as f32 * x[j] * col[j];
11012 }
11013 sum
11014 }
11015}
11016
11017#[cfg(target_arch = "aarch64")]
11018#[target_feature(enable = "neon")]
11019unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11020 unsafe {
11022 use core::arch::aarch64::*;
11023 let n = x.len();
11024 let wp = w.as_ptr() as *const i8;
11025 let xp = x.as_ptr();
11026 let cp = col.as_ptr();
11027 let (mut a0, mut a1, mut a2, mut a3) = (
11028 vdupq_n_f32(0.0),
11029 vdupq_n_f32(0.0),
11030 vdupq_n_f32(0.0),
11031 vdupq_n_f32(0.0),
11032 );
11033 let mut j = 0usize;
11034 while j + 16 <= n {
11035 let wb = vld1q_s8(wp.add(j));
11036 let lo = vmovl_s8(vget_low_s8(wb));
11037 let hi = vmovl_s8(vget_high_s8(wb));
11038 let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11039 let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11040 let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11041 let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11042 a0 = vfmaq_f32(
11043 a0,
11044 w0,
11045 vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
11046 );
11047 a1 = vfmaq_f32(
11048 a1,
11049 w1,
11050 vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
11051 );
11052 a2 = vfmaq_f32(
11053 a2,
11054 w2,
11055 vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
11056 );
11057 a3 = vfmaq_f32(
11058 a3,
11059 w3,
11060 vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
11061 );
11062 j += 16;
11063 }
11064 let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11065 while j < n {
11066 sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
11067 j += 1;
11068 }
11069 sum
11070 }
11071}
11072
11073#[cfg(target_arch = "aarch64")]
11074#[target_feature(enable = "neon")]
11075unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
11076 unsafe {
11078 use core::arch::aarch64::*;
11079 let n = x.len();
11080 let wp = w.as_ptr() as *const i8;
11081 let xp = x.as_ptr();
11082 let (mut a0, mut a1, mut a2, mut a3) = (
11083 vdupq_n_f32(0.0),
11084 vdupq_n_f32(0.0),
11085 vdupq_n_f32(0.0),
11086 vdupq_n_f32(0.0),
11087 );
11088 let mut j = 0usize;
11089 while j + 16 <= n {
11090 let wb = vld1q_s8(wp.add(j));
11091 let lo = vmovl_s8(vget_low_s8(wb));
11092 let hi = vmovl_s8(vget_high_s8(wb));
11093 let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11094 let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11095 let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11096 let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11097 a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
11098 a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
11099 a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
11100 a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
11101 j += 16;
11102 }
11103 let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11104 while j < n {
11105 sum += (*wp.add(j)) as f32 * *xp.add(j);
11106 j += 1;
11107 }
11108 sum
11109 }
11110}
11111
11112#[allow(clippy::too_many_arguments)]
11113fn qmatvec(
11114 q: &[u8],
11115 rep: &[u8],
11116 row_scale: &[f32],
11117 x: &[f32],
11118 col_field: &[f32],
11119 dtype: TensorDtype,
11120 rows: usize,
11121 cols: usize,
11122 out: &mut [f32],
11123 pool: Option<&Pool>,
11124) {
11125 debug_assert_eq!(out.len(), rows);
11126 #[cfg(not(target_arch = "aarch64"))]
11127 let _ = rep;
11128
11129 #[cfg(target_arch = "aarch64")]
11130 if sdot_enabled() {
11131 let act = if dtype == TensorDtype::Q8_2f {
11132 split_act_q8_2f(x, col_field)
11133 } else {
11134 split_act(x)
11135 };
11136 let out_addr = SendMut(out.as_mut_ptr());
11137 let run_range = |start: usize, end: usize| {
11138 q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
11139 };
11140 match pool {
11141 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11142 _ => run_range(0, rows),
11143 }
11144 return;
11145 }
11146 #[cfg(target_arch = "x86_64")]
11149 if avx2_a8w8_enabled() {
11150 let act = if dtype == TensorDtype::Q8_2f {
11151 split_act_q8_2f(x, col_field)
11152 } else {
11153 split_act(x)
11154 };
11155 let out_addr = SendMut(out.as_mut_ptr());
11156 let run_range = |start: usize, end: usize| {
11157 q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
11158 };
11159 match pool {
11160 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11161 _ => run_range(0, rows),
11162 }
11163 return;
11164 }
11165
11166 prescale_with(x, col_field, dtype, 1, |xs| {
11167 let out_addr = SendMut(out.as_mut_ptr());
11168 let run_range = move |start: usize, end: usize| {
11169 for o in start..end {
11170 let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
11171 unsafe { *out_addr.at(o) = v };
11173 }
11174 };
11175 match pool {
11176 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11177 _ => run_range(0, rows),
11178 }
11179 });
11180}
11181
11182#[allow(clippy::too_many_arguments)]
11183fn qmatvec2(
11184 q: &[u8],
11185 row_scale: &[f32],
11186 x1: &[f32],
11187 x2: &[f32],
11188 col_field: &[f32],
11189 dtype: TensorDtype,
11190 rows: usize,
11191 cols: usize,
11192 o1: &mut [f32],
11193 o2: &mut [f32],
11194 pool: Option<&Pool>,
11195) {
11196 #[cfg(target_arch = "aarch64")]
11197 if sdot_enabled() {
11198 let a1s = if dtype == TensorDtype::Q8_2f {
11199 split_act_q8_2f(x1, col_field)
11200 } else {
11201 split_act(x1)
11202 };
11203 let a2s = if dtype == TensorDtype::Q8_2f {
11204 split_act_q8_2f(x2, col_field)
11205 } else {
11206 split_act(x2)
11207 };
11208 let p1 = SendMut(o1.as_mut_ptr());
11209 let p2 = SendMut(o2.as_mut_ptr());
11210 let run_range = |start: usize, end: usize| {
11211 q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11212 };
11213 match pool {
11214 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11215 _ => run_range(0, rows),
11216 }
11217 return;
11218 }
11219 #[cfg(target_arch = "x86_64")]
11220 if avx2_a8w8_enabled() {
11221 let a1s = if dtype == TensorDtype::Q8_2f {
11222 split_act_q8_2f(x1, col_field)
11223 } else {
11224 split_act(x1)
11225 };
11226 let a2s = if dtype == TensorDtype::Q8_2f {
11227 split_act_q8_2f(x2, col_field)
11228 } else {
11229 split_act(x2)
11230 };
11231 let p1 = SendMut(o1.as_mut_ptr());
11232 let p2 = SendMut(o2.as_mut_ptr());
11233 let run_range = |start: usize, end: usize| {
11234 q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11235 };
11236 match pool {
11237 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11238 _ => run_range(0, rows),
11239 }
11240 return;
11241 }
11242
11243 prescale_with(x1, col_field, dtype, 1, |x1s| {
11244 prescale_with(x2, col_field, dtype, 2, |x2s| {
11245 let p1 = SendMut(o1.as_mut_ptr());
11246 let p2 = SendMut(o2.as_mut_ptr());
11247 let run_range = move |start: usize, end: usize| {
11248 for o in start..end {
11249 let row = &q[o * cols..(o + 1) * cols];
11250 let s1 = dot_i8_f32(row, x1s) * row_scale[o];
11251 let s2 = dot_i8_f32(row, x2s) * row_scale[o];
11252 unsafe {
11254 *p1.at(o) = s1;
11255 *p2.at(o) = s2;
11256 }
11257 }
11258 };
11259 match pool {
11260 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11261 _ => run_range(0, rows),
11262 }
11263 });
11264 });
11265}
11266
11267#[derive(Clone, Copy)]
11268struct SendMut(*mut f32);
11269unsafe impl Send for SendMut {}
11270unsafe impl Sync for SendMut {}
11271
11272impl SendMut {
11273 #[inline]
11274 fn at(self, i: usize) -> *mut f32 {
11275 unsafe { self.0.add(i) }
11276 }
11277}
11278
11279#[cfg(test)]
11280mod tests {
11281 #[test]
11285 fn q8_round_is_round_clamp() {
11286 let reference = |t: f32| t.round().clamp(-127.0, 127.0) as i8;
11287 let mut probe = vec![
11288 0.0f32,
11289 -0.0,
11290 f32::NAN,
11291 f32::INFINITY,
11292 f32::NEG_INFINITY,
11293 f32::MAX,
11294 f32::MIN,
11295 1e30,
11296 -1e30,
11297 f32::MIN_POSITIVE,
11298 -f32::MIN_POSITIVE,
11299 ];
11300 for k in -300i32..=300 {
11301 let h = k as f32 * 0.5;
11302 let up = f32::from_bits(h.to_bits() + 1);
11303 let down = f32::from_bits(h.to_bits().wrapping_sub(1));
11304 for t in [h, up, down] {
11305 probe.push(t);
11306 probe.push(-t);
11307 }
11308 }
11309 let mut t = -140.0f32;
11310 while t < 140.0 {
11311 probe.push(t);
11312 t += 0.000_731;
11313 }
11314 for t in probe {
11315 assert_eq!(q8_round(t), reference(t), "t = {t:e} ({:#x})", t.to_bits());
11316 }
11317 }
11318
11319 use super::*;
11320
11321 #[test]
11322 fn q2tp_i8_dot_matches_exact_on_grid() {
11323 let (rows, cols) = (5, 64);
11327 let gpr = cols / GROUP_SIZE;
11328 let chunks: Vec<u8> = (0..rows * gpr * Q2TP_CHUNK)
11332 .map(|i| (i as u32).wrapping_mul(2654435761) as u8)
11333 .collect();
11334 let scales: Vec<f32> = (0..gpr).map(|g| 0.5 + g as f32 * 0.25).collect();
11335 let x: Vec<f32> = (0..cols)
11336 .map(|i| if i % 3 == 0 { -1.0 } else { 1.0 })
11337 .collect();
11338 let act = split_act(&x);
11339 assert!(
11340 act.outliers.is_empty(),
11341 "on-grid input must have no outliers"
11342 );
11343 let gsum = q1_group_sums(&act.xq, gpr);
11344 for r in 0..rows {
11345 let exact = q2tp_row_exact(&chunks, r, gpr, &x, &scales);
11346 let fast = dot_q2tp_row_i8(&chunks, r, gpr, &act.xq, &gsum, &scales) * act.sx;
11347 assert!(
11348 (exact - fast).abs() <= exact.abs() * 1e-5 + 1e-5,
11349 "row {r}: exact {exact} vs i8 {fast}"
11350 );
11351 }
11352 }
11353
11354 #[test]
11355 fn q2tp_affine_fuses_half_scale_correction_without_changing_raw_decode() {
11356 let (rows, cols) = (1usize, GROUP_SIZE);
11357 let mut bytes = vec![0u8; Q2TP_CHUNK + 4 + 1];
11358 bytes[..Q2TP_CHUNK].fill(0x24); bytes[Q2TP_CHUNK..Q2TP_CHUNK + 2].copy_from_slice(&0u16.to_le_bytes());
11362 bytes[Q2TP_CHUNK + 2..Q2TP_CHUNK + 4].copy_from_slice(&0u16.to_le_bytes());
11363 bytes[Q2TP_CHUNK + 4] = 1; let x = vec![1.0f32; cols];
11365 let mut raw = vec![0.0f32; rows];
11366 let mut affine = vec![0.0f32; rows];
11367 q2tp_matvec_for_test(&bytes, &x, rows, cols, &mut raw);
11368 q2tp_affine_matvec_for_test(&bytes, &x, rows, cols, &mut affine);
11369 assert_eq!(raw, vec![-24.0]);
11370 assert_eq!(affine, vec![-8.0]);
11371 assert!((affine[0] - (raw[0] + 0.5 * cols as f32)).abs() < 1e-6);
11372 }
11373
11374 #[cfg(target_arch = "x86_64")]
11375 #[test]
11376 fn q2tp_avx2_dot_matches_scalar_for_random_patterns() {
11377 if !std::arch::is_x86_feature_detected!("avx2") {
11382 return;
11383 }
11384 let mut seed = 0x9e3779b9u32;
11385 let mut next = || {
11386 seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
11387 seed
11388 };
11389 for _ in 0..20_000 {
11390 let mut ch = [0u8; Q2TP_CHUNK];
11391 let mut x = [0i8; GROUP_SIZE];
11392 for b in &mut ch {
11393 *b = next() as u8;
11394 }
11395 for v in &mut x {
11396 *v = (next() >> 24) as i8;
11397 }
11398 let mut reference = 0i32;
11399 for (k, &b) in ch.iter().enumerate() {
11400 reference += (b & 3) as i32 * x[k * 4] as i32;
11401 reference += ((b >> 2) & 3) as i32 * x[k * 4 + 1] as i32;
11402 reference += ((b >> 4) & 3) as i32 * x[k * 4 + 2] as i32;
11403 reference += ((b >> 6) & 3) as i32 * x[k * 4 + 3] as i32;
11404 }
11405 let got = unsafe { q2tp_code_dot_avx2(&ch, &x) };
11408 assert_eq!(got, reference, "packed q2 lane mismatch");
11409 }
11410 }
11411
11412 #[test]
11413 fn q8_row_dot_fast_matches_scalar() {
11414 let cols = 96;
11417 let row: Vec<u8> = (0..cols)
11418 .map(|i| ((i * 37 % 251) - 125) as i8 as u8)
11419 .collect();
11420 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.13).sin()).collect();
11421 let act = split_act(&x);
11422 let fast = q8_row_dot(&row, &act);
11423 let scalar = q8_row_dot_scalar(&row, &act);
11424 assert!(
11425 (fast - scalar).abs() <= scalar.abs() * 1e-5 + 1e-5,
11426 "fast {fast} vs scalar {scalar}"
11427 );
11428 }
11429
11430 #[test]
11431 fn f32_matvec_matches_matvec_rows_bitexact() {
11432 let (rows, cols) = (300, 40);
11433 let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
11434 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
11435 let qt = QTensor::from_f32(w.clone(), rows, cols);
11436
11437 let mut a = vec![0.0f32; rows];
11438 matvec_rows(None, &w, &x, &mut a);
11439 let mut b = vec![0.0f32; rows];
11440 qt.matvec(&x, &mut b, None);
11441 assert_eq!(a, b);
11442 }
11443
11444 #[test]
11445 fn sdot_kernel_exact_on_grid() {
11446 eprintln!("sdot_enabled = {}", sdot_enabled());
11451 let (rows, cols) = (9, 80); let w: Vec<u8> = (0..rows * cols)
11453 .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11454 .collect();
11455 let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
11456 let x: Vec<f32> = (0..cols)
11457 .map(|i| match i % 3 {
11458 0 => 1.0,
11459 1 => -1.0,
11460 _ => 0.0,
11461 })
11462 .collect();
11463 let mut a = vec![0.0f32; rows];
11464 qmatvec(
11465 &w,
11466 &[],
11467 &scales,
11468 &x,
11469 &[],
11470 TensorDtype::Q8Row,
11471 rows,
11472 cols,
11473 &mut a,
11474 None,
11475 );
11476 for o in 0..rows {
11477 let mut acc = 0.0f32;
11478 for j in 0..cols {
11479 acc += (w[o * cols + j] as i8) as f32 * x[j];
11480 }
11481 let expect = acc * scales[o];
11482 assert!(
11483 (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11484 "row {o}: {} vs {expect}",
11485 a[o]
11486 );
11487 }
11488 }
11489
11490 #[test]
11491 fn q1_tbl_fast_path_matches_reference() {
11492 let (rows, cols) = (5, 256);
11497 let gpr = cols / GROUP_SIZE;
11498 let mut bytes = Vec::new();
11499 for t in 0..rows * gpr {
11500 let s = 0.007 + (t % 11) as f32 * 0.004;
11501 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11502 for j in 0..4 {
11503 bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
11504 }
11505 }
11506 let x: Vec<f32> = (0..cols)
11507 .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
11508 .collect();
11509 let mut w = vec![0.0f32; rows * cols];
11510 cortiq_core::quant::dequant_q1(&bytes, &mut w);
11511 let mut got = vec![0.0f32; rows];
11512 q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11513 for o in 0..rows {
11514 let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11515 assert!(
11516 (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11517 "row {o}: {} vs {expect}",
11518 got[o]
11519 );
11520 }
11521 let b = 5usize;
11524 let mut xs_all = Vec::new();
11525 for bi in 0..b {
11526 xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
11527 }
11528 let mut mm = vec![0.0f32; b * rows];
11529 q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
11530 for bi in 0..b {
11531 let mut single = vec![0.0f32; rows];
11532 q1_matvec(
11533 &bytes,
11534 &xs_all[bi * cols..(bi + 1) * cols],
11535 rows,
11536 cols,
11537 &mut single,
11538 None,
11539 );
11540 assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
11541 }
11542 }
11543
11544 #[test]
11545 fn q1_kernels_match_exact_reference() {
11546 let (rows, cols) = (7, 96);
11548 let gpr = cols / GROUP_SIZE;
11549 let mut bytes = Vec::new();
11550 for t in 0..rows * gpr {
11551 let s = 0.01 + (t % 13) as f32 * 0.003;
11552 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11553 for j in 0..4 {
11554 bytes.push(((t * 31 + j * 97) % 251) as u8);
11555 }
11556 }
11557 let x: Vec<f32> = (0..cols)
11559 .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
11560 .collect();
11561 let mut w = vec![0.0f32; rows * cols];
11563 cortiq_core::quant::dequant_q1(&bytes, &mut w);
11564 let mut expect = vec![0.0f32; rows];
11565 for o in 0..rows {
11566 expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11567 }
11568 let mut got = vec![0.0f32; rows];
11569 q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11570 for o in 0..rows {
11571 assert!(
11572 (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
11573 "row {o}: {} vs {}",
11574 got[o],
11575 expect[o]
11576 );
11577 }
11578 let x2: Vec<f32> = x.iter().map(|v| -v).collect();
11580 let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
11581 q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
11582 assert_eq!(a1, got);
11583 let mut xs = x.clone();
11584 xs.extend_from_slice(&x2);
11585 let mut mm = vec![0.0f32; 2 * rows];
11586 q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
11587 assert_eq!(&mm[..rows], got.as_slice());
11588 assert_eq!(&mm[rows..], a2.as_slice());
11589 }
11590
11591 #[test]
11592 fn repack_is_bit_identical() {
11593 let (rows, cols) = (267, 96); let w: Vec<u8> = (0..rows * cols)
11599 .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
11600 .collect();
11601 let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
11602 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
11603 let rep = q8_repack_layout(&w, rows, cols);
11604 for g in 0..rows / 4 {
11606 for c in 0..cols / 16 {
11607 for lane in 0..4 {
11608 assert_eq!(
11609 &rep[g * 4 * cols + c * 64 + lane * 16
11610 ..g * 4 * cols + c * 64 + lane * 16 + 16],
11611 &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
11612 );
11613 }
11614 }
11615 }
11616 let mut a = vec![0.0f32; rows];
11617 qmatvec(
11618 &w,
11619 &[],
11620 &scales,
11621 &x,
11622 &[],
11623 TensorDtype::Q8Row,
11624 rows,
11625 cols,
11626 &mut a,
11627 None,
11628 );
11629 let mut b = vec![0.0f32; rows];
11630 qmatvec(
11631 &w,
11632 &rep,
11633 &scales,
11634 &x,
11635 &[],
11636 TensorDtype::Q8Row,
11637 rows,
11638 cols,
11639 &mut b,
11640 None,
11641 );
11642 assert_eq!(a, b, "full-range repack output diverged");
11643
11644 #[cfg(target_arch = "aarch64")]
11645 if sdot_enabled() {
11646 let act = split_act(&x);
11648 let mut c1 = vec![0.0f32; rows];
11649 let mut c2 = vec![0.0f32; rows];
11650 q8_range_sdot(
11651 &w,
11652 &[],
11653 &scales,
11654 &act,
11655 cols,
11656 SendMut(c1.as_mut_ptr()),
11657 3,
11658 rows - 2,
11659 );
11660 q8_range_sdot(
11661 &w,
11662 &rep,
11663 &scales,
11664 &act,
11665 cols,
11666 SendMut(c2.as_mut_ptr()),
11667 3,
11668 rows - 2,
11669 );
11670 assert_eq!(c1, c2, "unaligned-range repack output diverged");
11671 }
11672 }
11673
11674 #[test]
11675 fn sdot_a8w8_noise_is_bounded() {
11676 let (rows, cols) = (16, 512);
11680 let w: Vec<u8> = (0..rows * cols)
11681 .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11682 .collect();
11683 let scales = vec![0.01f32; rows];
11684 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
11685 let mut a = vec![0.0f32; rows];
11686 qmatvec(
11687 &w,
11688 &[],
11689 &scales,
11690 &x,
11691 &[],
11692 TensorDtype::Q8Row,
11693 rows,
11694 cols,
11695 &mut a,
11696 None,
11697 );
11698 let (mut num, mut den) = (0f64, 0f64);
11699 for o in 0..rows {
11700 let mut acc = 0.0f32;
11701 for j in 0..cols {
11702 acc += (w[o * cols + j] as i8) as f32 * x[j];
11703 }
11704 let expect = acc * scales[o];
11705 num += ((a[o] - expect) as f64).powi(2);
11706 den += (expect as f64).powi(2);
11707 }
11708 let rel = (num / den.max(1e-12)).sqrt();
11709 assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
11710 }
11711
11712 #[test]
11713 fn i8_dot_neon_matches_scalar() {
11714 let n = 100;
11715 let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
11716 let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
11717 let mut scalar = 0.0f32;
11718 for j in 0..n {
11719 scalar += (w[j] as i8) as f32 * x[j];
11720 }
11721 let fast = dot_i8_f32(&w, &x);
11722 assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
11723 }
11724
11725 #[test]
11727 fn vbitmatvec_matches_full_dequant() {
11728 let (rows, cols) = (6, 64);
11729 let ng = cols / GROUP_SIZE;
11730 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
11732 let mut bytes = bits.clone();
11733 for g in 0..rows * ng {
11734 let s = 0.02 + 0.001 * g as f32;
11735 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11736 }
11737 for r in 0..rows {
11738 let b = bits[r] as usize;
11739 let (mut acc, mut nb) = (0u64, 0usize);
11740 let mut rowbytes = Vec::new();
11741 for i in 0..cols {
11742 let v = ((i * 7 + r * 13) % (1 << b)) as u64;
11743 acc = (acc << b) | v;
11744 nb += b;
11745 while nb >= 8 {
11746 nb -= 8;
11747 rowbytes.push(((acc >> nb) & 0xFF) as u8);
11748 }
11749 }
11750 if nb > 0 {
11751 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11752 }
11753 bytes.extend_from_slice(&rowbytes);
11754 }
11755 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
11756
11757 let mut reference = vec![0f32; rows * cols];
11758 cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
11759 let mut expect = vec![0f32; rows];
11760 for r in 0..rows {
11761 expect[r] = reference[r * cols..(r + 1) * cols]
11762 .iter()
11763 .zip(&x)
11764 .map(|(w, xv)| w * xv)
11765 .sum();
11766 }
11767 let mut got = vec![0f32; rows];
11768 let offsets = vbit_row_offsets(&bytes, rows, cols);
11769 vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
11770 let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
11774 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
11775 for r in 0..rows {
11776 assert!(
11777 (got[r] - expect[r]).abs() < tol * scale,
11778 "row {r}: {} vs {}",
11779 got[r],
11780 expect[r]
11781 );
11782 }
11783 }
11784
11785 #[test]
11790 #[cfg(target_arch = "x86_64")]
11791 fn vbit_matmat_blocked_matches_per_row() {
11792 let (rows, cols, b) = (64usize, 128usize, 9usize);
11793 let ng = cols / GROUP_SIZE;
11794 let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
11795 let mut bytes = bits.clone();
11796 for g in 0..rows * ng {
11797 let sc = 0.02 + 0.0005 * g as f32;
11798 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11799 }
11800 for r in 0..rows {
11801 let bw = bits[r] as usize;
11802 let (mut acc, mut nb) = (0u64, 0usize);
11803 let mut rowbytes = Vec::new();
11804 for i in 0..cols {
11805 let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
11806 acc = (acc << bw) | v;
11807 nb += bw;
11808 while nb >= 8 {
11809 nb -= 8;
11810 rowbytes.push(((acc >> nb) & 0xFF) as u8);
11811 }
11812 }
11813 if nb > 0 {
11814 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11815 }
11816 bytes.extend_from_slice(&rowbytes);
11817 }
11818 let x: Vec<f32> = (0..b * cols)
11819 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11820 .collect();
11821 let offsets = vbit_row_offsets(&bytes, rows, cols);
11822 let mut y_a = vec![0f32; b * rows];
11823 let mut y_b = vec![0f32; b * rows];
11824 unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11825 vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
11826 unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11827 vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
11828 unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11829 let max_d = y_a
11830 .iter()
11831 .zip(&y_b)
11832 .map(|(p, q)| (p - q).abs())
11833 .fold(0.0f32, f32::max);
11834 assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
11835 }
11836
11837 #[test]
11845 fn q4t_matmat_blocked_matches_per_row() {
11846 let (rows, cols, b) = (16usize, 64usize, 9usize);
11847 let gpr = cols / GROUP_SIZE;
11848 let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
11849 for r in 0..rows {
11850 for g in 0..gpr {
11851 let t = (r * gpr + g) * Q4_TILE;
11852 let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
11853 bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11854 for k in 0..16 {
11855 bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11856 }
11857 }
11858 }
11859 let x: Vec<f32> = (0..b * cols)
11860 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11861 .collect();
11862 let mut y_blk = vec![0f32; b * rows];
11863 let mut y_row = vec![0f32; b * rows];
11864 unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11865 q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
11866 unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11867 q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
11868 unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11869 assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
11870 }
11871
11872 fn synth_q4tp(rows: usize, cols: usize) -> Vec<u8> {
11879 use cortiq_core::quant::{f32_to_f16, q4tp_code_stride, q4tp_put_code};
11880 let gpr = cols / GROUP_SIZE;
11881 let stride = q4tp_code_stride(gpr);
11882 let (params_off, codes_off, _) = q4tp_sections(rows, cols);
11883 let mut b = vec![0u8; codes_off + rows * stride];
11884 for r in 0..rows {
11885 for g in 0..gpr {
11886 let t = (r * gpr + g) * Q4TP_NIB;
11887 for k in 0..16 {
11888 b[t + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11889 }
11890 }
11891 let lo = -6.0 - 0.03 * (r % 17) as f32;
11892 let step = 0.01 + 0.004 * (r % 11) as f32;
11893 let p = params_off + r * 4;
11894 b[p..p + 2].copy_from_slice(&f32_to_f16(lo).to_le_bytes());
11895 b[p + 2..p + 4].copy_from_slice(&f32_to_f16(step).to_le_bytes());
11896 let crow = &mut b[codes_off + r * stride..codes_off + (r + 1) * stride];
11897 for g in 0..gpr {
11898 q4tp_put_code(crow, g, (r * 5 + g * 3) % 32);
11899 }
11900 }
11901 b
11902 }
11903
11904 fn q4tp_as_q4t(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
11908 let gpr = cols / GROUP_SIZE;
11909 let v = Q4tpView::new(bytes, rows, cols);
11910 let mut out = vec![0u8; rows * gpr * Q4_TILE];
11911 let mut sc = vec![0f32; gpr];
11912 for r in 0..rows {
11913 v.scales_into(r, gpr, &mut sc);
11914 for g in 0..gpr {
11915 let t = (r * gpr + g) * Q4_TILE;
11916 let s = sc[g];
11917 out[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11918 let src = (r * gpr + g) * Q4TP_NIB;
11919 out[t + 2..t + Q4_TILE].copy_from_slice(&v.nib[src..src + Q4TP_NIB]);
11920 }
11921 }
11922 out
11923 }
11924
11925 #[test]
11931 fn q4tp_exact_path_matches_dequant_reference() {
11932 let (rows, cols) = (256usize, 512usize);
11933 let gpr = cols / GROUP_SIZE;
11934 let bytes = synth_q4tp(rows, cols);
11935 let mut w = vec![0f32; rows * cols];
11936 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
11937
11938 let x: Vec<f32> = (0..cols)
11939 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11940 .collect();
11941 let v = Q4tpView::new(&bytes, rows, cols);
11942 let mut sc = vec![0f32; gpr];
11943 for r in 0..rows {
11944 v.scales_into(r, gpr, &mut sc);
11945 let got = q4tp_row_exact(v.nib, r, gpr, &x, &sc);
11946 let want: f32 = (0..cols).map(|c| w[r * cols + c] * x[c]).sum();
11947 let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
11951 assert!(
11952 (got - want).abs() <= 1e-5 * mag,
11953 "row {r}: kernel {got} vs dequant {want}"
11954 );
11955 }
11956 }
11957
11958 #[test]
11964 fn q4tp_matvec_matches_the_q4t_kernel_it_was_ported_from() {
11965 let (rows, cols) = (256usize, 512usize);
11966 let bytes = synth_q4tp(rows, cols);
11967 let twin = q4tp_as_q4t(&bytes, rows, cols);
11968 let x: Vec<f32> = (0..cols)
11969 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11970 .collect();
11971
11972 let mut got = vec![0f32; rows];
11973 q4tp_matvec(&bytes, &x, rows, cols, &mut got, None);
11974 let mut want = vec![0f32; rows];
11975 q4t_matvec(&twin, &x, rows, cols, &mut want, None);
11976
11977 let mut w = vec![0f32; rows * cols];
11980 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
11981 for r in 0..rows {
11982 let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
11983 assert!(
11984 (got[r] - want[r]).abs() <= 1e-3 * mag,
11985 "row {r}: q4tp {} vs q4t {}",
11986 got[r],
11987 want[r]
11988 );
11989 }
11990 }
11991
11992 #[test]
11997 fn q4tp_matmat_matches_the_q4t_kernel_it_was_ported_from() {
11998 let (rows, cols, b) = (256usize, 512usize, 5usize);
11999 let bytes = synth_q4tp(rows, cols);
12000 let twin = q4tp_as_q4t(&bytes, rows, cols);
12001 let xs: Vec<f32> = (0..b * cols)
12002 .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12003 .collect();
12004
12005 let mut got = vec![0f32; b * rows];
12006 q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, None);
12007 let mut want = vec![0f32; b * rows];
12008 q4t_matmat(&twin, &xs, b, rows, cols, &mut want, None);
12009
12010 let mut w = vec![0f32; rows * cols];
12011 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12012 for t in 0..b {
12013 for r in 0..rows {
12014 let mag: f32 = (0..cols)
12015 .map(|c| (w[r * cols + c] * xs[t * cols + c]).abs())
12016 .sum();
12017 let (g, wa) = (got[t * rows + r], want[t * rows + r]);
12018 assert!(
12019 (g - wa).abs() <= 1e-3 * mag,
12020 "batch {t} row {r}: q4tp {g} vs q4t {wa}"
12021 );
12022 }
12023 }
12024 }
12025
12026 #[test]
12027 fn q4tp_matvec2_matches_the_single_stream_kernel() {
12028 let (rows, cols) = (128usize, 256usize);
12029 let gpr = cols / GROUP_SIZE;
12030 let bytes = synth_q4tp(rows, cols);
12031 let xs: Vec<f32> = (0..2 * cols)
12032 .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12033 .collect();
12034
12035 let (mut o1, mut o2) = (vec![0f32; rows], vec![0f32; rows]);
12036 q4tp_matvec2(
12037 &bytes,
12038 &xs[..cols],
12039 &xs[cols..],
12040 rows,
12041 cols,
12042 &mut o1,
12043 &mut o2,
12044 None,
12045 );
12046
12047 let v = Q4tpView::new(&bytes, rows, cols);
12050 let mut sc = vec![0f32; gpr];
12051 for r in 0..rows {
12052 v.scales_into(r, gpr, &mut sc);
12053 assert_eq!(o1[r], q4tp_row_exact(v.nib, r, gpr, &xs[..cols], &sc));
12054 assert_eq!(o2[r], q4tp_row_exact(v.nib, r, gpr, &xs[cols..], &sc));
12055 }
12056 }
12057
12058 #[test]
12065 fn q4tp_matvec_keeps_pace_with_q4t() {
12066 let (rows, cols) = (4096usize, 3072usize);
12067 let bytes = synth_q4tp(rows, cols);
12068 let twin = q4tp_as_q4t(&bytes, rows, cols);
12069 let x: Vec<f32> = (0..cols).map(|i| (i % 97) as f32 / 97.0 - 0.5).collect();
12070 let mut o = vec![0f32; rows];
12071 let n = 12;
12072 let mut best = (f64::MAX, f64::MAX);
12073 for _ in 0..3 {
12076 let t0 = std::time::Instant::now();
12077 for _ in 0..n {
12078 q4t_matvec(&twin, &x, rows, cols, &mut o, None);
12079 }
12080 best.0 = best.0.min(t0.elapsed().as_secs_f64());
12081 let t0 = std::time::Instant::now();
12082 for _ in 0..n {
12083 q4tp_matvec(&bytes, &x, rows, cols, &mut o, None);
12084 }
12085 best.1 = best.1.min(t0.elapsed().as_secs_f64());
12086 }
12087 let ratio = best.1 / best.0;
12088 println!(
12089 "q4t {:.3} ms | q4tp {:.3} ms | {ratio:.2}x",
12090 best.0 * 1e3 / n as f64,
12091 best.1 * 1e3 / n as f64
12092 );
12093 assert!(ratio < 2.0, "q4tp matvec {ratio:.2}x slower than q4t");
12094 }
12095
12096 #[cfg(target_os = "macos")]
12097 #[test]
12098 fn q4t_matmat_accel_matches_dequant_reference() {
12099 if !accel_gemm_enabled() {
12100 return; }
12102 let (rows, cols, b) = (512usize, 1024usize, 8usize); let gpr = cols / GROUP_SIZE;
12104 let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
12105 for r in 0..rows {
12106 for g in 0..gpr {
12107 let t = (r * gpr + g) * Q4_TILE;
12108 let sc = 0.02 + 0.0005 * ((r * gpr + g) % 64) as f32;
12109 bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
12110 for k in 0..16 {
12111 bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12112 }
12113 }
12114 }
12115 let x: Vec<f32> = (0..b * cols)
12116 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12117 .collect();
12118 let mut got = vec![0f32; b * rows];
12119 q4t_matmat(&bytes, &x, b, rows, cols, &mut got, None);
12120 let mut w = vec![0f32; rows * cols];
12122 for r in 0..rows {
12123 for g in 0..gpr {
12124 let t = (r * gpr + g) * Q4_TILE;
12125 let s = f16_to_f32(u16::from_le_bytes([bytes[t], bytes[t + 1]]));
12126 for (k, &bb) in bytes[t + 2..t + Q4_TILE].iter().enumerate() {
12127 w[r * cols + g * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
12128 w[r * cols + g * GROUP_SIZE + k * 2 + 1] =
12129 (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
12130 }
12131 }
12132 }
12133 for bi in 0..b {
12134 for r in 0..rows {
12135 let want: f32 = (0..cols).map(|j| x[bi * cols + j] * w[r * cols + j]).sum();
12136 let d = (got[bi * rows + r] - want).abs();
12137 assert!(
12138 d <= want.abs().max(1.0) * 1e-4,
12139 "accel q4t GEMM diverged at ({bi},{r}): {} vs {want}",
12140 got[bi * rows + r]
12141 );
12142 }
12143 }
12144 }
12145
12146 #[test]
12147 fn q4matvec_matches_full_dequant() {
12148 let (rows, cols) = (8, 64);
12149 let groups = rows * cols / GROUP_SIZE;
12150 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12152 for i in 0..groups * 16 {
12153 bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12154 }
12155 for g in 0..groups {
12156 let s = 0.01 + 0.003 * g as f32;
12157 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12158 }
12159 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12160
12161 let mut reference = vec![0.0f32; rows * cols];
12162 cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12163 let mut expect = vec![0.0f32; rows];
12164 for r in 0..rows {
12165 expect[r] = reference[r * cols..(r + 1) * cols]
12166 .iter()
12167 .zip(&x)
12168 .map(|(w, xv)| w * xv)
12169 .sum();
12170 }
12171
12172 let mut got = vec![0.0f32; rows];
12173 q4matvec(&bytes, &x, rows, cols, &mut got, None);
12174 let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
12178 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12179 for r in 0..rows {
12180 assert!(
12181 (got[r] - expect[r]).abs() < tol * scale,
12182 "row {r}: {} vs {}",
12183 got[r],
12184 expect[r]
12185 );
12186 }
12187 }
12188
12189 #[test]
12192 fn vbitmatvec2_equals_two_singles() {
12193 let (rows, cols) = (6, 64);
12194 let ng = cols / GROUP_SIZE;
12195 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
12196 let mut bytes = bits.clone();
12197 for g in 0..rows * ng {
12198 let s = 0.02 + 0.001 * g as f32;
12199 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12200 }
12201 for r in 0..rows {
12202 let b = bits[r] as usize;
12203 let (mut acc, mut nb) = (0u64, 0usize);
12204 let mut rowbytes = Vec::new();
12205 for i in 0..cols {
12206 let v = ((i * 7 + r * 13) % (1 << b)) as u64;
12207 acc = (acc << b) | v;
12208 nb += b;
12209 while nb >= 8 {
12210 nb -= 8;
12211 rowbytes.push(((acc >> nb) & 0xFF) as u8);
12212 }
12213 }
12214 if nb > 0 {
12215 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12216 }
12217 bytes.extend_from_slice(&rowbytes);
12218 }
12219 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
12220 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
12221 let offsets = vbit_row_offsets(&bytes, rows, cols);
12222
12223 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12224 vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
12225 vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
12226 let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12227 vbitmatvec2(
12228 &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
12229 );
12230 assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
12231 assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
12232 }
12233
12234 #[test]
12236 fn q4matvec2_equals_two_singles() {
12237 let (rows, cols) = (8, 128);
12238 let groups = rows * cols / GROUP_SIZE;
12239 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12240 for i in 0..groups * 16 {
12241 bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12242 }
12243 for g in 0..groups {
12244 let s = 0.01 + 0.003 * g as f32;
12245 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12246 }
12247 let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12250 x1[9] = 250.0;
12251 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12252
12253 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12254 q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
12255 q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
12256 let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12257 q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
12258 assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
12259 assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
12260 }
12261
12262 #[test]
12265 fn matvec_many_equals_separate_matvecs() {
12266 use crate::pool::Pool;
12267 let (r1, r2, cols) = (300, 200, 64);
12268 let mk = |salt: usize, rows: usize| {
12269 QTensor::from_f32(
12270 (0..rows * cols)
12271 .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
12272 .collect(),
12273 rows,
12274 cols,
12275 )
12276 };
12277 let (a, b) = (mk(1, r1), mk(5, r2));
12278 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
12279 let pool = Pool::new(3);
12280
12281 let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
12282 a.matvec(&x, &mut ea, Some(&pool));
12283 b.matvec(&x, &mut eb, Some(&pool));
12284 let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
12285 QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12286 assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
12287 assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
12288 }
12289
12290 #[test]
12295 fn q4tp_matvec_many_equals_separate_matvecs() {
12296 use crate::pool::Pool;
12297 use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12298
12299 let (r1, r2, cols) = (300usize, 200usize, 64usize);
12300 let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12301 "arch_name": "tiny-q4tp",
12302 "hidden_size": cols,
12303 "intermediate_size": cols * 2,
12304 "num_layers": 1,
12305 "num_attention_heads": 2,
12306 "num_kv_heads": 1,
12307 "head_dim": 32,
12308 "vocab_size": r1,
12309 "layer_types": ["FullAttention"],
12310 "rms_norm_eps": 1e-6,
12311 "max_position_embeddings": 8,
12312 "linear_conv_kernel_dim": 0,
12313 "linear_num_key_heads": 0,
12314 "linear_num_value_heads": 0
12315 }))
12316 .unwrap();
12317 let header = CmfHeader {
12318 format: "cmf".into(),
12319 version: CMF_VERSION,
12320 arch,
12321 quant_type: QuantType::Q4Block,
12322 provenance: None,
12323 tokenizer_config: None,
12324 section_hashes: None,
12325 skills: Vec::new(),
12326 shard: None,
12327 calibration: None,
12328 routing: None,
12329 };
12330 let specs = [
12331 TensorSpec {
12332 name: "q".into(),
12333 dtype: TensorDtype::Q4TiledP,
12334 shape: vec![r1, cols],
12335 data: synth_q4tp(r1, cols),
12336 },
12337 TensorSpec {
12338 name: "kv".into(),
12339 dtype: TensorDtype::Q4TiledP,
12340 shape: vec![r2, cols],
12341 data: synth_q4tp(r2, cols),
12342 },
12343 ];
12344 let dir = std::env::temp_dir().join(format!("cmf-q4tp-many-{}", std::process::id()));
12345 std::fs::create_dir_all(&dir).unwrap();
12346 let path = dir.join("m.cmf");
12347 CmfModel::write(&path, &header, &specs, None, None).unwrap();
12348 let model = Arc::new(CmfModel::open(&path).unwrap());
12349 let (a, b) = (
12350 QTensor::from_model(&model, "q").unwrap(),
12351 QTensor::from_model(&model, "kv").unwrap(),
12352 );
12353 assert_eq!(a.model_dtype(), Some(TensorDtype::Q4TiledP));
12354 assert_eq!(b.model_dtype(), Some(TensorDtype::Q4TiledP));
12355 let x: Vec<f32> = (0..cols)
12356 .map(|i| ((i * 17 + 3) % 97) as f32 / 97.0 - 0.5)
12357 .collect();
12358 let pool = Pool::new(3);
12359 let (mut ea, mut eb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12360 a.matvec(&x, &mut ea, Some(&pool));
12361 b.matvec(&x, &mut eb, Some(&pool));
12362 let (mut ga, mut gb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12363 QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12364 assert_eq!(ea, ga, "Q4TP fused lane 1 must be bit-identical");
12365 assert_eq!(eb, gb, "Q4TP fused lane 2 must be bit-identical");
12366 let _ = std::fs::remove_dir_all(&dir);
12367 }
12368
12369 #[test]
12375 fn multi_token_moe_rows_equal_single_token_decode() {
12376 use crate::pool::Pool;
12377 use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12378
12379 let (h, inter, ne) = (64usize, 128usize, 3usize);
12380 let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12381 "arch_name": "tiny-q4tp-moe",
12382 "hidden_size": h,
12383 "intermediate_size": inter,
12384 "num_layers": 1,
12385 "num_attention_heads": 2,
12386 "num_kv_heads": 1,
12387 "head_dim": 32,
12388 "vocab_size": 8,
12389 "layer_types": ["FullAttention"],
12390 "rms_norm_eps": 1e-6,
12391 "max_position_embeddings": 8,
12392 "linear_conv_kernel_dim": 0,
12393 "linear_num_key_heads": 0,
12394 "linear_num_value_heads": 0
12395 }))
12396 .unwrap();
12397 let header = CmfHeader {
12398 format: "cmf".into(),
12399 version: CMF_VERSION,
12400 arch,
12401 quant_type: QuantType::Q4Block,
12402 provenance: None,
12403 tokenizer_config: None,
12404 section_hashes: None,
12405 skills: Vec::new(),
12406 shard: None,
12407 calibration: None,
12408 routing: None,
12409 };
12410 let mut specs = Vec::new();
12411 for e in 0..ne {
12412 for (k, (n, r, c)) in [("g", inter, h), ("u", inter, h), ("d", h, inter)]
12413 .into_iter()
12414 .enumerate()
12415 {
12416 let mut data = synth_q4tp(r, c);
12419 for (i, byte) in data[..r * (c / GROUP_SIZE) * Q4TP_NIB]
12420 .iter_mut()
12421 .enumerate()
12422 {
12423 *byte ^= ((i * (e * 3 + k + 1)) % 251) as u8;
12424 }
12425 specs.push(TensorSpec {
12426 name: format!("{n}{e}"),
12427 dtype: TensorDtype::Q4TiledP,
12428 shape: vec![r, c],
12429 data,
12430 });
12431 }
12432 }
12433 let dir = std::env::temp_dir().join(format!(
12434 "cmf-moe-rows-{}-{}",
12435 std::process::id(),
12436 FLOAT_ACTIVATIONS.get()
12437 ));
12438 std::fs::create_dir_all(&dir).unwrap();
12439 let path = dir.join("m.cmf");
12440 CmfModel::write(&path, &header, &specs, None, None).unwrap();
12441 let model = Arc::new(CmfModel::open(&path).unwrap());
12442 let t = |n: String| QTensor::from_model(&model, &n).unwrap();
12443 let g: Vec<QTensor> = (0..ne).map(|e| t(format!("g{e}"))).collect();
12444 let u: Vec<QTensor> = (0..ne).map(|e| t(format!("u{e}"))).collect();
12445 let d: Vec<QTensor> = (0..ne).map(|e| t(format!("d{e}"))).collect();
12446 let b = 4usize;
12447 let mut xs: Vec<f32> = (0..b * h)
12448 .map(|i| ((i * 31 + 7) % 89) as f32 / 89.0 - 0.5)
12449 .collect();
12450 xs[5] = 9.0; let routes: Vec<(Vec<usize>, Vec<f32>)> = vec![
12453 (vec![2, 0], vec![0.6, 0.4]),
12454 (vec![0, 1, 2], vec![0.2, 0.5, 0.3]),
12455 (vec![1], vec![1.0]),
12456 (vec![2, 1, 0], vec![0.25, 0.25, 0.5]),
12457 ];
12458 let pool = Pool::new(3);
12459 let mut want = vec![0f32; b * h];
12461 for (tk, (idx, w)) in routes.iter().enumerate() {
12462 let x = &xs[tk * h..(tk + 1) * h];
12463 let pairs: Vec<(&QTensor, &QTensor)> = idx.iter().map(|&e| (&g[e], &u[e])).collect();
12464 let mut gs: Vec<Vec<f32>> = idx.iter().map(|_| vec![0f32; inter]).collect();
12465 assert!(QTensor::moe_gate_up_many(&pairs, x, &mut gs, Some(&pool)));
12466 if FLOAT_ACTIVATIONS.get() {
12467 for (slot, &e) in idx.iter().enumerate() {
12468 let (mut gate, mut up) = (vec![0.0; inter], vec![0.0; inter]);
12469 g[e].matvec(x, &mut gate, Some(&pool));
12470 u[e].matvec(x, &mut up, Some(&pool));
12471 for (v, u) in gate.iter_mut().zip(up) {
12472 *v = (*v / (1.0 + (-*v).exp())) * u;
12473 }
12474 assert_eq!(gs[slot], gate, "float gate/up must equal ordinary matvecs");
12475 }
12476 }
12477 let downs: Vec<&QTensor> = idx.iter().map(|&e| &d[e]).collect();
12478 assert!(QTensor::moe_down_many(
12479 &downs,
12480 &gs,
12481 w,
12482 &mut want[tk * h..(tk + 1) * h],
12483 Some(&pool)
12484 ));
12485 }
12486 if FLOAT_ACTIVATIONS.get() {
12487 for (tk, (idx, w)) in routes.iter().enumerate() {
12488 let mut scalar = vec![0.0; h];
12489 for (&e, &weight) in idx.iter().zip(w) {
12490 let (mut gate, mut up, mut down) =
12491 (vec![0.0; inter], vec![0.0; inter], vec![0.0; h]);
12492 g[e].matvec(&xs[tk * h..(tk + 1) * h], &mut gate, Some(&pool));
12493 u[e].matvec(&xs[tk * h..(tk + 1) * h], &mut up, Some(&pool));
12494 for (v, u) in gate.iter_mut().zip(up) {
12495 *v = (*v / (1.0 + (-*v).exp())) * u;
12496 }
12497 d[e].matvec(&gate, &mut down, Some(&pool));
12498 for (v, d) in scalar.iter_mut().zip(down) {
12499 *v += weight * d;
12500 }
12501 }
12502 assert_eq!(
12503 &want[tk * h..(tk + 1) * h],
12504 scalar,
12505 "float many equals scalar experts"
12506 );
12507 }
12508 }
12509 let mut experts: Vec<usize> = Vec::new();
12511 let mut groups: Vec<Vec<usize>> = Vec::new();
12512 for (tk, (idx, _)) in routes.iter().enumerate() {
12513 for &e in idx {
12514 match experts.iter().position(|&x| x == e) {
12515 Some(k) => groups[k].push(tk),
12516 None => {
12517 experts.push(e);
12518 groups.push(vec![tk]);
12519 }
12520 }
12521 }
12522 }
12523 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
12524 let pairs: Vec<(&QTensor, &QTensor)> = experts.iter().map(|&e| (&g[e], &u[e])).collect();
12525 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
12526 assert!(QTensor::moe_gate_up_rows(
12527 &pairs,
12528 &groups,
12529 &xs,
12530 &mut gs,
12531 Some(&pool)
12532 ));
12533 let downs: Vec<&QTensor> = experts.iter().map(|&e| &d[e]).collect();
12534 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
12535 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; h]).collect();
12536 assert!(QTensor::moe_down_rows(
12537 &downs,
12538 &lens,
12539 &gs,
12540 &mut ds,
12541 Some(&pool)
12542 ));
12543 let slot = |tk: usize, e: usize| {
12544 let k = experts.iter().position(|&x| x == e).unwrap();
12545 groups[..k].iter().map(|g| g.len()).sum::<usize>()
12546 + groups[k].iter().position(|&x| x == tk).unwrap()
12547 };
12548 let mut got = vec![0f32; b * h];
12549 for (tk, (idx, w)) in routes.iter().enumerate() {
12550 for i in 0..h {
12551 let mut acc = 0f32;
12552 for (&e, &we) in idx.iter().zip(w) {
12553 acc += we * ds[slot(tk, e)][i];
12554 }
12555 got[tk * h + i] = acc;
12556 }
12557 }
12558 assert!(want.iter().any(|v| *v != 0.0));
12559 assert_eq!(
12560 want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12561 got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12562 "multi-token MoE must equal decode bit for bit"
12563 );
12564
12565 let b5 = 5usize;
12568 let x5: Vec<f32> = (0..b5 * h)
12569 .map(|i| ((i * 13 + 5) % 71) as f32 / 71.0 - 0.5)
12570 .collect();
12571 let mut mm = vec![0f32; b5 * inter];
12572 row_exact_scope(|| g[1].matmat(&x5, b5, &mut mm, Some(&pool)));
12573 for tk in 0..b5 {
12574 let mut mv = vec![0f32; inter];
12575 g[1].matvec(&x5[tk * h..(tk + 1) * h], &mut mv, Some(&pool));
12576 assert_eq!(
12577 mv.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12578 mm[tk * inter..(tk + 1) * inter]
12579 .iter()
12580 .map(|v| v.to_bits())
12581 .collect::<Vec<_>>(),
12582 "row-exact matmat token {tk}"
12583 );
12584 }
12585 let _ = std::fs::remove_dir_all(&dir);
12588 }
12589
12590 #[test]
12598 fn q4tp_matmat_fast_path_unchanged_outside_row_exact() {
12599 use crate::pool::Pool;
12600 use std::sync::atomic::Ordering::Relaxed;
12601 let _alt = Q4TP_ALT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
12602 Q4TP_ALT.store(2, Relaxed);
12605 let pool = Pool::new(3);
12606 for &(rows, cols, b) in &[(64usize, 256usize, 7usize), (320, 1024, 9)] {
12610 let bytes = synth_q4tp(rows, cols);
12611 let mut xs: Vec<f32> = (0..b * cols)
12612 .map(|i| ((i * 29 + 11) % 83) as f32 / 83.0 - 0.5)
12613 .collect();
12614 xs[3] = 7.5; let run = |exact: bool| {
12616 let mut out = vec![0f32; b * rows];
12617 q4tp_matmat_with(&bytes, &xs, b, rows, cols, &mut out, Some(&pool), exact);
12618 out
12619 };
12620 let (fast, exact) = (run(false), run(true));
12621 #[cfg(not(target_arch = "aarch64"))]
12622 let _ = fast;
12623 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
12624 let mut matvecs = vec![0f32; b * rows];
12625 for (bi, o) in matvecs.chunks_mut(rows).enumerate() {
12626 q4tp_matvec(
12627 &bytes,
12628 &xs[bi * cols..(bi + 1) * cols],
12629 rows,
12630 cols,
12631 o,
12632 Some(&pool),
12633 );
12634 }
12635 assert!(matvecs.iter().any(|v| *v != 0.0));
12636 assert_eq!(
12637 bits(&exact),
12638 bits(&matvecs),
12639 "{rows}x{cols} b={b}: row-exact matmat must equal per-token matvecs"
12640 );
12641 #[cfg(target_arch = "aarch64")]
12642 {
12643 let gpr = cols / GROUP_SIZE;
12645 let v = Q4tpView::new(&bytes, rows, cols);
12646 let mut old = vec![0f32; b * rows];
12647 let mut sc = vec![0f32; gpr];
12648 let a8w8 = a8w8_enabled();
12649 let blocked = sdot_enabled() && blocked_enabled();
12650 let acts: Vec<SplitAct> = (0..b)
12651 .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
12652 .collect();
12653 for r in 0..rows {
12654 v.scales_into(r, gpr, &mut sc);
12655 if !a8w8 {
12656 for bi in 0..b {
12657 let x = &xs[bi * cols..(bi + 1) * cols];
12658 old[bi * rows + r] = q4tp_row_exact(v.nib, r, gpr, x, &sc);
12659 }
12660 continue;
12661 }
12662 let finish = |d: f32, act: &SplitAct| {
12663 let mut acc = d * act.sx;
12664 for &(j, xv) in &act.outliers {
12665 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
12666 acc += w * s * xv;
12667 }
12668 acc
12669 };
12670 let mut bi = 0usize;
12671 while blocked && bi + 4 <= b {
12672 let xs4 = [
12673 acts[bi].xq.as_slice(),
12674 acts[bi + 1].xq.as_slice(),
12675 acts[bi + 2].xq.as_slice(),
12676 acts[bi + 3].xq.as_slice(),
12677 ];
12678 let d = unsafe { dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs4, &sc) };
12679 for k in 0..4 {
12680 old[(bi + k) * rows + r] = finish(d[k], &acts[bi + k]);
12681 }
12682 bi += 4;
12683 }
12684 for (bi, act) in acts.iter().enumerate().skip(bi) {
12685 let d = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc);
12686 old[bi * rows + r] = finish(d, act);
12687 }
12688 }
12689 assert_eq!(
12690 bits(&fast),
12691 bits(&old),
12692 "{rows}x{cols} b={b}: the fast path changed outside row_exact"
12693 );
12694 if blocked {
12697 assert_ne!(
12698 bits(&fast),
12699 bits(&matvecs),
12700 "{rows}x{cols} b={b}: the tuned tile no longer runs outside row_exact"
12701 );
12702 }
12703 }
12704 }
12705 Q4TP_ALT.store(0, Relaxed);
12706 }
12707
12708 #[test]
12709 fn row_exact_scopes_survive_overlap_nesting_and_unwind() {
12710 use std::sync::{Barrier, atomic::{AtomicUsize, Ordering}};
12711 let active = AtomicUsize::new(0);
12714 counted_row_exact_scope(&active, || {
12715 assert_eq!(active.load(Ordering::Acquire), 1);
12716 counted_row_exact_scope(&active, || {
12717 assert_eq!(active.load(Ordering::Acquire), 2);
12718 });
12719 assert_eq!(active.load(Ordering::Acquire), 1);
12720 });
12721 assert_eq!(active.load(Ordering::Acquire), 0);
12722
12723 let both_entered = Barrier::new(2);
12724 let release_last = Barrier::new(2);
12725 std::thread::scope(|s| {
12726 let first = s.spawn(|| counted_row_exact_scope(&active, || {
12727 both_entered.wait();
12728 }));
12729 let last = s.spawn(|| counted_row_exact_scope(&active, || {
12730 both_entered.wait();
12731 release_last.wait();
12732 }));
12733 first.join().unwrap();
12734 let after_first = active.load(Ordering::Acquire);
12735 release_last.wait();
12736 last.join().unwrap();
12737 assert_eq!(after_first, 1, "second request must remain exact");
12738 });
12739 assert_eq!(active.load(Ordering::Acquire), 0);
12740 let panic = std::panic::catch_unwind(|| {
12741 counted_row_exact_scope(&active, || panic!("scope unwind"));
12742 });
12743 assert!(panic.is_err());
12744 assert_eq!(active.load(Ordering::Acquire), 0);
12745 }
12746
12747 #[test]
12748 #[cfg(target_arch = "x86_64")]
12749 fn q4tp_float_avx2_is_bitwise_scalar() {
12750 if !avx2_enabled() {
12751 return;
12752 }
12753 for cols in [32, 64, 96, 2048, 4096] {
12754 let rows = 9;
12755 let bytes = synth_q4tp(rows, cols);
12756 let v = Q4tpView::new(&bytes, rows, cols);
12757 let gpr = cols / GROUP_SIZE;
12758 let mut sc = vec![0.0; gpr];
12759 for seed in 1..=5 {
12760 let xs: Vec<f32> = (0..cols)
12761 .map(|i| (((i * 104729 + seed * 8191) % 100003) as f32 - 50001.0) / 7919.0)
12762 .collect();
12763 for r in 0..rows {
12764 v.scales_into(r, gpr, &mut sc);
12765 let scalar = q4tp_row_float_scalar(v.nib, r, gpr, &xs, &sc);
12766 let vector = unsafe { q4tp_row_float_avx2(v.nib, r, gpr, &xs, &sc) };
12767 assert_eq!(
12768 scalar.to_bits(),
12769 vector.to_bits(),
12770 "cols={cols} row={r} seed={seed}"
12771 );
12772 }
12773 }
12774 }
12775 }
12776
12777 #[test]
12778 fn multi_token_moe_rows_float_equal_single_token_decode() {
12779 float_activations_scope(multi_token_moe_rows_equal_single_token_decode);
12780 }
12781
12782 #[test]
12783 fn full_gpu_q8_scope_is_nested_and_thread_local() {
12784 assert!(!FULL_GPU_Q8.get());
12785 let before = gpu_split_frac();
12786 {
12787 let _guard = enter_full_gpu_q8_scope();
12788 assert_eq!(gpu_split_frac(), 1.0);
12789 {
12790 let _nested = enter_full_gpu_q8_scope();
12791 }
12792 assert_eq!(gpu_split_frac(), 1.0);
12793 std::thread::spawn(|| assert!(!FULL_GPU_Q8.get())).join().unwrap();
12794 }
12795 assert!(!FULL_GPU_Q8.get());
12796 assert_eq!(gpu_split_frac(), before);
12797 }
12798
12799 #[test]
12800 fn float_activation_scope_is_nested_thread_local_and_unwind_safe() {
12801 assert!(!FLOAT_ACTIVATIONS.get());
12802 let before = a8w8_enabled();
12803 float_activations_scope(|| {
12804 assert!(!a8w8_enabled());
12805 float_activations_scope(|| assert!(!a8w8_enabled()));
12806 assert!(FLOAT_ACTIVATIONS.get());
12807 std::thread::spawn(|| assert!(!FLOAT_ACTIVATIONS.get()))
12808 .join()
12809 .unwrap();
12810 });
12811 assert!(!FLOAT_ACTIVATIONS.get());
12812 assert_eq!(a8w8_enabled(), before);
12813 let _ = std::panic::catch_unwind(|| float_activations_scope(|| panic!("test unwind")));
12814 assert!(!FLOAT_ACTIVATIONS.get());
12815 }
12816
12817 #[test]
12820 fn batched_matmat_equals_per_position_matvec() {
12821 let (rows, cols, b) = (8, 64, 5);
12822 let groups = rows * cols / GROUP_SIZE;
12824 let mut q4 = Vec::new();
12825 for i in 0..groups * 16 {
12826 q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12827 }
12828 for g in 0..groups {
12829 q4.extend_from_slice(
12830 &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12831 );
12832 }
12833 let ng = cols / GROUP_SIZE;
12835 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
12836 let mut vb = bits.clone();
12837 for g in 0..rows * ng {
12838 vb.extend_from_slice(
12839 &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
12840 );
12841 }
12842 for r in 0..rows {
12843 let bw = bits[r] as usize;
12844 let (mut acc, mut nb) = (0u64, 0usize);
12845 let mut rowbytes = Vec::new();
12846 for i in 0..cols {
12847 let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
12848 acc = (acc << bw) | v;
12849 nb += bw;
12850 while nb >= 8 {
12851 nb -= 8;
12852 rowbytes.push(((acc >> nb) & 0xFF) as u8);
12853 }
12854 }
12855 if nb > 0 {
12856 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12857 }
12858 vb.extend_from_slice(&rowbytes);
12859 }
12860 let offsets = vbit_row_offsets(&vb, rows, cols);
12861
12862 let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
12863
12864 let mut got = vec![0f32; b * rows];
12866 q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
12867 for bi in 0..b {
12868 let mut expect = vec![0f32; rows];
12869 q4matvec(
12870 &q4,
12871 &xs[bi * cols..(bi + 1) * cols],
12872 rows,
12873 cols,
12874 &mut expect,
12875 None,
12876 );
12877 assert_eq!(
12878 &got[bi * rows..(bi + 1) * rows],
12879 &expect[..],
12880 "q4 batch pos {bi}"
12881 );
12882 }
12883
12884 let mut got = vec![0f32; b * rows];
12886 vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
12887 for bi in 0..b {
12888 let mut expect = vec![0f32; rows];
12889 vbitmatvec(
12890 &vb,
12891 &offsets,
12892 &xs[bi * cols..(bi + 1) * cols],
12893 rows,
12894 cols,
12895 &mut expect,
12896 None,
12897 );
12898 assert_eq!(
12899 &got[bi * rows..(bi + 1) * rows],
12900 &expect[..],
12901 "vbit batch pos {bi}"
12902 );
12903 }
12904 }
12905
12906 #[test]
12910 fn q4_tiled_matches_q4_block_bitexact() {
12911 let (rows, cols, b) = (8usize, 128usize, 3usize);
12912 let groups = rows * cols / GROUP_SIZE;
12913 let mut split = Vec::with_capacity(groups * 18);
12914 for i in 0..groups * 16 {
12915 split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12916 }
12917 for g in 0..groups {
12918 split.extend_from_slice(
12919 &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12920 );
12921 }
12922 let (packed, scales) = split.split_at(groups * 16);
12924 let mut tiled = Vec::with_capacity(groups * Q4_TILE);
12925 for g in 0..groups {
12926 tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
12927 tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
12928 }
12929
12930 let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12931 x1[9] = 250.0; let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12933
12934 let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
12935 q4matvec(&split, &x1, rows, cols, &mut a, None);
12936 q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
12937 assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
12938
12939 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12940 let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
12941 q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
12942 q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
12943 assert_eq!(a1, t1);
12944 assert_eq!(a2, t2);
12945
12946 let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
12947 let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
12948 q4matmat(&split, &xs, b, rows, cols, &mut am, None);
12949 q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
12950 assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
12951 }
12952
12953 #[test]
12960 fn q4matvec_sdot_outlier_exact() {
12961 let (rows, cols) = (4, 128);
12962 let groups = rows * cols / GROUP_SIZE;
12963 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12964 for i in 0..groups * 16 {
12965 bytes.push(((i * 11 + 5) % 256) as u8);
12966 }
12967 for g in 0..groups {
12968 let s = 0.02 + 0.002 * g as f32;
12969 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12970 }
12971 let mut x: Vec<f32> = (0..cols)
12972 .map(|i| match i % 3 {
12973 0 => 1.0,
12974 1 => -1.0,
12975 _ => 0.0,
12976 })
12977 .collect();
12978 x[17] = 300.0; let mut reference = vec![0.0f32; rows * cols];
12981 cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12982 let mut expect = vec![0.0f32; rows];
12983 for r in 0..rows {
12984 expect[r] = reference[r * cols..(r + 1) * cols]
12985 .iter()
12986 .zip(&x)
12987 .map(|(w, xv)| w * xv)
12988 .sum();
12989 }
12990 let mut got = vec![0.0f32; rows];
12991 q4matvec(&bytes, &x, rows, cols, &mut got, None);
12992 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12993 for r in 0..rows {
12994 assert!(
12995 (got[r] - expect[r]).abs() < 2e-3 * scale,
12996 "row {r}: {} vs {} (outlier term must be exact)",
12997 got[r],
12998 expect[r]
12999 );
13000 }
13001 }
13002
13003 #[test]
13007 fn q1t_matvec_matches_reference() {
13008 use cortiq_core::quant::{dequant_q1t, f32_to_f16};
13009 let (rows, cols) = (3usize, 64usize); let gpr = cols / GROUP_SIZE;
13011 let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
13012 let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
13014 let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
13015 let mut bytes = Vec::new();
13016 for r in 0..rows {
13017 for g in 0..gpr {
13018 bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
13019 let mut c = [0u8; 7];
13020 for k in 0..GROUP_SIZE {
13021 let code = if is_out(r * cols + g * GROUP_SIZE + k) {
13023 0
13024 } else {
13025 ((k + r * 3 + g) % 3) as u8 };
13027 cortiq_core::quant::q1t_pack(&mut c, k, code);
13028 }
13029 bytes.extend_from_slice(&c);
13030 }
13031 }
13032 let mut row_ptr = vec![0u32; rows + 1];
13035 for &(idx, _) in &outliers {
13036 row_ptr[idx as usize / cols + 1] += 1;
13037 }
13038 for r in 0..rows {
13039 row_ptr[r + 1] += row_ptr[r];
13040 }
13041 for &p in &row_ptr {
13042 bytes.extend_from_slice(&p.to_le_bytes());
13043 }
13044 for &(idx, v) in &outliers {
13045 bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
13046 bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
13047 }
13048
13049 let mut refw = vec![0f32; rows * cols];
13050 dequant_q1t(&bytes, rows, cols, &mut refw);
13051 let x: Vec<f32> = (0..cols)
13054 .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13055 .collect();
13056 let mut expect = vec![0f32; rows];
13057 for r in 0..rows {
13058 let mut a = 0.0f32;
13059 for j in 0..cols {
13060 a += refw[r * cols + j] * x[j];
13061 }
13062 expect[r] = a;
13063 }
13064 let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
13065 let mut got = vec![0f32; rows];
13066 q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
13067 for r in 0..rows {
13068 assert!(
13069 (got[r] - expect[r]).abs() < tol(expect[r]),
13070 "row {r}: {} vs {}",
13071 got[r],
13072 expect[r]
13073 );
13074 }
13075 let x2: Vec<f32> = x.iter().chain(x.iter()).copied().collect();
13077 let mut gm = vec![0f32; 2 * rows];
13078 q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
13079 for r in 0..rows {
13080 assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
13081 assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
13082 }
13083 let xb: Vec<f32> = (0..cols)
13087 .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
13088 .collect();
13089 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13090 q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
13091 q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
13092 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13093 q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
13094 assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
13095 assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
13096 }
13097
13098 #[test]
13101 fn q1t_matvec2_odd_gpr_matches_singles() {
13102 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13103 let (rows, cols) = (5usize, 96usize); let gpr = cols / GROUP_SIZE;
13105 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13106 for r in 0..rows {
13107 for g in 0..gpr {
13108 bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
13109 let mut c = [0u8; 7];
13110 for k in 0..GROUP_SIZE {
13111 q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
13112 }
13113 bytes.extend_from_slice(&c);
13114 }
13115 }
13116 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13117 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13118 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13119 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13120 q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13121 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13122 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13123 assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
13124 assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
13125 }
13126
13127 #[test]
13131 #[ignore]
13132 fn q1t_matvec2_speed() {
13133 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13134 use std::time::Instant;
13135 let (rows, cols) = (8192usize, 4096usize);
13136 let gpr = cols / GROUP_SIZE;
13137 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13138 for r in 0..rows {
13139 for g in 0..gpr {
13140 let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13141 bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13142 let mut c = [0u8; 7];
13143 for k in 0..GROUP_SIZE {
13144 q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13145 }
13146 bytes.extend_from_slice(&c);
13147 }
13148 }
13149 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13150 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13151 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13152 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13153 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13155 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13156 let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
13157 for _ in 0..8 {
13158 let t0 = Instant::now();
13159 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13160 t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
13161 let t1 = Instant::now();
13162 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13163 q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13164 t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
13165 }
13166 assert_eq!(p1, s1);
13167 assert_eq!(p2, s2);
13168 println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
13169 }
13170
13171 #[test]
13175 #[ignore]
13176 fn q1t_matvec_speed() {
13177 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
13178 use std::time::Instant;
13179 let (rows, cols) = (8192usize, 4096usize); let gpr = cols / GROUP_SIZE;
13181 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
13182 for r in 0..rows {
13183 for g in 0..gpr {
13184 let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13185 bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13186 let mut c = [0u8; 7];
13187 for k in 0..GROUP_SIZE {
13188 q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13189 }
13190 bytes.extend_from_slice(&c);
13191 }
13192 }
13193 let (n, stride) = (rows * cols, 40usize); let mut row_ptr = vec![0u32; rows + 1];
13195 let mut idx = 0usize;
13196 while idx < n {
13197 row_ptr[idx / cols + 1] += 1;
13198 idx += stride;
13199 }
13200 for r in 0..rows {
13201 row_ptr[r + 1] += row_ptr[r];
13202 }
13203 for &p in &row_ptr {
13204 bytes.extend_from_slice(&p.to_le_bytes());
13205 }
13206 let mut idx = 0usize;
13207 while idx < n {
13208 bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
13209 bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
13210 idx += stride;
13211 }
13212 let x: Vec<f32> = (0..cols)
13215 .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13216 .collect();
13217 let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
13218
13219 let slow = |out: &mut [f32]| {
13221 let mut buf = vec![0f32; cols];
13222 for r in 0..rows {
13223 for g in 0..gpr {
13224 let off = (r * gpr + g) * Q1T_TILE;
13225 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
13226 let codes = &bytes[off + 2..off + Q1T_TILE];
13227 for k in 0..GROUP_SIZE {
13228 buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
13229 1 => s,
13230 2 => -s,
13231 _ => 0.0,
13232 };
13233 }
13234 }
13235 out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
13236 + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
13237 }
13238 };
13239 let iters = 5;
13240 let mut a = vec![0f32; rows];
13241 slow(&mut a); let t = Instant::now();
13243 for _ in 0..iters {
13244 slow(&mut a);
13245 }
13246 let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13247
13248 let mut b = vec![0f32; rows];
13249 q1t_matvec(&bytes, &x, rows, cols, &mut b, None); let t = Instant::now();
13251 for _ in 0..iters {
13252 q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
13253 }
13254 let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13255
13256 for r in 0..rows {
13257 assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
13258 }
13259 println!(
13260 "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms fused-LUT {fast_ms:.2} ms => {:.2}x",
13261 slow_ms / fast_ms
13262 );
13263 }
13264}
13265
13266#[cfg(test)]
13267mod gemm_bench {
13268 #[test]
13279 #[ignore]
13280 fn q4tp_matmat_throughput() {
13281 let _alt = super::Q4TP_ALT_TEST_LOCK
13282 .lock()
13283 .unwrap_or_else(|e| e.into_inner());
13284 let b: usize = std::env::var("CMF_BENCH_B")
13288 .ok()
13289 .and_then(|v| v.parse().ok())
13290 .unwrap_or(296);
13291 let (rows, cols) = (9216usize, 2304usize);
13292 let (_, _, _) = (rows, cols, b);
13293 let total =
13294 cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4TiledP, &[rows, cols])
13295 .unwrap();
13296 let (params_off, codes_off, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13301 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13302 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13303 let step = cortiq_core::quant::f32_to_f16(0.1);
13304 for r in 0..rows {
13305 let o = params_off + r * 4;
13306 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13307 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13308 }
13309 let _ = codes_off;
13310 let xs: Vec<f32> = (0..b * cols)
13311 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13312 .collect();
13313 let mut out = vec![0f32; b * rows];
13314 let pool = crate::pool::Pool::from_env();
13315 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13321 let reps: usize = std::env::var("CMF_BENCH_REPS")
13322 .ok()
13323 .and_then(|v| v.parse().ok())
13324 .unwrap_or(10);
13325 let mut best = [f64::MAX; 2];
13326 let mut sums = [0f32; 2];
13327 for _ in 0..reps {
13328 for (k, w) in [(0usize, 1u8), (1usize, 2u8)] {
13329 super::Q4TP_ALT.store(w, std::sync::atomic::Ordering::Relaxed);
13330 let t = std::time::Instant::now();
13331 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13332 best[k] = best[k].min(t.elapsed().as_secs_f64());
13333 sums[k] = out.iter().take(64).sum::<f32>();
13334 }
13335 }
13336 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13337 for (k, name) in ["previous", "tuned "].iter().enumerate() {
13338 println!(
13339 "q4tp matmat {rows}x{cols} b={b} {name}: {:.1} ms {:.1} GFLOP/s (checksum {:.3})",
13340 best[k] * 1e3,
13341 flops / best[k] / 1e9,
13342 sums[k]
13343 );
13344 }
13345 assert!(
13346 (sums[0] - sums[1]).abs() < 1e-2,
13347 "the tuned kernel changed the result: {} vs {}",
13348 sums[0],
13349 sums[1]
13350 );
13351 }
13352
13353 #[test]
13359 fn q4tp_matmat_blocked_matches_scalar() {
13360 use std::sync::atomic::Ordering::Relaxed;
13361 let _alt = super::Q4TP_ALT_TEST_LOCK
13362 .lock()
13363 .unwrap_or_else(|e| e.into_inner());
13364 for &(rows, cols, b) in &[
13372 (64usize, 128usize, 7usize),
13373 (33, 96, 4),
13374 (16, 256, 9),
13375 (192, 2304, 37),
13376 ] {
13377 let total = cortiq_core::quant::expected_nbytes(
13378 cortiq_core::TensorDtype::Q4TiledP,
13379 &[rows, cols],
13380 )
13381 .unwrap();
13382 let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13383 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 61 % 251) as u8).collect();
13384 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13385 let step = cortiq_core::quant::f32_to_f16(0.1);
13386 for r in 0..rows {
13387 let o = params_off + r * 4;
13388 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13389 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13390 }
13391 let xs: Vec<f32> = (0..b * cols)
13392 .map(|i| ((i % 89) as f32 - 44.0) / 44.0)
13393 .collect();
13394 let mut got = vec![0f32; b * rows];
13395 let mut want = vec![0f32; b * rows];
13396 let gpr = cols / 32;
13397 let view = super::Q4tpView::new(&bytes, rows, cols);
13398 let pool = crate::pool::Pool::from_env();
13399 super::Q4TP_ALT.store(2, Relaxed);
13400 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, pool.as_deref());
13401 super::Q4TP_ALT.store(1, Relaxed);
13402 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut want, pool.as_deref());
13403 super::Q4TP_ALT.store(0, Relaxed);
13404 let scale = want.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
13411 let (mut worst, mut at) = (0f32, 0usize);
13412 for (i, (g, w)) in got.iter().zip(&want).enumerate() {
13413 if (g - w).abs() > worst {
13414 worst = (g - w).abs();
13415 at = i;
13416 }
13417 }
13418 assert!(
13419 worst <= 1e-4 * scale,
13420 "{rows}x{cols} b={b}: blocked and scalar disagree by {worst:.3e} \
13421 (scale {scale:.3e}) at cell {at}: {} vs {}",
13422 got[at],
13423 want[at]
13424 );
13425
13426 let (mut e_blocked, mut e_scalar) = (0f64, 0f64);
13433 for bi in 0..b {
13434 let act = super::split_act(&xs[bi * cols..(bi + 1) * cols]);
13435 for r in 0..rows {
13436 let mut sc = vec![0f32; gpr];
13437 view.scales_into(r, gpr, &mut sc);
13438 let mut exact = 0f64;
13439 for j in 0..cols {
13440 let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13441 exact += w as f64 * sq as f64 * act.xq[j] as f64;
13442 }
13443 exact *= act.sx as f64;
13444 for &(j, xv) in &act.outliers {
13445 let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13446 exact += w as f64 * sq as f64 * xv as f64;
13447 }
13448 let i = bi * rows + r;
13449 e_blocked = e_blocked.max((got[i] as f64 - exact).abs());
13450 e_scalar = e_scalar.max((want[i] as f64 - exact).abs());
13451 }
13452 }
13453 println!(
13454 "{rows}x{cols} b={b}: worst error vs f64 — blocked {e_blocked:.3e}, \
13455 per-column {e_scalar:.3e}"
13456 );
13457 assert!(
13462 e_blocked <= 1e-5 * scale as f64 && e_scalar <= 1e-5 * scale as f64,
13463 "{rows}x{cols} b={b}: error against f64 too large — blocked \
13464 {e_blocked:.3e}, per-column {e_scalar:.3e}, scale {scale:.3e}"
13465 );
13466 }
13467 }
13468
13469 #[test]
13475 #[ignore]
13476 fn q4tp_matmat_row_exact_cost() {
13477 let _alt = super::Q4TP_ALT_TEST_LOCK
13478 .lock()
13479 .unwrap_or_else(|e| e.into_inner());
13480 let pool = crate::pool::Pool::from_env();
13481 let reps: usize = std::env::var("CMF_BENCH_REPS")
13482 .ok()
13483 .and_then(|v| v.parse().ok())
13484 .unwrap_or(30);
13485 for &(rows, cols) in &[(2048usize, 4096usize), (4096, 2048), (4096, 4096)] {
13486 let total = cortiq_core::quant::expected_nbytes(
13487 cortiq_core::TensorDtype::Q4TiledP,
13488 &[rows, cols],
13489 )
13490 .unwrap();
13491 let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13492 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13493 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13494 let step = cortiq_core::quant::f32_to_f16(0.1);
13495 for r in 0..rows {
13496 let o = params_off + r * 4;
13497 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13498 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13499 }
13500 for &b in &[2usize, 4, 5, 8] {
13501 let xs: Vec<f32> = (0..b * cols)
13502 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13503 .collect();
13504 let mut out = vec![0f32; b * rows];
13505 let mut best = [f64::MAX; 3];
13506 for _ in 0..reps {
13507 for (k, best_k) in best.iter_mut().enumerate() {
13508 let t = std::time::Instant::now();
13509 match k {
13510 0 | 1 => super::q4tp_matmat_with(
13511 &bytes,
13512 &xs,
13513 b,
13514 rows,
13515 cols,
13516 &mut out,
13517 pool.as_deref(),
13518 k == 1,
13519 ),
13520 _ => {
13521 for (bi, o) in out.chunks_mut(rows).enumerate() {
13522 super::q4tp_matvec(
13523 &bytes,
13524 &xs[bi * cols..(bi + 1) * cols],
13525 rows,
13526 cols,
13527 o,
13528 pool.as_deref(),
13529 );
13530 }
13531 }
13532 }
13533 *best_k = best_k.min(t.elapsed().as_secs_f64());
13534 }
13535 }
13536 println!(
13537 "q4tp {rows}x{cols} b={b}: fast {:.3} ms, row-exact {:.3} ms, \
13538 {b} matvecs {:.3} ms",
13539 best[0] * 1e3,
13540 best[1] * 1e3,
13541 best[2] * 1e3
13542 );
13543 }
13544 }
13545 }
13546
13547 #[test]
13551 #[ignore]
13552 fn q4t_matmat_throughput() {
13553 let (rows, cols, b) = (9216usize, 2304usize, 296usize);
13554 let total =
13555 cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4Tiled, &[rows, cols])
13556 .unwrap();
13557 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13560 let sc = cortiq_core::quant::f32_to_f16(0.02);
13561 for t in bytes.chunks_mut(super::Q4_TILE) {
13562 t[..2].copy_from_slice(&sc.to_le_bytes());
13563 }
13564 let xs: Vec<f32> = (0..b * cols)
13565 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13566 .collect();
13567 let mut out = vec![0f32; b * rows];
13568 let pool = crate::pool::Pool::from_env();
13569 super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13570 let reps: usize = std::env::var("CMF_BENCH_REPS")
13571 .ok()
13572 .and_then(|v| v.parse().ok())
13573 .unwrap_or(10);
13574 let mut best = f64::MAX;
13575 for _ in 0..reps {
13576 let t = std::time::Instant::now();
13577 super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13578 best = best.min(t.elapsed().as_secs_f64());
13579 }
13580 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13581 println!(
13582 "q4t matmat {rows}x{cols} b={b}: {:.1} ms {:.1} GFLOP/s (checksum {:.3})",
13583 best * 1e3,
13584 flops / best / 1e9,
13585 out.iter().take(64).sum::<f32>()
13586 );
13587 }
13588}