1use crate::pool::{Pool, matvec_rows, matvec_rows2};
15use cortiq_core::quant::{
16 GROUP_SIZE, Q1_TILE, Q2TP_CHUNK, Q4_TILE, Q4TP_NIB, f16_to_f32, q2tp_ladder, q2tp_sections,
17 q4tp_code, q4tp_ladder, q4tp_sections,
18};
19use cortiq_core::{CmfModel, TensorDtype};
20use std::cell::UnsafeCell;
21use std::sync::Arc;
22
23pub enum QTensor {
24 F32 {
25 data: Vec<f32>,
26 rows: usize,
27 cols: usize,
28 },
29 Mapped {
30 model: Arc<CmfModel>,
31 idx: usize,
33 dtype: TensorDtype,
34 rows: usize,
35 cols: usize,
36 row_scale: Vec<f32>,
38 col_field: Vec<f32>,
40 vbit_offsets: Vec<usize>,
44 repack: Vec<u8>,
52 },
53}
54
55fn repack_enabled() -> bool {
62 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
63 *ON.get_or_init(|| {
64 std::env::var("CMF_REPACK")
65 .map(|v| v == "1")
66 .unwrap_or(cfg!(target_os = "android"))
67 })
68}
69
70fn q8_repack(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
74 #[cfg(target_arch = "aarch64")]
75 let arch_ok = sdot_enabled();
76 #[cfg(not(target_arch = "aarch64"))]
77 let arch_ok = false;
78 if !arch_ok || !repack_enabled() || rows < 256 || cols % 16 != 0 {
79 return Vec::new();
80 }
81 q8_repack_layout(bytes, rows, cols)
82}
83
84fn q8_repack_layout(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
87 let groups = rows / 4;
88 let mut rep = vec![0u8; groups * 4 * cols];
89 for g in 0..groups {
90 let dst = &mut rep[g * 4 * cols..(g + 1) * 4 * cols];
91 for c in 0..cols / 16 {
92 for lane in 0..4 {
93 let src = (g * 4 + lane) * cols + c * 16;
94 dst[c * 64 + lane * 16..c * 64 + lane * 16 + 16]
95 .copy_from_slice(&bytes[src..src + 16]);
96 }
97 }
98 }
99 rep
100}
101
102fn vbit_row_offsets(bytes: &[u8], rows: usize, cols: usize) -> Vec<usize> {
105 let ng = cols / GROUP_SIZE;
106 let bits = &bytes[..rows];
107 let mut offsets = Vec::with_capacity(rows + 1);
108 let mut off = rows + rows * ng * 2;
109 for r in 0..rows {
110 offsets.push(off);
111 off += (cols * bits[r] as usize).div_ceil(8);
112 }
113 offsets.push(off);
114 offsets
115}
116
117fn blocked_enabled() -> bool {
123 use std::sync::atomic::Ordering::Relaxed;
124 match BLOCKED_OVERRIDE.load(Relaxed) {
125 1 => false,
126 2 => true,
127 _ => {
128 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
129 *ON.get_or_init(|| {
130 std::env::var("CMF_X86_BLOCKED")
131 .map(|v| v != "0")
132 .unwrap_or(true)
133 })
134 }
135 }
136}
137
138static BLOCKED_OVERRIDE: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
139
140pub fn set_blocked_override(on: Option<bool>) {
147 let v = match on {
148 None => 0,
149 Some(false) => 1,
150 Some(true) => 2,
151 };
152 BLOCKED_OVERRIDE.store(v, std::sync::atomic::Ordering::Relaxed);
153}
154
155fn gpu_lmhead_enabled() -> bool {
156 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
157 *ON.get_or_init(|| {
158 std::env::var("CMF_GPU_LMHEAD")
159 .map(|v| v != "0")
160 .unwrap_or(true)
161 })
162}
163
164fn gpu_split_frac() -> f32 {
165 if FULL_GPU_Q8.get() {
168 return 1.0;
169 }
170 static F: std::sync::OnceLock<f32> = std::sync::OnceLock::new();
171 *F.get_or_init(|| {
172 std::env::var("CMF_GPU_SPLIT")
173 .ok()
174 .and_then(|v| v.parse::<f32>().ok())
175 .unwrap_or(0.5)
176 .clamp(0.0, 1.0)
177 })
178}
179
180impl QTensor {
181 pub fn from_f32(data: Vec<f32>, rows: usize, cols: usize) -> Self {
182 debug_assert_eq!(data.len(), rows * cols);
183 Self::F32 { data, rows, cols }
184 }
185
186 pub fn from_model(model: &Arc<CmfModel>, name: &str) -> Result<Self, String> {
189 let idx = model
192 .tensor_index(name)
193 .ok_or_else(|| format!("tensor '{name}' not found in CMF directory"))?;
194 let entry = &model.tensors[idx];
195 if entry.shape.len() != 2 {
196 return Err(format!("QTensor::from_model needs 2-D, got '{name}'"));
197 }
198 let (rows, cols) = (entry.shape[0], entry.shape[1]);
199 let bytes = model.entry_bytes(entry);
200
201 match entry.dtype {
202 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
203 let n = rows * cols;
204 let scales_off = n;
205 let row_scale: Vec<f32> = (0..rows)
206 .map(|o| {
207 f16_to_f32(u16::from_le_bytes([
208 bytes[scales_off + o * 2],
209 bytes[scales_off + o * 2 + 1],
210 ]))
211 })
212 .collect();
213 let col_field: Vec<f32> = if entry.dtype == TensorDtype::Q8_2f {
214 let col_off = n + rows * 2;
215 (0..cols)
216 .map(|i| {
217 f16_to_f32(u16::from_le_bytes([
218 bytes[col_off + i * 2],
219 bytes[col_off + i * 2 + 1],
220 ]))
221 })
222 .collect()
223 } else {
224 Vec::new()
225 };
226 Ok(Self::Mapped {
227 model: model.clone(),
228 idx,
229 dtype: entry.dtype,
230 rows,
231 cols,
232 row_scale,
233 col_field,
234 vbit_offsets: Vec::new(),
235 repack: q8_repack(bytes, rows, cols),
236 })
237 }
238 TensorDtype::Vbit if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
240 model: model.clone(),
241 idx,
242 dtype: entry.dtype,
243 rows,
244 cols,
245 row_scale: Vec::new(),
246 col_field: Vec::new(),
247 vbit_offsets: vbit_row_offsets(bytes, rows, cols),
248 repack: Vec::new(),
249 }),
250 TensorDtype::VbitRo if cols % GROUP_SIZE == 0 => {
254 let (_, off_off, packed_off) = cortiq_core::quant::vbit_ro_sections(rows, cols);
255 let offsets: Vec<usize> = (0..=rows)
256 .map(|r| packed_off + cortiq_core::quant::vbit_ro_offset(bytes, off_off, r))
257 .collect();
258 Ok(Self::Mapped {
259 model: model.clone(),
260 idx,
261 dtype: entry.dtype,
262 rows,
263 cols,
264 row_scale: Vec::new(),
265 col_field: Vec::new(),
266 vbit_offsets: offsets,
267 repack: Vec::new(),
268 })
269 }
270 TensorDtype::Q4Tiled if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
276 model: model.clone(),
277 idx,
278 dtype: entry.dtype,
279 rows,
280 cols,
281 row_scale: Vec::new(),
282 col_field: Vec::new(),
283 vbit_offsets: Vec::new(),
284 repack: Vec::new(),
285 }),
286 TensorDtype::Q4TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
289 model: model.clone(),
290 idx,
291 dtype: entry.dtype,
292 rows,
293 cols,
294 row_scale: Vec::new(),
295 col_field: Vec::new(),
296 vbit_offsets: Vec::new(),
297 repack: Vec::new(),
298 }),
299 TensorDtype::Q2TiledP if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
301 model: model.clone(),
302 idx,
303 dtype: entry.dtype,
304 rows,
305 cols,
306 row_scale: Vec::new(),
307 col_field: Vec::new(),
308 vbit_offsets: Vec::new(),
309 repack: Vec::new(),
310 }),
311 TensorDtype::Q4Block if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
312 model: model.clone(),
313 idx,
314 dtype: entry.dtype,
315 rows,
316 cols,
317 row_scale: Vec::new(),
318 col_field: Vec::new(),
319 vbit_offsets: Vec::new(),
320 repack: Vec::new(),
321 }),
322 TensorDtype::Q1 if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
324 model: model.clone(),
325 idx,
326 dtype: entry.dtype,
327 rows,
328 cols,
329 row_scale: Vec::new(),
330 col_field: Vec::new(),
331 vbit_offsets: Vec::new(),
332 repack: Vec::new(),
333 }),
334 TensorDtype::Q1T if cols % GROUP_SIZE == 0 => Ok(Self::Mapped {
338 model: model.clone(),
339 idx,
340 dtype: entry.dtype,
341 rows,
342 cols,
343 row_scale: Vec::new(),
344 col_field: Vec::new(),
345 vbit_offsets: Vec::new(),
346 repack: Vec::new(),
347 }),
348 _ => {
350 let mut data = vec![0.0f32; rows * cols];
351 cortiq_core::quant::dequant_tensor(entry, bytes, &mut data)?;
352 Ok(Self::from_f32(data, rows, cols))
353 }
354 }
355 }
356
357 pub(crate) fn is_q1(&self) -> bool {
360 matches!(
361 self,
362 Self::Mapped {
363 dtype: TensorDtype::Q1,
364 ..
365 }
366 )
367 }
368
369 pub(crate) fn f32_parts(&self) -> Option<(&[f32], usize, usize)> {
372 match self {
373 Self::F32 { data, rows, cols } => Some((data, *rows, *cols)),
374 _ => None,
375 }
376 }
377
378 pub(crate) fn q1_parts(&self) -> Option<(usize, usize, usize)> {
385 if self.has_prism_contract() {
386 return None;
387 }
388 match self {
389 #[cfg(target_os = "macos")]
390 Self::Mapped {
391 dtype: TensorDtype::Q1T,
392 ..
393 } if !crate::gpu::metal_q1t_enabled() => None,
394 Self::Mapped {
395 idx,
396 dtype:
397 TensorDtype::Q1
398 | TensorDtype::Q1T
399 | TensorDtype::Q4Block
400 | TensorDtype::Q4Tiled
401 | TensorDtype::Q4TiledP
405 | TensorDtype::Q8Row
406 | TensorDtype::Q8_2f,
407 rows,
408 cols,
409 ..
410 } => Some((*idx, *rows, *cols)),
411 _ => None,
412 }
413 }
414
415 #[cfg(target_os = "macos")]
422 pub(crate) fn metal_graph_parts(&self) -> Option<(usize, usize, usize)> {
423 if let Some((model, idx, kind, _)) = self.graph_weight_descriptor() {
424 let name = &model.tensors[idx].name;
425 let forward = kind == 9 && crate::prism::is_forward_weight(model, name);
426 let affine = kind == 9 && crate::prism::is_affine_target(model, name);
427 if forward && affine {
428 let e = model.tensors.get(idx)?;
429 return Some((idx, *e.shape.first()?, *e.shape.get(1)?));
430 }
431 }
432 self.q1_parts()
433 }
434
435 pub(crate) fn q4t_parts(&self) -> Option<(usize, usize, usize)> {
441 if self.has_prism_contract() {
442 return None;
443 }
444 match self {
445 Self::Mapped {
446 idx,
447 dtype: TensorDtype::Q4Tiled,
448 rows,
449 cols,
450 ..
451 } => Some((*idx, *rows, *cols)),
452 _ => None,
453 }
454 }
455
456 pub(crate) fn q4tp_parts(&self) -> Option<(usize, usize, usize)> {
460 if self.has_prism_contract() {
461 return None;
462 }
463 match self {
464 Self::Mapped {
465 idx,
466 dtype: TensorDtype::Q4TiledP,
467 rows,
468 cols,
469 ..
470 } => Some((*idx, *rows, *cols)),
471 _ => None,
472 }
473 }
474
475 pub(crate) fn q8_row_parts(&self) -> Option<(usize, usize, usize, &[f32])> {
480 if self.has_prism_contract() {
481 return None;
482 }
483 match self {
484 Self::Mapped {
485 idx,
486 dtype: TensorDtype::Q8Row,
487 rows,
488 cols,
489 row_scale,
490 col_field,
491 ..
492 } if col_field.is_empty() => Some((*idx, *rows, *cols, row_scale)),
493 _ => None,
494 }
495 }
496
497 pub fn model_dtype(&self) -> Option<cortiq_core::TensorDtype> {
501 match self {
502 Self::Mapped { dtype, .. } => Some(*dtype),
503 _ => None,
504 }
505 }
506
507 pub fn model_idx(&self) -> Option<usize> {
512 match self {
513 Self::Mapped { idx, .. } => Some(*idx),
514 _ => None,
515 }
516 }
517
518 pub fn model_arc(&self) -> Option<std::sync::Arc<cortiq_core::CmfModel>> {
523 match self {
524 Self::Mapped { model, .. } => Some(model.clone()),
525 _ => None,
526 }
527 }
528
529 pub(crate) fn has_prism_contract(&self) -> bool {
534 matches!(self, Self::Mapped { model, .. } if crate::prism::has_contract(model))
535 }
536
537 pub fn rows(&self) -> usize {
538 match self {
539 Self::F32 { rows, .. } | Self::Mapped { rows, .. } => *rows,
540 }
541 }
542
543 pub(crate) fn mapped_q4t(&self) -> Option<(&Arc<CmfModel>, usize)> {
546 if self.has_prism_contract() {
547 return None;
548 }
549 match self {
550 Self::Mapped {
551 model,
552 idx,
553 dtype: TensorDtype::Q4Tiled,
554 ..
555 } => Some((model, *idx)),
556 _ => None,
557 }
558 }
559
560 pub fn mapped_q4tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
563 if self.has_prism_contract() {
564 return None;
565 }
566 match self {
567 Self::Mapped {
568 model,
569 idx,
570 dtype: TensorDtype::Q4TiledP,
571 ..
572 } => Some((model, *idx)),
573 _ => None,
574 }
575 }
576
577 pub fn mapped_device_gemm(&self) -> Option<(&Arc<CmfModel>, usize)> {
585 if self.has_prism_contract() {
586 return None;
587 }
588 match self {
589 Self::Mapped {
590 model,
591 idx,
592 dtype: TensorDtype::Q4TiledP | TensorDtype::Q8Row | TensorDtype::Q8_2f,
593 ..
594 } => Some((model, *idx)),
595 _ => None,
596 }
597 }
598
599 pub fn mapped_q2tp(&self) -> Option<(&Arc<CmfModel>, usize)> {
602 if self.has_prism_contract() {
603 return None;
604 }
605 match self {
606 Self::Mapped {
607 model,
608 idx,
609 dtype: TensorDtype::Q2TiledP,
610 ..
611 } => Some((model, *idx)),
612 _ => None,
613 }
614 }
615
616 pub fn cols(&self) -> usize {
617 match self {
618 Self::F32 { cols, .. } | Self::Mapped { cols, .. } => *cols,
619 }
620 }
621
622 pub fn mapped_q1(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
625 if self.has_prism_contract() {
626 return None;
627 }
628 match self {
629 Self::Mapped {
630 model,
631 idx,
632 dtype: TensorDtype::Q1,
633 ..
634 } => Some((model, *idx)),
635 _ => None,
636 }
637 }
638
639 pub fn graph_weight(&self) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
649 if self.has_prism_contract() {
650 return None;
651 }
652 self.graph_weight_descriptor()
653 }
654
655 pub(crate) fn graph_weight_descriptor(
660 &self,
661 ) -> Option<(&std::sync::Arc<CmfModel>, usize, u8, &[f32])> {
662 match self {
663 Self::Mapped {
664 model,
665 idx,
666 dtype: TensorDtype::Q8Row,
667 row_scale,
668 ..
669 } => Some((model, *idx, 0, row_scale.as_slice())),
670 Self::Mapped {
671 model,
672 idx,
673 dtype: TensorDtype::Q1,
674 ..
675 } => Some((model, *idx, 1, &[])),
676 Self::Mapped {
681 model,
682 idx,
683 dtype: TensorDtype::Q4Tiled,
684 ..
685 } => Some((model, *idx, 5, &[])),
686 Self::Mapped {
690 model,
691 idx,
692 dtype: TensorDtype::Q4TiledP,
693 ..
694 } => Some((model, *idx, 6, &[])),
695 Self::Mapped {
696 model,
697 idx,
698 dtype: TensorDtype::Q4Block,
699 ..
700 } => Some((model, *idx, 2, &[])),
701 Self::Mapped {
706 model,
707 idx,
708 dtype: TensorDtype::Q8_2f,
709 ..
710 } => Some((model, *idx, 7, &[])),
711 Self::Mapped {
712 model,
713 idx,
714 dtype: TensorDtype::Q1T,
715 ..
716 } => Some((model, *idx, 3, &[])),
717 Self::Mapped {
721 model,
722 idx,
723 dtype: TensorDtype::Q2TiledP,
724 ..
725 } => Some((model, *idx, 9, &[])),
726 _ => None,
727 }
728 }
729
730 pub fn as_f32(&self) -> Option<&[f32]> {
733 match self {
734 Self::F32 { data, .. } => Some(data),
735 Self::Mapped { .. } => None,
736 }
737 }
738
739 fn quant_bytes(&self) -> &[u8] {
740 match self {
741 Self::Mapped { model, idx, .. } => model.entry_bytes(&model.tensors[*idx]),
742 Self::F32 { .. } => unreachable!("quant_bytes on F32"),
743 }
744 }
745
746 pub fn row_f32(&self, r: usize, dst: &mut [f32]) {
748 let cols = self.cols();
749 debug_assert_eq!(dst.len(), cols);
750 match self {
751 Self::F32 { data, .. } => dst.copy_from_slice(&data[r * cols..(r + 1) * cols]),
752 Self::Mapped {
753 model,
754 idx,
755 dtype,
756 row_scale,
757 col_field,
758 vbit_offsets,
759 ..
760 } => {
761 if *dtype == TensorDtype::Q4Tiled {
762 let bytes = self.quant_bytes();
763 let gpr = cols / GROUP_SIZE;
764 for gi in 0..gpr {
765 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
766 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
767 for (k, &b) in tile[2..].iter().enumerate() {
768 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
769 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
770 }
771 }
772 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
773 crate::prism::inverse_embedding(model, dst);
774 }
775 return;
776 }
777 if *dtype == TensorDtype::Q4TiledP {
778 let bytes = self.quant_bytes();
779 let gpr = cols / GROUP_SIZE;
780 let v = Q4tpView::new(bytes, self.rows(), cols);
781 let mut sc = vec![0f32; gpr];
782 v.scales_into(r, gpr, &mut sc);
783 for gi in 0..gpr {
784 let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
785 let s = sc[gi];
786 for (k, &b) in tile.iter().enumerate() {
787 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
788 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
789 }
790 }
791 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
792 crate::prism::inverse_embedding(model, dst);
793 }
794 return;
795 }
796 if *dtype == TensorDtype::Q2TiledP {
797 let bytes = self.quant_bytes();
798 let gpr = cols / GROUP_SIZE;
799 let v = Q4tpView::new_q2(bytes, self.rows(), cols);
800 let mut sc = vec![0f32; gpr];
801 v.scales_into(r, gpr, &mut sc);
802 for gi in 0..gpr {
803 let ch =
804 &v.nib[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
805 let s = sc[gi];
806 for (k, &b) in ch.iter().enumerate() {
807 for j in 0..4 {
808 let center = if crate::prism::is_affine_target(
809 model,
810 &model.tensors[*idx].name,
811 ) {
812 1.0
813 } else {
814 1.5
815 };
816 dst[gi * GROUP_SIZE + k * 4 + j] =
817 (((b >> (2 * j)) & 3) as f32 - center) * s;
818 }
819 }
820 }
821 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
822 crate::prism::inverse_embedding(model, dst);
823 }
824 return;
825 }
826 if *dtype == TensorDtype::Q4Block {
827 let (packed, scales) = q4_split(self.quant_bytes(), self.rows(), cols);
828 let gpr = cols / GROUP_SIZE;
829 for gi in 0..gpr {
830 let g = r * gpr + gi;
831 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
832 for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
833 dst[gi * GROUP_SIZE + k * 2] = ((b & 0x0F) as f32 - 8.0) * s;
834 dst[gi * GROUP_SIZE + k * 2 + 1] = (((b >> 4) & 0x0F) as f32 - 8.0) * s;
835 }
836 }
837 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
838 crate::prism::inverse_embedding(model, dst);
839 }
840 return;
841 }
842 if *dtype == TensorDtype::Q1 {
843 let bytes = self.quant_bytes();
844 let gpr = cols / GROUP_SIZE;
845 for gi in 0..gpr {
846 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
847 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
848 for (j, &b) in tile[2..].iter().enumerate() {
849 for k in 0..8 {
850 dst[gi * GROUP_SIZE + j * 8 + k] =
851 (((b >> k) & 1) as f32 * 2.0 - 1.0) * s;
852 }
853 }
854 }
855 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
856 crate::prism::inverse_embedding(model, dst);
857 }
858 return;
859 }
860 if *dtype == TensorDtype::Q1T {
861 let bytes = self.quant_bytes();
862 let gpr = cols / GROUP_SIZE;
863 let base_len = self.rows() * gpr * cortiq_core::quant::Q1T_TILE;
864 for gi in 0..gpr {
865 let off = (r * gpr + gi) * cortiq_core::quant::Q1T_TILE;
866 let s = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
867 bytes[off],
868 bytes[off + 1],
869 ]));
870 let codes = &bytes[off + 2..off + cortiq_core::quant::Q1T_TILE];
871 for k in 0..GROUP_SIZE {
872 dst[gi * GROUP_SIZE + k] = match cortiq_core::quant::q1t_code(codes, k)
873 {
874 1 => s,
875 2 => -s,
876 _ => 0.0,
877 };
878 }
879 }
880 let rows = self.rows();
882 let entries = base_len + (rows + 1) * 4;
883 if entries <= bytes.len() {
884 let ptrs = &bytes[base_len..base_len + (rows + 1) * 4];
885 let r0 = u32::from_le_bytes([
886 ptrs[r * 4],
887 ptrs[r * 4 + 1],
888 ptrs[r * 4 + 2],
889 ptrs[r * 4 + 3],
890 ]) as usize;
891 let r1 = u32::from_le_bytes([
892 ptrs[(r + 1) * 4],
893 ptrs[(r + 1) * 4 + 1],
894 ptrs[(r + 1) * 4 + 2],
895 ptrs[(r + 1) * 4 + 3],
896 ]) as usize;
897 let off = entries + r0 * 4;
898 for i in 0..r1 - r0 {
899 let item = &bytes[off + i * 4..off + i * 4 + 4];
900 let c = u16::from_le_bytes([item[0], item[1]]) as usize;
901 let v = cortiq_core::quant::f16_to_f32(u16::from_le_bytes([
902 item[2], item[3],
903 ]));
904 if c < cols {
905 dst[c] = v;
906 }
907 }
908 }
909 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
910 crate::prism::inverse_embedding(model, dst);
911 }
912 return;
913 }
914 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
915 let bytes = self.quant_bytes();
916 let rows = self.rows();
917 let ng = cols / GROUP_SIZE;
918 let bits = &bytes[..rows];
919 let sc_off = rows;
920 let off = vbit_offsets[r];
923 let b = bits[r] as usize;
924 let l = ((1usize << (b - 1)) - 1) as f32;
925 let data = &bytes[off..];
926 let (mut acc, mut nbits, mut byte_idx) = (0u64, 0usize, 0usize);
927 for (i, d) in dst.iter_mut().enumerate() {
928 while nbits < b {
929 acc = (acc << 8) | data[byte_idx] as u64;
930 byte_idx += 1;
931 nbits += 8;
932 }
933 let u = ((acc >> (nbits - b)) & ((1u64 << b) - 1)) as f32;
934 nbits -= b;
935 let so = (r * ng + i / GROUP_SIZE) * 2;
936 let sv = f16_to_f32(u16::from_le_bytes([
937 bytes[sc_off + so],
938 bytes[sc_off + so + 1],
939 ]));
940 *d = (u - l) * sv;
941 }
942 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
943 crate::prism::inverse_embedding(model, dst);
944 }
945 return;
946 }
947 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
948 let s = row_scale[r];
949 match dtype {
950 TensorDtype::Q8Row => {
951 for (d, &b) in dst.iter_mut().zip(q) {
952 *d = (b as i8) as f32 * s;
953 }
954 }
955 TensorDtype::Q8_2f => {
956 for (i, (d, &b)) in dst.iter_mut().zip(q).enumerate() {
957 *d = (b as i8) as f32 * s * col_field[i];
958 }
959 }
960 _ => unreachable!(),
961 }
962 if crate::prism::is_inverse_embedding(model, &model.tensors[*idx].name) {
963 crate::prism::inverse_embedding(model, dst);
964 }
965 }
966 }
967 }
968
969 pub fn sparse_col_ok(&self) -> bool {
974 match self {
975 Self::F32 { .. } => true,
976 Self::Mapped { dtype, .. } => {
977 matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
978 }
979 }
980 }
981
982 pub fn add_col_scaled(&self, c: usize, w: f32, out: &mut [f32]) {
986 let inter = self.cols();
987 let hidden = self.rows();
988 debug_assert_eq!(out.len(), hidden);
989 match self {
990 Self::F32 { data, .. } => {
991 for (k, o) in out.iter_mut().enumerate() {
992 *o += w * data[k * inter + c];
993 }
994 }
995 Self::Mapped {
996 dtype,
997 row_scale,
998 col_field,
999 ..
1000 } => {
1001 let q = self.quant_bytes();
1002 let colf = if *dtype == TensorDtype::Q8_2f {
1003 col_field[c]
1004 } else {
1005 1.0
1006 };
1007 let wc = w * colf;
1008 for (k, o) in out.iter_mut().enumerate() {
1009 let b = q[k * inter + c] as i8 as f32;
1010 *o += wc * b * row_scale[k];
1011 }
1012 }
1013 }
1014 }
1015
1016 #[inline]
1025 pub fn prefetch_row(&self, r: usize) {
1026 let Self::Mapped { dtype, .. } = self else {
1027 return;
1028 };
1029 if !matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f) {
1030 return;
1031 }
1032 let cols = self.cols();
1033 let q = self.quant_bytes();
1034 let (a, b) = (r * cols, (r + 1) * cols);
1035 if b > q.len() {
1036 return;
1037 }
1038 let mut j = a;
1039 while j < b {
1040 unsafe { std::ptr::read_volatile(q.as_ptr().add(j)) };
1041 j += 512;
1042 }
1043 }
1044
1045 pub fn add_row_scaled(&self, r: usize, w: f32, out: &mut [f32], scratch: &mut [f32]) {
1054 let cols = self.cols();
1055 debug_assert_eq!(out.len(), cols);
1056 match self {
1057 Self::F32 { data, .. } => {
1058 let row = &data[r * cols..(r + 1) * cols];
1059 for (o, v) in out.iter_mut().zip(row) {
1060 *o += w * v;
1061 }
1062 }
1063 Self::Mapped {
1064 dtype,
1065 row_scale,
1066 col_field,
1067 ..
1068 } => match dtype {
1069 TensorDtype::Q8Row => {
1070 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1071 let ws = w * row_scale[r];
1072 let row: &[i8] =
1073 unsafe { std::slice::from_raw_parts(q.as_ptr() as *const i8, q.len()) };
1074 axpy_i8_f32(out, row, ws);
1075 }
1076 TensorDtype::Q8_2f => {
1077 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1078 let ws = w * row_scale[r];
1079 for ((o, b), c) in out.iter_mut().zip(q).zip(col_field) {
1080 *o += ws * c * (*b as i8 as f32);
1081 }
1082 }
1083 _ => {
1084 self.row_f32(r, scratch);
1085 for (o, v) in out.iter_mut().zip(scratch.iter()) {
1086 *o += w * v;
1087 }
1088 }
1089 },
1090 }
1091 }
1092
1093 pub fn row_dot(&self, r: usize, x: &[f32], scratch: &mut [f32]) -> f32 {
1097 let cols = self.cols();
1098 match self {
1099 Self::F32 { data, .. } => {
1100 let row = &data[r * cols..(r + 1) * cols];
1101 row.iter().zip(x).map(|(w, v)| w * v).sum()
1102 }
1103 Self::Mapped {
1104 model,
1105 idx,
1106 dtype,
1107 row_scale,
1108 col_field,
1109 ..
1110 } => {
1111 let prism_forward =
1112 crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1113 if prism_forward {
1114 let transformed = crate::prism::forward(model, &x[..cols]);
1115 let gpr = cols / GROUP_SIZE;
1116 match dtype {
1117 TensorDtype::Q2TiledP => {
1118 let v = Q4tpView::new_q2(self.quant_bytes(), self.rows(), cols);
1119 let mut sc = vec![0f32; gpr];
1120 v.scales_into(r, gpr, &mut sc);
1121 if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1122 return q2tp_affine_row_exact(v.nib, r, gpr, &transformed, &sc);
1123 }
1124 return q2tp_row_exact(v.nib, r, gpr, &transformed, &sc);
1125 }
1126 _ => {
1127 self.row_f32(r, scratch);
1128 return scratch.iter().zip(&transformed).map(|(w, v)| w * v).sum();
1129 }
1130 }
1131 }
1132 match dtype {
1133 TensorDtype::Q8Row => {
1134 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1135 dot_i8_f32(q, x) * row_scale[r]
1136 }
1137 TensorDtype::Q8_2f => {
1138 let q = &self.quant_bytes()[r * cols..(r + 1) * cols];
1139 dot_i8_col_f32(q, x, col_field) * row_scale[r]
1140 }
1141 _ => {
1142 self.row_f32(r, scratch);
1143 scratch.iter().zip(x).map(|(w, v)| w * v).sum()
1144 }
1145 }
1146 }
1147 }
1148 }
1149
1150 pub fn matvec(&self, x: &[f32], out: &mut [f32], pool: Option<&Pool>) {
1153 match self {
1154 Self::F32 { data, .. } => {
1158 if !crate::f32_backend::matvec(data, x, out) {
1159 matvec_rows(pool, data, x, out);
1160 }
1161 }
1162 Self::Mapped {
1163 model,
1164 idx,
1165 dtype,
1166 rows,
1167 cols,
1168 row_scale,
1169 col_field,
1170 vbit_offsets,
1171 repack,
1172 } => {
1173 let _ = (model, idx);
1174 assert!(
1184 out.len() >= *rows && x.len() >= *cols,
1185 "matvec {rows}x{cols}: out {} (need {rows}), x {} (need {cols})",
1186 out.len(),
1187 x.len(),
1188 );
1189 let prism_forward =
1190 crate::prism::is_forward_weight(model, &model.tensors[*idx].name);
1191 if *dtype == TensorDtype::Q2TiledP
1192 && std::env::var("CMF_Q2TP_TRACE").as_deref() == Ok("1")
1193 {
1194 use std::sync::atomic::{AtomicUsize, Ordering};
1195 static N: AtomicUsize = AtomicUsize::new(0);
1196 let n = N.fetch_add(1, Ordering::Relaxed);
1197 if n < 128 {
1198 eprintln!(
1199 "q2tp-dispatch #{n} name={} prism={} rows={} cols={} gpu={} optin={} layer={}",
1200 model.tensors[*idx].name,
1201 prism_forward,
1202 rows,
1203 cols,
1204 crate::gpu::enabled_here(),
1205 crate::gpu::q2tp_gpu_opt_in(),
1206 crate::gpu::cur_layer(),
1207 );
1208 }
1209 }
1210 if prism_forward {
1215 let transformed = crate::prism::forward(model, &x[..*cols]);
1216 match dtype {
1217 TensorDtype::Q4Block => {
1218 q4matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1219 }
1220 TensorDtype::Q4Tiled => {
1221 q4t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1222 }
1223 TensorDtype::Q4TiledP => {
1224 q4tp_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1225 }
1226 TensorDtype::Q2TiledP => {
1227 let affine =
1228 crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1229 if *rows * *cols >= 8_388_608
1230 && crate::gpu::enabled_here()
1231 && crate::gpu::q2tp_gpu_opt_in()
1232 {
1233 let gpu_ok = if affine {
1234 crate::gpu::q2tp_affine_matvec(
1235 model,
1236 *idx,
1237 &transformed,
1238 *rows,
1239 *cols,
1240 out,
1241 )
1242 } else {
1243 crate::gpu::q2tp_matvec(
1244 model,
1245 *idx,
1246 &transformed,
1247 *rows,
1248 *cols,
1249 out,
1250 )
1251 };
1252 if gpu_ok {
1253 return;
1254 }
1255 }
1256 if affine {
1257 q2tp_affine_matvec(
1258 self.quant_bytes(),
1259 &transformed,
1260 *rows,
1261 *cols,
1262 out,
1263 pool,
1264 )
1265 } else {
1266 q2tp_matvec(
1267 self.quant_bytes(),
1268 &transformed,
1269 *rows,
1270 *cols,
1271 out,
1272 pool,
1273 )
1274 }
1275 }
1276 TensorDtype::Q1 => {
1277 q1_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1278 }
1279 TensorDtype::Q1T => {
1280 q1t_matvec(self.quant_bytes(), &transformed, *rows, *cols, out, pool)
1281 }
1282 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec(
1283 self.quant_bytes(),
1284 vbit_offsets,
1285 &transformed,
1286 *rows,
1287 *cols,
1288 out,
1289 pool,
1290 ),
1291 TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec(
1292 self.quant_bytes(),
1293 repack,
1294 row_scale,
1295 &transformed,
1296 col_field,
1297 *dtype,
1298 *rows,
1299 *cols,
1300 out,
1301 pool,
1302 ),
1303 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1304 }
1305 return;
1306 }
1307 if *dtype == TensorDtype::Q4Block {
1308 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1312 let t0 = std::time::Instant::now();
1313 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1314 crate::gpu::ProbeArm::Gpu => {
1315 if crate::gpu::q4b_matvec(model, *idx, x, *rows, *cols, out) {
1316 crate::gpu::probe_record(
1317 crate::gpu::OpClass::Matvec,
1318 true,
1319 t0.elapsed(),
1320 );
1321 return;
1322 }
1323 }
1324 crate::gpu::ProbeArm::CpuTimed => {
1325 q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1326 crate::gpu::probe_record(
1327 crate::gpu::OpClass::Matvec,
1328 false,
1329 t0.elapsed(),
1330 );
1331 return;
1332 }
1333 crate::gpu::ProbeArm::Cpu => {}
1334 }
1335 }
1336 q4matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1337 return;
1338 }
1339 if *dtype == TensorDtype::Q4Tiled {
1340 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1345 let t0 = std::time::Instant::now();
1346 let cls = crate::gpu::matvec_class(*rows, *cols);
1347 match crate::gpu::probe_arm(cls) {
1348 crate::gpu::ProbeArm::Gpu => {
1349 if crate::gpu::q4t_matvec(model, *idx, x, *rows, *cols, out) {
1350 crate::gpu::probe_record(cls, true, t0.elapsed());
1351 return;
1352 }
1353 }
1354 crate::gpu::ProbeArm::CpuTimed => {
1355 q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1356 crate::gpu::probe_record(cls, false, t0.elapsed());
1357 return;
1358 }
1359 crate::gpu::ProbeArm::Cpu => {}
1360 }
1361 }
1362 q4t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1363 return;
1364 }
1365 if *dtype == TensorDtype::Q4TiledP {
1366 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1372 let t0 = std::time::Instant::now();
1373 let cls = crate::gpu::matvec_class(*rows, *cols);
1374 match crate::gpu::probe_arm(cls) {
1375 crate::gpu::ProbeArm::Gpu => {
1376 if crate::gpu::q4tp_matvec(model, *idx, x, *rows, *cols, out) {
1377 crate::gpu::probe_record(cls, true, t0.elapsed());
1378 return;
1379 }
1380 }
1381 crate::gpu::ProbeArm::CpuTimed => {
1382 q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1383 crate::gpu::probe_record(cls, false, t0.elapsed());
1384 return;
1385 }
1386 crate::gpu::ProbeArm::Cpu => {}
1387 }
1388 }
1389 q4tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1390 return;
1391 }
1392 if *dtype == TensorDtype::Q2TiledP {
1393 q2tp_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1394 return;
1395 }
1396 if *dtype == TensorDtype::Q1 {
1397 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1402 let t0 = std::time::Instant::now();
1403 let arm = if crate::gpu::q1_force() {
1404 crate::gpu::ProbeArm::Gpu
1405 } else {
1406 crate::gpu::probe_arm(crate::gpu::OpClass::Matvec)
1407 };
1408 match arm {
1409 crate::gpu::ProbeArm::Gpu => {
1410 if crate::gpu::q1_matvec(model, *idx, x, *rows, *cols, out) {
1411 crate::gpu::probe_record(
1412 crate::gpu::OpClass::Matvec,
1413 true,
1414 t0.elapsed(),
1415 );
1416 return;
1417 }
1418 }
1419 crate::gpu::ProbeArm::CpuTimed => {
1420 q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1421 crate::gpu::probe_record(
1422 crate::gpu::OpClass::Matvec,
1423 false,
1424 t0.elapsed(),
1425 );
1426 return;
1427 }
1428 crate::gpu::ProbeArm::Cpu => {}
1429 }
1430 }
1431 q1_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1432 return;
1433 }
1434 if *dtype == TensorDtype::Q1T {
1435 if *rows * *cols >= 8_388_608 && crate::gpu::enabled_here() {
1439 let t0 = std::time::Instant::now();
1440 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1441 crate::gpu::ProbeArm::Gpu => {
1442 if crate::gpu::q1t_matvec(model, *idx, x, *rows, *cols, out) {
1443 q1t_add_overlay(self.quant_bytes(), x, *rows, *cols, out, pool);
1444 crate::gpu::probe_record(
1445 crate::gpu::OpClass::Matvec,
1446 true,
1447 t0.elapsed(),
1448 );
1449 return;
1450 }
1451 }
1452 crate::gpu::ProbeArm::CpuTimed => {
1453 q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1454 crate::gpu::probe_record(
1455 crate::gpu::OpClass::Matvec,
1456 false,
1457 t0.elapsed(),
1458 );
1459 return;
1460 }
1461 crate::gpu::ProbeArm::Cpu => {}
1462 }
1463 }
1464 q1t_matvec(self.quant_bytes(), x, *rows, *cols, out, pool);
1465 return;
1466 }
1467 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1468 vbitmatvec(self.quant_bytes(), vbit_offsets, x, *rows, *cols, out, pool);
1469 return;
1470 }
1471 let xs = prescale(x, col_field, *dtype);
1472 if *rows >= crate::gpu::min_rows()
1477 && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
1478 && gpu_lmhead_enabled()
1479 && crate::gpu::enabled_here()
1480 {
1481 let t0 = std::time::Instant::now();
1484 match crate::gpu::probe_arm(crate::gpu::OpClass::Matvec) {
1485 crate::gpu::ProbeArm::Gpu => {}
1486 crate::gpu::ProbeArm::CpuTimed => {
1487 qmatvec(
1488 self.quant_bytes(),
1489 repack,
1490 row_scale,
1491 x,
1492 col_field,
1493 *dtype,
1494 *rows,
1495 *cols,
1496 out,
1497 pool,
1498 );
1499 crate::gpu::probe_record(
1500 crate::gpu::OpClass::Matvec,
1501 false,
1502 t0.elapsed(),
1503 );
1504 return;
1505 }
1506 crate::gpu::ProbeArm::Cpu => {
1507 qmatvec(
1508 self.quant_bytes(),
1509 repack,
1510 row_scale,
1511 x,
1512 col_field,
1513 *dtype,
1514 *rows,
1515 *cols,
1516 out,
1517 pool,
1518 );
1519 return;
1520 }
1521 }
1522 let frac = gpu_split_frac();
1523 let cpu_rows = ((*rows as f32) * (1.0 - frac)) as usize;
1524 let (out_cpu, out_gpu) = out.split_at_mut(cpu_rows);
1525 let bytes = self.quant_bytes();
1526 let ok = std::thread::scope(|sc| {
1527 let g = sc.spawn(|| {
1528 crate::gpu::q8_matvec_range(
1529 model,
1530 *idx,
1531 cpu_rows,
1532 &row_scale[cpu_rows..],
1533 &xs,
1534 *rows - cpu_rows,
1535 *cols,
1536 out_gpu,
1537 )
1538 });
1539 if cpu_rows > 0 {
1540 let rep_cpu = if repack.is_empty() {
1543 &[][..]
1544 } else {
1545 &repack[..(cpu_rows / 4) * 4 * *cols]
1546 };
1547 qmatvec(
1548 &bytes[..cpu_rows * *cols],
1549 rep_cpu,
1550 &row_scale[..cpu_rows],
1551 x,
1552 col_field,
1553 *dtype,
1554 cpu_rows,
1555 *cols,
1556 out_cpu,
1557 pool,
1558 );
1559 }
1560 g.join().unwrap_or(false)
1561 });
1562 if ok {
1563 crate::gpu::probe_record(crate::gpu::OpClass::Matvec, true, t0.elapsed());
1564 return;
1565 }
1566 qmatvec(
1569 &bytes[cpu_rows * *cols..(*rows) * *cols],
1570 &[],
1571 &row_scale[cpu_rows..],
1572 x,
1573 col_field,
1574 *dtype,
1575 *rows - cpu_rows,
1576 *cols,
1577 out_gpu,
1578 pool,
1579 );
1580 return;
1581 }
1582 qmatvec(
1583 self.quant_bytes(),
1584 repack,
1585 row_scale,
1586 x,
1587 col_field,
1588 *dtype,
1589 *rows,
1590 *cols,
1591 out,
1592 pool,
1593 );
1594 }
1595 }
1596 }
1597
1598 pub fn matvec2(
1600 &self,
1601 x1: &[f32],
1602 x2: &[f32],
1603 o1: &mut [f32],
1604 o2: &mut [f32],
1605 pool: Option<&Pool>,
1606 ) {
1607 match self {
1608 Self::F32 { data, .. } => matvec_rows2(pool, data, x1, x2, o1, o2),
1609 Self::Mapped {
1610 model,
1611 idx,
1612 dtype,
1613 rows,
1614 cols,
1615 row_scale,
1616 col_field,
1617 vbit_offsets,
1618 ..
1619 } => {
1620 if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1621 let tx1 = crate::prism::forward(model, &x1[..*cols]);
1622 let tx2 = crate::prism::forward(model, &x2[..*cols]);
1623 match dtype {
1624 TensorDtype::Q4Block => {
1625 q4matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1626 }
1627 TensorDtype::Q4Tiled => {
1628 q4t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1629 }
1630 TensorDtype::Q4TiledP => {
1631 q4tp_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1632 }
1633 TensorDtype::Q2TiledP => {
1634 if crate::prism::is_affine_target(model, &model.tensors[*idx].name) {
1635 q2tp_affine_matvec2(
1636 self.quant_bytes(),
1637 &tx1,
1638 &tx2,
1639 *rows,
1640 *cols,
1641 o1,
1642 o2,
1643 pool,
1644 )
1645 } else {
1646 q2tp_matvec2(
1647 self.quant_bytes(),
1648 &tx1,
1649 &tx2,
1650 *rows,
1651 *cols,
1652 o1,
1653 o2,
1654 pool,
1655 )
1656 }
1657 }
1658 TensorDtype::Q1 => {
1659 q1_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1660 }
1661 TensorDtype::Q1T => {
1662 q1t_matvec2(self.quant_bytes(), &tx1, &tx2, *rows, *cols, o1, o2, pool)
1663 }
1664 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatvec2(
1665 self.quant_bytes(),
1666 vbit_offsets,
1667 &tx1,
1668 &tx2,
1669 *rows,
1670 *cols,
1671 o1,
1672 o2,
1673 pool,
1674 ),
1675 TensorDtype::Q8Row | TensorDtype::Q8_2f => qmatvec2(
1676 self.quant_bytes(),
1677 row_scale,
1678 &tx1,
1679 &tx2,
1680 col_field,
1681 *dtype,
1682 *rows,
1683 *cols,
1684 o1,
1685 o2,
1686 pool,
1687 ),
1688 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1689 }
1690 return;
1691 }
1692 if *dtype == TensorDtype::Q4Block {
1693 q4matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1694 return;
1695 }
1696 if *dtype == TensorDtype::Q4Tiled {
1697 q4t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1698 return;
1699 }
1700 if *dtype == TensorDtype::Q4TiledP {
1701 q4tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1702 return;
1703 }
1704 if *dtype == TensorDtype::Q2TiledP {
1705 q2tp_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1706 return;
1707 }
1708 if *dtype == TensorDtype::Q1 {
1709 q1_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1710 return;
1711 }
1712 if *dtype == TensorDtype::Q1T {
1713 q1t_matvec2(self.quant_bytes(), x1, x2, *rows, *cols, o1, o2, pool);
1719 return;
1720 }
1721 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
1722 vbitmatvec2(
1723 self.quant_bytes(),
1724 vbit_offsets,
1725 x1,
1726 x2,
1727 *rows,
1728 *cols,
1729 o1,
1730 o2,
1731 pool,
1732 );
1733 return;
1734 }
1735 qmatvec2(
1736 self.quant_bytes(),
1737 row_scale,
1738 x1,
1739 x2,
1740 col_field,
1741 *dtype,
1742 *rows,
1743 *cols,
1744 o1,
1745 o2,
1746 pool,
1747 );
1748 }
1749 }
1750 }
1751}
1752
1753impl QTensor {
1754 pub fn q4tp_mapped(&self) -> Option<(&std::sync::Arc<CmfModel>, usize)> {
1762 if self.has_prism_contract() {
1763 return None;
1764 }
1765 match self {
1766 Self::Mapped {
1767 model, idx, dtype, ..
1768 } if *dtype == TensorDtype::Q4TiledP => Some((model, *idx)),
1769 _ => None,
1770 }
1771 }
1772
1773 pub fn matmat(&self, xs_all: &[f32], b: usize, out: &mut [f32], pool: Option<&Pool>) {
1774 let cols = self.cols();
1775 let rows = self.rows();
1776 debug_assert_eq!(xs_all.len(), b * cols);
1777 debug_assert_eq!(out.len(), b * rows);
1778 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::Matmat);
1779 if crate::gptq_capture::capturing() {
1783 if let Self::Mapped { model, idx, .. } = self {
1784 crate::gptq_capture::accumulate(&model.tensors[*idx].name, xs_all, b, cols);
1785 }
1786 }
1787 match self {
1788 Self::F32 { data, .. } => {
1789 if crate::f32_backend::matmat(data, xs_all, b, rows, cols, out) {
1790 return;
1791 }
1792 let out_addr = SendMut(out.as_mut_ptr());
1793 let run = |start: usize, end: usize| {
1794 for o in start..end {
1795 let row = &data[o * cols..(o + 1) * cols];
1796 for bi in 0..b {
1797 let x = &xs_all[bi * cols..(bi + 1) * cols];
1798 let mut acc = 0f32;
1799 for j in 0..cols {
1800 acc += row[j] * x[j];
1801 }
1802 unsafe { *out_addr.at(bi * rows + o) = acc };
1803 }
1804 }
1805 };
1806 dispatch_rows(pool, rows, &run);
1807 }
1808 Self::Mapped {
1809 model,
1810 idx,
1811 dtype,
1812 row_scale,
1813 col_field,
1814 vbit_offsets,
1815 ..
1816 } => {
1817 if crate::prism::is_forward_weight(model, &model.tensors[*idx].name) {
1818 let mut transformed = Vec::with_capacity(xs_all.len());
1819 for bi in 0..b {
1820 transformed.extend_from_slice(&crate::prism::forward(
1821 model,
1822 &xs_all[bi * cols..(bi + 1) * cols],
1823 ));
1824 }
1825 match dtype {
1826 TensorDtype::Q4Block => {
1827 q4matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1828 }
1829 TensorDtype::Q4Tiled => {
1830 q4t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1831 }
1832 TensorDtype::Q4TiledP => {
1833 q4tp_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1834 }
1835 TensorDtype::Q2TiledP => {
1836 let affine =
1837 crate::prism::is_affine_target(model, &model.tensors[*idx].name);
1838 let gpu_batch_ok = if affine {
1844 b >= 2
1845 } else {
1846 b >= 32 && b * rows * cols >= 128_000_000
1847 };
1848 if gpu_batch_ok
1849 && cols % 32 == 0
1850 && crate::gpu::enabled_here()
1851 && crate::gpu::q2tp_gpu_opt_in()
1852 {
1853 let gpu_ok = if affine {
1854 crate::gpu::q2tp_affine_matmat(
1855 model,
1856 *idx,
1857 &transformed,
1858 b,
1859 rows,
1860 cols,
1861 out,
1862 )
1863 } else {
1864 crate::gpu::q2tp_matmat(
1865 model,
1866 *idx,
1867 &transformed,
1868 b,
1869 rows,
1870 cols,
1871 out,
1872 )
1873 };
1874 if gpu_ok {
1875 return;
1876 }
1877 }
1878 if affine
1882 && b == 1
1883 && cols % 32 == 0
1884 && crate::gpu::enabled_here()
1885 && crate::gpu::q2tp_gpu_opt_in()
1886 && crate::gpu::q2tp_affine_matvec(
1887 model,
1888 *idx,
1889 &transformed[..cols],
1890 rows,
1891 cols,
1892 &mut out[..rows],
1893 )
1894 {
1895 return;
1896 }
1897 if affine {
1898 q2tp_affine_matmat(
1899 self.quant_bytes(),
1900 &transformed,
1901 b,
1902 rows,
1903 cols,
1904 out,
1905 pool,
1906 )
1907 } else {
1908 q2tp_matmat(
1909 self.quant_bytes(),
1910 &transformed,
1911 b,
1912 rows,
1913 cols,
1914 out,
1915 pool,
1916 )
1917 }
1918 }
1919 TensorDtype::Q1 => {
1920 q1_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1921 }
1922 TensorDtype::Q1T => {
1923 q1t_matmat(self.quant_bytes(), &transformed, b, rows, cols, out, pool)
1924 }
1925 TensorDtype::Vbit | TensorDtype::VbitRo => vbitmatmat(
1926 self.quant_bytes(),
1927 vbit_offsets,
1928 &transformed,
1929 b,
1930 rows,
1931 cols,
1932 out,
1933 pool,
1934 ),
1935 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
1936 let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
1937 .map(|bi| {
1938 prescale(
1939 &transformed[bi * cols..(bi + 1) * cols],
1940 col_field,
1941 *dtype,
1942 )
1943 })
1944 .collect();
1945 qmatmat(self.quant_bytes(), row_scale, &pre, rows, cols, out, pool)
1946 }
1947 _ => unreachable!("unsupported mapped Prism dtype {dtype:?}"),
1948 }
1949 return;
1950 }
1951 if *dtype == TensorDtype::Q4Block {
1952 q4matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
1953 return;
1954 }
1955 if *dtype == TensorDtype::Q4TiledP {
1956 if b >= 32
1970 && b * rows * cols >= 128_000_000
1971 && cols % 32 == 0
1972 && !row_exact()
1973 && !crate::gpu::mm_killed()
1974 && crate::gpu::enabled_here()
1975 {
1976 let class = if b >= 128 {
1977 crate::gpu::OpClass::MatmatWide
1978 } else {
1979 crate::gpu::OpClass::Matmat
1980 };
1981 if let Self::Mapped { model, idx, .. } = self {
1982 if crate::mm_ab::on() {
1994 let mut g = vec![0f32; b * rows];
1995 let t = std::time::Instant::now();
1996 let took = crate::gpu::q4tp_matmat(
1997 model, *idx, xs_all, b, rows, cols, &mut g,
1998 );
1999 let dg = t.elapsed();
2000 let t = std::time::Instant::now();
2001 q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2002 let dc = t.elapsed();
2003 crate::mm_ab::record(b, rows, cols, took, dg, dc, &g, out);
2004 return;
2005 }
2006 let t0 = std::time::Instant::now();
2007 let resident = crate::gpu::weight_is_resident(model, *idx);
2011 match crate::gpu::probe_arm_cold_prefers_gpu(class, resident) {
2012 crate::gpu::ProbeArm::Gpu => {
2013 if crate::gpu::q4tp_matmat(
2014 model, *idx, xs_all, b, rows, cols, out,
2015 ) {
2016 let el = t0.elapsed();
2017 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2027 let budget = std::time::Duration::from_secs_f64(
2028 flops / 1.5e12 * 8.0 + 0.020,
2029 );
2030 crate::gpu::mm_budget_check(
2031 "q4tp matmat",
2032 el,
2033 budget,
2034 crate::gpu::probe_was_cold() || !resident,
2035 );
2036 crate::gpu::probe_record(class, true, el);
2037 return;
2038 }
2039 }
2040 crate::gpu::ProbeArm::CpuTimed => {
2041 q4tp_matmat(
2042 self.quant_bytes(),
2043 xs_all,
2044 b,
2045 rows,
2046 cols,
2047 out,
2048 pool,
2049 );
2050 crate::gpu::probe_record(class, false, t0.elapsed());
2051 return;
2052 }
2053 crate::gpu::ProbeArm::Cpu => {}
2054 }
2055 }
2056 }
2057 q4tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2058 return;
2059 }
2060 if *dtype == TensorDtype::Q2TiledP {
2061 if b >= 32
2067 && b * rows * cols >= 128_000_000
2068 && cols % 32 == 0
2069 && !crate::gpu::mm_killed()
2070 && crate::gpu::enabled_here()
2071 {
2072 let class = if b >= 128 {
2073 crate::gpu::OpClass::MatmatWide
2074 } else {
2075 crate::gpu::OpClass::Matmat
2076 };
2077 if let Self::Mapped { model, idx, .. } = self {
2078 let t0 = std::time::Instant::now();
2079 match crate::gpu::probe_arm(class) {
2080 crate::gpu::ProbeArm::Gpu => {
2081 if crate::gpu::q2tp_matmat(
2082 model, *idx, xs_all, b, rows, cols, out,
2083 ) {
2084 crate::gpu::probe_record(class, true, t0.elapsed());
2085 return;
2086 }
2087 }
2088 crate::gpu::ProbeArm::CpuTimed => {
2089 q2tp_matmat(
2090 self.quant_bytes(),
2091 xs_all,
2092 b,
2093 rows,
2094 cols,
2095 out,
2096 pool,
2097 );
2098 crate::gpu::probe_record(class, false, t0.elapsed());
2099 return;
2100 }
2101 crate::gpu::ProbeArm::Cpu => {}
2102 }
2103 }
2104 }
2105 q2tp_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2110 return;
2111 }
2112 if *dtype == TensorDtype::Q4Tiled {
2113 if b >= 32
2125 && b * rows * cols >= 128_000_000
2126 && cols % 32 == 0
2127 && !crate::gpu::mm_killed()
2128 && crate::gpu::enabled_here()
2129 {
2130 let class = if b >= 128 {
2131 crate::gpu::OpClass::MatmatWide
2132 } else {
2133 crate::gpu::OpClass::Matmat
2134 };
2135 if let Self::Mapped { model, idx, .. } = self {
2136 let t0 = std::time::Instant::now();
2137 match crate::gpu::probe_arm(class) {
2138 crate::gpu::ProbeArm::Gpu => {
2139 if crate::gpu::q4t_matmat(
2140 model, *idx, xs_all, b, rows, cols, out,
2141 ) {
2142 let el = t0.elapsed();
2143 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
2153 let budget = std::time::Duration::from_secs_f64(
2154 flops / 1.5e12 * 8.0 + 0.020,
2155 );
2156 crate::gpu::mm_budget_check(
2157 "q4t matmat",
2158 el,
2159 budget,
2160 crate::gpu::probe_was_cold(),
2161 );
2162 crate::gpu::probe_record(class, true, el);
2163 return;
2164 }
2165 }
2166 crate::gpu::ProbeArm::CpuTimed => {
2167 q4t_matmat(
2168 self.quant_bytes(),
2169 xs_all,
2170 b,
2171 rows,
2172 cols,
2173 out,
2174 pool,
2175 );
2176 crate::gpu::probe_record(class, false, t0.elapsed());
2177 return;
2178 }
2179 crate::gpu::ProbeArm::Cpu => {}
2180 }
2181 }
2182 }
2183 q4t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2184 return;
2185 }
2186 if *dtype == TensorDtype::Q1 {
2187 if b >= 32
2190 && b * rows * cols >= 128_000_000
2191 && cols % 64 == 0
2192 && crate::gpu::enabled_here()
2193 {
2194 if let Self::Mapped { model, idx, .. } = self {
2195 let t0 = std::time::Instant::now();
2196 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2197 crate::gpu::ProbeArm::Gpu => {
2198 if crate::gpu::q1_matmat(
2199 model, *idx, xs_all, b, rows, cols, out,
2200 ) {
2201 crate::gpu::probe_record(
2202 crate::gpu::OpClass::Matmat,
2203 true,
2204 t0.elapsed(),
2205 );
2206 return;
2207 }
2208 }
2209 crate::gpu::ProbeArm::CpuTimed => {
2210 q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2211 crate::gpu::probe_record(
2212 crate::gpu::OpClass::Matmat,
2213 false,
2214 t0.elapsed(),
2215 );
2216 return;
2217 }
2218 crate::gpu::ProbeArm::Cpu => {}
2219 }
2220 }
2221 }
2222 q1_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2223 return;
2224 }
2225 if *dtype == TensorDtype::Q1T {
2226 if b >= 32 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2229 if let Self::Mapped { model, idx, .. } = self {
2230 let t0 = std::time::Instant::now();
2231 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2232 crate::gpu::ProbeArm::Gpu => {
2233 if crate::gpu::q1t_matmat(
2234 model, *idx, xs_all, b, rows, cols, out,
2235 ) {
2236 crate::gpu::probe_record(
2237 crate::gpu::OpClass::Matmat,
2238 true,
2239 t0.elapsed(),
2240 );
2241 return;
2242 }
2243 }
2244 crate::gpu::ProbeArm::CpuTimed => {
2245 q1t_matmat(
2246 self.quant_bytes(),
2247 xs_all,
2248 b,
2249 rows,
2250 cols,
2251 out,
2252 pool,
2253 );
2254 crate::gpu::probe_record(
2255 crate::gpu::OpClass::Matmat,
2256 false,
2257 t0.elapsed(),
2258 );
2259 return;
2260 }
2261 crate::gpu::ProbeArm::Cpu => {}
2262 }
2263 }
2264 }
2265 q1t_matmat(self.quant_bytes(), xs_all, b, rows, cols, out, pool);
2266 return;
2267 }
2268 if matches!(dtype, TensorDtype::Vbit | TensorDtype::VbitRo) {
2269 vbitmatmat(
2270 self.quant_bytes(),
2271 vbit_offsets,
2272 xs_all,
2273 b,
2274 rows,
2275 cols,
2276 out,
2277 pool,
2278 );
2279 return;
2280 }
2281 let pre: Vec<std::borrow::Cow<'_, [f32]>> = (0..b)
2282 .map(|bi| prescale(&xs_all[bi * cols..(bi + 1) * cols], col_field, *dtype))
2283 .collect();
2284 if row_exact()
2290 && (1..=4).contains(&b)
2291 && matches!(dtype, TensorDtype::Q8Row | TensorDtype::Q8_2f)
2292 && crate::gpu::enabled_here()
2293 && crate::gpu::wgpu_active()
2294 {
2295 let flat: Vec<f32> = pre.iter().flat_map(|v| v.iter().copied()).collect();
2296 if crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out) {
2297 return;
2298 }
2299 }
2300 if b >= 8 && b * rows * cols >= 128_000_000 && crate::gpu::enabled_here() {
2306 if let Self::Mapped { model, idx, .. } = self {
2307 let t0 = std::time::Instant::now();
2308 match crate::gpu::probe_arm(crate::gpu::OpClass::Matmat) {
2309 crate::gpu::ProbeArm::Gpu
2310 if crate::gpu::probe_deciding(crate::gpu::OpClass::Matmat)
2311 && !crate::gpu::q8_resident_or_upload(model, *idx) =>
2312 {
2313 let q = self.quant_bytes();
2317 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2318 return;
2319 }
2320 crate::gpu::ProbeArm::Gpu => {
2321 let flat: Vec<f32> =
2322 pre.iter().flat_map(|v| v.iter().copied()).collect();
2323 if crate::gpu::q8_matmat(
2324 model, *idx, row_scale, &flat, b, rows, cols, out,
2325 ) {
2326 crate::gpu::probe_record(
2327 crate::gpu::OpClass::Matmat,
2328 true,
2329 t0.elapsed(),
2330 );
2331 return;
2332 }
2333 }
2334 crate::gpu::ProbeArm::CpuTimed => {
2335 let q = self.quant_bytes();
2336 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2337 crate::gpu::probe_record(
2338 crate::gpu::OpClass::Matmat,
2339 false,
2340 t0.elapsed(),
2341 );
2342 return;
2343 }
2344 crate::gpu::ProbeArm::Cpu => {}
2345 }
2346 }
2347 }
2348 let q = self.quant_bytes();
2349 qmatmat(q, row_scale, &pre, rows, cols, out, pool);
2350 }
2351 }
2352 }
2353}
2354
2355impl QTensor {
2356 pub fn device_matmat(&self, xs: &[f32], b: usize, out: &mut [f32]) -> bool {
2366 let (rows, cols) = (self.rows(), self.cols());
2367 let Self::Mapped {
2368 model,
2369 idx,
2370 dtype,
2371 row_scale,
2372 col_field,
2373 ..
2374 } = self
2375 else {
2376 return false;
2377 };
2378 if crate::prism::has_contract(model) {
2379 return false;
2380 }
2381 match *dtype {
2382 TensorDtype::Q4TiledP => crate::gpu::q4tp_matmat(model, *idx, xs, b, rows, cols, out),
2383 TensorDtype::Q8Row | TensorDtype::Q8_2f => {
2387 if *dtype == TensorDtype::Q8_2f
2390 && std::env::var("CMF_Q8_2F_DEV").as_deref() != Ok("0")
2391 && crate::gpu::q8_matmat_2f(
2392 model, *idx, row_scale, col_field, xs, b, rows, cols, out,
2393 )
2394 {
2395 return true;
2396 }
2397 let flat: Vec<f32> = (0..b)
2398 .flat_map(|bi| {
2399 prescale(&xs[bi * cols..(bi + 1) * cols], col_field, *dtype).into_owned()
2400 })
2401 .collect();
2402 crate::gpu::q8_matmat(model, *idx, row_scale, &flat, b, rows, cols, out)
2403 }
2404 _ => false,
2405 }
2406 }
2407
2408 pub fn matvec_many<const N: usize>(
2415 ts: [&QTensor; N],
2416 x: &[f32],
2417 mut outs: [&mut [f32]; N],
2418 pool: Option<&Pool>,
2419 ) {
2420 let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2421 if ts.iter().any(|t| t.has_prism_contract()) {
2422 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2426 t.matvec(x, o, pool);
2427 }
2428 return;
2429 }
2430 let uniform_q8 = ts.iter().all(|t| {
2431 matches!(
2432 t,
2433 Self::Mapped {
2434 dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2435 ..
2436 }
2437 )
2438 });
2439 let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2440 if uniform_f32 && crate::f32_backend::active() {
2441 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2442 t.matvec(x, o, pool);
2443 }
2444 return;
2445 }
2446 let uniform_q4 = ts.iter().all(|t| {
2447 matches!(
2448 t,
2449 Self::Mapped {
2450 dtype: TensorDtype::Q4Block,
2451 ..
2452 }
2453 )
2454 });
2455 let uniform_vbit = ts.iter().all(|t| {
2456 matches!(
2457 t,
2458 Self::Mapped {
2459 dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2460 ..
2461 }
2462 )
2463 });
2464 let uniform_q1 = ts.iter().all(|t| {
2465 matches!(
2466 t,
2467 Self::Mapped {
2468 dtype: TensorDtype::Q1,
2469 ..
2470 }
2471 )
2472 });
2473 let uniform_q1t = ts.iter().all(|t| {
2474 matches!(
2475 t,
2476 Self::Mapped {
2477 dtype: TensorDtype::Q1T,
2478 ..
2479 }
2480 )
2481 });
2482 let uniform_q4tp = ts.iter().all(|t| {
2487 matches!(
2488 t,
2489 Self::Mapped {
2490 dtype: TensorDtype::Q4TiledP,
2491 ..
2492 }
2493 )
2494 }) && ts
2495 .iter()
2496 .all(|t| t.cols() == ts[0].cols() && t.cols() % GROUP_SIZE == 0);
2497 let Some(pool) = pool else {
2498 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2499 t.matvec(x, o, None);
2500 }
2501 return;
2502 };
2503 if total_rows < 256
2504 || !(uniform_q8
2505 || uniform_f32
2506 || uniform_q4
2507 || uniform_vbit
2508 || uniform_q1
2509 || uniform_q1t
2510 || uniform_q4tp)
2511 {
2512 for (t, o) in ts.iter().zip(outs.iter_mut()) {
2513 t.matvec(x, o, Some(pool));
2514 }
2515 return;
2516 }
2517
2518 if uniform_q4tp {
2519 let cols = ts[0].cols();
2525 let gpr = cols / GROUP_SIZE;
2526 let views: [Q4tpView; N] =
2527 std::array::from_fn(|i| Q4tpView::new(ts[i].quant_bytes(), ts[i].rows(), cols));
2528 let rows_of: [usize; N] = std::array::from_fn(|i| ts[i].rows());
2529 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2530 let locate = |flat: usize| -> (usize, usize) {
2532 let mut acc = 0;
2533 for (i, &r) in rows_of.iter().enumerate() {
2534 if flat < acc + r {
2535 return (i, flat - acc);
2536 }
2537 acc += r;
2538 }
2539 (rows_of.len() - 1, 0)
2540 };
2541 let (views, outs_addr) = (&views, &outs_addr);
2542 if a8w8_enabled() {
2543 let act = split_act(x);
2544 let act = &act;
2545 let run = |start: usize, end: usize| {
2546 with_krow(gpr, |sc| {
2547 for flat in start..end {
2548 let (t, r) = locate(flat);
2549 let v = &views[t];
2550 v.scales_into(r, gpr, sc);
2551 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
2552 for &(j, xv) in &act.outliers {
2553 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
2554 acc += w * s * xv;
2555 }
2556 unsafe { *outs_addr[t].at(r) = acc };
2558 }
2559 });
2560 };
2561 pool.run_rows(total_rows, &run);
2562 } else {
2563 let run = |start: usize, end: usize| {
2564 with_krow(gpr, |sc| {
2565 for flat in start..end {
2566 let (t, r) = locate(flat);
2567 let v = &views[t];
2568 v.scales_into(r, gpr, sc);
2569 unsafe { *outs_addr[t].at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
2571 }
2572 });
2573 };
2574 pool.run_rows(total_rows, &run);
2575 }
2576 return;
2577 }
2578
2579 if uniform_q1 {
2580 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2583 if a8w8_enabled() {
2584 let act = split_act(x);
2585 let gsum = q1_group_sums(&act.xq, ts[0].cols() / GROUP_SIZE);
2586 let (act, gsum) = (&act, &gsum);
2587 let closures: [_; N] = std::array::from_fn(|i| {
2588 let (bytes, gpr, out) =
2589 (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2590 move |s: usize, e: usize| q1_range_a8w8(bytes, gpr, act, gsum, out, s, e)
2591 });
2592 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2593 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2594 pool.run_many(&parts);
2595 } else {
2596 let closures: [_; N] = std::array::from_fn(|i| {
2597 let (bytes, gpr, out) =
2598 (ts[i].quant_bytes(), ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2599 move |s: usize, e: usize| q1_range_f32(bytes, gpr, x, out, s, e)
2600 });
2601 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2602 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2603 pool.run_many(&parts);
2604 }
2605 return;
2606 }
2607
2608 if uniform_q1t {
2609 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2613 const TILE: usize = cortiq_core::quant::Q1T_TILE;
2614 if a8w8_enabled() {
2615 let act = split_act(x);
2616 let act = &act;
2617 let x_ref = x;
2618 let closures: [_; N] = std::array::from_fn(|i| {
2619 let bytes = ts[i].quant_bytes();
2620 let (rows, cols) = (ts[i].rows(), ts[i].cols());
2621 let gpr = cols / GROUP_SIZE;
2622 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2623 let out = outs_addr[i];
2624 move |s: usize, e: usize| {
2625 q1t_range_a8w8(bytes, gpr, rp_off, ent_off, has_ov, act, x_ref, out, s, e)
2626 }
2627 });
2628 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2629 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2630 pool.run_many(&parts);
2631 } else {
2632 let x_ref = x;
2633 let closures: [_; N] = std::array::from_fn(|i| {
2634 let bytes = ts[i].quant_bytes();
2635 let (rows, cols) = (ts[i].rows(), ts[i].cols());
2636 let gpr = cols / GROUP_SIZE;
2637 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
2638 let out = outs_addr[i];
2639 move |s: usize, e: usize| {
2640 q1t_range_f32_batch(bytes, gpr, rp_off, ent_off, has_ov, x_ref, out, s, e)
2641 }
2642 });
2643 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2644 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2645 pool.run_many(&parts);
2646 }
2647 return;
2648 }
2649
2650 if uniform_q4 || uniform_vbit {
2651 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2652 if a8w8_enabled() {
2654 let act = split_act(x);
2655 let act = &act;
2656 if uniform_q4 {
2657 let closures: [_; N] = std::array::from_fn(|i| {
2658 let (packed, scales) =
2659 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2660 let (gpr, cols, out) =
2661 (ts[i].cols() / GROUP_SIZE, ts[i].cols(), outs_addr[i]);
2662 move |s: usize, e: usize| {
2663 q4_range_a8w8(packed, scales, gpr, cols, act, out, s, e)
2664 }
2665 });
2666 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2667 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2668 pool.run_many(&parts);
2669 } else {
2670 let closures: [_; N] = std::array::from_fn(|i| {
2671 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2672 unreachable!()
2673 };
2674 let (bytes, rows, cols, out) = (
2675 ts[i].quant_bytes(),
2676 ts[i].rows(),
2677 ts[i].cols(),
2678 outs_addr[i],
2679 );
2680 move |s: usize, e: usize| {
2681 vbit_range_a8w8(bytes, vbit_offsets, x, act, rows, cols, out, s, e)
2682 }
2683 });
2684 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2685 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2686 pool.run_many(&parts);
2687 }
2688 return;
2689 }
2690 if uniform_q4 {
2691 let closures: [_; N] = std::array::from_fn(|i| {
2692 let (packed, scales) =
2693 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2694 let (gpr, out) = (ts[i].cols() / GROUP_SIZE, outs_addr[i]);
2695 move |s: usize, e: usize| q4_range_f32(packed, scales, gpr, x, out, s, e)
2696 });
2697 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2698 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2699 pool.run_many(&parts);
2700 } else {
2701 let closures: [_; N] = std::array::from_fn(|i| {
2702 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2703 unreachable!()
2704 };
2705 let (bytes, rows, cols, out) = (
2706 ts[i].quant_bytes(),
2707 ts[i].rows(),
2708 ts[i].cols(),
2709 outs_addr[i],
2710 );
2711 move |s: usize, e: usize| {
2712 vbit_range_f32(bytes, vbit_offsets, x, rows, cols, out, s, e)
2713 }
2714 });
2715 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2716 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2717 pool.run_many(&parts);
2718 }
2719 return;
2720 }
2721
2722 if uniform_f32 {
2723 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2724 let closures: [_; N] = std::array::from_fn(|i| {
2725 let Self::F32 { data, cols, .. } = ts[i] else {
2726 unreachable!()
2727 };
2728 let out = outs_addr[i];
2729 move |start: usize, end: usize| {
2730 for o in start..end {
2731 let row = &data[o * cols..(o + 1) * cols];
2732 let mut sum = 0.0f32;
2733 for j in 0..*cols {
2734 sum += row[j] * x[j];
2735 }
2736 unsafe { *out.at(o) = sum };
2738 }
2739 }
2740 });
2741 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2742 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2743 pool.run_many(&parts);
2744 return;
2745 }
2746
2747 struct Ctx<'a> {
2750 bytes: &'a [u8],
2751 #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
2752 rep: &'a [u8],
2753 row_scale: &'a [f32],
2754 cols: usize,
2755 xs: std::borrow::Cow<'a, [f32]>,
2756 }
2757 let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
2758 let Self::Mapped {
2759 dtype,
2760 cols,
2761 row_scale,
2762 col_field,
2763 repack,
2764 ..
2765 } = ts[i]
2766 else {
2767 unreachable!()
2768 };
2769 Ctx {
2770 bytes: ts[i].quant_bytes(),
2771 rep: repack,
2772 row_scale,
2773 cols: *cols,
2774 xs: prescale(x, col_field, *dtype),
2775 }
2776 });
2777 let outs_addr: [SendMut; N] = std::array::from_fn(|i| SendMut(outs[i].as_mut_ptr()));
2778 #[cfg(target_arch = "aarch64")]
2779 if sdot_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_sdot(c.bytes, c.rep, 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 #[cfg(target_arch = "x86_64")]
2793 if avx2_a8w8_enabled() {
2794 let acts: [SplitAct; N] = std::array::from_fn(|i| split_act(&ctxs[i].xs));
2795 let closures: [_; N] = std::array::from_fn(|i| {
2796 let (c, act, out) = (&ctxs[i], &acts[i], outs_addr[i]);
2797 move |start: usize, end: usize| {
2798 q8_range_avx2(c.bytes, c.row_scale, act, c.cols, out, start, end)
2799 }
2800 });
2801 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2802 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2803 pool.run_many(&parts);
2804 return;
2805 }
2806 let closures: [_; N] = std::array::from_fn(|i| {
2807 let (c, out) = (&ctxs[i], outs_addr[i]);
2808 move |start: usize, end: usize| {
2809 q8_range_f32(c.bytes, c.row_scale, &c.xs, c.cols, out, start, end)
2810 }
2811 });
2812 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2813 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2814 pool.run_many(&parts);
2815 }
2816}
2817
2818impl QTensor {
2819 #[allow(clippy::needless_range_loop)]
2824 pub fn matvec2_many<const N: usize>(
2825 ts: [&QTensor; N],
2826 x1: &[f32],
2827 x2: &[f32],
2828 mut o1s: [&mut [f32]; N],
2829 mut o2s: [&mut [f32]; N],
2830 pool: Option<&Pool>,
2831 ) {
2832 let total_rows: usize = ts.iter().map(|t| t.rows()).sum();
2833 if ts.iter().any(|t| t.has_prism_contract()) {
2834 for i in 0..N {
2835 ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2836 }
2837 return;
2838 }
2839 let uniform_q8 = ts.iter().all(|t| {
2840 matches!(
2841 t,
2842 Self::Mapped {
2843 dtype: TensorDtype::Q8Row | TensorDtype::Q8_2f,
2844 ..
2845 }
2846 )
2847 });
2848 let uniform_f32 = ts.iter().all(|t| matches!(t, Self::F32 { .. }));
2849 let uniform_q4 = ts.iter().all(|t| {
2850 matches!(
2851 t,
2852 Self::Mapped {
2853 dtype: TensorDtype::Q4Block,
2854 ..
2855 }
2856 )
2857 });
2858 let uniform_vbit = ts.iter().all(|t| {
2859 matches!(
2860 t,
2861 Self::Mapped {
2862 dtype: TensorDtype::Vbit | TensorDtype::VbitRo,
2863 ..
2864 }
2865 )
2866 });
2867 let fusable = pool.is_some()
2868 && total_rows >= 256
2869 && (uniform_q8 || uniform_f32 || uniform_q4 || uniform_vbit);
2870 if !fusable {
2871 for i in 0..N {
2872 ts[i].matvec2(x1, x2, o1s[i], o2s[i], pool);
2873 }
2874 return;
2875 }
2876 let pool = pool.unwrap();
2877
2878 if uniform_q4 || uniform_vbit {
2879 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2880 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2881 if a8w8_enabled() {
2883 let a1 = split_act(x1);
2884 let a2 = split_act(x2);
2885 let (a1, a2) = (&a1, &a2);
2886 if uniform_q4 {
2887 let closures: [_; N] = std::array::from_fn(|i| {
2888 let (packed, scales) =
2889 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2890 let (gpr, cols, o1, o2) =
2891 (ts[i].cols() / GROUP_SIZE, ts[i].cols(), p1[i], p2[i]);
2892 move |s: usize, e: usize| {
2893 q4_range2_a8w8(packed, scales, gpr, cols, a1, a2, o1, o2, s, e)
2894 }
2895 });
2896 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2897 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2898 pool.run_many(&parts);
2899 } else {
2900 let closures: [_; N] = std::array::from_fn(|i| {
2901 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2902 unreachable!()
2903 };
2904 let (bytes, rows, cols, o1, o2) = (
2905 ts[i].quant_bytes(),
2906 ts[i].rows(),
2907 ts[i].cols(),
2908 p1[i],
2909 p2[i],
2910 );
2911 move |s: usize, e: usize| {
2912 vbit_range2_a8w8(
2913 bytes,
2914 vbit_offsets,
2915 x1,
2916 x2,
2917 a1,
2918 a2,
2919 rows,
2920 cols,
2921 o1,
2922 o2,
2923 s,
2924 e,
2925 )
2926 }
2927 });
2928 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2929 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2930 pool.run_many(&parts);
2931 }
2932 return;
2933 }
2934 if uniform_q4 {
2935 let closures: [_; N] = std::array::from_fn(|i| {
2936 let (packed, scales) =
2937 q4_split(ts[i].quant_bytes(), ts[i].rows(), ts[i].cols());
2938 let (gpr, o1, o2) = (ts[i].cols() / GROUP_SIZE, p1[i], p2[i]);
2939 move |s: usize, e: usize| {
2940 q4_range2_f32(packed, scales, gpr, x1, x2, o1, o2, s, e)
2941 }
2942 });
2943 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2944 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2945 pool.run_many(&parts);
2946 } else {
2947 let closures: [_; N] = std::array::from_fn(|i| {
2948 let Self::Mapped { vbit_offsets, .. } = ts[i] else {
2949 unreachable!()
2950 };
2951 let (bytes, rows, cols, o1, o2) = (
2952 ts[i].quant_bytes(),
2953 ts[i].rows(),
2954 ts[i].cols(),
2955 p1[i],
2956 p2[i],
2957 );
2958 move |s: usize, e: usize| {
2959 vbit_range2_f32(bytes, vbit_offsets, x1, x2, rows, cols, o1, o2, s, e)
2960 }
2961 });
2962 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2963 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2964 pool.run_many(&parts);
2965 }
2966 return;
2967 }
2968
2969 if uniform_f32 {
2970 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
2971 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
2972 let closures: [_; N] = std::array::from_fn(|i| {
2973 let Self::F32 { data, cols, .. } = ts[i] else {
2974 unreachable!()
2975 };
2976 let (o1, o2) = (p1[i], p2[i]);
2977 move |start: usize, end: usize| {
2978 for o in start..end {
2979 let row = &data[o * cols..(o + 1) * cols];
2980 let (mut s1, mut s2) = (0.0f32, 0.0f32);
2981 for j in 0..*cols {
2982 s1 += row[j] * x1[j];
2983 s2 += row[j] * x2[j];
2984 }
2985 unsafe {
2987 *o1.at(o) = s1;
2988 *o2.at(o) = s2;
2989 }
2990 }
2991 }
2992 });
2993 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
2994 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
2995 pool.run_many(&parts);
2996 return;
2997 }
2998
2999 struct Ctx<'a> {
3000 bytes: &'a [u8],
3001 row_scale: &'a [f32],
3002 cols: usize,
3003 xs1: std::borrow::Cow<'a, [f32]>,
3004 xs2: std::borrow::Cow<'a, [f32]>,
3005 }
3006 let ctxs: [Ctx<'_>; N] = std::array::from_fn(|i| {
3007 let Self::Mapped {
3008 dtype,
3009 cols,
3010 row_scale,
3011 col_field,
3012 ..
3013 } = ts[i]
3014 else {
3015 unreachable!()
3016 };
3017 Ctx {
3018 bytes: ts[i].quant_bytes(),
3019 row_scale,
3020 cols: *cols,
3021 xs1: prescale(x1, col_field, *dtype),
3022 xs2: prescale(x2, col_field, *dtype),
3023 }
3024 });
3025 let p1: [SendMut; N] = std::array::from_fn(|i| SendMut(o1s[i].as_mut_ptr()));
3026 let p2: [SendMut; N] = std::array::from_fn(|i| SendMut(o2s[i].as_mut_ptr()));
3027 #[cfg(target_arch = "aarch64")]
3028 if sdot_enabled() {
3029 let acts: [(SplitAct, SplitAct); N] =
3030 std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3031 let closures: [_; N] = std::array::from_fn(|i| {
3032 let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3033 move |start: usize, end: usize| {
3034 q8_range2_sdot(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3035 }
3036 });
3037 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3038 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3039 pool.run_many(&parts);
3040 return;
3041 }
3042 #[cfg(target_arch = "x86_64")]
3043 if avx2_a8w8_enabled() {
3044 let acts: [(SplitAct, SplitAct); N] =
3045 std::array::from_fn(|i| (split_act(&ctxs[i].xs1), split_act(&ctxs[i].xs2)));
3046 let closures: [_; N] = std::array::from_fn(|i| {
3047 let (c, a, o1, o2) = (&ctxs[i], &acts[i], p1[i], p2[i]);
3048 move |start: usize, end: usize| {
3049 q8_range2_avx2(c.bytes, c.row_scale, &a.0, &a.1, c.cols, o1, o2, start, end)
3050 }
3051 });
3052 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3053 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3054 pool.run_many(&parts);
3055 return;
3056 }
3057 let closures: [_; N] = std::array::from_fn(|i| {
3058 let (c, o1, o2) = (&ctxs[i], p1[i], p2[i]);
3059 move |start: usize, end: usize| {
3060 q8_range2_f32(
3061 c.bytes,
3062 c.row_scale,
3063 &c.xs1,
3064 &c.xs2,
3065 c.cols,
3066 o1,
3067 o2,
3068 start,
3069 end,
3070 )
3071 }
3072 });
3073 let parts: [(usize, &(dyn Fn(usize, usize) + Sync)); N] =
3074 std::array::from_fn(|i| (ts[i].rows(), &closures[i] as _));
3075 pool.run_many(&parts);
3076 }
3077
3078 pub fn matvec_silu_mul(
3083 gate: &QTensor,
3084 up: &QTensor,
3085 x: &[f32],
3086 out: &mut [f32],
3087 pool: Option<&Pool>,
3088 ) -> bool {
3089 Self::matvec_silu_mul_limited(gate, up, x, out, 0.0, pool)
3090 }
3091
3092 pub fn matvec_silu_mul_limited(
3098 gate: &QTensor,
3099 up: &QTensor,
3100 x: &[f32],
3101 out: &mut [f32],
3102 limit: f32,
3103 pool: Option<&Pool>,
3104 ) -> bool {
3105 if gate.has_prism_contract() || up.has_prism_contract() {
3106 return false;
3110 }
3111 let inter = gate.rows();
3112 debug_assert_eq!(up.rows(), inter);
3113 debug_assert_eq!(out.len(), inter);
3114 debug_assert_eq!(gate.cols(), up.cols());
3115 if !a8w8_enabled() {
3116 return false;
3117 }
3118 let act = split_act(x);
3119 let act = &act;
3120 let x_ref = x;
3121 let out_addr = SendMut(out.as_mut_ptr());
3122
3123 match (gate, up) {
3124 (
3126 Self::Mapped {
3127 dtype: TensorDtype::Q4Block,
3128 ..
3129 },
3130 Self::Mapped {
3131 dtype: TensorDtype::Q4Block,
3132 ..
3133 },
3134 ) => {
3135 let (gp, gs) = q4_split(gate.quant_bytes(), gate.rows(), gate.cols());
3136 let (up_p, up_s) = q4_split(up.quant_bytes(), up.rows(), up.cols());
3137 let gpr = gate.cols() / GROUP_SIZE;
3138 let cols = gate.cols();
3139 let run = move |start: usize, end: usize| {
3140 for r in start..end {
3141 let mut gv = dot_q4_row_i8(gp, gs, r * gpr, gpr, &act.xq) * act.sx;
3142 let mut uv = dot_q4_row_i8(up_p, up_s, r * gpr, gpr, &act.xq) * act.sx;
3143 for &(j, xv) in &act.outliers {
3144 let flat = r * cols + j;
3145 let gb = gp[flat / 2];
3146 let gn = if flat & 1 == 0 { gb & 0x0F } else { gb >> 4 };
3147 let gsc = f16_to_f32(u16::from_le_bytes([
3148 gs[(flat / GROUP_SIZE) * 2],
3149 gs[(flat / GROUP_SIZE) * 2 + 1],
3150 ]));
3151 gv += ((gn as i32 - 8) as f32) * gsc * xv;
3152 let ub = up_p[flat / 2];
3153 let un = if flat & 1 == 0 { ub & 0x0F } else { ub >> 4 };
3154 let usc = f16_to_f32(u16::from_le_bytes([
3155 up_s[(flat / GROUP_SIZE) * 2],
3156 up_s[(flat / GROUP_SIZE) * 2 + 1],
3157 ]));
3158 uv += ((un as i32 - 8) as f32) * usc * xv;
3159 }
3160 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3162 }
3163 };
3164 dispatch_rows(pool, inter, &run);
3165 true
3166 }
3167 (
3171 Self::Mapped {
3172 dtype: TensorDtype::Q4Tiled,
3173 ..
3174 },
3175 Self::Mapped {
3176 dtype: TensorDtype::Q4Tiled,
3177 ..
3178 },
3179 ) => {
3180 let g_bytes = gate.quant_bytes();
3181 let u_bytes = up.quant_bytes();
3182 let gpr = gate.cols() / GROUP_SIZE;
3183 let run = move |start: usize, end: usize| {
3184 for r in start..end {
3185 let mut gv = dot_q4t_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3186 let mut uv = dot_q4t_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3187 for &(j, xv) in &act.outliers {
3188 let (w, s) = q4t_outlier(g_bytes, r, gpr, j);
3189 gv += w * s * xv;
3190 let (w, s) = q4t_outlier(u_bytes, r, gpr, j);
3191 uv += w * s * xv;
3192 }
3193 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3195 }
3196 };
3197 dispatch_rows(pool, inter, &run);
3198 true
3199 }
3200 (
3203 Self::Mapped {
3204 dtype: TensorDtype::Q4TiledP,
3205 ..
3206 },
3207 Self::Mapped {
3208 dtype: TensorDtype::Q4TiledP,
3209 ..
3210 },
3211 ) => {
3212 let cols = gate.cols();
3213 let gpr = cols / GROUP_SIZE;
3214 let gv_view = Q4tpView::new(gate.quant_bytes(), inter, cols);
3215 let uv_view = Q4tpView::new(up.quant_bytes(), inter, cols);
3216 let run = |start: usize, end: usize| {
3217 with_krows(gpr, |gsc, usc| {
3218 for r in start..end {
3219 gv_view.scales_into(r, gpr, gsc);
3220 uv_view.scales_into(r, gpr, usc);
3221 let mut gv =
3222 dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx;
3223 let mut uv =
3224 dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx;
3225 for &(j, xv) in &act.outliers {
3226 let (w, s) = q4tp_outlier(gv_view.nib, r, gpr, j, gsc);
3227 gv += w * s * xv;
3228 let (w, s) = q4tp_outlier(uv_view.nib, r, gpr, j, usc);
3229 uv += w * s * xv;
3230 }
3231 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3233 }
3234 });
3235 };
3236 dispatch_rows(pool, inter, &run);
3237 true
3238 }
3239 (
3245 Self::Mapped {
3246 dtype: TensorDtype::Q1,
3247 ..
3248 },
3249 Self::Mapped {
3250 dtype: TensorDtype::Q1,
3251 ..
3252 },
3253 ) => {
3254 let g_bytes = gate.quant_bytes();
3255 let u_bytes = up.quant_bytes();
3256 let gpr = gate.cols() / GROUP_SIZE;
3257 let gsum = q1_group_sums(&act.xq, gpr);
3258 let gsum = &gsum;
3259 let run = move |start: usize, end: usize| {
3260 for r in start..end {
3261 let mut gv = dot_q1_row_i8(g_bytes, r, gpr, &act.xq, gsum) * act.sx;
3262 let mut uv = dot_q1_row_i8(u_bytes, r, gpr, &act.xq, gsum) * act.sx;
3263 for &(j, xv) in &act.outliers {
3264 let (w, s) = q1_outlier(g_bytes, r, gpr, j);
3265 gv += w * s * xv;
3266 let (w, s) = q1_outlier(u_bytes, r, gpr, j);
3267 uv += w * s * xv;
3268 }
3269 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3271 }
3272 };
3273 dispatch_rows(pool, inter, &run);
3274 true
3275 }
3276 (
3280 Self::Mapped {
3281 dtype: TensorDtype::Q2TiledP,
3282 ..
3283 },
3284 Self::Mapped {
3285 dtype: TensorDtype::Q2TiledP,
3286 ..
3287 },
3288 ) => {
3289 let cols = gate.cols();
3290 let gpr = cols / GROUP_SIZE;
3291 let gv_view = Q4tpView::new_q2(gate.quant_bytes(), inter, cols);
3292 let uv_view = Q4tpView::new_q2(up.quant_bytes(), inter, cols);
3293 let gsum = q1_group_sums(&act.xq, gpr);
3294 let gsum = &gsum;
3295 let run = move |start: usize, end: usize| {
3296 with_krows(gpr, |gsc, usc| {
3297 for r in start..end {
3298 gv_view.scales_into(r, gpr, gsc);
3299 uv_view.scales_into(r, gpr, usc);
3300 let mut gv =
3301 dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx;
3302 let mut uv =
3303 dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx;
3304 for &(j, xv) in &act.outliers {
3305 let (w, s) = q2tp_outlier(gv_view.nib, r, gpr, j, gsc);
3306 gv += w * s * xv;
3307 let (w, s) = q2tp_outlier(uv_view.nib, r, gpr, j, usc);
3308 uv += w * s * xv;
3309 }
3310 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3312 }
3313 });
3314 };
3315 dispatch_rows(pool, inter, &run);
3316 true
3317 }
3318 (
3323 Self::Mapped {
3324 dtype: TensorDtype::Q8Row,
3325 row_scale: g_rs,
3326 ..
3327 },
3328 Self::Mapped {
3329 dtype: TensorDtype::Q8Row,
3330 row_scale: u_rs,
3331 ..
3332 },
3333 ) => {
3334 let g_bytes = gate.quant_bytes();
3335 let u_bytes = up.quant_bytes();
3336 let cols = gate.cols();
3337 let run = move |start: usize, end: usize| {
3338 for r in start..end {
3339 let gv = q8_row_dot(&g_bytes[r * cols..(r + 1) * cols], act) * g_rs[r];
3340 let uv = q8_row_dot(&u_bytes[r * cols..(r + 1) * cols], act) * u_rs[r];
3341 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3343 }
3344 };
3345 dispatch_rows(pool, inter, &run);
3346 true
3347 }
3348 (
3350 Self::Mapped {
3351 dtype: TensorDtype::Q1T,
3352 ..
3353 },
3354 Self::Mapped {
3355 dtype: TensorDtype::Q1T,
3356 ..
3357 },
3358 ) => {
3359 const TILE: usize = cortiq_core::quant::Q1T_TILE;
3360 let g_bytes = gate.quant_bytes();
3361 let u_bytes = up.quant_bytes();
3362 let gpr = gate.cols() / GROUP_SIZE;
3363 let (g_rp, g_ent, g_ov) = q1t_overlay(g_bytes, inter * gpr * TILE, inter);
3364 let (u_rp, u_ent, u_ov) = q1t_overlay(u_bytes, inter * gpr * TILE, inter);
3365 let run = move |start: usize, end: usize| {
3366 for r in start..end {
3367 let mut gv = q1t_dot_row_i8(g_bytes, r, gpr, &act.xq) * act.sx;
3368 let mut uv = q1t_dot_row_i8(u_bytes, r, gpr, &act.xq) * act.sx;
3369 for &(j, xv) in &act.outliers {
3370 gv += q1t_base_weight(g_bytes, r, gpr, j) * xv;
3371 uv += q1t_base_weight(u_bytes, r, gpr, j) * xv;
3372 }
3373 gv += q1t_row_outlier_correction(g_bytes, r, g_rp, g_ent, g_ov, x_ref);
3374 uv += q1t_row_outlier_correction(u_bytes, r, u_rp, u_ent, u_ov, x_ref);
3375 unsafe { *out_addr.at(r) = silu_mul_limited(gv, uv, limit) };
3377 }
3378 };
3379 dispatch_rows(pool, inter, &run);
3380 true
3381 }
3382 _ => false,
3383 }
3384 }
3385
3386 pub fn moe_gate_up_many(
3401 pairs: &[(&QTensor, &QTensor)],
3402 x: &[f32],
3403 outs: &mut [Vec<f32>],
3404 pool: Option<&Pool>,
3405 ) -> bool {
3406 Self::moe_gate_up_many_limited(pairs, x, outs, 0.0, pool)
3407 }
3408
3409 pub fn moe_gate_up_many_limited(
3413 pairs: &[(&QTensor, &QTensor)],
3414 x: &[f32],
3415 outs: &mut [Vec<f32>],
3416 limit: f32,
3417 pool: Option<&Pool>,
3418 ) -> bool {
3419 if pairs.is_empty() || pairs.len() != outs.len() {
3420 return false;
3421 }
3422 if !a8w8_enabled() {
3423 if limit > 0.0 {
3424 return false;
3427 }
3428 let groups = vec![vec![0]; pairs.len()];
3429 return Self::moe_gate_up_rows(pairs, &groups, x, outs, pool);
3430 }
3431 let inter = pairs[0].0.rows();
3432 let cols = pairs[0].0.cols();
3433 if cols % GROUP_SIZE != 0 {
3434 return false;
3435 }
3436 let gpr = cols / GROUP_SIZE;
3437 let q2 = matches!(
3440 pairs[0].0,
3441 Self::Mapped {
3442 dtype: TensorDtype::Q2TiledP,
3443 ..
3444 }
3445 );
3446 let want = if q2 {
3447 TensorDtype::Q2TiledP
3448 } else {
3449 TensorDtype::Q4TiledP
3450 };
3451 let mut views = Vec::with_capacity(pairs.len() * 2);
3452 for ((g, u), o) in pairs.iter().zip(outs.iter()) {
3453 let both = matches!(g, Self::Mapped { dtype, .. } if *dtype == want)
3454 && matches!(u, Self::Mapped { dtype, .. } if *dtype == want);
3455 if !both
3456 || g.rows() != inter
3457 || u.rows() != inter
3458 || g.cols() != cols
3459 || u.cols() != cols
3460 || o.len() != inter
3461 {
3462 return false;
3463 }
3464 let mk = if q2 { Q4tpView::new_q2 } else { Q4tpView::new };
3465 views.push(mk(g.quant_bytes(), inter, cols));
3466 views.push(mk(u.quant_bytes(), inter, cols));
3467 }
3468 let act = split_act(x);
3469 let gsum = if q2 {
3470 q1_group_sums(&act.xq, gpr)
3471 } else {
3472 Vec::new()
3473 };
3474 let (act, gsum) = (&act, &gsum);
3475 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3476 let (views, ptrs) = (&views, &ptrs);
3477 let run = |start: usize, end: usize| {
3478 with_krows(gpr, |gsc, usc| {
3479 for flat in start..end {
3480 let (e, r) = (flat / inter, flat % inter);
3481 let gv_view = &views[e * 2];
3482 let uv_view = &views[e * 2 + 1];
3483 gv_view.scales_into(r, gpr, gsc);
3484 uv_view.scales_into(r, gpr, usc);
3485 let (mut gv, mut uv) = if q2 {
3486 (
3487 dot_q2tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsum, gsc) * act.sx,
3488 dot_q2tp_row_i8(uv_view.nib, r, gpr, &act.xq, gsum, usc) * act.sx,
3489 )
3490 } else {
3491 (
3492 dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, gsc) * act.sx,
3493 dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, usc) * act.sx,
3494 )
3495 };
3496 for &(j, xv) in &act.outliers {
3497 let (og, ou) = if q2 {
3498 (
3499 q2tp_outlier(gv_view.nib, r, gpr, j, gsc),
3500 q2tp_outlier(uv_view.nib, r, gpr, j, usc),
3501 )
3502 } else {
3503 (
3504 q4tp_outlier(gv_view.nib, r, gpr, j, gsc),
3505 q4tp_outlier(uv_view.nib, r, gpr, j, usc),
3506 )
3507 };
3508 gv += og.0 * og.1 * xv;
3509 uv += ou.0 * ou.1 * xv;
3510 }
3511 unsafe { *ptrs[e].at(r) = silu_mul_limited(gv, uv, limit) };
3513 }
3514 });
3515 };
3516 dispatch_rows(pool, pairs.len() * inter, &run);
3517 true
3518 }
3519
3520 pub fn moe_down_many(
3529 downs: &[&QTensor],
3530 gs: &[Vec<f32>],
3531 weights: &[f32],
3532 out: &mut [f32],
3533 pool: Option<&Pool>,
3534 ) -> bool {
3535 if downs.is_empty() || downs.len() != gs.len() || downs.len() != weights.len() {
3536 return false;
3537 }
3538 if !a8w8_enabled() {
3539 let mut terms = vec![vec![0.0; out.len()]; downs.len()];
3540 if !Self::moe_down_rows(downs, &vec![1; downs.len()], gs, &mut terms, pool) {
3541 return false;
3542 }
3543 out.fill(0.0);
3544 for (row, &w) in terms.iter().zip(weights) {
3545 for (o, &v) in out.iter_mut().zip(row) {
3546 *o += w * v;
3547 }
3548 }
3549 return true;
3550 }
3551 let rows = out.len();
3552 let cols = downs[0].cols();
3553 if cols % GROUP_SIZE != 0 {
3554 return false;
3555 }
3556 let gpr = cols / GROUP_SIZE;
3557 let mut views = Vec::with_capacity(downs.len());
3558 for (d, g) in downs.iter().zip(gs.iter()) {
3559 if !matches!(
3560 d,
3561 Self::Mapped {
3562 dtype: TensorDtype::Q4TiledP,
3563 ..
3564 }
3565 ) || d.rows() != rows
3566 || d.cols() != cols
3567 || g.len() != cols
3568 {
3569 return false;
3570 }
3571 views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3572 }
3573 let acts: Vec<SplitAct> = gs.iter().map(|g| split_act(g)).collect();
3575 let out_addr = SendMut(out.as_mut_ptr());
3583 let (views, acts, weights) = (&views, &acts, &weights);
3584 let run = |start: usize, end: usize| {
3585 with_krow(gpr, |sc| {
3586 for r in start..end {
3587 let mut acc = 0f32;
3588 for (e, v) in views.iter().enumerate() {
3589 v.scales_into(r, gpr, sc);
3590 let a = &acts[e];
3591 let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, sc) * a.sx;
3592 for &(j, xv) in &a.outliers {
3593 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
3594 d += w * s * xv;
3595 }
3596 acc += weights[e] * d;
3597 }
3598 unsafe { *out_addr.at(r) = acc };
3600 }
3601 });
3602 };
3603 dispatch_rows(pool, rows, &run);
3604 true
3605 }
3606
3607 pub fn moe_gate_up_rows(
3619 pairs: &[(&QTensor, &QTensor)],
3620 groups: &[Vec<usize>],
3621 xs: &[f32],
3622 outs: &mut [Vec<f32>],
3623 pool: Option<&Pool>,
3624 ) -> bool {
3625 if pairs.is_empty() || pairs.len() != groups.len() {
3626 return false;
3627 }
3628 let inter = pairs[0].0.rows();
3629 let cols = pairs[0].0.cols();
3630 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
3631 if cols == 0 || cols % GROUP_SIZE != 0 || outs.len() != n_pairs || xs.len() % cols != 0 {
3632 return false;
3633 }
3634 let b = xs.len() / cols;
3635 let gpr = cols / GROUP_SIZE;
3636 let mut views = Vec::with_capacity(pairs.len() * 2);
3637 for (g, u) in pairs {
3638 let q4tp = |t: &QTensor| {
3639 matches!(
3640 t,
3641 Self::Mapped {
3642 dtype: TensorDtype::Q4TiledP,
3643 ..
3644 }
3645 )
3646 };
3647 if g.has_prism_contract()
3648 || u.has_prism_contract()
3649 || !q4tp(g)
3650 || !q4tp(u)
3651 || g.rows() != inter
3652 || u.rows() != inter
3653 || g.cols() != cols
3654 || u.cols() != cols
3655 {
3656 return false;
3657 }
3658 views.push(Q4tpView::new(g.quant_bytes(), inter, cols));
3659 views.push(Q4tpView::new(u.quant_bytes(), inter, cols));
3660 }
3661 if outs.iter().any(|o| o.len() != inter) || groups.iter().flatten().any(|&t| t >= b) {
3662 return false;
3663 }
3664 let quantized = a8w8_enabled();
3665 let acts: Vec<SplitAct> = if quantized {
3666 (0..b)
3667 .map(|t| split_act(&xs[t * cols..(t + 1) * cols]))
3668 .collect()
3669 } else {
3670 Vec::new()
3671 };
3672 let mut offs = Vec::with_capacity(groups.len());
3673 let mut o = 0usize;
3674 for g in groups {
3675 offs.push(o);
3676 o += g.len();
3677 }
3678 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3679 let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3680 let run = |start: usize, end: usize| {
3681 let (mut gsc, mut usc) = (vec![0f32; gpr], vec![0f32; gpr]);
3682 for flat in start..end {
3683 let (e, r) = (flat / inter, flat % inter);
3684 let (gv_view, uv_view) = (&views[e * 2], &views[e * 2 + 1]);
3685 gv_view.scales_into(r, gpr, &mut gsc);
3686 uv_view.scales_into(r, gpr, &mut usc);
3687 for (k, &t) in groups[e].iter().enumerate() {
3688 if !quantized {
3689 let x = &xs[t * cols..(t + 1) * cols];
3690 let gv = q4tp_row_exact(gv_view.nib, r, gpr, x, &gsc);
3691 let uv = q4tp_row_exact(uv_view.nib, r, gpr, x, &usc);
3692 unsafe { *ptrs[offs[e] + k].at(r) = (gv / (1.0 + (-gv).exp())) * uv };
3693 continue;
3694 }
3695 let act = &acts[t];
3696 let mut gv = dot_q4tp_row_i8(gv_view.nib, r, gpr, &act.xq, &gsc) * act.sx;
3697 let mut uv = dot_q4tp_row_i8(uv_view.nib, r, gpr, &act.xq, &usc) * act.sx;
3698 for &(j, xv) in &act.outliers {
3699 let og = q4tp_outlier(gv_view.nib, r, gpr, j, &gsc);
3700 let ou = q4tp_outlier(uv_view.nib, r, gpr, j, &usc);
3701 gv += og.0 * og.1 * xv;
3702 uv += ou.0 * ou.1 * xv;
3703 }
3704 let silu_g = gv / (1.0 + (-gv).exp());
3705 unsafe { *ptrs[offs[e] + k].at(r) = silu_g * uv };
3708 }
3709 }
3710 };
3711 dispatch_rows(pool, pairs.len() * inter, &run);
3712 true
3713 }
3714
3715 pub fn moe_down_rows(
3722 downs: &[&QTensor],
3723 group_lens: &[usize],
3724 gs: &[Vec<f32>],
3725 outs: &mut [Vec<f32>],
3726 pool: Option<&Pool>,
3727 ) -> bool {
3728 if downs.is_empty() || downs.len() != group_lens.len() {
3729 return false;
3730 }
3731 let rows = downs[0].rows();
3732 let cols = downs[0].cols();
3733 let n_pairs: usize = group_lens.iter().sum();
3734 if cols == 0 || cols % GROUP_SIZE != 0 || gs.len() != n_pairs || outs.len() != n_pairs {
3735 return false;
3736 }
3737 let gpr = cols / GROUP_SIZE;
3738 let mut views = Vec::with_capacity(downs.len());
3739 for d in downs {
3740 if d.has_prism_contract()
3741 || !matches!(
3742 d,
3743 Self::Mapped {
3744 dtype: TensorDtype::Q4TiledP,
3745 ..
3746 }
3747 )
3748 || d.rows() != rows
3749 || d.cols() != cols
3750 {
3751 return false;
3752 }
3753 views.push(Q4tpView::new(d.quant_bytes(), rows, cols));
3754 }
3755 if gs.iter().any(|g| g.len() != cols) || outs.iter().any(|o| o.len() != rows) {
3756 return false;
3757 }
3758 let quantized = a8w8_enabled();
3759 let acts: Vec<SplitAct> = if quantized {
3760 gs.iter().map(|g| split_act(g)).collect()
3761 } else {
3762 Vec::new()
3763 };
3764 let mut offs = Vec::with_capacity(group_lens.len());
3765 let mut o = 0usize;
3766 for &l in group_lens {
3767 offs.push(o);
3768 o += l;
3769 }
3770 let ptrs: Vec<SendMut> = outs.iter_mut().map(|o| SendMut(o.as_mut_ptr())).collect();
3771 let (views, ptrs, acts, offs) = (&views, &ptrs, &acts, &offs);
3772 let run = |start: usize, end: usize| {
3773 let mut sc = vec![0f32; gpr];
3774 for flat in start..end {
3775 let (e, r) = (flat / rows, flat % rows);
3776 let v = &views[e];
3777 v.scales_into(r, gpr, &mut sc);
3778 for k in 0..group_lens[e] {
3779 if !quantized {
3780 let d = q4tp_row_exact(v.nib, r, gpr, &gs[offs[e] + k], &sc);
3781 unsafe { *ptrs[offs[e] + k].at(r) = d };
3782 continue;
3783 }
3784 let a = &acts[offs[e] + k];
3785 let mut d = dot_q4tp_row_i8(v.nib, r, gpr, &a.xq, &sc) * a.sx;
3786 for &(j, xv) in &a.outliers {
3787 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
3788 d += w * s * xv;
3789 }
3790 unsafe { *ptrs[offs[e] + k].at(r) = d };
3792 }
3793 }
3794 };
3795 dispatch_rows(pool, downs.len() * rows, &run);
3796 true
3797 }
3798}
3799
3800#[cfg(target_os = "macos")]
3805mod accel_blas {
3806 #[link(name = "Accelerate", kind = "framework")]
3807 unsafe extern "C" {
3808 pub fn cblas_sgemm(
3809 order: i32,
3810 trans_a: i32,
3811 trans_b: i32,
3812 m: i32,
3813 n: i32,
3814 k: i32,
3815 alpha: f32,
3816 a: *const f32,
3817 lda: i32,
3818 b: *const f32,
3819 ldb: i32,
3820 beta: f32,
3821 c: *mut f32,
3822 ldc: i32,
3823 );
3824 }
3825}
3826
3827#[cfg(target_os = "macos")]
3828pub(crate) fn accel_gemm_enabled() -> bool {
3829 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
3830 *ON.get_or_init(|| std::env::var("CMF_ACCEL").map(|v| v != "0").unwrap_or(true))
3831}
3832
3833#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3836pub(crate) fn accel_gemm_enabled() -> bool {
3837 true
3838}
3839
3840#[cfg(target_arch = "aarch64")]
3847#[allow(clippy::too_many_arguments)]
3848pub(crate) fn neon_gemm_rm(
3849 m: usize,
3850 n: usize,
3851 k: usize,
3852 alpha: f32,
3853 a: &[f32],
3854 lda: usize,
3855 b_mat: &[f32],
3856 ldb: usize,
3857 b_rows_are_n: bool,
3858 c: &mut [f32],
3859 ldc: usize,
3860) {
3861 debug_assert!(a.len() >= (m - 1) * lda + k);
3862 debug_assert!(c.len() >= (m - 1) * ldc + n);
3863 unsafe {
3865 use core::arch::aarch64::*;
3866 let mut i = 0usize;
3867 while i < m {
3868 let mi = (m - i).min(4);
3869 let mut j = 0usize;
3870 while j < n {
3871 let nj = (n - j).min(8);
3872 if mi == 4 && nj == 8 {
3873 let (mut c0a, mut c0b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3874 let (mut c1a, mut c1b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3875 let (mut c2a, mut c2b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3876 let (mut c3a, mut c3b) = (vdupq_n_f32(0.0), vdupq_n_f32(0.0));
3877 for p in 0..k {
3878 let (b0, b1) = if b_rows_are_n {
3879 let base = b_mat.as_ptr().add(j * ldb + p);
3882 let g = |o: usize| *base.add(o * ldb);
3883 ([g(0), g(1), g(2), g(3)], [g(4), g(5), g(6), g(7)])
3884 } else {
3885 let base = b_mat.as_ptr().add(p * ldb + j);
3886 (
3887 [*base, *base.add(1), *base.add(2), *base.add(3)],
3888 [*base.add(4), *base.add(5), *base.add(6), *base.add(7)],
3889 )
3890 };
3891 let bv0 = vld1q_f32(b0.as_ptr());
3892 let bv1 = vld1q_f32(b1.as_ptr());
3893 let a0 = vdupq_n_f32(*a.as_ptr().add(i * lda + p));
3894 let a1 = vdupq_n_f32(*a.as_ptr().add((i + 1) * lda + p));
3895 let a2 = vdupq_n_f32(*a.as_ptr().add((i + 2) * lda + p));
3896 let a3 = vdupq_n_f32(*a.as_ptr().add((i + 3) * lda + p));
3897 c0a = vfmaq_f32(c0a, a0, bv0);
3898 c0b = vfmaq_f32(c0b, a0, bv1);
3899 c1a = vfmaq_f32(c1a, a1, bv0);
3900 c1b = vfmaq_f32(c1b, a1, bv1);
3901 c2a = vfmaq_f32(c2a, a2, bv0);
3902 c2b = vfmaq_f32(c2b, a2, bv1);
3903 c3a = vfmaq_f32(c3a, a3, bv0);
3904 c3b = vfmaq_f32(c3b, a3, bv1);
3905 }
3906 let al = vdupq_n_f32(alpha);
3907 for (r, (ca, cb)) in [(c0a, c0b), (c1a, c1b), (c2a, c2b), (c3a, c3b)]
3908 .iter()
3909 .enumerate()
3910 {
3911 let dst = c.as_mut_ptr().add((i + r) * ldc + j);
3912 vst1q_f32(dst, vmulq_f32(*ca, al));
3913 vst1q_f32(dst.add(4), vmulq_f32(*cb, al));
3914 }
3915 } else {
3916 for r in 0..mi {
3917 for q in 0..nj {
3918 let mut acc = 0f32;
3919 for p in 0..k {
3920 let bv = if b_rows_are_n {
3921 b_mat[(j + q) * ldb + p]
3922 } else {
3923 b_mat[p * ldb + j + q]
3924 };
3925 acc += a[(i + r) * lda + p] * bv;
3926 }
3927 c[(i + r) * ldc + j + q] = acc * alpha;
3928 }
3929 }
3930 }
3931 j += nj;
3932 }
3933 i += mi;
3934 }
3935 }
3936}
3937
3938#[cfg(all(target_arch = "aarch64", not(target_os = "macos")))]
3940#[allow(clippy::too_many_arguments)]
3941pub(crate) fn sgemm_rm(
3942 m: usize,
3943 n: usize,
3944 k: usize,
3945 alpha: f32,
3946 a: &[f32],
3947 lda: usize,
3948 b_mat: &[f32],
3949 ldb: usize,
3950 b_rows_are_n: bool,
3951 c: &mut [f32],
3952 ldc: usize,
3953) {
3954 neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3955}
3956
3957#[allow(clippy::too_many_arguments)]
3961pub fn sgemm_public(
3962 m: usize,
3963 n: usize,
3964 k: usize,
3965 alpha: f32,
3966 a: &[f32],
3967 lda: usize,
3968 b_mat: &[f32],
3969 ldb: usize,
3970 b_rows_are_n: bool,
3971 c: &mut [f32],
3972 ldc: usize,
3973) {
3974 #[cfg(any(target_os = "macos", target_arch = "aarch64"))]
3975 {
3976 sgemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
3977 }
3978 #[cfg(not(any(target_os = "macos", target_arch = "aarch64")))]
3983 {
3984 for i in 0..m {
3985 for j in 0..n {
3986 let mut acc = 0f32;
3987 for p in 0..k {
3988 let bv = if b_rows_are_n {
3989 b_mat[j * ldb + p]
3990 } else {
3991 b_mat[p * ldb + j]
3992 };
3993 acc += a[i * lda + p] * bv;
3994 }
3995 c[i * ldc + j] = alpha * acc;
3996 }
3997 }
3998 }
3999}
4000
4001#[cfg(target_os = "macos")]
4004#[allow(clippy::too_many_arguments)]
4005pub(crate) fn sgemm_rm(
4006 m: usize,
4007 n: usize,
4008 k: usize,
4009 alpha: f32,
4010 a: &[f32],
4011 lda: usize,
4012 b_mat: &[f32],
4013 ldb: usize,
4014 b_rows_are_n: bool,
4015 c: &mut [f32],
4016 ldc: usize,
4017) {
4018 debug_assert!(a.len() >= (m - 1) * lda + k);
4019 debug_assert!(c.len() >= (m - 1) * ldc + n);
4020 #[cfg(target_arch = "aarch64")]
4025 if std::env::var("CMF_FORCE_NEON_GEMM")
4026 .map(|v| v == "1")
4027 .unwrap_or(false)
4028 {
4029 return neon_gemm_rm(m, n, k, alpha, a, lda, b_mat, ldb, b_rows_are_n, c, ldc);
4030 }
4031 unsafe {
4032 accel_blas::cblas_sgemm(
4033 101, 111, if b_rows_are_n { 112 } else { 111 },
4036 m as i32,
4037 n as i32,
4038 k as i32,
4039 alpha,
4040 a.as_ptr(),
4041 lda as i32,
4042 b_mat.as_ptr(),
4043 ldb as i32,
4044 0.0,
4045 c.as_mut_ptr(),
4046 ldc as i32,
4047 );
4048 }
4049}
4050
4051#[cfg(target_os = "macos")]
4058fn qmatmat_accel(
4059 q: &[u8],
4060 row_scale: &[f32],
4061 pre: &[std::borrow::Cow<'_, [f32]>],
4062 rows: usize,
4063 cols: usize,
4064 out: &mut [f32],
4065 pool: Option<&Pool>,
4066) {
4067 const TR: usize = 2048;
4072 let b = pre.len();
4073 thread_local! {
4074 static XPANEL: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4075 static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
4076 }
4077 XPANEL.with(|xp| {
4078 WTILE.with(|wt| {
4079 let mut xpanel = xp.borrow_mut();
4080 xpanel.clear();
4081 for x in pre {
4082 xpanel.extend_from_slice(x);
4083 }
4084 let mut wtile = wt.borrow_mut();
4085 wtile.resize(TR * cols, 0.0);
4086 let mut r0 = 0usize;
4087 while r0 < rows {
4088 let tr = TR.min(rows - r0);
4089 let wt_addr = SendMut(wtile.as_mut_ptr());
4091 let run = |start: usize, end: usize| {
4092 for r in start..end {
4093 let row = &q[(r0 + r) * cols..(r0 + r + 1) * cols];
4094 let s = row_scale[r0 + r];
4095 let dst =
4097 unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
4098 for (d, &v) in dst.iter_mut().zip(row) {
4099 *d = (v as i8) as f32 * s;
4100 }
4101 }
4102 };
4103 dispatch_rows(pool, tr, &run);
4104 unsafe {
4106 accel_blas::cblas_sgemm(
4107 101, 111, 112, b as i32,
4111 tr as i32,
4112 cols as i32,
4113 1.0,
4114 xpanel.as_ptr(),
4115 cols as i32,
4116 wtile.as_ptr(),
4117 cols as i32,
4118 0.0,
4119 out.as_mut_ptr().add(r0),
4120 rows as i32,
4121 );
4122 }
4123 r0 += tr;
4124 }
4125 })
4126 });
4127}
4128
4129fn qmatmat(
4130 q: &[u8],
4131 row_scale: &[f32],
4132 pre: &[std::borrow::Cow<'_, [f32]>],
4133 rows: usize,
4134 cols: usize,
4135 out: &mut [f32],
4136 pool: Option<&Pool>,
4137) {
4138 let b = pre.len();
4139 debug_assert_eq!(out.len(), b * rows);
4140 #[cfg(target_os = "macos")]
4145 if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
4146 qmatmat_accel(q, row_scale, pre, rows, cols, out, pool);
4147 return;
4148 }
4149 #[cfg(target_arch = "aarch64")]
4150 if sdot_enabled() {
4151 let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4152 let out_addr = SendMut(out.as_mut_ptr());
4153 let blocked_ok = blocked_enabled();
4156 let use_i8mm = i8mm_enabled();
4157 if blocked_ok {
4158 let run = |start: usize, end: usize| {
4159 let mut o = start;
4160 while o < end {
4161 if o + 2 <= end {
4162 let r0 = &q[o * cols..(o + 1) * cols];
4163 let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4164 let mut bi = 0usize;
4165 while bi + 4 <= acts.len() {
4166 let xs = [
4167 acts[bi].xq.as_slice(),
4168 acts[bi + 1].xq.as_slice(),
4169 acts[bi + 2].xq.as_slice(),
4170 acts[bi + 3].xq.as_slice(),
4171 ];
4172 let d = if use_i8mm {
4173 unsafe { dot_i8_smmla_2x4(r0, r1, xs) }
4174 } else {
4175 unsafe { dot_i8_sdot_2x4(r0, r1, xs) }
4176 };
4177 for (r, row) in [r0, r1].into_iter().enumerate() {
4178 for k in 0..4 {
4179 let act = &acts[bi + k];
4180 let mut v = d[r][k] as f32 * act.sx;
4181 for &(j, xv) in &act.outliers {
4182 v += (row[j] as i8) as f32 * xv;
4183 }
4184 unsafe {
4185 *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4186 };
4187 }
4188 }
4189 bi += 4;
4190 }
4191 while bi < acts.len() {
4192 for (r, row) in [r0, r1].into_iter().enumerate() {
4193 let v = row_dot_sdot(row, &acts[bi]) * row_scale[o + r];
4194 unsafe { *out_addr.at(bi * rows + o + r) = v };
4195 }
4196 bi += 1;
4197 }
4198 o += 2;
4199 } else {
4200 let row = &q[o * cols..(o + 1) * cols];
4201 for (bi, act) in acts.iter().enumerate() {
4202 let v = row_dot_sdot(row, act) * row_scale[o];
4203 unsafe { *out_addr.at(bi * rows + o) = v };
4204 }
4205 o += 1;
4206 }
4207 }
4208 };
4209 dispatch_rows(pool, rows, &run);
4210 return;
4211 }
4212 let run = |start: usize, end: usize| {
4213 for o in start..end {
4214 let row = &q[o * cols..(o + 1) * cols];
4215 for (bi, act) in acts.iter().enumerate() {
4216 let v = row_dot_sdot(row, act) * row_scale[o];
4217 unsafe { *out_addr.at(bi * rows + o) = v };
4218 }
4219 }
4220 };
4221 dispatch_rows(pool, rows, &run);
4222 return;
4223 }
4224 #[cfg(target_arch = "x86_64")]
4229 if avx2_a8w8_enabled() {
4230 let acts: Vec<SplitAct> = pre.iter().map(|x| split_act(x)).collect();
4231 let out_addr = SendMut(out.as_mut_ptr());
4232 let blocked_ok = blocked_enabled();
4235 if !avx512vnni_enabled() && blocked_ok && !row_exact() {
4236 let run = |start: usize, end: usize| {
4237 let mut o = start;
4238 while o < end {
4239 if o + 2 <= end {
4240 let r0 = &q[o * cols..(o + 1) * cols];
4241 let r1 = &q[(o + 1) * cols..(o + 2) * cols];
4242 let mut bi = 0usize;
4243 while bi + 4 <= acts.len() {
4244 let xs = [
4245 acts[bi].xq.as_slice(),
4246 acts[bi + 1].xq.as_slice(),
4247 acts[bi + 2].xq.as_slice(),
4248 acts[bi + 3].xq.as_slice(),
4249 ];
4250 let d = unsafe { dot_i8_i8_avx2_2x4(r0, r1, xs) };
4251 for (r, row) in [r0, r1].into_iter().enumerate() {
4252 for k in 0..4 {
4253 let act = &acts[bi + k];
4254 let mut v = d[r][k] as f32 * act.sx;
4255 for &(j, xv) in &act.outliers {
4256 v += (row[j] as i8) as f32 * xv;
4257 }
4258 unsafe {
4259 *out_addr.at((bi + k) * rows + o + r) = v * row_scale[o + r]
4260 };
4261 }
4262 }
4263 bi += 4;
4264 }
4265 while bi < acts.len() {
4266 for (r, row) in [r0, r1].into_iter().enumerate() {
4267 let v = row_dot_avx2(row, &acts[bi]) * row_scale[o + r];
4268 unsafe { *out_addr.at(bi * rows + o + r) = v };
4269 }
4270 bi += 1;
4271 }
4272 o += 2;
4273 } else {
4274 let row = &q[o * cols..(o + 1) * cols];
4275 for (bi, act) in acts.iter().enumerate() {
4276 let v = row_dot_avx2(row, act) * row_scale[o];
4277 unsafe { *out_addr.at(bi * rows + o) = v };
4278 }
4279 o += 1;
4280 }
4281 }
4282 };
4283 dispatch_rows(pool, rows, &run);
4284 return;
4285 }
4286 let run = |start: usize, end: usize| {
4287 for o in start..end {
4288 let row = &q[o * cols..(o + 1) * cols];
4289 for (bi, act) in acts.iter().enumerate() {
4290 let v = row_dot_avx2(row, act) * row_scale[o];
4291 unsafe { *out_addr.at(bi * rows + o) = v };
4292 }
4293 }
4294 };
4295 dispatch_rows(pool, rows, &run);
4296 return;
4297 }
4298 let out_addr = SendMut(out.as_mut_ptr());
4299 let run = |start: usize, end: usize| {
4300 for o in start..end {
4301 let row = &q[o * cols..(o + 1) * cols];
4302 for (bi, x) in pre.iter().enumerate() {
4303 let mut acc = 0f32;
4304 for j in 0..cols {
4305 acc += (row[j] as i8) as f32 * x[j];
4306 }
4307 unsafe { *out_addr.at(bi * rows + o) = acc * row_scale[o] };
4308 }
4309 }
4310 };
4311 dispatch_rows(pool, rows, &run);
4312}
4313
4314fn dispatch_rows(pool: Option<&Pool>, rows: usize, run: &(dyn Fn(usize, usize) + Sync)) {
4317 match pool {
4318 Some(pool) if rows >= 256 => pool.run_rows(rows, run),
4319 _ => run(0, rows),
4320 }
4321}
4322
4323fn q4_split(bytes: &[u8], rows: usize, cols: usize) -> (&[u8], &[u8]) {
4325 let groups = rows * cols / GROUP_SIZE;
4326 bytes.split_at(groups * 16)
4327}
4328
4329#[inline]
4334fn vbit_fill4(data: &[u8], buf: &mut [u8]) {
4335 #[cfg(target_arch = "aarch64")]
4336 unsafe {
4337 return vbit_fill4_neon(data, buf);
4338 }
4339 #[cfg(target_arch = "x86_64")]
4340 if avx2_enabled() {
4341 return unsafe { vbit_fill4_avx2(data, buf) };
4342 }
4343 #[allow(unreachable_code)]
4344 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4345 let u = unpack8::<4>(&data[blk * 4..]);
4346 for k in 0..8 {
4347 chunk[k] = (u[k] - 7) as i8 as u8;
4348 }
4349 }
4350}
4351
4352#[cfg(target_arch = "aarch64")]
4353#[target_feature(enable = "neon")]
4354unsafe fn vbit_fill4_neon(data: &[u8], buf: &mut [u8]) {
4355 unsafe {
4358 use core::arch::aarch64::*;
4359 let n = buf.len();
4360 let mask = vdupq_n_u8(0x0F);
4361 let seven = vdupq_n_s8(7);
4362 let mut g = 0usize;
4363 while g * 32 + 32 <= n {
4364 let b = vld1q_u8(data.as_ptr().add(g * 16));
4365 let hi = vshrq_n_u8::<4>(b);
4366 let lo = vandq_u8(b, mask);
4367 let z0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(hi, lo)), seven);
4368 let z1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(hi, lo)), seven);
4369 vst1q_u8(buf.as_mut_ptr().add(g * 32), vreinterpretq_u8_s8(z0));
4370 vst1q_u8(buf.as_mut_ptr().add(g * 32 + 16), vreinterpretq_u8_s8(z1));
4371 g += 1;
4372 }
4373 }
4374}
4375
4376#[cfg(target_arch = "x86_64")]
4377#[target_feature(enable = "avx2")]
4378unsafe fn vbit_fill4_avx2(data: &[u8], buf: &mut [u8]) {
4379 unsafe {
4381 use core::arch::x86_64::*;
4382 let n = buf.len();
4383 let mask = _mm_set1_epi8(0x0F);
4384 let seven = _mm256_set1_epi8(7);
4385 let mut g = 0usize;
4386 while g * 32 + 32 <= n {
4387 let b = _mm_loadu_si128(data.as_ptr().add(g * 16) as *const __m128i);
4388 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), mask);
4389 let lo = _mm_and_si128(b, mask);
4390 let z = _mm256_sub_epi8(
4391 _mm256_set_m128i(_mm_unpackhi_epi8(hi, lo), _mm_unpacklo_epi8(hi, lo)),
4392 seven,
4393 );
4394 _mm256_storeu_si256(buf.as_mut_ptr().add(g * 32) as *mut __m256i, z);
4395 g += 1;
4396 }
4397 }
4398}
4399
4400#[inline(always)]
4405fn unpack8<const B: usize>(data: &[u8]) -> [i32; 8] {
4406 let mut acc = 0u64;
4407 for i in 0..B {
4408 acc = (acc << 8) | data[i] as u64;
4409 }
4410 let mask = (1u64 << B) - 1;
4411 let mut out = [0i32; 8];
4412 for (k, o) in out.iter_mut().enumerate() {
4413 *o = ((acc >> ((7 - k) * B)) & mask) as i32;
4414 }
4415 out
4416}
4417
4418#[allow(clippy::too_many_arguments)]
4424fn vbitmatvec(
4425 bytes: &[u8],
4426 offsets: &[usize],
4427 x: &[f32],
4428 rows: usize,
4429 cols: usize,
4430 out: &mut [f32],
4431 pool: Option<&Pool>,
4432) {
4433 debug_assert_eq!(out.len(), rows);
4434 debug_assert_eq!(offsets.len(), rows + 1);
4435
4436 if a8w8_enabled() {
4440 let act = split_act(x);
4441 let out_addr = SendMut(out.as_mut_ptr());
4442 let run = move |start: usize, end: usize| {
4443 vbit_range_a8w8(bytes, offsets, x, &act, rows, cols, out_addr, start, end)
4444 };
4445 dispatch_rows(pool, rows, &run);
4446 return;
4447 }
4448
4449 let out_addr = SendMut(out.as_mut_ptr());
4450 let run = move |start: usize, end: usize| {
4451 vbit_range_f32(bytes, offsets, x, rows, cols, out_addr, start, end)
4452 };
4453 dispatch_rows(pool, rows, &run);
4454}
4455
4456#[allow(clippy::too_many_arguments)]
4460fn vbit_range_a8w8(
4461 bytes: &[u8],
4462 offsets: &[usize],
4463 x: &[f32],
4464 act: &SplitAct,
4465 rows: usize,
4466 cols: usize,
4467 out: SendMut,
4468 start: usize,
4469 end: usize,
4470) {
4471 let ng = cols / GROUP_SIZE;
4472 let bits = &bytes[..rows];
4473 let sc_off = rows;
4474 let row_dot = |r: usize| -> f32 {
4475 let b = bits[r] as usize;
4476 let l = (1i32 << (b - 1)) - 1;
4477 let mask = (1u64 << b) - 1;
4478 let data = &bytes[offsets[r]..offsets[r + 1]];
4479 if b == 8 {
4480 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4482 let mut dot = 0f32;
4483 for g in 0..ng {
4484 let so = (r * ng + g) * 2;
4485 let sgf = f16_to_f32(u16::from_le_bytes([
4486 bytes[sc_off + so],
4487 bytes[sc_off + so + 1],
4488 ]));
4489 let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4490 let mut gd = 0f32;
4491 for &xv in xg.iter() {
4492 if nbits < 8 {
4493 acc = (acc << 8) | data[idx] as u64;
4494 idx += 1;
4495 nbits += 8;
4496 }
4497 let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4498 nbits -= 8;
4499 gd += (u - l) as f32 * xv;
4500 }
4501 dot += gd * sgf;
4502 }
4503 return dot;
4504 }
4505 thread_local! {
4509 static VBIT_SCRATCH: std::cell::RefCell<Vec<u8>> =
4510 const { std::cell::RefCell::new(Vec::new()) };
4511 }
4512 #[inline(always)]
4513 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4514 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4515 let u = unpack8::<B>(&data[blk * B..]);
4516 for k in 0..8 {
4517 chunk[k] = (u[k] - l) as i8 as u8;
4518 }
4519 }
4520 }
4521 let _ = mask;
4522 VBIT_SCRATCH.with(|scratch| {
4523 let mut buf = scratch.borrow_mut();
4524 buf.resize(cols, 0);
4525 match b {
4526 3 => fill::<3>(data, l, &mut buf),
4527 4 => vbit_fill4(data, &mut buf),
4528 5 => fill::<5>(data, l, &mut buf),
4529 6 => fill::<6>(data, l, &mut buf),
4530 _ => unreachable!(),
4531 }
4532 let mut dot = 0f32;
4533 for g in 0..ng {
4534 let so = (r * ng + g) * 2;
4535 let s = f16_to_f32(u16::from_le_bytes([
4536 bytes[sc_off + so],
4537 bytes[sc_off + so + 1],
4538 ]));
4539 let d = dot_i8_i8(
4540 &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4541 &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
4542 ) as f32
4543 * act.sx;
4544 dot += d * s;
4545 }
4546 for &(j, xv) in &act.outliers {
4547 let so = (r * ng + j / GROUP_SIZE) * 2;
4548 let s = f16_to_f32(u16::from_le_bytes([
4549 bytes[sc_off + so],
4550 bytes[sc_off + so + 1],
4551 ]));
4552 dot += (buf[j] as i8) as f32 * s * xv;
4554 }
4555 dot
4556 })
4557 };
4558 for r in start..end {
4559 unsafe { *out.at(r) = row_dot(r) };
4561 }
4562}
4563
4564#[allow(clippy::too_many_arguments)]
4566fn vbit_range_f32(
4567 bytes: &[u8],
4568 offsets: &[usize],
4569 x: &[f32],
4570 rows: usize,
4571 cols: usize,
4572 out: SendMut,
4573 start: usize,
4574 end: usize,
4575) {
4576 let ng = cols / GROUP_SIZE;
4577 let bits = &bytes[..rows];
4578 let sc_off = rows;
4579 #[inline(always)]
4583 fn dot_row<const B: usize>(
4584 data: &[u8],
4585 bytes: &[u8],
4586 sc_off: usize,
4587 r: usize,
4588 ng: usize,
4589 x: &[f32],
4590 ) -> f32 {
4591 let l = ((1i32 << (B - 1)) - 1) as f32;
4592 let gbytes = GROUP_SIZE * B / 8;
4593 let mut dot = 0f32;
4594 for g in 0..ng {
4595 let so = (r * ng + g) * 2;
4596 let s = f16_to_f32(u16::from_le_bytes([
4597 bytes[sc_off + so],
4598 bytes[sc_off + so + 1],
4599 ]));
4600 let xg = &x[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4601 let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4602 let mut gd = 0f32;
4603 for blk in 0..GROUP_SIZE / 8 {
4604 let u = unpack8::<B>(&gd0[blk * B..]);
4605 let xb = &xg[blk * 8..blk * 8 + 8];
4606 for k in 0..8 {
4607 gd += (u[k] as f32 - l) * xb[k];
4608 }
4609 }
4610 dot += gd * s;
4611 }
4612 dot
4613 }
4614 for r in start..end {
4615 let data = &bytes[offsets[r]..offsets[r + 1]];
4616 let v = match bits[r] {
4617 3 => dot_row::<3>(data, bytes, sc_off, r, ng, x),
4618 4 => dot_row::<4>(data, bytes, sc_off, r, ng, x),
4619 5 => dot_row::<5>(data, bytes, sc_off, r, ng, x),
4620 6 => dot_row::<6>(data, bytes, sc_off, r, ng, x),
4621 8 => dot_row::<8>(data, bytes, sc_off, r, ng, x),
4622 b => unreachable!("vbit bit-width {b} (validated at load)"),
4623 };
4624 unsafe { *out.at(r) = v };
4626 }
4627}
4628
4629#[allow(clippy::too_many_arguments)]
4634fn vbitmatvec2(
4635 bytes: &[u8],
4636 offsets: &[usize],
4637 x1: &[f32],
4638 x2: &[f32],
4639 rows: usize,
4640 cols: usize,
4641 o1: &mut [f32],
4642 o2: &mut [f32],
4643 pool: Option<&Pool>,
4644) {
4645 debug_assert_eq!(o1.len(), rows);
4646 debug_assert_eq!(o2.len(), rows);
4647
4648 if a8w8_enabled() {
4649 let a1 = split_act(x1);
4650 let a2 = split_act(x2);
4651 let p1 = SendMut(o1.as_mut_ptr());
4652 let p2 = SendMut(o2.as_mut_ptr());
4653 let run = move |start: usize, end: usize| {
4654 vbit_range2_a8w8(
4655 bytes, offsets, x1, x2, &a1, &a2, rows, cols, p1, p2, start, end,
4656 )
4657 };
4658 dispatch_rows(pool, rows, &run);
4659 return;
4660 }
4661
4662 let p1 = SendMut(o1.as_mut_ptr());
4663 let p2 = SendMut(o2.as_mut_ptr());
4664 let run = move |start: usize, end: usize| {
4665 vbit_range2_f32(bytes, offsets, x1, x2, rows, cols, p1, p2, start, end)
4666 };
4667 dispatch_rows(pool, rows, &run);
4668}
4669
4670#[allow(clippy::too_many_arguments)]
4674fn vbit_range2_a8w8(
4675 bytes: &[u8],
4676 offsets: &[usize],
4677 x1: &[f32],
4678 x2: &[f32],
4679 a1: &SplitAct,
4680 a2: &SplitAct,
4681 rows: usize,
4682 cols: usize,
4683 p1: SendMut,
4684 p2: SendMut,
4685 start: usize,
4686 end: usize,
4687) {
4688 let ng = cols / GROUP_SIZE;
4689 let bits = &bytes[..rows];
4690 let sc_off = rows;
4691 let row_dots = |r: usize| -> (f32, f32) {
4692 let b = bits[r] as usize;
4693 let l = (1i32 << (b - 1)) - 1;
4694 let data = &bytes[offsets[r]..offsets[r + 1]];
4695 if b == 8 {
4696 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
4699 let (mut d1, mut d2) = (0f32, 0f32);
4700 for g in 0..ng {
4701 let so = (r * ng + g) * 2;
4702 let sgf = f16_to_f32(u16::from_le_bytes([
4703 bytes[sc_off + so],
4704 bytes[sc_off + so + 1],
4705 ]));
4706 let (mut g1, mut g2) = (0f32, 0f32);
4707 for k in 0..GROUP_SIZE {
4708 if nbits < 8 {
4709 acc = (acc << 8) | data[idx] as u64;
4710 idx += 1;
4711 nbits += 8;
4712 }
4713 let u = ((acc >> (nbits - 8)) & 0xFF) as i32;
4714 nbits -= 8;
4715 let w = (u - l) as f32;
4716 g1 += w * x1[g * GROUP_SIZE + k];
4717 g2 += w * x2[g * GROUP_SIZE + k];
4718 }
4719 d1 += g1 * sgf;
4720 d2 += g2 * sgf;
4721 }
4722 return (d1, d2);
4723 }
4724 thread_local! {
4725 static VBIT_SCRATCH2: std::cell::RefCell<Vec<u8>> =
4726 const { std::cell::RefCell::new(Vec::new()) };
4727 }
4728 #[inline(always)]
4729 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
4730 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
4731 let u = unpack8::<B>(&data[blk * B..]);
4732 for k in 0..8 {
4733 chunk[k] = (u[k] - l) as i8 as u8;
4734 }
4735 }
4736 }
4737 VBIT_SCRATCH2.with(|scratch| {
4738 let mut buf = scratch.borrow_mut();
4739 buf.resize(cols, 0);
4740 match b {
4741 3 => fill::<3>(data, l, &mut buf),
4742 4 => vbit_fill4(data, &mut buf),
4743 5 => fill::<5>(data, l, &mut buf),
4744 6 => fill::<6>(data, l, &mut buf),
4745 _ => unreachable!(),
4746 }
4747 let (mut d1, mut d2) = (0f32, 0f32);
4748 for g in 0..ng {
4749 let so = (r * ng + g) * 2;
4750 let s = f16_to_f32(u16::from_le_bytes([
4751 bytes[sc_off + so],
4752 bytes[sc_off + so + 1],
4753 ]));
4754 let wg = &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4755 let v1 = dot_i8_i8(wg, &a1.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a1.sx;
4756 let v2 = dot_i8_i8(wg, &a2.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE]) as f32 * a2.sx;
4757 d1 += v1 * s;
4758 d2 += v2 * s;
4759 }
4760 for &(j, xv) in &a1.outliers {
4761 let so = (r * ng + j / GROUP_SIZE) * 2;
4762 let s = f16_to_f32(u16::from_le_bytes([
4763 bytes[sc_off + so],
4764 bytes[sc_off + so + 1],
4765 ]));
4766 d1 += (buf[j] as i8) as f32 * s * xv;
4767 }
4768 for &(j, xv) in &a2.outliers {
4769 let so = (r * ng + j / GROUP_SIZE) * 2;
4770 let s = f16_to_f32(u16::from_le_bytes([
4771 bytes[sc_off + so],
4772 bytes[sc_off + so + 1],
4773 ]));
4774 d2 += (buf[j] as i8) as f32 * s * xv;
4775 }
4776 (d1, d2)
4777 })
4778 };
4779 for r in start..end {
4780 let (v1, v2) = row_dots(r);
4781 unsafe {
4783 *p1.at(r) = v1;
4784 *p2.at(r) = v2;
4785 }
4786 }
4787}
4788
4789#[allow(clippy::too_many_arguments)]
4793fn vbit_range2_f32(
4794 bytes: &[u8],
4795 offsets: &[usize],
4796 x1: &[f32],
4797 x2: &[f32],
4798 rows: usize,
4799 cols: usize,
4800 p1: SendMut,
4801 p2: SendMut,
4802 start: usize,
4803 end: usize,
4804) {
4805 let ng = cols / GROUP_SIZE;
4806 let bits = &bytes[..rows];
4807 let sc_off = rows;
4808 #[inline(always)]
4809 #[allow(clippy::too_many_arguments)]
4810 fn dot_row2<const B: usize>(
4811 data: &[u8],
4812 bytes: &[u8],
4813 sc_off: usize,
4814 r: usize,
4815 ng: usize,
4816 x1: &[f32],
4817 x2: &[f32],
4818 ) -> (f32, f32) {
4819 let l = ((1i32 << (B - 1)) - 1) as f32;
4820 let gbytes = GROUP_SIZE * B / 8;
4821 let (mut d1, mut d2) = (0f32, 0f32);
4822 for g in 0..ng {
4823 let so = (r * ng + g) * 2;
4824 let s = f16_to_f32(u16::from_le_bytes([
4825 bytes[sc_off + so],
4826 bytes[sc_off + so + 1],
4827 ]));
4828 let x1g = &x1[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4829 let x2g = &x2[g * GROUP_SIZE..(g + 1) * GROUP_SIZE];
4830 let gd0 = &data[g * gbytes..(g + 1) * gbytes];
4831 let (mut g1, mut g2) = (0f32, 0f32);
4832 for blk in 0..GROUP_SIZE / 8 {
4833 let u = unpack8::<B>(&gd0[blk * B..]);
4834 for k in 0..8 {
4835 let w = u[k] as f32 - l;
4836 g1 += w * x1g[blk * 8 + k];
4837 g2 += w * x2g[blk * 8 + k];
4838 }
4839 }
4840 d1 += g1 * s;
4841 d2 += g2 * s;
4842 }
4843 (d1, d2)
4844 }
4845 for r in start..end {
4846 let data = &bytes[offsets[r]..offsets[r + 1]];
4847 let (v1, v2) = match bits[r] {
4848 3 => dot_row2::<3>(data, bytes, sc_off, r, ng, x1, x2),
4849 4 => dot_row2::<4>(data, bytes, sc_off, r, ng, x1, x2),
4850 5 => dot_row2::<5>(data, bytes, sc_off, r, ng, x1, x2),
4851 6 => dot_row2::<6>(data, bytes, sc_off, r, ng, x1, x2),
4852 8 => dot_row2::<8>(data, bytes, sc_off, r, ng, x1, x2),
4853 b => unreachable!("vbit bit-width {b} (validated at load)"),
4854 };
4855 unsafe {
4857 *p1.at(r) = v1;
4858 *p2.at(r) = v2;
4859 }
4860 }
4861}
4862
4863#[inline]
4870#[allow(unreachable_code)]
4871fn dot_q4t_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4872 #[cfg(target_arch = "aarch64")]
4873 unsafe {
4874 return dot_q4t_row_sdot(bytes, r, gpr, xq);
4875 }
4876 #[cfg(target_arch = "x86_64")]
4877 unsafe {
4878 if vnni_tiles_enabled() {
4879 return dot_q4t_row_vnni(bytes, r, gpr, xq);
4880 }
4881 return dot_q4t_row_avx2(bytes, r, gpr, xq);
4882 }
4883 let mut acc = 0f32;
4884 for gi in 0..gpr {
4885 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
4886 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
4887 let mut d = 0i32;
4888 for (k, &b) in tile[2..].iter().enumerate() {
4889 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
4890 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
4891 }
4892 acc += d as f32 * s;
4893 }
4894 acc
4895}
4896
4897#[cfg(target_arch = "aarch64")]
4898#[target_feature(enable = "neon,dotprod")]
4899unsafe fn dot_q4t_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4900 unsafe {
4903 use core::arch::aarch64::*;
4904 use core::arch::asm;
4905 let lomask = vdupq_n_u8(0x0F);
4906 let eight = vdupq_n_s8(8);
4907 let mut acc = 0f32;
4908 for gi in 0..gpr {
4909 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4910 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4911 let b = vld1q_u8(t.add(2));
4912 let lo = vandq_u8(b, lomask);
4913 let hi = vshrq_n_u8::<4>(b);
4914 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
4915 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
4916 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
4917 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
4918 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
4919 asm!(
4920 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
4921 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
4922 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
4923 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
4924 options(pure, nomem, nostack),
4925 );
4926 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
4927 }
4928 acc
4929 }
4930}
4931
4932#[cfg(target_arch = "x86_64")]
4933#[target_feature(enable = "avx2")]
4934unsafe fn dot_q4t_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4935 unsafe {
4937 use core::arch::x86_64::*;
4938 let lomask = _mm_set1_epi8(0x0F);
4939 let eight = _mm256_set1_epi8(8);
4940 let ones = _mm256_set1_epi16(1);
4941 let mut acc = 0f32;
4942 for gi in 0..gpr {
4943 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4944 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4945 let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4946 let lo = _mm_and_si128(b, lomask);
4947 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4948 let w = _mm256_sub_epi8(
4949 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4950 eight,
4951 );
4952 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4953 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4954 let d = _mm256_madd_epi16(p16, ones);
4955 let hi128 = _mm256_extracti128_si256::<1>(d);
4956 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
4957 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
4958 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
4959 acc += _mm_cvtsi128_si32(s32) as f32 * s;
4960 }
4961 acc
4962 }
4963}
4964
4965#[cfg(target_arch = "x86_64")]
4969#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
4970unsafe fn dot_q4t_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
4971 unsafe {
4973 use core::arch::x86_64::*;
4974 let lomask = _mm_set1_epi8(0x0F);
4975 let eight = _mm256_set1_epi8(8);
4976 let mut acc = 0f32;
4977 for gi in 0..gpr {
4978 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
4979 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
4980 let b = _mm_loadu_si128(t.add(2) as *const __m128i);
4981 let lo = _mm_and_si128(b, lomask);
4982 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
4983 let w = _mm256_sub_epi8(
4984 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
4985 eight,
4986 );
4987 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
4988 let d = dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
4989 acc += d as f32 * s;
4990 }
4991 acc
4992 }
4993}
4994
4995#[cfg(target_arch = "x86_64")]
5000#[target_feature(enable = "avx2,fma")]
5005unsafe fn dot_q4t_row_1x4_avx2(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5006 unsafe {
5008 use core::arch::x86_64::*;
5009 let lomask = _mm_set1_epi8(0x0F);
5010 let eight = _mm256_set1_epi8(8);
5011 let ones = _mm256_set1_epi16(1);
5012 let mut f0 = _mm256_setzero_ps();
5025 let mut f1 = _mm256_setzero_ps();
5026 let mut f2 = _mm256_setzero_ps();
5027 let mut f3 = _mm256_setzero_ps();
5028 for gi in 0..gpr {
5029 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5030 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5031 let sv = _mm256_set1_ps(s);
5032 let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5033 let lo = _mm_and_si128(bb, lomask);
5034 let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5035 let w = _mm256_sub_epi8(
5036 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5037 eight,
5038 );
5039 let aw = _mm256_abs_epi8(w);
5040 let off = gi * GROUP_SIZE;
5041 let dot = |xq: &[i8]| {
5042 let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5043 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
5044 _mm256_cvtepi32_ps(_mm256_madd_epi16(p16, ones))
5045 };
5046 f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5047 f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5048 f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5049 f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5050 }
5051 [
5052 hsum256_ps(f0),
5053 hsum256_ps(f1),
5054 hsum256_ps(f2),
5055 hsum256_ps(f3),
5056 ]
5057 }
5058}
5059
5060#[cfg(target_arch = "x86_64")]
5063#[target_feature(enable = "avx2")]
5064#[inline]
5065unsafe fn hsum256_ps(v: core::arch::x86_64::__m256) -> f32 {
5066 unsafe {
5068 use core::arch::x86_64::*;
5069 let hi = _mm256_extractf128_ps::<1>(v);
5070 let s = _mm_add_ps(_mm256_castps256_ps128(v), hi);
5071 let s = _mm_add_ps(s, _mm_movehl_ps(s, s));
5072 let s = _mm_add_ss(s, _mm_shuffle_ps::<0x55>(s, s));
5073 _mm_cvtss_f32(s)
5074 }
5075}
5076
5077#[cfg(target_arch = "x86_64")]
5079#[target_feature(enable = "avx2,fma,avx512f,avx512bw,avx512vl,avx512vnni")]
5080unsafe fn dot_q4t_row_1x4_vnni(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5081 unsafe {
5083 use core::arch::x86_64::*;
5084 let lomask = _mm_set1_epi8(0x0F);
5085 let eight = _mm256_set1_epi8(8);
5086 let mut f0 = _mm256_setzero_ps();
5089 let mut f1 = _mm256_setzero_ps();
5090 let mut f2 = _mm256_setzero_ps();
5091 let mut f3 = _mm256_setzero_ps();
5092 for gi in 0..gpr {
5093 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5094 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5095 let sv = _mm256_set1_ps(s);
5096 let bb = _mm_loadu_si128(t.add(2) as *const __m128i);
5097 let lo = _mm_and_si128(bb, lomask);
5098 let hi = _mm_and_si128(_mm_srli_epi16::<4>(bb), lomask);
5099 let w = _mm256_sub_epi8(
5100 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5101 eight,
5102 );
5103 let aw = _mm256_abs_epi8(w);
5104 let off = gi * GROUP_SIZE;
5105 let dot = |xq: &[i8]| {
5106 let x = _mm256_loadu_si256(xq.as_ptr().add(off) as *const __m256i);
5107 _mm256_cvtepi32_ps(_mm256_dpbusd_epi32(
5108 _mm256_setzero_si256(),
5109 aw,
5110 _mm256_sign_epi8(x, w),
5111 ))
5112 };
5113 f0 = _mm256_fmadd_ps(dot(xs[0]), sv, f0);
5114 f1 = _mm256_fmadd_ps(dot(xs[1]), sv, f1);
5115 f2 = _mm256_fmadd_ps(dot(xs[2]), sv, f2);
5116 f3 = _mm256_fmadd_ps(dot(xs[3]), sv, f3);
5117 }
5118 let acc = [
5119 hsum256_ps(f0),
5120 hsum256_ps(f1),
5121 hsum256_ps(f2),
5122 hsum256_ps(f3),
5123 ];
5124 acc
5125 }
5126}
5127
5128#[cfg(target_arch = "aarch64")]
5133#[target_feature(enable = "neon,dotprod")]
5134unsafe fn dot_q4t_row_1x4_sdot(bytes: &[u8], r: usize, gpr: usize, xs: [&[i8]; 4]) -> [f32; 4] {
5135 unsafe {
5137 use core::arch::aarch64::*;
5138 use core::arch::asm;
5139 let lomask = vdupq_n_u8(0x0F);
5140 let eight = vdupq_n_s8(8);
5141 let mut acc = [0f32; 4];
5142 for gi in 0..gpr {
5143 let t = bytes.as_ptr().add((r * gpr + gi) * Q4_TILE);
5144 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
5145 let b = vld1q_u8(t.add(2));
5146 let lo = vandq_u8(b, lomask);
5147 let hi = vshrq_n_u8::<4>(b);
5148 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5149 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5150 for (k, xq) in xs.iter().enumerate() {
5151 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5152 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5153 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5154 asm!(
5155 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5156 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5157 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5158 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5159 options(pure, nomem, nostack),
5160 );
5161 acc[k] += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5162 }
5163 }
5164 acc
5165 }
5166}
5167
5168#[inline]
5170fn q4t_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
5171 let gi = j / GROUP_SIZE;
5172 let k = j % GROUP_SIZE;
5173 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5174 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5175 let byte = tile[2 + k / 2];
5176 let nib = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5177 ((nib as i32 - 8) as f32, s)
5178}
5179
5180#[inline]
5183fn q4t_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
5184 let mut acc = 0f32;
5185 for gi in 0..gpr {
5186 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
5187 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
5188 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5189 let mut ga = 0f32;
5190 for (k, &b) in tile[2..].iter().enumerate() {
5191 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5192 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5193 }
5194 acc += ga * s;
5195 }
5196 acc
5197}
5198
5199struct Q4tpView<'a> {
5203 nib: &'a [u8],
5204 params: &'a [u8],
5205 codes: &'a [u8],
5206 stride: usize,
5207 zero_rung: bool,
5209}
5210
5211impl<'a> Q4tpView<'a> {
5212 fn new(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5213 let (params_off, codes_off, stride) = q4tp_sections(rows, cols);
5214 Self {
5215 nib: &bytes[..params_off],
5216 params: &bytes[params_off..codes_off],
5217 codes: &bytes[codes_off..],
5218 stride,
5219 zero_rung: false,
5220 }
5221 }
5222
5223 fn new_q2(bytes: &'a [u8], rows: usize, cols: usize) -> Self {
5225 let (params_off, codes_off, stride) = q2tp_sections(rows, cols);
5226 Self {
5227 nib: &bytes[..params_off],
5228 params: &bytes[params_off..codes_off],
5229 codes: &bytes[codes_off..],
5230 stride,
5231 zero_rung: true,
5232 }
5233 }
5234
5235 #[inline]
5251 fn scales_into(&self, r: usize, gpr: usize, out: &mut [f32]) {
5252 let tab = if self.zero_rung {
5253 q2tp_ladder(self.params, r)
5254 } else {
5255 q4tp_ladder(self.params, r)
5256 };
5257 let codes = &self.codes[r * self.stride..(r + 1) * self.stride];
5258 let out = &mut out[..gpr];
5259 let mut chunks = out.chunks_exact_mut(8);
5260 let mut ci = 0usize;
5261 for c in &mut chunks {
5262 let w = u64::from(codes[ci])
5263 | u64::from(codes[ci + 1]) << 8
5264 | u64::from(codes[ci + 2]) << 16
5265 | u64::from(codes[ci + 3]) << 24
5266 | u64::from(codes[ci + 4]) << 32;
5267 for (k, o) in c.iter_mut().enumerate() {
5268 *o = tab[((w >> (5 * k)) & 31) as usize];
5269 }
5270 ci += 5;
5271 }
5272 let tail = &codes[ci..];
5275 for (k, o) in chunks.into_remainder().iter_mut().enumerate() {
5276 *o = tab[q4tp_code(tail, k)];
5277 }
5278 }
5279}
5280
5281#[inline]
5282fn dot_q4tp_row_i8(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5283 #[cfg(target_arch = "aarch64")]
5284 unsafe {
5285 return dot_q4tp_row_sdot(nib, r, gpr, xq, scales);
5286 }
5287 #[cfg(target_arch = "x86_64")]
5288 unsafe {
5289 if vnni_tiles_enabled() {
5290 return dot_q4tp_row_vnni(nib, r, gpr, xq, scales);
5291 }
5292 return dot_q4tp_row_avx2(nib, r, gpr, xq, scales);
5293 }
5294 #[allow(unreachable_code)]
5295 {
5296 let mut acc = 0f32;
5297 for gi in 0..gpr {
5298 let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5299 let s = scales[gi];
5300 let mut d = 0i32;
5301 for (k, &b) in tile.iter().enumerate() {
5302 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
5303 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
5304 }
5305 acc += d as f32 * s;
5306 }
5307 acc
5308 }
5309}
5310
5311#[cfg(target_arch = "aarch64")]
5314#[target_feature(enable = "neon,dotprod")]
5315unsafe fn dot_q4tp_row_sdot(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5316 unsafe {
5319 use core::arch::aarch64::*;
5320 use core::arch::asm;
5321 let lomask = vdupq_n_u8(0x0F);
5322 let eight = vdupq_n_s8(8);
5323 let mut acc = 0f32;
5324 for gi in 0..gpr {
5325 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5326 let s = *scales.get_unchecked(gi);
5327 let b = vld1q_u8(t);
5328 let lo = vandq_u8(b, lomask);
5329 let hi = vshrq_n_u8::<4>(b);
5330 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
5331 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
5332 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
5333 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
5334 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
5335 asm!(
5336 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
5337 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
5338 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
5339 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
5340 options(pure, nomem, nostack),
5341 );
5342 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
5343 }
5344 acc
5345 }
5346}
5347
5348#[cfg(target_arch = "x86_64")]
5349#[target_feature(enable = "avx2")]
5350unsafe fn dot_q4tp_row_avx2(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5351 unsafe {
5353 use core::arch::x86_64::*;
5354 let lomask = _mm_set1_epi8(0x0F);
5355 let eight = _mm256_set1_epi8(8);
5356 let ones = _mm256_set1_epi16(1);
5357 let mut acc = 0f32;
5358 for gi in 0..gpr {
5359 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5360 let s = *scales.get_unchecked(gi);
5361 let b = _mm_loadu_si128(t as *const __m128i);
5362 let lo = _mm_and_si128(b, lomask);
5363 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5364 let w = _mm256_sub_epi8(
5365 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5366 eight,
5367 );
5368 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5369 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
5370 let d = _mm256_madd_epi16(p16, ones);
5371 let hi128 = _mm256_extracti128_si256::<1>(d);
5372 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
5373 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
5374 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
5375 acc += _mm_cvtsi128_si32(s32) as f32 * s;
5376 }
5377 acc
5378 }
5379}
5380
5381#[cfg(target_arch = "x86_64")]
5384#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
5385unsafe fn dot_q4tp_row_vnni(nib: &[u8], r: usize, gpr: usize, xq: &[i8], scales: &[f32]) -> f32 {
5386 unsafe {
5388 use core::arch::x86_64::*;
5389 let lomask = _mm_set1_epi8(0x0F);
5390 let eight = _mm256_set1_epi8(8);
5391 let mut acc = 0f32;
5392 for gi in 0..gpr {
5393 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
5394 let s = *scales.get_unchecked(gi);
5395 let b = _mm_loadu_si128(t as *const __m128i);
5396 let lo = _mm_and_si128(b, lomask);
5397 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
5398 let w = _mm256_sub_epi8(
5399 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
5400 eight,
5401 );
5402 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
5403 acc += dpbusd_hsum(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w)) as f32 * s;
5404 }
5405 acc
5406 }
5407}
5408
5409#[inline]
5412fn q4tp_row_exact(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5413 #[cfg(target_arch = "x86_64")]
5414 if avx2_enabled() {
5415 return unsafe { q4tp_row_float_avx2(nib, r, gpr, x, scales) };
5417 }
5418 q4tp_row_float_scalar(nib, r, gpr, x, scales)
5419}
5420
5421#[inline]
5422fn q4tp_row_float_scalar(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5423 let mut acc = 0f32;
5424 for gi in 0..gpr {
5425 let tile = &nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
5426 let s = scales[gi];
5427 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5428 let mut ga = 0f32;
5429 for (k, &b) in tile.iter().enumerate() {
5430 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
5431 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
5432 }
5433 acc += ga * s;
5434 }
5435 acc
5436}
5437
5438#[cfg(target_arch = "x86_64")]
5442#[target_feature(enable = "avx2")]
5443unsafe fn q4tp_row_float_avx2(nib: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5444 unsafe {
5447 use core::arch::x86_64::*;
5448 let mask = _mm_set1_epi8(15);
5449 let eight = _mm_set1_epi8(8);
5450 let order = _mm256_setr_epi32(0, 1, 4, 5, 2, 3, 6, 7);
5451 let mut acc = 0.0f32;
5452 for gi in 0..gpr {
5453 let packed = _mm_loadu_si128(nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB).cast());
5454 let lo = _mm_and_si128(packed, mask);
5455 let hi = _mm_and_si128(_mm_srli_epi16::<4>(packed), mask);
5456 let w0 = _mm_sub_epi8(_mm_unpacklo_epi8(lo, hi), eight);
5457 let w1 = _mm_sub_epi8(_mm_unpackhi_epi8(lo, hi), eight);
5458 let xp = x.as_ptr().add(gi * GROUP_SIZE);
5459 let a = _mm256_mul_ps(
5460 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w0)),
5461 _mm256_loadu_ps(xp),
5462 );
5463 let b = _mm256_mul_ps(
5464 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w0))),
5465 _mm256_loadu_ps(xp.add(8)),
5466 );
5467 let c = _mm256_mul_ps(
5468 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(w1)),
5469 _mm256_loadu_ps(xp.add(16)),
5470 );
5471 let d = _mm256_mul_ps(
5472 _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(w1))),
5473 _mm256_loadu_ps(xp.add(24)),
5474 );
5475 let mut pairs = [0.0f32; 16];
5476 _mm256_storeu_ps(
5477 pairs.as_mut_ptr(),
5478 _mm256_permutevar8x32_ps(_mm256_hadd_ps(a, b), order),
5479 );
5480 _mm256_storeu_ps(
5481 pairs.as_mut_ptr().add(8),
5482 _mm256_permutevar8x32_ps(_mm256_hadd_ps(c, d), order),
5483 );
5484 let mut ga = 0.0f32;
5485 for v in pairs {
5486 ga += v;
5487 }
5488 acc += ga * scales[gi];
5489 }
5490 acc
5491 }
5492}
5493
5494#[inline]
5497fn q4tp_outlier(nib: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5498 let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5499 let byte = nib[(r * gpr + gi) * Q4TP_NIB + k / 2];
5500 let n = if k & 1 == 0 { byte & 0x0F } else { byte >> 4 };
5501 ((n as i32 - 8) as f32, scales[gi])
5502}
5503
5504fn q4tp_matvec(
5506 bytes: &[u8],
5507 x: &[f32],
5508 rows: usize,
5509 cols: usize,
5510 out: &mut [f32],
5511 pool: Option<&Pool>,
5512) {
5513 debug_assert_eq!(out.len(), rows);
5514 let gpr = cols / GROUP_SIZE;
5515 let v = Q4tpView::new(bytes, rows, cols);
5516 let out_addr = SendMut(out.as_mut_ptr());
5517 if a8w8_enabled() {
5518 let act = split_act(x);
5519 let run = |start: usize, end: usize| {
5520 with_krow(gpr, |sc| {
5522 for r in start..end {
5523 v.scales_into(r, gpr, sc);
5524 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, sc) * act.sx;
5525 for &(j, xv) in &act.outliers {
5526 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, sc);
5527 acc += w * s * xv;
5528 }
5529 unsafe { *out_addr.at(r) = acc };
5531 }
5532 })
5533 };
5534 dispatch_rows(pool, rows, &run);
5535 return;
5536 }
5537 let run = |start: usize, end: usize| {
5538 with_krow(gpr, |sc| {
5539 for r in start..end {
5540 v.scales_into(r, gpr, sc);
5541 unsafe { *out_addr.at(r) = q4tp_row_exact(v.nib, r, gpr, x, sc) };
5543 }
5544 })
5545 };
5546 dispatch_rows(pool, rows, &run);
5547}
5548
5549#[allow(clippy::too_many_arguments)]
5552fn q4tp_matvec2(
5553 bytes: &[u8],
5554 x1: &[f32],
5555 x2: &[f32],
5556 rows: usize,
5557 cols: usize,
5558 o1: &mut [f32],
5559 o2: &mut [f32],
5560 pool: Option<&Pool>,
5561) {
5562 let gpr = cols / GROUP_SIZE;
5563 let v = Q4tpView::new(bytes, rows, cols);
5564 let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5565 let run = |start: usize, end: usize| {
5566 let mut sc = vec![0f32; gpr];
5567 for r in start..end {
5568 v.scales_into(r, gpr, &mut sc);
5569 unsafe {
5571 *p1.at(r) = q4tp_row_exact(v.nib, r, gpr, x1, &sc);
5572 *p2.at(r) = q4tp_row_exact(v.nib, r, gpr, x2, &sc);
5573 }
5574 }
5575 };
5576 dispatch_rows(pool, rows, &run);
5577}
5578
5579#[inline]
5582fn q2tp_outlier(chunks: &[u8], r: usize, gpr: usize, j: usize, scales: &[f32]) -> (f32, f32) {
5583 let (gi, k) = (j / GROUP_SIZE, j % GROUP_SIZE);
5584 let byte = chunks[(r * gpr + gi) * Q2TP_CHUNK + k / 4];
5585 let c = (byte >> (2 * (k % 4))) & 3;
5586 (c as f32 - 1.5, scales[gi])
5587}
5588
5589#[cfg(target_arch = "x86_64")]
5590const Q2TP_DECODE_U32: [u32; 256] = {
5591 let mut tab = [0u32; 256];
5592 let mut b = 0usize;
5593 while b < 256 {
5594 tab[b] = ((b as u32) & 3)
5595 | ((((b as u32) >> 2) & 3) << 8)
5596 | ((((b as u32) >> 4) & 3) << 16)
5597 | ((((b as u32) >> 6) & 3) << 24);
5598 b += 1;
5599 }
5600 tab
5601};
5602
5603#[cfg(target_arch = "x86_64")]
5607#[target_feature(enable = "avx2")]
5608unsafe fn q2tp_code_dot_avx2(ch: &[u8], x: &[i8]) -> i32 {
5609 use core::arch::x86_64::*;
5610 debug_assert!(ch.len() >= Q2TP_CHUNK && x.len() >= GROUP_SIZE);
5611 let codes = _mm256_setr_epi32(
5612 Q2TP_DECODE_U32[ch[0] as usize] as i32,
5613 Q2TP_DECODE_U32[ch[1] as usize] as i32,
5614 Q2TP_DECODE_U32[ch[2] as usize] as i32,
5615 Q2TP_DECODE_U32[ch[3] as usize] as i32,
5616 Q2TP_DECODE_U32[ch[4] as usize] as i32,
5617 Q2TP_DECODE_U32[ch[5] as usize] as i32,
5618 Q2TP_DECODE_U32[ch[6] as usize] as i32,
5619 Q2TP_DECODE_U32[ch[7] as usize] as i32,
5620 );
5621 let xv = unsafe { _mm256_loadu_si256(x.as_ptr().cast()) };
5622 let pair = _mm256_maddubs_epi16(codes, xv);
5623 let quad = _mm256_madd_epi16(pair, _mm256_set1_epi16(1));
5624 let sum128 = _mm_add_epi32(
5625 _mm256_castsi256_si128(quad),
5626 _mm256_extracti128_si256(quad, 1),
5627 );
5628 let sum64 = _mm_hadd_epi32(sum128, sum128);
5629 _mm_cvtsi128_si32(_mm_hadd_epi32(sum64, sum64))
5630}
5631
5632#[cfg(target_arch = "x86_64")]
5633#[inline]
5634fn q2tp_avx2_enabled() -> bool {
5635 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
5636 *ON.get_or_init(|| std::arch::is_x86_feature_detected!("avx2"))
5637}
5638
5639#[inline]
5646fn dot_q2tp_row_i8(
5647 chunks: &[u8],
5648 r: usize,
5649 gpr: usize,
5650 xq: &[i8],
5651 gsum: &[i32],
5652 scales: &[f32],
5653) -> f32 {
5654 let mut acc = 0f32;
5655 let base = r * gpr * Q2TP_CHUNK;
5656 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5657 let mut codes = [0i8; GROUP_SIZE];
5658 #[cfg(target_arch = "x86_64")]
5659 let avx2 = q2tp_avx2_enabled();
5664 for gi in 0..gpr {
5665 let ch = &chunks[base + gi * Q2TP_CHUNK..base + (gi + 1) * Q2TP_CHUNK];
5666 let xg = &xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5667 #[cfg(target_arch = "aarch64")]
5668 let dot = unsafe {
5674 use core::arch::aarch64::*;
5675 let b = vld1_u8(ch.as_ptr());
5676 let three = vdup_n_u8(3);
5677 let c0 = vreinterpret_s8_u8(vand_u8(b, three));
5678 let c1 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 2), three));
5679 let c2 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 4), three));
5680 let c3 = vreinterpret_s8_u8(vand_u8(vshr_n_u8(b, 6), three));
5681 let x4 = vld4_s8(xg.as_ptr());
5682 let mut acc4 = vdupq_n_s32(0);
5683 acc4 = vpadalq_s16(acc4, vmull_s8(c0, x4.0));
5684 acc4 = vpadalq_s16(acc4, vmull_s8(c1, x4.1));
5685 acc4 = vpadalq_s16(acc4, vmull_s8(c2, x4.2));
5686 acc4 = vpadalq_s16(acc4, vmull_s8(c3, x4.3));
5687 vaddvq_s32(acc4)
5688 };
5689 #[cfg(target_arch = "x86_64")]
5690 let dot: i32 = if avx2 {
5691 unsafe { q2tp_code_dot_avx2(ch, xg) }
5694 } else {
5695 ch.iter()
5696 .enumerate()
5697 .map(|(k, &b)| {
5698 ((b & 3) as i32) * xg[k * 4] as i32
5699 + (((b >> 2) & 3) as i32) * xg[k * 4 + 1] as i32
5700 + (((b >> 4) & 3) as i32) * xg[k * 4 + 2] as i32
5701 + (((b >> 6) & 3) as i32) * xg[k * 4 + 3] as i32
5702 })
5703 .sum()
5704 };
5705 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
5706 let dot: i32 = {
5707 for (k, &b) in ch.iter().enumerate() {
5708 codes[k * 4] = (b & 3) as i8;
5709 codes[k * 4 + 1] = ((b >> 2) & 3) as i8;
5710 codes[k * 4 + 2] = ((b >> 4) & 3) as i8;
5711 codes[k * 4 + 3] = ((b >> 6) & 3) as i8;
5712 }
5713 codes
5714 .iter()
5715 .zip(xg)
5716 .map(|(&c, &x)| c as i32 * x as i32)
5717 .sum()
5718 };
5719 acc += scales[gi] * (dot as f32 - 1.5 * gsum[gi] as f32);
5720 }
5721 acc
5722}
5723
5724fn q2tp_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5728 q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.5)
5729}
5730
5731#[inline]
5735fn q2tp_affine_row_exact(chunks: &[u8], r: usize, gpr: usize, x: &[f32], scales: &[f32]) -> f32 {
5736 q2tp_row_exact_center(chunks, r, gpr, x, scales, 1.0)
5737}
5738
5739#[inline]
5740fn q2tp_row_exact_center(
5741 chunks: &[u8],
5742 r: usize,
5743 gpr: usize,
5744 x: &[f32],
5745 scales: &[f32],
5746 center: f32,
5747) -> f32 {
5748 let mut acc = 0f32;
5749 for gi in 0..gpr {
5750 let ch = &chunks[(r * gpr + gi) * Q2TP_CHUNK..(r * gpr + gi + 1) * Q2TP_CHUNK];
5751 let s = scales[gi];
5752 let xb = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
5753 let mut g = 0f32;
5754 for (k, &b) in ch.iter().enumerate() {
5755 g += ((b & 3) as f32 - center) * xb[k * 4]
5756 + (((b >> 2) & 3) as f32 - center) * xb[k * 4 + 1]
5757 + (((b >> 4) & 3) as f32 - center) * xb[k * 4 + 2]
5758 + (((b >> 6) & 3) as f32 - center) * xb[k * 4 + 3];
5759 }
5760 acc += s * g;
5761 }
5762 acc
5763}
5764
5765fn q2tp_matvec(
5766 bytes: &[u8],
5767 x: &[f32],
5768 rows: usize,
5769 cols: usize,
5770 out: &mut [f32],
5771 pool: Option<&Pool>,
5772) {
5773 q2tp_matvec_mode(bytes, x, rows, cols, out, pool, false);
5774}
5775
5776fn q2tp_affine_matvec(
5777 bytes: &[u8],
5778 x: &[f32],
5779 rows: usize,
5780 cols: usize,
5781 out: &mut [f32],
5782 pool: Option<&Pool>,
5783) {
5784 q2tp_matvec_mode(bytes, x, rows, cols, out, pool, true);
5785}
5786
5787fn q2tp_matvec_mode(
5788 bytes: &[u8],
5789 x: &[f32],
5790 rows: usize,
5791 cols: usize,
5792 out: &mut [f32],
5793 pool: Option<&Pool>,
5794 affine: bool,
5795) {
5796 debug_assert_eq!(out.len(), rows);
5797 let gpr = cols / GROUP_SIZE;
5798 let v = Q4tpView::new_q2(bytes, rows, cols);
5799 let out_addr = SendMut(out.as_mut_ptr());
5800 if !affine && a8w8_enabled() {
5805 let act = split_act(x);
5806 let gsum = q1_group_sums(&act.xq, gpr);
5807 let (act, gsum) = (&act, &gsum);
5808 let run = move |start: usize, end: usize| {
5809 with_krow(gpr, |sc| {
5810 for r in start..end {
5811 v.scales_into(r, gpr, sc);
5812 let mut acc = dot_q2tp_row_i8(v.nib, r, gpr, &act.xq, gsum, sc) * act.sx;
5813 for &(j, xv) in &act.outliers {
5814 let (w, s) = q2tp_outlier(v.nib, r, gpr, j, sc);
5815 acc += w * s * xv;
5816 }
5817 unsafe { *out_addr.at(r) = acc };
5819 }
5820 })
5821 };
5822 dispatch_rows(pool, rows, &run);
5823 return;
5824 }
5825 let run = |start: usize, end: usize| {
5826 with_krow(gpr, |sc| {
5827 for r in start..end {
5828 v.scales_into(r, gpr, sc);
5829 unsafe {
5831 *out_addr.at(r) = if affine {
5832 q2tp_affine_row_exact(v.nib, r, gpr, x, sc)
5833 } else {
5834 q2tp_row_exact(v.nib, r, gpr, x, sc)
5835 }
5836 };
5837 }
5838 })
5839 };
5840 dispatch_rows(pool, rows, &run);
5841}
5842
5843#[allow(clippy::too_many_arguments)]
5845fn q2tp_matvec2(
5846 bytes: &[u8],
5847 x1: &[f32],
5848 x2: &[f32],
5849 rows: usize,
5850 cols: usize,
5851 o1: &mut [f32],
5852 o2: &mut [f32],
5853 pool: Option<&Pool>,
5854) {
5855 q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, false);
5856}
5857
5858#[allow(clippy::too_many_arguments)]
5859fn q2tp_affine_matvec2(
5860 bytes: &[u8],
5861 x1: &[f32],
5862 x2: &[f32],
5863 rows: usize,
5864 cols: usize,
5865 o1: &mut [f32],
5866 o2: &mut [f32],
5867 pool: Option<&Pool>,
5868) {
5869 q2tp_matvec2_mode(bytes, x1, x2, rows, cols, o1, o2, pool, true);
5870}
5871
5872#[allow(clippy::too_many_arguments)]
5873fn q2tp_matvec2_mode(
5874 bytes: &[u8],
5875 x1: &[f32],
5876 x2: &[f32],
5877 rows: usize,
5878 cols: usize,
5879 o1: &mut [f32],
5880 o2: &mut [f32],
5881 pool: Option<&Pool>,
5882 affine: bool,
5883) {
5884 let gpr = cols / GROUP_SIZE;
5885 let v = Q4tpView::new_q2(bytes, rows, cols);
5886 let (p1, p2) = (SendMut(o1.as_mut_ptr()), SendMut(o2.as_mut_ptr()));
5887 let run = |start: usize, end: usize| {
5888 let mut sc = vec![0f32; gpr];
5889 for r in start..end {
5890 v.scales_into(r, gpr, &mut sc);
5891 unsafe {
5893 *p1.at(r) = if affine {
5894 q2tp_affine_row_exact(v.nib, r, gpr, x1, &sc)
5895 } else {
5896 q2tp_row_exact(v.nib, r, gpr, x1, &sc)
5897 };
5898 *p2.at(r) = if affine {
5899 q2tp_affine_row_exact(v.nib, r, gpr, x2, &sc)
5900 } else {
5901 q2tp_row_exact(v.nib, r, gpr, x2, &sc)
5902 };
5903 }
5904 }
5905 };
5906 dispatch_rows(pool, rows, &run);
5907}
5908
5909pub fn q2tp_matvec_for_test(bytes: &[u8], x: &[f32], rows: usize, cols: usize, out: &mut [f32]) {
5916 let gpr = cols / GROUP_SIZE;
5921 let v = Q4tpView::new_q2(bytes, rows, cols);
5922 with_krow(gpr, |sc| {
5923 for r in 0..rows {
5924 v.scales_into(r, gpr, sc);
5925 out[r] = q2tp_row_exact(v.nib, r, gpr, x, sc);
5926 }
5927 });
5928}
5929
5930pub fn q2tp_affine_matvec_for_test(
5933 bytes: &[u8],
5934 x: &[f32],
5935 rows: usize,
5936 cols: usize,
5937 out: &mut [f32],
5938) {
5939 q2tp_affine_matvec(bytes, x, rows, cols, out, None);
5940}
5941
5942pub fn q2tp_matmat_for_test(
5943 bytes: &[u8],
5944 xs_all: &[f32],
5945 b: usize,
5946 rows: usize,
5947 cols: usize,
5948 out: &mut [f32],
5949) {
5950 q2tp_matmat(bytes, xs_all, b, rows, cols, out, None);
5951}
5952
5953fn q2tp_matmat(
5954 bytes: &[u8],
5955 xs_all: &[f32],
5956 b: usize,
5957 rows: usize,
5958 cols: usize,
5959 out: &mut [f32],
5960 pool: Option<&Pool>,
5961) {
5962 q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, false);
5963}
5964
5965fn q2tp_affine_matmat(
5966 bytes: &[u8],
5967 xs_all: &[f32],
5968 b: usize,
5969 rows: usize,
5970 cols: usize,
5971 out: &mut [f32],
5972 pool: Option<&Pool>,
5973) {
5974 q2tp_matmat_mode(bytes, xs_all, b, rows, cols, out, pool, true);
5975}
5976
5977fn q2tp_matmat_mode(
5978 bytes: &[u8],
5979 xs_all: &[f32],
5980 b: usize,
5981 rows: usize,
5982 cols: usize,
5983 out: &mut [f32],
5984 pool: Option<&Pool>,
5985 affine: bool,
5986) {
5987 debug_assert_eq!(out.len(), b * rows);
5988 let gpr = cols / GROUP_SIZE;
5989 let v = Q4tpView::new_q2(bytes, rows, cols);
5990 let out_addr = SendMut(out.as_mut_ptr());
5991 let run = |start: usize, end: usize| {
5992 let mut sc = vec![0f32; gpr];
5993 for r in start..end {
5994 v.scales_into(r, gpr, &mut sc);
5995 for bi in 0..b {
5996 let x = &xs_all[bi * cols..(bi + 1) * cols];
5997 unsafe {
5999 *out_addr.at(bi * rows + r) = if affine {
6000 q2tp_affine_row_exact(v.nib, r, gpr, x, &sc)
6001 } else {
6002 q2tp_row_exact(v.nib, r, gpr, x, &sc)
6003 }
6004 };
6005 }
6006 }
6007 };
6008 dispatch_rows(pool, rows, &run);
6009}
6010
6011#[cfg(target_arch = "aarch64")]
6019#[target_feature(enable = "neon,dotprod")]
6020unsafe fn dot_q4tp_row_1x4_sdot_v1(
6021 nib: &[u8],
6022 r: usize,
6023 gpr: usize,
6024 xs: [&[i8]; 4],
6025 scales: &[f32],
6026) -> [f32; 4] {
6027 unsafe {
6028 use core::arch::aarch64::*;
6029 use core::arch::asm;
6030 let lomask = vdupq_n_u8(0x0F);
6031 let eight = vdupq_n_s8(8);
6032 let (mut f0, mut f1, mut f2, mut f3) = (0f32, 0f32, 0f32, 0f32);
6033 for gi in 0..gpr {
6034 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6035 let s = *scales.get_unchecked(gi);
6036 let bb = vld1q_u8(t);
6037 let lo = vandq_u8(bb, lomask);
6038 let hi = vshrq_n_u8::<4>(bb);
6039 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6040 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6041 let mut d = [0f32; 4];
6042 for (k, dk) in d.iter_mut().enumerate() {
6043 let x0 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE));
6044 let x1 = vld1q_s8(xs[k].as_ptr().add(gi * GROUP_SIZE + 16));
6045 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6046 asm!(
6047 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6048 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6049 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6050 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6051 options(pure, nomem, nostack),
6052 );
6053 *dk = vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
6054 }
6055 f0 += d[0];
6056 f1 += d[1];
6057 f2 += d[2];
6058 f3 += d[3];
6059 }
6060 [f0, f1, f2, f3]
6061 }
6062}
6063
6064#[allow(dead_code)]
6071static Q4TP_ALT: std::sync::atomic::AtomicU8 = std::sync::atomic::AtomicU8::new(0);
6072
6073#[cfg(test)]
6076static Q4TP_ALT_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
6077
6078#[cfg(target_arch = "x86_64")]
6084fn q4tp_blocked_x86() -> bool {
6085 match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6086 1 => false,
6087 2 => avx512vnni_enabled(),
6092 _ => blocked_enabled() && avx512vnni_enabled(),
6096 }
6097}
6098
6099#[cfg(target_arch = "aarch64")]
6101#[allow(dead_code)]
6102fn q4tp_v1() -> bool {
6103 match Q4TP_ALT.load(std::sync::atomic::Ordering::Relaxed) {
6104 1 => true,
6105 2 => false,
6106 _ => {
6107 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
6108 *ON.get_or_init(|| std::env::var("CMF_Q4TP_V1").is_ok_and(|v| v != "0"))
6109 }
6110 }
6111}
6112
6113#[cfg(target_arch = "x86_64")]
6125#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6126unsafe fn dot_q4tp_2x8_avx512(
6127 nib: &[u8],
6128 r0: usize,
6129 gpr: usize,
6130 xs: [&[i8]; 8],
6131 sc0: &[f32],
6132 sc1: &[f32],
6133) -> [[f32; 8]; 2] {
6134 unsafe {
6137 use core::arch::x86_64::*;
6138 let lomask = _mm256_set1_epi8(0x0F);
6139 let eight = _mm256_set1_epi8(8);
6140 let zero = _mm512_setzero_si512();
6141 let mut v0 = [_mm512_setzero_ps(); 8];
6142 let mut v1 = [_mm512_setzero_ps(); 8];
6143 let pairs = gpr / 2;
6144 let unpack = |r: usize, gi: usize| -> (__m512i, __mmask64) {
6145 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6146 let bb = _mm256_loadu_si256(t as *const __m256i);
6147 let lo = _mm256_and_si256(bb, lomask);
6148 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6149 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6150 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6151 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6152 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6153 (_mm512_abs_epi8(w), _mm512_movepi8_mask(w))
6154 };
6155 for gp in 0..pairs {
6156 let gi = gp * 2;
6157 let (wa0, neg0) = unpack(r0, gi);
6158 let (wa1, neg1) = unpack(r0 + 1, gi);
6159 let off = gi * GROUP_SIZE;
6160 let sv = |sc: &[f32]| {
6161 _mm512_insertf32x8::<1>(
6162 _mm512_castps256_ps512(_mm256_set1_ps(*sc.get_unchecked(gi))),
6163 _mm256_set1_ps(*sc.get_unchecked(gi + 1)),
6164 )
6165 };
6166 let s0 = sv(sc0);
6167 let s1 = sv(sc1);
6168 for k in 0..8 {
6169 let xv = _mm512_loadu_si512(xs[k].as_ptr().add(off) as *const __m512i);
6170 let d0 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6171 zero,
6172 wa0,
6173 _mm512_mask_sub_epi8(xv, neg0, zero, xv),
6174 ));
6175 let d1 = _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(
6176 zero,
6177 wa1,
6178 _mm512_mask_sub_epi8(xv, neg1, zero, xv),
6179 ));
6180 v0[k] = _mm512_fmadd_ps(d0, s0, v0[k]);
6181 v1[k] = _mm512_fmadd_ps(d1, s1, v1[k]);
6182 }
6183 }
6184 let mut acc = [[0f32; 8]; 2];
6185 for k in 0..8 {
6186 acc[0][k] = _mm512_reduce_add_ps(v0[k]);
6187 acc[1][k] = _mm512_reduce_add_ps(v1[k]);
6188 }
6189 if gpr % 2 == 1 {
6190 let off = (gpr - 1) * GROUP_SIZE;
6191 for j in off..off + GROUP_SIZE {
6192 let (w0, sa) = q4tp_outlier(nib, r0, gpr, j, sc0);
6193 let (w1, sb) = q4tp_outlier(nib, r0 + 1, gpr, j, sc1);
6194 for k in 0..8 {
6195 let x = *xs[k].get_unchecked(j) as f32;
6196 acc[0][k] += w0 * sa * x;
6197 acc[1][k] += w1 * sb * x;
6198 }
6199 }
6200 }
6201 acc
6202 }
6203}
6204
6205#[cfg(target_arch = "x86_64")]
6210#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6211unsafe fn dot_q4tp_row_1x8_avx512(
6212 nib: &[u8],
6213 r: usize,
6214 gpr: usize,
6215 xs: [&[i8]; 8],
6216 scales: &[f32],
6217) -> [f32; 8] {
6218 unsafe {
6220 use core::arch::x86_64::*;
6221 let lomask = _mm256_set1_epi8(0x0F);
6222 let eight = _mm256_set1_epi8(8);
6223 let zero = _mm512_setzero_si512();
6224 let (mut v0, mut v1, mut v2, mut v3) = (
6225 _mm512_setzero_ps(),
6226 _mm512_setzero_ps(),
6227 _mm512_setzero_ps(),
6228 _mm512_setzero_ps(),
6229 );
6230 let (mut v4, mut v5, mut v6, mut v7) = (
6231 _mm512_setzero_ps(),
6232 _mm512_setzero_ps(),
6233 _mm512_setzero_ps(),
6234 _mm512_setzero_ps(),
6235 );
6236 let pairs = gpr / 2;
6237 for gp in 0..pairs {
6238 let gi = gp * 2;
6239 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6240 let bb = _mm256_loadu_si256(t as *const __m256i);
6241 let lo = _mm256_and_si256(bb, lomask);
6242 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6243 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6248 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6249 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6250 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6251 let wabs = _mm512_abs_epi8(w);
6252 let neg = _mm512_movepi8_mask(w);
6253 let off = gi * GROUP_SIZE;
6254 let sv = _mm512_insertf32x8::<1>(
6255 _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6256 _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6257 );
6258 let dot = |x: &[i8]| -> __m512 {
6259 let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6260 let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6261 _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6262 };
6263 v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6264 v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6265 v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6266 v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6267 v4 = _mm512_fmadd_ps(dot(xs[4]), sv, v4);
6268 v5 = _mm512_fmadd_ps(dot(xs[5]), sv, v5);
6269 v6 = _mm512_fmadd_ps(dot(xs[6]), sv, v6);
6270 v7 = _mm512_fmadd_ps(dot(xs[7]), sv, v7);
6271 }
6272 let mut acc = [
6273 _mm512_reduce_add_ps(v0),
6274 _mm512_reduce_add_ps(v1),
6275 _mm512_reduce_add_ps(v2),
6276 _mm512_reduce_add_ps(v3),
6277 _mm512_reduce_add_ps(v4),
6278 _mm512_reduce_add_ps(v5),
6279 _mm512_reduce_add_ps(v6),
6280 _mm512_reduce_add_ps(v7),
6281 ];
6282 if gpr % 2 == 1 {
6285 let off = (gpr - 1) * GROUP_SIZE;
6286 for j in off..off + GROUP_SIZE {
6287 let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6288 let ws = w * s;
6289 for k in 0..8 {
6290 acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6291 }
6292 }
6293 }
6294 acc
6295 }
6296}
6297
6298#[cfg(target_arch = "x86_64")]
6311#[target_feature(enable = "avx512f,avx512bw,avx512vnni")]
6312unsafe fn dot_q4tp_row_1x4_avx512(
6313 nib: &[u8],
6314 r: usize,
6315 gpr: usize,
6316 xs: [&[i8]; 4],
6317 scales: &[f32],
6318) -> [f32; 4] {
6319 unsafe {
6321 use core::arch::x86_64::*;
6322 let lomask = _mm256_set1_epi8(0x0F);
6323 let eight = _mm256_set1_epi8(8);
6324 let zero = _mm512_setzero_si512();
6325 let (mut v0, mut v1, mut v2, mut v3) = (
6326 _mm512_setzero_ps(),
6327 _mm512_setzero_ps(),
6328 _mm512_setzero_ps(),
6329 _mm512_setzero_ps(),
6330 );
6331 let pairs = gpr / 2;
6332 for gp in 0..pairs {
6333 let gi = gp * 2;
6334 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6335 let bb = _mm256_loadu_si256(t as *const __m256i);
6336 let lo = _mm256_and_si256(bb, lomask);
6337 let hi = _mm256_and_si256(_mm256_srli_epi16::<4>(bb), lomask);
6338 let ul = _mm256_sub_epi8(_mm256_unpacklo_epi8(lo, hi), eight);
6343 let uh = _mm256_sub_epi8(_mm256_unpackhi_epi8(lo, hi), eight);
6344 let cat = _mm512_inserti64x4::<1>(_mm512_castsi256_si512(ul), uh);
6345 let w = _mm512_shuffle_i64x2::<0b11_01_10_00>(cat, cat);
6346 let wabs = _mm512_abs_epi8(w);
6347 let neg = _mm512_movepi8_mask(w);
6348 let off = gi * GROUP_SIZE;
6349 let sv = _mm512_insertf32x8::<1>(
6350 _mm512_castps256_ps512(_mm256_set1_ps(*scales.get_unchecked(gi))),
6351 _mm256_set1_ps(*scales.get_unchecked(gi + 1)),
6352 );
6353 let dot = |x: &[i8]| -> __m512 {
6354 let xv = _mm512_loadu_si512(x.as_ptr().add(off) as *const __m512i);
6355 let sx = _mm512_mask_sub_epi8(xv, neg, zero, xv);
6356 _mm512_cvtepi32_ps(_mm512_dpbusd_epi32(zero, wabs, sx))
6357 };
6358 v0 = _mm512_fmadd_ps(dot(xs[0]), sv, v0);
6359 v1 = _mm512_fmadd_ps(dot(xs[1]), sv, v1);
6360 v2 = _mm512_fmadd_ps(dot(xs[2]), sv, v2);
6361 v3 = _mm512_fmadd_ps(dot(xs[3]), sv, v3);
6362 }
6363 let mut acc = [
6364 _mm512_reduce_add_ps(v0),
6365 _mm512_reduce_add_ps(v1),
6366 _mm512_reduce_add_ps(v2),
6367 _mm512_reduce_add_ps(v3),
6368 ];
6369 if gpr % 2 == 1 {
6372 let off = (gpr - 1) * GROUP_SIZE;
6373 for j in off..off + GROUP_SIZE {
6374 let (w, s) = q4tp_outlier(nib, r, gpr, j, scales);
6375 let ws = w * s;
6376 for k in 0..4 {
6377 acc[k] += ws * *xs[k].get_unchecked(j) as f32;
6378 }
6379 }
6380 }
6381 acc
6382 }
6383}
6384
6385#[cfg(target_arch = "aarch64")]
6389#[target_feature(enable = "neon,dotprod")]
6390unsafe fn dot_q4tp_row_1x4_sdot(
6391 nib: &[u8],
6392 r: usize,
6393 gpr: usize,
6394 xs: [&[i8]; 4],
6395 scales: &[f32],
6396) -> [f32; 4] {
6397 unsafe {
6399 use core::arch::aarch64::*;
6400 use core::arch::asm;
6401 let lomask = vdupq_n_u8(0x0F);
6402 let eight = vdupq_n_s8(8);
6403 let (mut v0, mut v1, mut v2, mut v3) = (
6419 vdupq_n_f32(0.0),
6420 vdupq_n_f32(0.0),
6421 vdupq_n_f32(0.0),
6422 vdupq_n_f32(0.0),
6423 );
6424 for gi in 0..gpr {
6425 let t = nib.as_ptr().add((r * gpr + gi) * Q4TP_NIB);
6426 let s = *scales.get_unchecked(gi);
6427 let bb = vld1q_u8(t);
6428 let lo = vandq_u8(bb, lomask);
6429 let hi = vshrq_n_u8::<4>(bb);
6430 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
6431 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
6432 let off = gi * GROUP_SIZE;
6433 let dot4 = |x: &[i8]| -> int32x4_t {
6434 let x0 = vld1q_s8(x.as_ptr().add(off));
6435 let x1 = vld1q_s8(x.as_ptr().add(off + 16));
6436 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
6437 asm!(
6438 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
6439 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
6440 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
6441 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
6442 options(pure, nomem, nostack),
6443 );
6444 vaddq_s32(a0, a1)
6445 };
6446 v0 = vfmaq_n_f32(v0, vcvtq_f32_s32(dot4(xs[0])), s);
6447 v1 = vfmaq_n_f32(v1, vcvtq_f32_s32(dot4(xs[1])), s);
6448 v2 = vfmaq_n_f32(v2, vcvtq_f32_s32(dot4(xs[2])), s);
6449 v3 = vfmaq_n_f32(v3, vcvtq_f32_s32(dot4(xs[3])), s);
6450 }
6451 [
6452 vaddvq_f32(v0),
6453 vaddvq_f32(v1),
6454 vaddvq_f32(v2),
6455 vaddvq_f32(v3),
6456 ]
6457 }
6458}
6459
6460fn q4tp_matmat(
6468 bytes: &[u8],
6469 xs_all: &[f32],
6470 b: usize,
6471 rows: usize,
6472 cols: usize,
6473 out: &mut [f32],
6474 pool: Option<&Pool>,
6475) {
6476 q4tp_matmat_with(bytes, xs_all, b, rows, cols, out, pool, row_exact())
6477}
6478
6479#[allow(clippy::too_many_arguments)]
6483fn q4tp_matmat_with(
6484 bytes: &[u8],
6485 xs_all: &[f32],
6486 b: usize,
6487 rows: usize,
6488 cols: usize,
6489 out: &mut [f32],
6490 pool: Option<&Pool>,
6491 exact: bool,
6492) {
6493 debug_assert_eq!(out.len(), b * rows);
6494 let gpr = cols / GROUP_SIZE;
6495 let v = Q4tpView::new(bytes, rows, cols);
6496
6497 #[cfg(target_os = "macos")]
6501 if !exact && b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6502 dequant_matmat_accel(
6503 &|r, dst| {
6504 let mut sc = [0f32; 32];
6505 let mut scv;
6506 let s: &[f32] = if gpr <= 32 {
6507 v.scales_into(r, gpr, &mut sc);
6508 &sc[..gpr]
6509 } else {
6510 scv = vec![0f32; gpr];
6511 v.scales_into(r, gpr, &mut scv);
6512 &scv
6513 };
6514 for gi in 0..gpr {
6515 let tile = &v.nib[(r * gpr + gi) * Q4TP_NIB..(r * gpr + gi + 1) * Q4TP_NIB];
6516 for (k, &bb) in tile.iter().enumerate() {
6517 dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s[gi];
6518 dst[gi * GROUP_SIZE + k * 2 + 1] =
6519 (((bb >> 4) & 0x0F) as f32 - 8.0) * s[gi];
6520 }
6521 }
6522 },
6523 xs_all,
6524 b,
6525 rows,
6526 cols,
6527 out,
6528 pool,
6529 );
6530 return;
6531 }
6532
6533 let out_addr = SendMut(out.as_mut_ptr());
6534 if a8w8_enabled() {
6535 let acts: Vec<SplitAct> = (0..b)
6536 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6537 .collect();
6538 let acts = &acts;
6539 #[cfg(target_arch = "aarch64")]
6546 let blocked_ok = sdot_enabled() && blocked_enabled();
6547 #[cfg(target_arch = "x86_64")]
6554 let blocked_ok = q4tp_blocked_x86() && !exact;
6555 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
6556 let blocked_ok = {
6557 let _ = exact;
6558 false
6559 };
6560 let panel_cols: usize = std::env::var("CMF_Q4TP_PANEL")
6570 .ok()
6571 .and_then(|v| v.parse().ok())
6572 .filter(|v| *v > 0)
6573 .unwrap_or(256);
6574 let run = |start: usize, end: usize| {
6575 for abase in (0..acts.len()).step_by(panel_cols) {
6576 let alen = (acts.len() - abase).min(panel_cols);
6577 let mut sc = vec![0f32; gpr];
6578 #[cfg(target_arch = "x86_64")]
6579 let mut r_lo = start;
6580 #[cfg(target_arch = "x86_64")]
6581 if blocked_ok && alen >= 8 {
6582 let mut sc1 = vec![0f32; gpr];
6583 while r_lo + 2 <= end {
6584 v.scales_into(r_lo, gpr, &mut sc);
6585 v.scales_into(r_lo + 1, gpr, &mut sc1);
6586 let mut bi = 0usize;
6587 while bi + 8 <= alen {
6588 let xs = [
6589 acts[abase + bi].xq.as_slice(),
6590 acts[abase + bi + 1].xq.as_slice(),
6591 acts[abase + bi + 2].xq.as_slice(),
6592 acts[abase + bi + 3].xq.as_slice(),
6593 acts[abase + bi + 4].xq.as_slice(),
6594 acts[abase + bi + 5].xq.as_slice(),
6595 acts[abase + bi + 6].xq.as_slice(),
6596 acts[abase + bi + 7].xq.as_slice(),
6597 ];
6598 let d = unsafe { dot_q4tp_2x8_avx512(v.nib, r_lo, gpr, xs, &sc, &sc1) };
6599 for (row, dr, scr) in [(r_lo, &d[0], &sc), (r_lo + 1, &d[1], &sc1)] {
6600 for k in 0..8 {
6601 let act = &acts[abase + bi + k];
6602 let mut acc = dr[k] * act.sx;
6603 for &(j, xv) in &act.outliers {
6604 let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6605 acc += w * s * xv;
6606 }
6607 unsafe { *out_addr.at((abase + bi + k) * rows + row) = acc };
6609 }
6610 }
6611 bi += 8;
6612 }
6613 for row in [r_lo, r_lo + 1] {
6616 let scr: &[f32] = if row == r_lo { &sc } else { &sc1 };
6617 for b2 in bi..alen {
6618 let act = &acts[abase + b2];
6619 let xs4 = [
6620 act.xq.as_slice(),
6621 act.xq.as_slice(),
6622 act.xq.as_slice(),
6623 act.xq.as_slice(),
6624 ];
6625 let d =
6626 unsafe { dot_q4tp_row_1x4_avx512(v.nib, row, gpr, xs4, scr) };
6627 let mut acc = d[0] * act.sx;
6628 for &(j, xv) in &act.outliers {
6629 let (w, s) = q4tp_outlier(v.nib, row, gpr, j, scr);
6630 acc += w * s * xv;
6631 }
6632 unsafe { *out_addr.at((abase + b2) * rows + row) = acc };
6634 }
6635 }
6636 r_lo += 2;
6637 }
6638 }
6639 #[cfg(target_arch = "x86_64")]
6640 let row_start = r_lo;
6641 #[cfg(not(target_arch = "x86_64"))]
6642 let row_start = start;
6643 for r in row_start..end {
6644 v.scales_into(r, gpr, &mut sc);
6645 let mut bi = 0usize;
6646 #[cfg(target_arch = "x86_64")]
6647 if blocked_ok {
6648 while bi + 8 <= alen {
6649 let xs = [
6650 acts[abase + bi].xq.as_slice(),
6651 acts[abase + bi + 1].xq.as_slice(),
6652 acts[abase + bi + 2].xq.as_slice(),
6653 acts[abase + bi + 3].xq.as_slice(),
6654 acts[abase + bi + 4].xq.as_slice(),
6655 acts[abase + bi + 5].xq.as_slice(),
6656 acts[abase + bi + 6].xq.as_slice(),
6657 acts[abase + bi + 7].xq.as_slice(),
6658 ];
6659 let d = unsafe { dot_q4tp_row_1x8_avx512(v.nib, r, gpr, xs, &sc) };
6660 for k in 0..8 {
6661 let act = &acts[abase + bi + k];
6662 let mut acc = d[k] * act.sx;
6663 for &(j, xv) in &act.outliers {
6664 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6665 acc += w * s * xv;
6666 }
6667 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6669 }
6670 bi += 8;
6671 }
6672 while bi + 4 <= alen {
6673 let xs = [
6674 acts[abase + bi].xq.as_slice(),
6675 acts[abase + bi + 1].xq.as_slice(),
6676 acts[abase + bi + 2].xq.as_slice(),
6677 acts[abase + bi + 3].xq.as_slice(),
6678 ];
6679 let d = unsafe { dot_q4tp_row_1x4_avx512(v.nib, r, gpr, xs, &sc) };
6680 for k in 0..4 {
6681 let act = &acts[abase + bi + k];
6682 let mut acc = d[k] * act.sx;
6683 for &(j, xv) in &act.outliers {
6684 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6685 acc += w * s * xv;
6686 }
6687 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6689 }
6690 bi += 4;
6691 }
6692 }
6693 #[cfg(target_arch = "aarch64")]
6694 if blocked_ok {
6695 while bi + 4 <= alen {
6696 let xs = [
6697 acts[abase + bi].xq.as_slice(),
6698 acts[abase + bi + 1].xq.as_slice(),
6699 acts[abase + bi + 2].xq.as_slice(),
6700 acts[abase + bi + 3].xq.as_slice(),
6701 ];
6702 let d = unsafe {
6703 if exact || q4tp_v1() {
6704 dot_q4tp_row_1x4_sdot_v1(v.nib, r, gpr, xs, &sc)
6705 } else {
6706 dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs, &sc)
6707 }
6708 };
6709 for k in 0..4 {
6710 let act = &acts[abase + bi + k];
6711 let mut acc = d[k] * act.sx;
6712 for &(j, xv) in &act.outliers {
6713 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6714 acc += w * s * xv;
6715 }
6716 unsafe { *out_addr.at((abase + bi + k) * rows + r) = acc };
6718 }
6719 bi += 4;
6720 }
6721 }
6722 let _ = blocked_ok;
6723 while bi < alen {
6724 let act = &acts[abase + bi];
6725 let mut acc = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc) * act.sx;
6726 for &(j, xv) in &act.outliers {
6727 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
6728 acc += w * s * xv;
6729 }
6730 unsafe { *out_addr.at((abase + bi) * rows + r) = acc };
6732 bi += 1;
6733 }
6734 }
6735 }
6736 };
6737 dispatch_rows(pool, rows, &run);
6738 return;
6739 }
6740
6741 let run = |start: usize, end: usize| {
6742 let mut sc = vec![0f32; gpr];
6743 for r in start..end {
6744 v.scales_into(r, gpr, &mut sc);
6745 for bi in 0..b {
6746 let x = &xs_all[bi * cols..(bi + 1) * cols];
6747 unsafe { *out_addr.at(bi * rows + r) = q4tp_row_exact(v.nib, r, gpr, x, &sc) };
6749 }
6750 }
6751 };
6752 dispatch_rows(pool, rows, &run);
6753}
6754
6755fn q4t_matvec(
6757 bytes: &[u8],
6758 x: &[f32],
6759 rows: usize,
6760 cols: usize,
6761 out: &mut [f32],
6762 pool: Option<&Pool>,
6763) {
6764 debug_assert_eq!(out.len(), rows);
6765 let gpr = cols / GROUP_SIZE;
6766 let out_addr = SendMut(out.as_mut_ptr());
6767 if a8w8_enabled() {
6768 let act = split_act(x);
6769 let run = move |start: usize, end: usize| {
6770 for r in start..end {
6771 let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
6772 for &(j, xv) in &act.outliers {
6773 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6774 acc += w * s * xv;
6775 }
6776 unsafe { *out_addr.at(r) = acc };
6778 }
6779 };
6780 dispatch_rows(pool, rows, &run);
6781 return;
6782 }
6783 let run = move |start: usize, end: usize| {
6784 for r in start..end {
6785 unsafe { *out_addr.at(r) = q4t_row_exact(bytes, r, gpr, x) };
6787 }
6788 };
6789 dispatch_rows(pool, rows, &run);
6790}
6791
6792#[allow(clippy::too_many_arguments)]
6794fn q4t_matvec2(
6795 bytes: &[u8],
6796 x1: &[f32],
6797 x2: &[f32],
6798 rows: usize,
6799 cols: usize,
6800 o1: &mut [f32],
6801 o2: &mut [f32],
6802 pool: Option<&Pool>,
6803) {
6804 let gpr = cols / GROUP_SIZE;
6805 let p1 = SendMut(o1.as_mut_ptr());
6806 let p2 = SendMut(o2.as_mut_ptr());
6807 if a8w8_enabled() {
6808 let a1 = split_act(x1);
6809 let a2 = split_act(x2);
6810 let run = move |start: usize, end: usize| {
6811 for r in start..end {
6812 let mut v1 = dot_q4t_row_i8(bytes, r, gpr, &a1.xq) * a1.sx;
6813 let mut v2 = dot_q4t_row_i8(bytes, r, gpr, &a2.xq) * a2.sx;
6814 for &(j, xv) in &a1.outliers {
6815 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6816 v1 += w * s * xv;
6817 }
6818 for &(j, xv) in &a2.outliers {
6819 let (w, s) = q4t_outlier(bytes, r, gpr, j);
6820 v2 += w * s * xv;
6821 }
6822 unsafe {
6824 *p1.at(r) = v1;
6825 *p2.at(r) = v2;
6826 }
6827 }
6828 };
6829 dispatch_rows(pool, rows, &run);
6830 return;
6831 }
6832 let run = move |start: usize, end: usize| {
6833 for r in start..end {
6834 unsafe {
6836 *p1.at(r) = q4t_row_exact(bytes, r, gpr, x1);
6837 *p2.at(r) = q4t_row_exact(bytes, r, gpr, x2);
6838 }
6839 }
6840 };
6841 dispatch_rows(pool, rows, &run);
6842}
6843
6844#[allow(clippy::too_many_arguments)]
6846#[cfg(target_os = "macos")]
6852fn dequant_matmat_accel(
6853 dequant_row: &(dyn Fn(usize, &mut [f32]) + Sync),
6854 xs_all: &[f32],
6855 b: usize,
6856 rows: usize,
6857 cols: usize,
6858 out: &mut [f32],
6859 pool: Option<&Pool>,
6860) {
6861 const TR: usize = 2048;
6862 thread_local! {
6863 static WTILE: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
6864 }
6865 WTILE.with(|wt| {
6866 let mut wtile = wt.borrow_mut();
6867 wtile.resize(TR * cols, 0.0);
6868 let mut r0 = 0usize;
6869 while r0 < rows {
6870 let tr = TR.min(rows - r0);
6871 let wt_addr = SendMut(wtile.as_mut_ptr());
6872 let run = |start: usize, end: usize| {
6873 for r in start..end {
6874 let dst = unsafe { std::slice::from_raw_parts_mut(wt_addr.at(r * cols), cols) };
6876 dequant_row(r0 + r, dst);
6877 }
6878 };
6879 dispatch_rows(pool, tr, &run);
6880 unsafe {
6881 accel_blas::cblas_sgemm(
6882 101, 111, 112, b as i32,
6886 tr as i32,
6887 cols as i32,
6888 1.0,
6889 xs_all.as_ptr(),
6890 cols as i32,
6891 wtile.as_ptr(),
6892 cols as i32,
6893 0.0,
6894 out.as_mut_ptr().add(r0),
6895 rows as i32,
6896 );
6897 }
6898 r0 += tr;
6899 }
6900 });
6901}
6902
6903fn q4t_matmat(
6904 bytes: &[u8],
6905 xs_all: &[f32],
6906 b: usize,
6907 rows: usize,
6908 cols: usize,
6909 out: &mut [f32],
6910 pool: Option<&Pool>,
6911) {
6912 debug_assert_eq!(out.len(), b * rows);
6913 let gpr = cols / GROUP_SIZE;
6914 #[cfg(target_os = "macos")]
6918 if b >= 8 && rows * cols >= 500_000 && accel_gemm_enabled() {
6919 dequant_matmat_accel(
6920 &|r, dst| {
6921 for gi in 0..gpr {
6922 let tile = &bytes[(r * gpr + gi) * Q4_TILE..(r * gpr + gi + 1) * Q4_TILE];
6923 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
6924 for (k, &bb) in tile[2..].iter().enumerate() {
6925 dst[gi * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
6926 dst[gi * GROUP_SIZE + k * 2 + 1] = (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
6927 }
6928 }
6929 },
6930 xs_all,
6931 b,
6932 rows,
6933 cols,
6934 out,
6935 pool,
6936 );
6937 return;
6938 }
6939 let out_addr = SendMut(out.as_mut_ptr());
6940 if a8w8_enabled() {
6941 let acts: Vec<SplitAct> = (0..b)
6942 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
6943 .collect();
6944 let acts = &acts;
6945 #[cfg(target_arch = "x86_64")]
6946 let blocked_ok = avx2_enabled() && blocked_enabled();
6947 #[cfg(target_arch = "aarch64")]
6948 let blocked_ok = sdot_enabled() && blocked_enabled();
6949 #[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
6950 let blocked_ok = false;
6951 let run = move |start: usize, end: usize| {
6952 for r in start..end {
6953 let mut bi = 0usize;
6954 #[cfg(target_arch = "aarch64")]
6955 if blocked_ok {
6956 while bi + 4 <= acts.len() {
6957 let xs = [
6958 acts[bi].xq.as_slice(),
6959 acts[bi + 1].xq.as_slice(),
6960 acts[bi + 2].xq.as_slice(),
6961 acts[bi + 3].xq.as_slice(),
6962 ];
6963 let d = unsafe { dot_q4t_row_1x4_sdot(bytes, r, gpr, xs) };
6964 for k in 0..4 {
6965 let act = &acts[bi + k];
6966 let mut acc = d[k] * act.sx;
6967 for &(j, xv) in &act.outliers {
6968 let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6969 acc += w * sc * xv;
6970 }
6971 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
6973 }
6974 bi += 4;
6975 }
6976 }
6977 #[cfg(target_arch = "x86_64")]
6978 if blocked_ok {
6979 while bi + 4 <= acts.len() {
6980 let xs = [
6981 acts[bi].xq.as_slice(),
6982 acts[bi + 1].xq.as_slice(),
6983 acts[bi + 2].xq.as_slice(),
6984 acts[bi + 3].xq.as_slice(),
6985 ];
6986 let d = unsafe {
6987 if vnni_tiles_enabled() {
6988 dot_q4t_row_1x4_vnni(bytes, r, gpr, xs)
6989 } else {
6990 dot_q4t_row_1x4_avx2(bytes, r, gpr, xs)
6991 }
6992 };
6993 for k in 0..4 {
6994 let act = &acts[bi + k];
6995 let mut acc = d[k] * act.sx;
6996 for &(j, xv) in &act.outliers {
6997 let (w, sc) = q4t_outlier(bytes, r, gpr, j);
6998 acc += w * sc * xv;
6999 }
7000 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
7002 }
7003 bi += 4;
7004 }
7005 }
7006 let _ = blocked_ok;
7007 while bi < acts.len() {
7008 let act = &acts[bi];
7009 let mut acc = dot_q4t_row_i8(bytes, r, gpr, &act.xq) * act.sx;
7010 for &(j, xv) in &act.outliers {
7011 let (w, s) = q4t_outlier(bytes, r, gpr, j);
7012 acc += w * s * xv;
7013 }
7014 unsafe { *out_addr.at(bi * rows + r) = acc };
7016 bi += 1;
7017 }
7018 }
7019 };
7020 dispatch_rows(pool, rows, &run);
7021 return;
7022 }
7023 let run = move |start: usize, end: usize| {
7024 for r in start..end {
7025 for bi in 0..b {
7026 let x = &xs_all[bi * cols..(bi + 1) * cols];
7027 unsafe { *out_addr.at(bi * rows + r) = q4t_row_exact(bytes, r, gpr, x) };
7029 }
7030 }
7031 };
7032 dispatch_rows(pool, rows, &run);
7033}
7034
7035fn q1_group_sums(xq: &[i8], gpr: usize) -> Vec<i32> {
7044 (0..gpr)
7045 .map(|gi| {
7046 xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE]
7047 .iter()
7048 .map(|&v| v as i32)
7049 .sum()
7050 })
7051 .collect()
7052}
7053
7054#[inline]
7058#[allow(unreachable_code)]
7059#[cfg(target_arch = "x86_64")]
7064#[target_feature(enable = "avx2")]
7065unsafe fn dot_q1_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7066 unsafe {
7068 use core::arch::x86_64::*;
7069 let expand = _mm256_setr_epi8(
7071 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,
7072 3, 3, 3,
7073 );
7074 let bitsel = _mm256_setr_epi8(
7075 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7076 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7077 );
7078 let ones8 = _mm256_set1_epi8(1);
7079 let ones16 = _mm256_set1_epi16(1);
7080 let mut acc = 0f32;
7081 for gi in 0..gpr {
7082 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7083 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7084 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7085 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7086 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7087 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7088 let sel = _mm256_and_si256(x, mask);
7089 let p16 = _mm256_maddubs_epi16(ones8, sel);
7091 let d32 = _mm256_madd_epi16(p16, ones16);
7092 let hi128 = _mm256_extracti128_si256::<1>(d32);
7093 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7094 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7095 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7096 let msum = _mm_cvtsi128_si32(s32);
7097 let d = 2 * msum - gsum[gi];
7100 acc += d as f32 * s;
7101 }
7102 acc
7103 }
7104}
7105
7106#[cfg(target_arch = "x86_64")]
7109#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7110unsafe fn dot_q1_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7111 unsafe {
7113 use core::arch::x86_64::*;
7114 let expand = _mm256_setr_epi8(
7115 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,
7116 3, 3, 3,
7117 );
7118 let bitsel = _mm256_setr_epi8(
7119 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7120 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7121 );
7122 let ones8 = _mm256_set1_epi8(1);
7123 let mut acc = 0f32;
7124 for gi in 0..gpr {
7125 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7126 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7127 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7128 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7129 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7130 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7131 let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7132 let d = 2 * msum - gsum[gi];
7133 acc += d as f32 * s;
7134 }
7135 acc
7136 }
7137}
7138
7139#[cfg(target_arch = "x86_64")]
7141#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7142unsafe fn dot_q1_row_1x4_vnni(
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 mut acc = [0f32; 4];
7162 for gi in 0..gpr {
7163 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7164 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7165 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7166 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7167 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7168 for (k, xq) in xs.iter().enumerate() {
7169 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7170 let msum = dpbusd_hsum(ones8, _mm256_and_si256(x, mask));
7171 let d = 2 * msum - gsums[k][gi];
7172 acc[k] += d as f32 * s;
7173 }
7174 }
7175 acc
7176 }
7177}
7178
7179#[cfg(target_arch = "x86_64")]
7182#[target_feature(enable = "avx2")]
7183unsafe fn dot_q1_row_1x4_avx2(
7184 bytes: &[u8],
7185 r: usize,
7186 gpr: usize,
7187 xs: [&[i8]; 4],
7188 gsums: [&[i32]; 4],
7189) -> [f32; 4] {
7190 unsafe {
7192 use core::arch::x86_64::*;
7193 let expand = _mm256_setr_epi8(
7194 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,
7195 3, 3, 3,
7196 );
7197 let bitsel = _mm256_setr_epi8(
7198 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64, -128, 1, 2, 4, 8, 16, 32, 64,
7199 -128, 1, 2, 4, 8, 16, 32, 64, -128,
7200 );
7201 let ones8 = _mm256_set1_epi8(1);
7202 let ones16 = _mm256_set1_epi16(1);
7203 let mut acc = [0f32; 4];
7204 for gi in 0..gpr {
7205 let t = bytes.as_ptr().add((r * gpr + gi) * Q1_TILE);
7206 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7207 let bits = u32::from_le_bytes([*t.add(2), *t.add(3), *t.add(4), *t.add(5)]);
7208 let bc = _mm256_shuffle_epi8(_mm256_set1_epi32(bits as i32), expand);
7209 let mask = _mm256_cmpeq_epi8(_mm256_and_si256(bc, bitsel), bitsel);
7210 for (k, xq) in xs.iter().enumerate() {
7211 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
7212 let sel = _mm256_and_si256(x, mask);
7213 let p16 = _mm256_maddubs_epi16(ones8, sel);
7214 let d32 = _mm256_madd_epi16(p16, ones16);
7215 let hi128 = _mm256_extracti128_si256::<1>(d32);
7216 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d32), hi128);
7217 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7218 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7219 let msum = _mm_cvtsi128_si32(s32);
7220 let d = 2 * msum - gsums[k][gi];
7221 acc[k] += d as f32 * s;
7222 }
7223 }
7224 acc
7225 }
7226}
7227
7228#[allow(unreachable_code)]
7229fn dot_q1_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7230 #[cfg(target_arch = "aarch64")]
7231 unsafe {
7232 return dot_q1_row_sdot(bytes, r, gpr, xq, gsum);
7233 }
7234 #[cfg(target_arch = "x86_64")]
7235 if avx2_enabled() {
7236 unsafe {
7237 if vnni_tiles_enabled() {
7238 return dot_q1_row_vnni(bytes, r, gpr, xq, gsum);
7239 }
7240 return dot_q1_row_avx2(bytes, r, gpr, xq, gsum);
7241 }
7242 }
7243 let _ = gsum;
7244 let mut acc = 0f32;
7245 for gi in 0..gpr {
7246 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7247 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7248 let mut d = 0i32;
7249 for (j, &b) in tile[2..].iter().enumerate() {
7250 for k in 0..8 {
7251 let w = ((b >> k) & 1) as i32 * 2 - 1;
7252 d += w * xq[gi * GROUP_SIZE + j * 8 + k] as i32;
7253 }
7254 }
7255 acc += d as f32 * s;
7256 }
7257 acc
7258}
7259
7260#[cfg(target_arch = "aarch64")]
7269#[target_feature(enable = "neon,dotprod")]
7270unsafe fn dot_q1_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8], gsum: &[i32]) -> f32 {
7271 unsafe {
7274 use core::arch::aarch64::*;
7275 use core::arch::asm;
7276 const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7277 let m = vld1q_u8(MASKS.as_ptr());
7278 macro_rules! tile_dot {
7280 ($t:expr, $x:expr) => {{
7281 let v0 = vcombine_u8(vdup_n_u8(*$t.add(2)), vdup_n_u8(*$t.add(3)));
7282 let v1 = vcombine_u8(vdup_n_u8(*$t.add(4)), vdup_n_u8(*$t.add(5)));
7283 let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7284 let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7285 let x0 = vld1q_s8($x);
7286 let x1 = vld1q_s8($x.add(16));
7287 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7288 asm!(
7289 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7290 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7291 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7292 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7293 options(pure, nomem, nostack),
7294 );
7295 vaddq_s32(a0, a1)
7296 }};
7297 }
7298 const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7307 const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7308 const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7309 const IW11: [u8; 16] = [
7310 10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7311 ];
7312 const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7313 let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7314 let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7315 let isc = vld1_u8(ISC.as_ptr());
7316 macro_rules! tile_dot_tbl {
7318 ($ld:expr, $i0:expr, $i1:expr, $x:expr) => {{
7319 let w0 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i0), m));
7320 let w1 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8($ld, $i1), m));
7321 let x0 = vld1q_s8($x);
7322 let x1 = vld1q_s8($x.add(16));
7323 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7324 asm!(
7325 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7326 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7327 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7328 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7329 options(pure, nomem, nostack),
7330 );
7331 vaddq_s32(a0, a1)
7332 }};
7333 }
7334 let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7335 let row_base = r * gpr * Q1_TILE;
7336 let abs_end = bytes.len();
7337 let xp = xq.as_ptr();
7338 let gp = gsum.as_ptr();
7339 let mut accv = vdupq_n_f32(0.0);
7340 let mut gi = 0;
7341 while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7344 let t0 = base.add(gi * Q1_TILE);
7345 let ld_a = vld1q_u8(t0);
7346 let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7347 let d0 = tile_dot_tbl!(ld_a, iw00, iw01, xp.add(gi * GROUP_SIZE));
7348 let d1 = tile_dot_tbl!(ld_a, iw10, iw11, xp.add((gi + 1) * GROUP_SIZE));
7349 let d2 = tile_dot_tbl!(ld_b, iw00, iw01, xp.add((gi + 2) * GROUP_SIZE));
7350 let d3 = tile_dot_tbl!(ld_b, iw10, iw11, xp.add((gi + 3) * GROUP_SIZE));
7351 let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7353 let g = vld1q_s32(gp.add(gi));
7354 let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7355 let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7356 let scf: float32x4_t;
7357 asm!(
7358 "fcvtl {o:v}.4s, {i:v}.4h",
7359 o = out(vreg) scf, i = in(vreg) sc16,
7360 options(pure, nomem, nostack),
7361 );
7362 accv = vfmaq_f32(accv, vcvtq_f32_s32(dots), scf);
7363 gi += 4;
7364 }
7365 let mut acc = vaddvq_f32(accv);
7366 while gi < gpr {
7367 let t = base.add(gi * Q1_TILE);
7368 let s = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7369 let d = vaddvq_s32(tile_dot!(t, xp.add(gi * GROUP_SIZE)));
7370 acc += (-(2 * d + *gp.add(gi))) as f32 * s;
7371 gi += 1;
7372 }
7373 acc
7374 }
7375}
7376
7377#[cfg(target_arch = "aarch64")]
7382#[target_feature(enable = "neon,dotprod")]
7383unsafe fn dot_q1_row_1x4_sdot(
7384 bytes: &[u8],
7385 r: usize,
7386 gpr: usize,
7387 xs: [&[i8]; 4],
7388 gs: [&[i32]; 4],
7389) -> [f32; 4] {
7390 unsafe {
7392 use core::arch::aarch64::*;
7393 use core::arch::asm;
7394 const MASKS: [u8; 16] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128];
7395 const IW00: [u8; 16] = [2, 2, 2, 2, 2, 2, 2, 2, 3, 3, 3, 3, 3, 3, 3, 3];
7396 const IW01: [u8; 16] = [4, 4, 4, 4, 4, 4, 4, 4, 5, 5, 5, 5, 5, 5, 5, 5];
7397 const IW10: [u8; 16] = [8, 8, 8, 8, 8, 8, 8, 8, 9, 9, 9, 9, 9, 9, 9, 9];
7398 const IW11: [u8; 16] = [
7399 10, 10, 10, 10, 10, 10, 10, 10, 11, 11, 11, 11, 11, 11, 11, 11,
7400 ];
7401 const ISC: [u8; 8] = [0, 1, 6, 7, 16, 17, 22, 23];
7402 let m = vld1q_u8(MASKS.as_ptr());
7403 let (iw00, iw01) = (vld1q_u8(IW00.as_ptr()), vld1q_u8(IW01.as_ptr()));
7404 let (iw10, iw11) = (vld1q_u8(IW10.as_ptr()), vld1q_u8(IW11.as_ptr()));
7405 let isc = vld1_u8(ISC.as_ptr());
7406 macro_rules! sdot2 {
7407 ($w0:expr, $w1:expr, $x:expr) => {{
7408 let x0 = vld1q_s8($x);
7409 let x1 = vld1q_s8($x.add(16));
7410 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7411 asm!(
7412 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7413 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7414 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7415 w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
7416 options(pure, nomem, nostack),
7417 );
7418 vaddq_s32(a0, a1)
7419 }};
7420 }
7421 let base = bytes.as_ptr().add(r * gpr * Q1_TILE);
7422 let row_base = r * gpr * Q1_TILE;
7423 let abs_end = bytes.len();
7424 let mut accv = [vdupq_n_f32(0.0); 4];
7425 let mut gi = 0;
7426 while gi + 4 <= gpr && row_base + (gi + 4) * Q1_TILE + 4 <= abs_end {
7427 let t0 = base.add(gi * Q1_TILE);
7428 let ld_a = vld1q_u8(t0);
7429 let ld_b = vld1q_u8(t0.add(2 * Q1_TILE));
7430 let w00 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw00), m));
7432 let w01 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw01), m));
7433 let w10 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw10), m));
7434 let w11 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_a, iw11), m));
7435 let w20 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw00), m));
7436 let w21 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw01), m));
7437 let w30 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw10), m));
7438 let w31 = vreinterpretq_s8_u8(vtstq_u8(vqtbl1q_u8(ld_b, iw11), m));
7439 let sc16 = vqtbl2_u8(uint8x16x2_t(ld_a, ld_b), isc);
7440 let scf: float32x4_t;
7441 asm!(
7442 "fcvtl {o:v}.4s, {i:v}.4h",
7443 o = out(vreg) scf, i = in(vreg) sc16,
7444 options(pure, nomem, nostack),
7445 );
7446 for k in 0..4 {
7447 let xp = xs[k].as_ptr();
7448 let d0 = sdot2!(w00, w01, xp.add(gi * GROUP_SIZE));
7449 let d1 = sdot2!(w10, w11, xp.add((gi + 1) * GROUP_SIZE));
7450 let d2 = sdot2!(w20, w21, xp.add((gi + 2) * GROUP_SIZE));
7451 let d3 = sdot2!(w30, w31, xp.add((gi + 3) * GROUP_SIZE));
7452 let neg = vpaddq_s32(vpaddq_s32(d0, d1), vpaddq_s32(d2, d3));
7453 let g = vld1q_s32(gs[k].as_ptr().add(gi));
7454 let dots = vnegq_s32(vaddq_s32(vshlq_n_s32::<1>(neg), g));
7455 accv[k] = vfmaq_f32(accv[k], vcvtq_f32_s32(dots), scf);
7456 }
7457 gi += 4;
7458 }
7459 let mut acc = [
7460 vaddvq_f32(accv[0]),
7461 vaddvq_f32(accv[1]),
7462 vaddvq_f32(accv[2]),
7463 vaddvq_f32(accv[3]),
7464 ];
7465 while gi < gpr {
7466 let t = base.add(gi * Q1_TILE);
7467 let sc = f16_to_f32(u16::from_le_bytes([*t, *t.add(1)]));
7468 let v0 = vcombine_u8(vdup_n_u8(*t.add(2)), vdup_n_u8(*t.add(3)));
7469 let v1 = vcombine_u8(vdup_n_u8(*t.add(4)), vdup_n_u8(*t.add(5)));
7470 let w0 = vreinterpretq_s8_u8(vtstq_u8(v0, m));
7471 let w1 = vreinterpretq_s8_u8(vtstq_u8(v1, m));
7472 for k in 0..4 {
7473 let d = vaddvq_s32(sdot2!(w0, w1, xs[k].as_ptr().add(gi * GROUP_SIZE)));
7474 acc[k] += (-(2 * d + *gs[k].as_ptr().add(gi))) as f32 * sc;
7475 }
7476 gi += 1;
7477 }
7478 acc
7479 }
7480}
7481
7482#[inline]
7484fn q1_outlier(bytes: &[u8], r: usize, gpr: usize, j: usize) -> (f32, f32) {
7485 let gi = j / GROUP_SIZE;
7486 let k = j % GROUP_SIZE;
7487 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7488 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7489 let bit = (tile[2 + k / 8] >> (k % 8)) & 1;
7490 ((bit as i32 * 2 - 1) as f32, s)
7491}
7492
7493#[inline]
7495fn q1_row_exact(bytes: &[u8], r: usize, gpr: usize, x: &[f32]) -> f32 {
7496 let mut acc = 0f32;
7497 for gi in 0..gpr {
7498 let tile = &bytes[(r * gpr + gi) * Q1_TILE..(r * gpr + gi + 1) * Q1_TILE];
7499 let s = f16_to_f32(u16::from_le_bytes([tile[0], tile[1]]));
7500 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
7501 let mut ga = 0f32;
7502 for (j, &b) in tile[2..].iter().enumerate() {
7503 for k in 0..8 {
7504 ga += (((b >> k) & 1) as f32 * 2.0 - 1.0) * xg[j * 8 + k];
7505 }
7506 }
7507 acc += ga * s;
7508 }
7509 acc
7510}
7511
7512#[allow(clippy::too_many_arguments)]
7515fn q1_range_a8w8(
7516 bytes: &[u8],
7517 gpr: usize,
7518 act: &SplitAct,
7519 gsum: &[i32],
7520 out: SendMut,
7521 start: usize,
7522 end: usize,
7523) {
7524 for r in start..end {
7525 let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
7526 for &(j, xv) in &act.outliers {
7527 let (w, s) = q1_outlier(bytes, r, gpr, j);
7528 acc += w * s * xv;
7529 }
7530 unsafe { *out.at(r) = acc };
7532 }
7533}
7534
7535fn q1_range_f32(bytes: &[u8], gpr: usize, x: &[f32], out: SendMut, start: usize, end: usize) {
7537 for r in start..end {
7538 unsafe { *out.at(r) = q1_row_exact(bytes, r, gpr, x) };
7540 }
7541}
7542
7543fn q1t_overlay(bytes: &[u8], base_len: usize, rows: usize) -> (usize, usize, bool) {
7548 let entries = base_len + (rows + 1) * 4;
7549 (base_len, entries, entries <= bytes.len())
7550}
7551
7552#[inline]
7554fn q1t_rowptr(bytes: &[u8], rp_off: usize, r: usize) -> usize {
7555 let o = rp_off + r * 4;
7556 u32::from_le_bytes([bytes[o], bytes[o + 1], bytes[o + 2], bytes[o + 3]]) as usize
7557}
7558
7559const SIGN5: [[f32; 5]; 256] = {
7563 let mut lut = [[0.0f32; 5]; 256];
7564 let pow3 = [1u16, 3, 9, 27, 81];
7565 let mut byte = 0usize;
7566 while byte < 256 {
7567 let mut i = 0usize;
7568 while i < 5 {
7569 let code = (byte as u16 / pow3[i]) % 3;
7570 lut[byte][i] = if code == 1 {
7571 1.0
7572 } else if code == 2 {
7573 -1.0
7574 } else {
7575 0.0
7576 };
7577 i += 1;
7578 }
7579 byte += 1;
7580 }
7581 lut
7582};
7583
7584const SIGN5_I8: [[i8; 5]; 256] = {
7586 let mut lut = [[0i8; 5]; 256];
7587 let pow3 = [1u16, 3, 9, 27, 81];
7588 let mut byte = 0usize;
7589 while byte < 256 {
7590 let mut i = 0usize;
7591 while i < 5 {
7592 let code = (byte as u16 / pow3[i]) % 3;
7593 lut[byte][i] = if code == 1 {
7594 1
7595 } else if code == 2 {
7596 -1
7597 } else {
7598 0
7599 };
7600 i += 1;
7601 }
7602 byte += 1;
7603 }
7604 lut
7605};
7606
7607const SIGN5_U64: [u64; 256] = {
7613 let mut lut = [0u64; 256];
7614 let pow3 = [1u16, 3, 9, 27, 81];
7615 let mut byte = 0usize;
7616 while byte < 256 {
7617 let mut v = 0u64;
7618 let mut i = 0usize;
7619 while i < 5 {
7620 let code = (byte as u16 / pow3[i]) % 3;
7621 let s: u8 = if code == 1 {
7622 1
7623 } else if code == 2 {
7624 0xFF
7625 } else {
7626 0
7627 };
7628 v |= (s as u64) << (i * 8);
7629 i += 1;
7630 }
7631 lut[byte] = v;
7632 byte += 1;
7633 }
7634 lut
7635};
7636
7637#[inline]
7642fn q1t_base_weight(bytes: &[u8], r: usize, gpr: usize, j: usize) -> f32 {
7643 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7644 let off = (r * gpr + j / GROUP_SIZE) * TILE;
7645 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7646 let within = j % GROUP_SIZE;
7647 SIGN5[bytes[off + 2 + within / 5] as usize][within % 5] * s
7648}
7649
7650#[cfg(target_arch = "aarch64")]
7653#[target_feature(enable = "neon,dotprod")]
7654#[inline]
7655unsafe fn sdot32_i8(w: *const i8, x: *const i8) -> i32 {
7656 unsafe {
7658 use core::arch::aarch64::*;
7659 use core::arch::asm;
7660 let w0 = vld1q_s8(w);
7661 let w1 = vld1q_s8(w.add(16));
7662 let x0 = vld1q_s8(x);
7663 let x1 = vld1q_s8(x.add(16));
7664 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7665 asm!(
7666 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7667 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7668 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7669 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7670 options(pure, nomem, nostack),
7671 );
7672 vaddvq_s32(vaddq_s32(a0, a1))
7673 }
7674}
7675
7676#[cfg(target_arch = "x86_64")]
7679#[target_feature(enable = "avx2")]
7680#[inline]
7681unsafe fn i8dot32_avx2(w: *const i8, x: *const i8) -> i32 {
7682 unsafe {
7684 use core::arch::x86_64::*;
7685 let wv = _mm256_loadu_si256(w as *const __m256i);
7686 let xv = _mm256_loadu_si256(x as *const __m256i);
7687 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7688 let d = _mm256_madd_epi16(p16, _mm256_set1_epi16(1));
7689 let hi128 = _mm256_extracti128_si256::<1>(d);
7690 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
7691 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
7692 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
7693 _mm_cvtsi128_si32(s32)
7694 }
7695}
7696
7697#[inline]
7702fn q1t_unpack_group_i8(codes: *const u8, dst: &mut [i8]) {
7703 debug_assert!(dst.len() >= 40);
7704 unsafe {
7707 let p = dst.as_mut_ptr();
7708 for bi in 0..7 {
7709 core::ptr::write_unaligned(
7710 p.add(bi * 5) as *mut u64,
7711 SIGN5_U64[*codes.add(bi) as usize],
7712 );
7713 }
7714 }
7715}
7716
7717#[inline]
7722fn q1t_i8dot32(w: *const i8, x: *const i8) -> i32 {
7723 #[cfg(target_arch = "aarch64")]
7724 unsafe {
7725 return sdot32_i8(w, x);
7726 }
7727 #[cfg(target_arch = "x86_64")]
7728 unsafe {
7729 return i8dot32_avx2(w, x);
7730 }
7731 #[allow(unreachable_code)]
7732 unsafe {
7733 let mut s = 0i32;
7734 for k in 0..GROUP_SIZE {
7735 s += *w.add(k) as i32 * *x.add(k) as i32;
7736 }
7737 s
7738 }
7739}
7740
7741#[inline]
7742unsafe fn q1t_unpack_reg_u64s(codes: *const u8) -> (u64, u64, u64, u64) {
7743 let (s0, s1, s2, s3, s4, s5, s6) = unsafe {
7744 (
7745 SIGN5_U64[*codes as usize],
7746 SIGN5_U64[*codes.add(1) as usize],
7747 SIGN5_U64[*codes.add(2) as usize],
7748 SIGN5_U64[*codes.add(3) as usize],
7749 SIGN5_U64[*codes.add(4) as usize],
7750 SIGN5_U64[*codes.add(5) as usize],
7751 SIGN5_U64[*codes.add(6) as usize],
7752 )
7753 };
7754
7755 let u0 = s0 | (s1 << 40);
7756 let u1 = (s1 >> 24) | (s2 << 16) | (s3 << 56);
7757 let u2 = (s3 >> 8) | (s4 << 32);
7758 let u3 = (s4 >> 32) | (s5 << 8) | (s6 << 48);
7759
7760 (u0, u1, u2, u3)
7761}
7762
7763#[cfg(target_arch = "aarch64")]
7767#[target_feature(enable = "neon,dotprod")]
7768unsafe fn q1t_dot_row_sdot(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7769 use core::arch::aarch64::*;
7770 use core::arch::asm;
7771 unsafe {
7772 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7773 let mut acc = 0f32;
7774 let bytes_ptr = bytes.as_ptr();
7775 let xq_ptr = xq.as_ptr();
7776 let row_off = r * gpr * TILE;
7777
7778 let gpr2 = gpr & !1;
7779 let mut gi = 0;
7780 while gi < gpr2 {
7781 let off0 = row_off + gi * TILE;
7782 let off1 = off0 + TILE;
7783 let s0 = f16_to_f32(u16::from_le_bytes([
7784 *bytes_ptr.add(off0),
7785 *bytes_ptr.add(off0 + 1),
7786 ]));
7787 let s1 = f16_to_f32(u16::from_le_bytes([
7788 *bytes_ptr.add(off1),
7789 *bytes_ptr.add(off1 + 1),
7790 ]));
7791
7792 let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
7793 let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
7794
7795 let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
7796 let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
7797 let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
7798 let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
7799
7800 let x0_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7801 let x1_0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7802 let x0_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE));
7803 let x1_1 = vld1q_s8(xq_ptr.add((gi + 1) * GROUP_SIZE + 16));
7804
7805 let (mut a0_0, mut a1_0) = (vdupq_n_s32(0), vdupq_n_s32(0));
7806 let (mut a0_1, mut a1_1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7807 asm!(
7808 "sdot {a0_0:v}.4s, {w0_0:v}.16b, {x0_0:v}.16b",
7809 "sdot {a1_0:v}.4s, {w1_0:v}.16b, {x1_0:v}.16b",
7810 "sdot {a0_1:v}.4s, {w0_1:v}.16b, {x0_1:v}.16b",
7811 "sdot {a1_1:v}.4s, {w1_1:v}.16b, {x1_1:v}.16b",
7812 a0_0 = inout(vreg) a0_0, a1_0 = inout(vreg) a1_0,
7813 a0_1 = inout(vreg) a0_1, a1_1 = inout(vreg) a1_1,
7814 w0_0 = in(vreg) w0_0, x0_0 = in(vreg) x0_0, w1_0 = in(vreg) w1_0, x1_0 = in(vreg) x1_0,
7815 w0_1 = in(vreg) w0_1, x0_1 = in(vreg) x0_1, w1_1 = in(vreg) w1_1, x1_1 = in(vreg) x1_1,
7816 options(pure, nomem, nostack),
7817 );
7818 let d0 = vaddvq_s32(vaddq_s32(a0_0, a1_0));
7819 let d1 = vaddvq_s32(vaddq_s32(a0_1, a1_1));
7820 acc += d0 as f32 * s0 + d1 as f32 * s1;
7821 gi += 2;
7822 }
7823
7824 if gi < gpr {
7825 let off = row_off + gi * TILE;
7826 let s = f16_to_f32(u16::from_le_bytes([
7827 *bytes_ptr.add(off),
7828 *bytes_ptr.add(off + 1),
7829 ]));
7830 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7831 let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
7832 let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
7833 let x0 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE));
7834 let x1 = vld1q_s8(xq_ptr.add(gi * GROUP_SIZE + 16));
7835 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
7836 asm!(
7837 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
7838 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
7839 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
7840 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
7841 options(pure, nomem, nostack),
7842 );
7843 let d = vaddvq_s32(vaddq_s32(a0, a1));
7844 acc += d as f32 * s;
7845 }
7846 acc
7847 }
7848}
7849
7850#[cfg(target_arch = "x86_64")]
7852#[target_feature(enable = "avx2")]
7853unsafe fn q1t_dot_row_avx2(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7854 use core::arch::x86_64::*;
7855 unsafe {
7856 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7857 let mut acc = 0f32;
7858 let bytes_ptr = bytes.as_ptr();
7859 let xq_ptr = xq.as_ptr();
7860 let row_off = r * gpr * TILE;
7861
7862 let ones = _mm256_set1_epi16(1);
7863 for gi in 0..gpr {
7864 let off = row_off + gi * TILE;
7865 let s = f16_to_f32(u16::from_le_bytes([
7866 *bytes_ptr.add(off),
7867 *bytes_ptr.add(off + 1),
7868 ]));
7869 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7870 let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7871 let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7872 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7873 let d256 = _mm256_madd_epi16(p16, ones);
7874 let d128 = _mm_add_epi32(
7875 _mm256_castsi256_si128(d256),
7876 _mm256_extracti128_si256(d256, 1),
7877 );
7878 let d64 = _mm_add_epi32(d128, _mm_shuffle_epi32(d128, 0xee));
7879 let d32 = _mm_cvtsi128_si32(_mm_add_epi32(d64, _mm_shuffle_epi32(d64, 0x55)));
7880 acc += d32 as f32 * s;
7881 }
7882 acc
7883 }
7884}
7885
7886#[cfg(target_arch = "x86_64")]
7888#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
7889unsafe fn q1t_dot_row_vnni(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7890 use core::arch::x86_64::*;
7891 unsafe {
7893 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7894 let mut acc = 0f32;
7895 let bytes_ptr = bytes.as_ptr();
7896 let xq_ptr = xq.as_ptr();
7897 let row_off = r * gpr * TILE;
7898 for gi in 0..gpr {
7899 let off = row_off + gi * TILE;
7900 let s = f16_to_f32(u16::from_le_bytes([
7901 *bytes_ptr.add(off),
7902 *bytes_ptr.add(off + 1),
7903 ]));
7904 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
7905 let wv = _mm256_set_epi64x(u3 as i64, u2 as i64, u1 as i64, u0 as i64);
7906 let xv = _mm256_loadu_si256(xq_ptr.add(gi * GROUP_SIZE) as *const __m256i);
7907 let d = dpbusd_hsum(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
7908 acc += d as f32 * s;
7909 }
7910 acc
7911 }
7912}
7913
7914#[inline]
7918fn q1t_dot_row_i8(bytes: &[u8], r: usize, gpr: usize, xq: &[i8]) -> f32 {
7919 #[cfg(target_arch = "aarch64")]
7920 unsafe {
7921 return q1t_dot_row_sdot(bytes, r, gpr, xq);
7922 }
7923 #[cfg(target_arch = "x86_64")]
7924 unsafe {
7925 if vnni_tiles_enabled() {
7926 return q1t_dot_row_vnni(bytes, r, gpr, xq);
7927 }
7928 return q1t_dot_row_avx2(bytes, r, gpr, xq);
7929 }
7930 #[allow(unreachable_code)]
7931 {
7932 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7933 let mut acc = 0f32;
7934 let mut sg = [0i8; GROUP_SIZE + 8]; for gi in 0..gpr {
7936 let off = (r * gpr + gi) * TILE;
7937 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7938 q1t_unpack_group_i8(bytes.as_ptr().wrapping_add(off + 2), &mut sg);
7939 let mut d = 0i32;
7940 for k in 0..GROUP_SIZE {
7941 d += sg[k] as i32 * xq[gi * GROUP_SIZE + k] as i32;
7942 }
7943 acc += d as f32 * s;
7944 }
7945 acc
7946 }
7947}
7948
7949fn q1t_row_outlier_correction(
7956 bytes: &[u8],
7957 r: usize,
7958 rp_off: usize,
7959 entries_off: usize,
7960 has_ov: bool,
7961 x: &[f32],
7962) -> f32 {
7963 if !has_ov {
7964 return 0.0;
7965 }
7966 let (c0, c1) = (
7967 q1t_rowptr(bytes, rp_off, r),
7968 q1t_rowptr(bytes, rp_off, r + 1),
7969 );
7970 let mut corr = 0f32;
7971 for p in c0..c1 {
7972 let e = entries_off + p * 4;
7973 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
7974 let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
7975 corr += val * x[col];
7976 }
7977 corr
7978}
7979
7980fn q1t_dequant_row(
7984 bytes: &[u8],
7985 r: usize,
7986 gpr: usize,
7987 rp_off: usize,
7988 entries_off: usize,
7989 has_ov: bool,
7990 buf: &mut [f32],
7991) {
7992 const TILE: usize = cortiq_core::quant::Q1T_TILE;
7993 for g in 0..gpr {
7994 let off = (r * gpr + g) * TILE;
7995 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
7996 let codes = &bytes[off + 2..off + TILE];
7997 let bc = g * GROUP_SIZE;
7998 for bi in 0..6 {
8000 let lut = &SIGN5[codes[bi] as usize];
8001 let d = &mut buf[bc + bi * 5..bc + bi * 5 + 5];
8002 for i in 0..5 {
8003 d[i] = lut[i] * s;
8004 }
8005 }
8006 let lut = &SIGN5[codes[6] as usize];
8007 buf[bc + 30] = lut[0] * s;
8008 buf[bc + 31] = lut[1] * s;
8009 }
8010 if !has_ov {
8011 return;
8012 }
8013 let (c0, c1) = (
8014 q1t_rowptr(bytes, rp_off, r),
8015 q1t_rowptr(bytes, rp_off, r + 1),
8016 );
8017 for p in c0..c1 {
8018 let e = entries_off + p * 4;
8019 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8020 buf[col] = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8021 }
8022}
8023
8024fn q1t_add_overlay(
8028 bytes: &[u8],
8029 x: &[f32],
8030 rows: usize,
8031 cols: usize,
8032 out: &mut [f32],
8033 pool: Option<&Pool>,
8034) {
8035 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8036 let gpr = cols / GROUP_SIZE;
8037 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8038 if !has_ov {
8039 return;
8040 }
8041 let out_addr = SendMut(out.as_mut_ptr());
8042 let run = move |start: usize, end: usize| {
8043 for r in start..end {
8044 let corr = q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8045 unsafe { *out_addr.at(r) += corr };
8047 }
8048 };
8049 dispatch_rows(pool, rows, &run);
8050}
8051
8052#[allow(clippy::too_many_arguments)]
8055fn q1t_range_a8w8(
8056 bytes: &[u8],
8057 gpr: usize,
8058 rp_off: usize,
8059 ent_off: usize,
8060 has_ov: bool,
8061 act: &SplitAct,
8062 x: &[f32],
8063 out: SendMut,
8064 start: usize,
8065 end: usize,
8066) {
8067 for r in start..end {
8068 let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8069 for &(j, xv) in &act.outliers {
8070 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
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
8078#[allow(clippy::too_many_arguments)]
8081fn q1t_range_f32_batch(
8082 bytes: &[u8],
8083 gpr: usize,
8084 rp_off: usize,
8085 ent_off: usize,
8086 has_ov: bool,
8087 x: &[f32],
8088 out: SendMut,
8089 start: usize,
8090 end: usize,
8091) {
8092 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8093 let mut sg = [0f32; GROUP_SIZE];
8094 for r in start..end {
8095 let mut acc = 0f32;
8096 for g in 0..gpr {
8097 let off = (r * gpr + g) * TILE;
8098 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8099 let codes = &bytes[off + 2..off + TILE];
8100 let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8101 for bi in 0..6 {
8102 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8103 }
8104 let lut = &SIGN5[codes[6] as usize];
8105 sg[30] = lut[0];
8106 sg[31] = lut[1];
8107 let mut gsum = 0f32;
8108 for k in 0..GROUP_SIZE {
8109 gsum += sg[k] * xg[k];
8110 }
8111 acc += s * gsum;
8112 }
8113 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8114 unsafe { *out.at(r) = acc };
8116 }
8117}
8118
8119fn q1t_matvec(
8123 bytes: &[u8],
8124 x: &[f32],
8125 rows: usize,
8126 cols: usize,
8127 out: &mut [f32],
8128 pool: Option<&Pool>,
8129) {
8130 debug_assert_eq!(out.len(), rows);
8131 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8132 let gpr = cols / GROUP_SIZE;
8133 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8134 let out_addr = SendMut(out.as_mut_ptr());
8135 if a8w8_enabled() {
8139 let act = split_act(x);
8140 let act = &act;
8141 let run = move |start: usize, end: usize| {
8142 for r in start..end {
8143 let mut acc = q1t_dot_row_i8(bytes, r, gpr, &act.xq) * act.sx;
8144 for &(j, xv) in &act.outliers {
8145 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8146 }
8147 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8148 unsafe { *out_addr.at(r) = acc };
8150 }
8151 };
8152 dispatch_rows(pool, rows, &run);
8153 return;
8154 }
8155 let run = move |start: usize, end: usize| {
8156 let mut sg = [0f32; GROUP_SIZE];
8160 for r in start..end {
8161 let mut acc = 0f32;
8162 for g in 0..gpr {
8163 let off = (r * gpr + g) * TILE;
8164 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8165 let codes = &bytes[off + 2..off + TILE];
8166 let xg = &x[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8167 for bi in 0..6 {
8168 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8169 }
8170 let lut = &SIGN5[codes[6] as usize];
8171 sg[30] = lut[0];
8172 sg[31] = lut[1];
8173 let mut gsum = 0f32;
8174 for k in 0..GROUP_SIZE {
8175 gsum += sg[k] * xg[k];
8176 }
8177 acc += s * gsum;
8178 }
8179 acc += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x);
8180 unsafe { *out_addr.at(r) = acc };
8181 }
8182 };
8183 dispatch_rows(pool, rows, &run);
8184}
8185
8186#[cfg(target_arch = "aarch64")]
8192#[target_feature(enable = "neon,dotprod")]
8193unsafe fn q1t_dot_row_sdot2(bytes: &[u8], r: usize, gpr: usize, xa: &[i8], xb: &[i8]) -> [f32; 2] {
8194 use core::arch::aarch64::*;
8195 use core::arch::asm;
8196 unsafe {
8198 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8199 let bytes_ptr = bytes.as_ptr();
8200 let row_off = r * gpr * TILE;
8201 let xp = [xa.as_ptr(), xb.as_ptr()];
8202 let mut acc = [0f32; 2];
8203 macro_rules! sdot2 {
8204 ($w0:expr, $w1:expr, $x:expr) => {{
8205 let x0 = vld1q_s8($x);
8206 let x1 = vld1q_s8($x.add(16));
8207 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
8208 asm!(
8209 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
8210 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
8211 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
8212 w0 = in(vreg) $w0, x0 = in(vreg) x0, w1 = in(vreg) $w1, x1 = in(vreg) x1,
8213 options(pure, nomem, nostack),
8214 );
8215 vaddvq_s32(vaddq_s32(a0, a1))
8216 }};
8217 }
8218 let gpr2 = gpr & !1;
8219 let mut gi = 0;
8220 while gi < gpr2 {
8221 let off0 = row_off + gi * TILE;
8222 let off1 = off0 + TILE;
8223 let s0 = f16_to_f32(u16::from_le_bytes([
8224 *bytes_ptr.add(off0),
8225 *bytes_ptr.add(off0 + 1),
8226 ]));
8227 let s1 = f16_to_f32(u16::from_le_bytes([
8228 *bytes_ptr.add(off1),
8229 *bytes_ptr.add(off1 + 1),
8230 ]));
8231 let (u0_0, u1_0, u2_0, u3_0) = q1t_unpack_reg_u64s(bytes_ptr.add(off0 + 2));
8232 let (u0_1, u1_1, u2_1, u3_1) = q1t_unpack_reg_u64s(bytes_ptr.add(off1 + 2));
8233 let w0_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_0), vcreate_u64(u1_0)));
8234 let w1_0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_0), vcreate_u64(u3_0)));
8235 let w0_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0_1), vcreate_u64(u1_1)));
8236 let w1_1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2_1), vcreate_u64(u3_1)));
8237 for k in 0..2 {
8238 let d0 = sdot2!(w0_0, w1_0, xp[k].add(gi * GROUP_SIZE));
8239 let d1 = sdot2!(w0_1, w1_1, xp[k].add((gi + 1) * GROUP_SIZE));
8240 acc[k] += d0 as f32 * s0 + d1 as f32 * s1;
8241 }
8242 gi += 2;
8243 }
8244 if gi < gpr {
8245 let off = row_off + gi * TILE;
8246 let s = f16_to_f32(u16::from_le_bytes([
8247 *bytes_ptr.add(off),
8248 *bytes_ptr.add(off + 1),
8249 ]));
8250 let (u0, u1, u2, u3) = q1t_unpack_reg_u64s(bytes_ptr.add(off + 2));
8251 let w0 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u0), vcreate_u64(u1)));
8252 let w1 = vreinterpretq_s8_u64(vcombine_u64(vcreate_u64(u2), vcreate_u64(u3)));
8253 for k in 0..2 {
8254 let d = sdot2!(w0, w1, xp[k].add(gi * GROUP_SIZE));
8255 acc[k] += d as f32 * s;
8256 }
8257 }
8258 acc
8259 }
8260}
8261
8262fn q1t_matvec2(
8268 bytes: &[u8],
8269 x1: &[f32],
8270 x2: &[f32],
8271 rows: usize,
8272 cols: usize,
8273 o1: &mut [f32],
8274 o2: &mut [f32],
8275 pool: Option<&Pool>,
8276) {
8277 debug_assert_eq!(o1.len(), rows);
8278 debug_assert_eq!(o2.len(), rows);
8279 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8280 let gpr = cols / GROUP_SIZE;
8281 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8282 let out1 = SendMut(o1.as_mut_ptr());
8283 let out2 = SendMut(o2.as_mut_ptr());
8284 if a8w8_enabled() {
8285 let a1 = split_act(x1);
8286 let a2 = split_act(x2);
8287 let (a1, a2) = (&a1, &a2);
8288 let run = move |start: usize, end: usize| {
8289 for r in start..end {
8290 #[cfg(target_arch = "aarch64")]
8291 let ds = unsafe { q1t_dot_row_sdot2(bytes, r, gpr, &a1.xq, &a2.xq) };
8294 #[cfg(not(target_arch = "aarch64"))]
8295 let ds = [
8296 q1t_dot_row_i8(bytes, r, gpr, &a1.xq),
8297 q1t_dot_row_i8(bytes, r, gpr, &a2.xq),
8298 ];
8299 let mut acc1 = ds[0] * a1.sx;
8300 for &(j, xv) in &a1.outliers {
8301 acc1 += q1t_base_weight(bytes, r, gpr, j) * xv;
8302 }
8303 acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8304 let mut acc2 = ds[1] * a2.sx;
8305 for &(j, xv) in &a2.outliers {
8306 acc2 += q1t_base_weight(bytes, r, gpr, j) * xv;
8307 }
8308 acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8309 unsafe {
8311 *out1.at(r) = acc1;
8312 *out2.at(r) = acc2;
8313 }
8314 }
8315 };
8316 dispatch_rows(pool, rows, &run);
8317 return;
8318 }
8319 let run = move |start: usize, end: usize| {
8320 let mut sg = [0f32; GROUP_SIZE];
8323 for r in start..end {
8324 let mut acc1 = 0f32;
8325 let mut acc2 = 0f32;
8326 for g in 0..gpr {
8327 let off = (r * gpr + g) * TILE;
8328 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8329 let codes = &bytes[off + 2..off + TILE];
8330 for bi in 0..6 {
8331 sg[bi * 5..bi * 5 + 5].copy_from_slice(&SIGN5[codes[bi] as usize]);
8332 }
8333 let lut = &SIGN5[codes[6] as usize];
8334 sg[30] = lut[0];
8335 sg[31] = lut[1];
8336 let xg1 = &x1[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8337 let xg2 = &x2[g * GROUP_SIZE..g * GROUP_SIZE + GROUP_SIZE];
8338 let mut gsum1 = 0f32;
8339 for k in 0..GROUP_SIZE {
8340 gsum1 += sg[k] * xg1[k];
8341 }
8342 acc1 += s * gsum1;
8343 let mut gsum2 = 0f32;
8344 for k in 0..GROUP_SIZE {
8345 gsum2 += sg[k] * xg2[k];
8346 }
8347 acc2 += s * gsum2;
8348 }
8349 acc1 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x1);
8350 acc2 += q1t_row_outlier_correction(bytes, r, rp_off, ent_off, has_ov, x2);
8351 unsafe {
8353 *out1.at(r) = acc1;
8354 *out2.at(r) = acc2;
8355 }
8356 }
8357 };
8358 dispatch_rows(pool, rows, &run);
8359}
8360
8361fn q1t_matmat(
8364 bytes: &[u8],
8365 xs: &[f32],
8366 b: usize,
8367 rows: usize,
8368 cols: usize,
8369 out: &mut [f32],
8370 pool: Option<&Pool>,
8371) {
8372 debug_assert_eq!(out.len(), b * rows);
8373 const TILE: usize = cortiq_core::quant::Q1T_TILE;
8374 let gpr = cols / GROUP_SIZE;
8375 let (rp_off, ent_off, has_ov) = q1t_overlay(bytes, rows * gpr * TILE, rows);
8376 let out_addr = SendMut(out.as_mut_ptr());
8377 if a8w8_enabled() {
8381 let acts: Vec<SplitAct> = (0..b)
8382 .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
8383 .collect();
8384 let acts = &acts;
8385 let run = move |start: usize, end: usize| {
8386 let mut sg = vec![0i8; cols + 8]; let mut sc = vec![0f32; gpr]; let mut accs = vec![0f32; b]; for r in start..end {
8390 for g in 0..gpr {
8391 let off = (r * gpr + g) * TILE;
8392 sc[g] = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
8393 q1t_unpack_group_i8(
8394 bytes.as_ptr().wrapping_add(off + 2),
8395 &mut sg[g * GROUP_SIZE..],
8396 );
8397 }
8398 for bi in 0..b {
8399 let act = &acts[bi];
8400 let mut isum = 0f32;
8401 for g in 0..gpr {
8402 let d = q1t_i8dot32(
8403 sg.as_ptr().wrapping_add(g * GROUP_SIZE),
8404 act.xq.as_ptr().wrapping_add(g * GROUP_SIZE),
8405 );
8406 isum += d as f32 * sc[g];
8407 }
8408 let mut acc = isum * act.sx;
8409 for &(j, xv) in &act.outliers {
8410 acc += q1t_base_weight(bytes, r, gpr, j) * xv;
8411 }
8412 accs[bi] = acc;
8413 }
8414 if has_ov {
8418 let (c0, c1) = (
8419 q1t_rowptr(bytes, rp_off, r),
8420 q1t_rowptr(bytes, rp_off, r + 1),
8421 );
8422 for p in c0..c1 {
8423 let e = ent_off + p * 4;
8424 let col = u16::from_le_bytes([bytes[e], bytes[e + 1]]) as usize;
8425 let val = f16_to_f32(u16::from_le_bytes([bytes[e + 2], bytes[e + 3]]));
8426 for bi in 0..b {
8427 accs[bi] += val * xs[bi * cols + col];
8428 }
8429 }
8430 }
8431 for bi in 0..b {
8432 unsafe { *out_addr.at(bi * rows + r) = accs[bi] };
8433 }
8434 }
8435 };
8436 dispatch_rows(pool, rows, &run);
8437 return;
8438 }
8439 let run = move |start: usize, end: usize| {
8440 let mut buf = vec![0f32; cols];
8441 for r in start..end {
8442 q1t_dequant_row(bytes, r, gpr, rp_off, ent_off, has_ov, &mut buf);
8443 for bi in 0..b {
8444 let xr = &xs[bi * cols..(bi + 1) * cols];
8445 let mut acc = 0f32;
8446 for j in 0..cols {
8447 acc += buf[j] * xr[j];
8448 }
8449 unsafe { *out_addr.at(bi * rows + r) = acc };
8450 }
8451 }
8452 };
8453 dispatch_rows(pool, rows, &run);
8454}
8455
8456fn q1_matvec(
8457 bytes: &[u8],
8458 x: &[f32],
8459 rows: usize,
8460 cols: usize,
8461 out: &mut [f32],
8462 pool: Option<&Pool>,
8463) {
8464 debug_assert_eq!(out.len(), rows);
8465 let gpr = cols / GROUP_SIZE;
8466 let out_addr = SendMut(out.as_mut_ptr());
8467 if a8w8_enabled() {
8468 let act = split_act(x);
8469 let gsum = q1_group_sums(&act.xq, gpr);
8470 let (act, gsum) = (&act, &gsum);
8471 let run = move |start: usize, end: usize| {
8472 q1_range_a8w8(bytes, gpr, act, gsum, out_addr, start, end)
8473 };
8474 dispatch_rows(pool, rows, &run);
8475 return;
8476 }
8477 let run = move |start: usize, end: usize| q1_range_f32(bytes, gpr, x, out_addr, start, end);
8478 dispatch_rows(pool, rows, &run);
8479}
8480
8481#[allow(clippy::too_many_arguments)]
8483fn q1_matvec2(
8484 bytes: &[u8],
8485 x1: &[f32],
8486 x2: &[f32],
8487 rows: usize,
8488 cols: usize,
8489 o1: &mut [f32],
8490 o2: &mut [f32],
8491 pool: Option<&Pool>,
8492) {
8493 let gpr = cols / GROUP_SIZE;
8494 let p1 = SendMut(o1.as_mut_ptr());
8495 let p2 = SendMut(o2.as_mut_ptr());
8496 if a8w8_enabled() {
8497 let a1 = split_act(x1);
8498 let a2 = split_act(x2);
8499 let g1 = q1_group_sums(&a1.xq, gpr);
8500 let g2 = q1_group_sums(&a2.xq, gpr);
8501 let (a1, a2, g1, g2) = (&a1, &a2, &g1, &g2);
8502 let run = move |start: usize, end: usize| {
8503 for r in start..end {
8504 let mut v1 = dot_q1_row_i8(bytes, r, gpr, &a1.xq, g1) * a1.sx;
8505 let mut v2 = dot_q1_row_i8(bytes, r, gpr, &a2.xq, g2) * a2.sx;
8506 for &(j, xv) in &a1.outliers {
8507 let (w, s) = q1_outlier(bytes, r, gpr, j);
8508 v1 += w * s * xv;
8509 }
8510 for &(j, xv) in &a2.outliers {
8511 let (w, s) = q1_outlier(bytes, r, gpr, j);
8512 v2 += w * s * xv;
8513 }
8514 unsafe {
8516 *p1.at(r) = v1;
8517 *p2.at(r) = v2;
8518 }
8519 }
8520 };
8521 dispatch_rows(pool, rows, &run);
8522 return;
8523 }
8524 let run = move |start: usize, end: usize| {
8525 for r in start..end {
8526 unsafe {
8528 *p1.at(r) = q1_row_exact(bytes, r, gpr, x1);
8529 *p2.at(r) = q1_row_exact(bytes, r, gpr, x2);
8530 }
8531 }
8532 };
8533 dispatch_rows(pool, rows, &run);
8534}
8535
8536#[allow(clippy::too_many_arguments)]
8538fn q1_matmat(
8539 bytes: &[u8],
8540 xs_all: &[f32],
8541 b: usize,
8542 rows: usize,
8543 cols: usize,
8544 out: &mut [f32],
8545 pool: Option<&Pool>,
8546) {
8547 debug_assert_eq!(out.len(), b * rows);
8548 let gpr = cols / GROUP_SIZE;
8549 let out_addr = SendMut(out.as_mut_ptr());
8550 if a8w8_enabled() {
8551 let acts: Vec<(SplitAct, Vec<i32>)> = (0..b)
8552 .map(|bi| {
8553 let act = split_act(&xs_all[bi * cols..(bi + 1) * cols]);
8554 let gsum = q1_group_sums(&act.xq, gpr);
8555 (act, gsum)
8556 })
8557 .collect();
8558 let acts = &acts;
8559 #[cfg(target_arch = "x86_64")]
8560 let blocked_ok = avx2_enabled() && blocked_enabled();
8561 #[cfg(target_arch = "aarch64")]
8562 let blocked_ok = sdot_enabled() && blocked_enabled();
8563 let run = move |start: usize, end: usize| {
8564 for r in start..end {
8565 let mut bi = 0usize;
8566 #[cfg(target_arch = "aarch64")]
8569 if blocked_ok {
8570 while bi + 4 <= acts.len() {
8571 let xs = [
8572 acts[bi].0.xq.as_slice(),
8573 acts[bi + 1].0.xq.as_slice(),
8574 acts[bi + 2].0.xq.as_slice(),
8575 acts[bi + 3].0.xq.as_slice(),
8576 ];
8577 let gs = [
8578 acts[bi].1.as_slice(),
8579 acts[bi + 1].1.as_slice(),
8580 acts[bi + 2].1.as_slice(),
8581 acts[bi + 3].1.as_slice(),
8582 ];
8583 let d = unsafe { dot_q1_row_1x4_sdot(bytes, r, gpr, xs, gs) };
8584 for k in 0..4 {
8585 let (act, _) = &acts[bi + k];
8586 let mut acc = d[k] * act.sx;
8587 for &(j, xv) in &act.outliers {
8588 let (w, sc) = q1_outlier(bytes, r, gpr, j);
8589 acc += w * sc * xv;
8590 }
8591 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8593 }
8594 bi += 4;
8595 }
8596 }
8597 #[cfg(target_arch = "x86_64")]
8598 if blocked_ok {
8599 while bi + 4 <= acts.len() {
8600 let xs = [
8601 acts[bi].0.xq.as_slice(),
8602 acts[bi + 1].0.xq.as_slice(),
8603 acts[bi + 2].0.xq.as_slice(),
8604 acts[bi + 3].0.xq.as_slice(),
8605 ];
8606 let gs = [
8607 acts[bi].1.as_slice(),
8608 acts[bi + 1].1.as_slice(),
8609 acts[bi + 2].1.as_slice(),
8610 acts[bi + 3].1.as_slice(),
8611 ];
8612 let d = unsafe {
8613 if vnni_tiles_enabled() {
8614 dot_q1_row_1x4_vnni(bytes, r, gpr, xs, gs)
8615 } else {
8616 dot_q1_row_1x4_avx2(bytes, r, gpr, xs, gs)
8617 }
8618 };
8619 for k in 0..4 {
8620 let (act, _) = &acts[bi + k];
8621 let mut acc = d[k] * act.sx;
8622 for &(j, xv) in &act.outliers {
8623 let (w, sc) = q1_outlier(bytes, r, gpr, j);
8624 acc += w * sc * xv;
8625 }
8626 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
8628 }
8629 bi += 4;
8630 }
8631 }
8632 while bi < acts.len() {
8633 let (act, gsum) = &acts[bi];
8634 let mut acc = dot_q1_row_i8(bytes, r, gpr, &act.xq, gsum) * act.sx;
8635 for &(j, xv) in &act.outliers {
8636 let (w, s) = q1_outlier(bytes, r, gpr, j);
8637 acc += w * s * xv;
8638 }
8639 unsafe { *out_addr.at(bi * rows + r) = acc };
8641 bi += 1;
8642 }
8643 }
8644 };
8645 dispatch_rows(pool, rows, &run);
8646 return;
8647 }
8648 let run = move |start: usize, end: usize| {
8649 for r in start..end {
8650 for bi in 0..b {
8651 let x = &xs_all[bi * cols..(bi + 1) * cols];
8652 unsafe { *out_addr.at(bi * rows + r) = q1_row_exact(bytes, r, gpr, x) };
8654 }
8655 }
8656 };
8657 dispatch_rows(pool, rows, &run);
8658}
8659
8660fn q4matvec(
8666 bytes: &[u8],
8667 x: &[f32],
8668 rows: usize,
8669 cols: usize,
8670 out: &mut [f32],
8671 pool: Option<&Pool>,
8672) {
8673 debug_assert_eq!(out.len(), rows);
8674 let (packed, scales) = q4_split(bytes, rows, cols);
8675 let gpr = cols / GROUP_SIZE;
8676 let out_addr = SendMut(out.as_mut_ptr());
8677
8678 if a8w8_enabled() {
8679 let act = split_act(x);
8680 let run = move |start: usize, end: usize| {
8681 q4_range_a8w8(packed, scales, gpr, cols, &act, out_addr, start, end)
8682 };
8683 dispatch_rows(pool, rows, &run);
8684 return;
8685 }
8686
8687 let run =
8688 move |start: usize, end: usize| q4_range_f32(packed, scales, gpr, x, out_addr, start, end);
8689 dispatch_rows(pool, rows, &run);
8690}
8691
8692#[inline]
8695#[allow(unreachable_code)]
8696#[cfg(target_arch = "x86_64")]
8701#[target_feature(enable = "avx2")]
8702unsafe fn dot_q4b_row_1x4_avx2(
8703 buf: &[u8],
8704 scales: &[u8],
8705 g0: usize,
8706 gpr: usize,
8707 xs: [&[i8]; 4],
8708) -> [f32; 4] {
8709 unsafe {
8711 use core::arch::x86_64::*;
8712 let ones = _mm256_set1_epi16(1);
8713 let mut acc = [0f32; 4];
8714 for gi in 0..gpr {
8715 let s = f16_to_f32(u16::from_le_bytes([
8716 scales[(g0 + gi) * 2],
8717 scales[(g0 + gi) * 2 + 1],
8718 ]));
8719 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8720 let aw = _mm256_abs_epi8(w);
8721 for (k, xq) in xs.iter().enumerate() {
8722 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8723 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8724 let d = _mm256_madd_epi16(p16, ones);
8725 let hi128 = _mm256_extracti128_si256::<1>(d);
8726 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8727 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8728 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8729 acc[k] += _mm_cvtsi128_si32(s32) as f32 * s;
8730 }
8731 }
8732 acc
8733 }
8734}
8735
8736#[cfg(target_arch = "x86_64")]
8738#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8739unsafe fn dot_q4b_row_1x4_vnni(
8740 buf: &[u8],
8741 scales: &[u8],
8742 g0: usize,
8743 gpr: usize,
8744 xs: [&[i8]; 4],
8745) -> [f32; 4] {
8746 unsafe {
8748 use core::arch::x86_64::*;
8749 let mut acc = [0f32; 4];
8750 for gi in 0..gpr {
8751 let s = f16_to_f32(u16::from_le_bytes([
8752 scales[(g0 + gi) * 2],
8753 scales[(g0 + gi) * 2 + 1],
8754 ]));
8755 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8756 let aw = _mm256_abs_epi8(w);
8757 for (k, xq) in xs.iter().enumerate() {
8758 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8759 let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8760 acc[k] += d as f32 * s;
8761 }
8762 }
8763 acc
8764 }
8765}
8766
8767#[cfg(target_arch = "x86_64")]
8773#[target_feature(enable = "avx2")]
8774unsafe fn dot_q4b_row_1x4_sx_avx2(
8775 buf: &[u8],
8776 scales: &[u8],
8777 g0: usize,
8778 gpr: usize,
8779 xs: [&[i8]; 4],
8780 sxs: [f32; 4],
8781) -> [f32; 4] {
8782 unsafe {
8784 use core::arch::x86_64::*;
8785 let ones = _mm256_set1_epi16(1);
8786 let mut acc = [0f32; 4];
8787 for gi in 0..gpr {
8788 let s = f16_to_f32(u16::from_le_bytes([
8789 scales[(g0 + gi) * 2],
8790 scales[(g0 + gi) * 2 + 1],
8791 ]));
8792 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8793 let aw = _mm256_abs_epi8(w);
8794 for (k, xq) in xs.iter().enumerate() {
8795 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8796 let p16 = _mm256_maddubs_epi16(aw, _mm256_sign_epi8(x, w));
8797 let d = _mm256_madd_epi16(p16, ones);
8798 let hi128 = _mm256_extracti128_si256::<1>(d);
8799 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
8800 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
8801 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
8802 acc[k] += (_mm_cvtsi128_si32(s32) as f32 * sxs[k]) * s;
8803 }
8804 }
8805 acc
8806 }
8807}
8808
8809#[cfg(target_arch = "x86_64")]
8812#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
8813unsafe fn dot_q4b_row_1x4_sx_vnni(
8814 buf: &[u8],
8815 scales: &[u8],
8816 g0: usize,
8817 gpr: usize,
8818 xs: [&[i8]; 4],
8819 sxs: [f32; 4],
8820) -> [f32; 4] {
8821 unsafe {
8823 use core::arch::x86_64::*;
8824 let mut acc = [0f32; 4];
8825 for gi in 0..gpr {
8826 let s = f16_to_f32(u16::from_le_bytes([
8827 scales[(g0 + gi) * 2],
8828 scales[(g0 + gi) * 2 + 1],
8829 ]));
8830 let w = _mm256_loadu_si256(buf.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8831 let aw = _mm256_abs_epi8(w);
8832 for (k, xq) in xs.iter().enumerate() {
8833 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
8834 let d = dpbusd_hsum(aw, _mm256_sign_epi8(x, w));
8835 acc[k] += (d as f32 * sxs[k]) * s;
8836 }
8837 }
8838 acc
8839 }
8840}
8841
8842#[allow(unreachable_code)]
8843fn dot_q4_row_i8(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
8844 #[cfg(target_arch = "aarch64")]
8845 unsafe {
8846 return dot_q4_row_sdot(packed, scales, g0, gpr, xq);
8847 }
8848 #[cfg(target_arch = "x86_64")]
8849 unsafe {
8850 return dot_q4_row_avx2(packed, scales, g0, gpr, xq);
8851 }
8852 let mut acc = 0f32;
8853 for gi in 0..gpr {
8854 let g = g0 + gi;
8855 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8856 let mut d = 0i32;
8857 for (k, &b) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
8858 d += ((b & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2] as i32
8859 + (((b >> 4) & 0x0F) as i32 - 8) * xq[gi * GROUP_SIZE + k * 2 + 1] as i32;
8860 }
8861 acc += d as f32 * s;
8862 }
8863 acc
8864}
8865
8866#[inline]
8868#[allow(unreachable_code)]
8869fn dot_q4_row_i8_2(
8870 packed: &[u8],
8871 scales: &[u8],
8872 g0: usize,
8873 gpr: usize,
8874 xq1: &[i8],
8875 xq2: &[i8],
8876) -> (f32, f32) {
8877 #[cfg(target_arch = "aarch64")]
8878 unsafe {
8879 return dot_q4_row_sdot2(packed, scales, g0, gpr, xq1, xq2);
8880 }
8881 #[cfg(target_arch = "x86_64")]
8882 unsafe {
8883 return dot_q4_row_avx2_2(packed, scales, g0, gpr, xq1, xq2);
8884 }
8885 (
8886 dot_q4_row_i8(packed, scales, g0, gpr, xq1),
8887 dot_q4_row_i8(packed, scales, g0, gpr, xq2),
8888 )
8889}
8890
8891#[allow(clippy::too_many_arguments)]
8894fn q4_range_a8w8(
8895 packed: &[u8],
8896 scales: &[u8],
8897 gpr: usize,
8898 cols: usize,
8899 act: &SplitAct,
8900 out: SendMut,
8901 start: usize,
8902 end: usize,
8903) {
8904 for r in start..end {
8905 let mut acc = dot_q4_row_i8(packed, scales, r * gpr, gpr, &act.xq) * act.sx;
8906 for &(j, xv) in &act.outliers {
8908 let flat = r * cols + j;
8909 let byte = packed[flat / 2];
8910 let nib = if flat & 1 == 0 {
8911 byte & 0x0F
8912 } else {
8913 byte >> 4
8914 };
8915 let s = f16_to_f32(u16::from_le_bytes([
8916 scales[(flat / GROUP_SIZE) * 2],
8917 scales[(flat / GROUP_SIZE) * 2 + 1],
8918 ]));
8919 acc += ((nib as i32 - 8) as f32) * s * xv;
8920 }
8921 unsafe { *out.at(r) = acc };
8923 }
8924}
8925
8926#[allow(clippy::too_many_arguments)]
8929fn q4_range2_a8w8(
8930 packed: &[u8],
8931 scales: &[u8],
8932 gpr: usize,
8933 cols: usize,
8934 a1: &SplitAct,
8935 a2: &SplitAct,
8936 p1: SendMut,
8937 p2: SendMut,
8938 start: usize,
8939 end: usize,
8940) {
8941 for r in start..end {
8942 let (s1, s2) = dot_q4_row_i8_2(packed, scales, r * gpr, gpr, &a1.xq, &a2.xq);
8943 let mut acc1 = s1 * a1.sx;
8944 let mut acc2 = s2 * a2.sx;
8945 let fix = |outliers: &[(usize, f32)], acc: &mut f32| {
8947 for &(j, xv) in outliers {
8948 let flat = r * cols + j;
8949 let byte = packed[flat / 2];
8950 let nib = if flat & 1 == 0 {
8951 byte & 0x0F
8952 } else {
8953 byte >> 4
8954 };
8955 let s = f16_to_f32(u16::from_le_bytes([
8956 scales[(flat / GROUP_SIZE) * 2],
8957 scales[(flat / GROUP_SIZE) * 2 + 1],
8958 ]));
8959 *acc += ((nib as i32 - 8) as f32) * s * xv;
8960 }
8961 };
8962 fix(&a1.outliers, &mut acc1);
8963 fix(&a2.outliers, &mut acc2);
8964 unsafe {
8966 *p1.at(r) = acc1;
8967 *p2.at(r) = acc2;
8968 }
8969 }
8970}
8971
8972fn q4_range_f32(
8974 packed: &[u8],
8975 scales: &[u8],
8976 gpr: usize,
8977 x: &[f32],
8978 out: SendMut,
8979 start: usize,
8980 end: usize,
8981) {
8982 for r in start..end {
8983 let mut acc = 0f32;
8984 for gi in 0..gpr {
8985 let g = r * gpr + gi;
8986 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
8987 let pk = &packed[g * 16..(g + 1) * 16];
8988 let xg = &x[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
8989 let mut ga = 0f32;
8990 for (k, &b) in pk.iter().enumerate() {
8991 ga += ((b & 0x0F) as f32 - 8.0) * xg[k * 2]
8992 + (((b >> 4) & 0x0F) as f32 - 8.0) * xg[k * 2 + 1];
8993 }
8994 acc += ga * s;
8995 }
8996 unsafe { *out.at(r) = acc };
8998 }
8999}
9000
9001#[allow(clippy::too_many_arguments)]
9005fn q4matvec2(
9006 bytes: &[u8],
9007 x1: &[f32],
9008 x2: &[f32],
9009 rows: usize,
9010 cols: usize,
9011 o1: &mut [f32],
9012 o2: &mut [f32],
9013 pool: Option<&Pool>,
9014) {
9015 debug_assert_eq!(o1.len(), rows);
9016 debug_assert_eq!(o2.len(), rows);
9017 let (packed, scales) = q4_split(bytes, rows, cols);
9018 let gpr = cols / GROUP_SIZE;
9019
9020 if a8w8_enabled() {
9021 let a1 = split_act(x1);
9022 let a2 = split_act(x2);
9023 let p1 = SendMut(o1.as_mut_ptr());
9024 let p2 = SendMut(o2.as_mut_ptr());
9025 let run = move |start: usize, end: usize| {
9026 q4_range2_a8w8(packed, scales, gpr, cols, &a1, &a2, p1, p2, start, end)
9027 };
9028 dispatch_rows(pool, rows, &run);
9029 return;
9030 }
9031
9032 let p1 = SendMut(o1.as_mut_ptr());
9033 let p2 = SendMut(o2.as_mut_ptr());
9034 let run = move |start: usize, end: usize| {
9035 q4_range2_f32(packed, scales, gpr, x1, x2, p1, p2, start, end)
9036 };
9037 dispatch_rows(pool, rows, &run);
9038}
9039
9040#[allow(clippy::too_many_arguments)]
9042fn q4_range2_f32(
9043 packed: &[u8],
9044 scales: &[u8],
9045 gpr: usize,
9046 x1: &[f32],
9047 x2: &[f32],
9048 p1: SendMut,
9049 p2: SendMut,
9050 start: usize,
9051 end: usize,
9052) {
9053 for r in start..end {
9054 let (mut acc1, mut acc2) = (0f32, 0f32);
9055 for gi in 0..gpr {
9056 let g = r * gpr + gi;
9057 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9058 let pk = &packed[g * 16..(g + 1) * 16];
9059 let x1g = &x1[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9060 let x2g = &x2[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE];
9061 let (mut g1, mut g2) = (0f32, 0f32);
9062 for (k, &b) in pk.iter().enumerate() {
9063 let wl = (b & 0x0F) as f32 - 8.0;
9064 let wh = ((b >> 4) & 0x0F) as f32 - 8.0;
9065 g1 += wl * x1g[k * 2] + wh * x1g[k * 2 + 1];
9066 g2 += wl * x2g[k * 2] + wh * x2g[k * 2 + 1];
9067 }
9068 acc1 += g1 * s;
9069 acc2 += g2 * s;
9070 }
9071 unsafe {
9073 *p1.at(r) = acc1;
9074 *p2.at(r) = acc2;
9075 }
9076 }
9077}
9078
9079thread_local! {
9080 static ROW_I8: std::cell::RefCell<Vec<u8>> = const { std::cell::RefCell::new(Vec::new()) };
9083 static ROW_F32: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9084}
9085
9086#[allow(clippy::too_many_arguments)]
9092fn q4matmat(
9093 bytes: &[u8],
9094 xs_all: &[f32],
9095 b: usize,
9096 rows: usize,
9097 cols: usize,
9098 out: &mut [f32],
9099 pool: Option<&Pool>,
9100) {
9101 debug_assert_eq!(xs_all.len(), b * cols);
9102 debug_assert_eq!(out.len(), b * rows);
9103 let (packed, scales) = q4_split(bytes, rows, cols);
9104 let gpr = cols / GROUP_SIZE;
9105 let gscale = |g: usize| f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
9106
9107 if a8w8_enabled() {
9108 let acts: Vec<SplitAct> = (0..b)
9109 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9110 .collect();
9111 let acts = &acts;
9112 let out_addr = SendMut(out.as_mut_ptr());
9113 let run = move |start: usize, end: usize| {
9114 ROW_I8.with(|rb| {
9115 let mut buf = rb.borrow_mut();
9116 buf.resize(cols, 0);
9117 for r in start..end {
9118 for gi in 0..gpr {
9122 let g = r * gpr + gi;
9123 for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9124 buf[gi * GROUP_SIZE + k * 2] = ((bt & 0x0F) as i32 - 8) as i8 as u8;
9125 buf[gi * GROUP_SIZE + k * 2 + 1] =
9126 (((bt >> 4) & 0x0F) as i32 - 8) as i8 as u8;
9127 }
9128 }
9129 let mut bi = 0usize;
9130 #[cfg(target_arch = "x86_64")]
9131 if avx2_enabled() && blocked_enabled() {
9132 while bi + 4 <= acts.len() {
9133 let xs = [
9134 acts[bi].xq.as_slice(),
9135 acts[bi + 1].xq.as_slice(),
9136 acts[bi + 2].xq.as_slice(),
9137 acts[bi + 3].xq.as_slice(),
9138 ];
9139 let d = unsafe {
9140 if vnni_tiles_enabled() {
9141 dot_q4b_row_1x4_vnni(&buf, scales, r * gpr, gpr, xs)
9142 } else {
9143 dot_q4b_row_1x4_avx2(&buf, scales, r * gpr, gpr, xs)
9144 }
9145 };
9146 for k in 0..4 {
9147 let act = &acts[bi + k];
9148 let mut acc = d[k] * act.sx;
9149 for &(j, xv) in &act.outliers {
9150 acc += (buf[j] as i8) as f32
9151 * gscale((r * cols + j) / GROUP_SIZE)
9152 * xv;
9153 }
9154 unsafe { *out_addr.at((bi + k) * rows + r) = acc };
9156 }
9157 bi += 4;
9158 }
9159 }
9160 while bi < acts.len() {
9161 let act = &acts[bi];
9162 let mut acc = 0f32;
9163 for gi in 0..gpr {
9164 let d = dot_i8_i8(
9165 &buf[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9166 &act.xq[gi * GROUP_SIZE..(gi + 1) * GROUP_SIZE],
9167 );
9168 acc += d as f32 * gscale(r * gpr + gi);
9169 }
9170 acc *= act.sx;
9171 for &(j, xv) in &act.outliers {
9173 acc += (buf[j] as i8) as f32 * gscale((r * cols + j) / GROUP_SIZE) * xv;
9174 }
9175 unsafe { *out_addr.at(bi * rows + r) = acc };
9177 bi += 1;
9178 }
9179 }
9180 })
9181 };
9182 dispatch_rows(pool, rows, &run);
9183 return;
9184 }
9185
9186 let out_addr = SendMut(out.as_mut_ptr());
9187 let run = move |start: usize, end: usize| {
9188 ROW_F32.with(|rb| {
9189 let mut buf = rb.borrow_mut();
9190 buf.resize(cols, 0.0);
9191 for r in start..end {
9192 for gi in 0..gpr {
9195 let g = r * gpr + gi;
9196 for (k, &bt) in packed[g * 16..(g + 1) * 16].iter().enumerate() {
9197 buf[gi * GROUP_SIZE + k * 2] = (bt & 0x0F) as f32 - 8.0;
9198 buf[gi * GROUP_SIZE + k * 2 + 1] = ((bt >> 4) & 0x0F) as f32 - 8.0;
9199 }
9200 }
9201 for bi in 0..b {
9202 let x = &xs_all[bi * cols..(bi + 1) * cols];
9203 let mut acc = 0f32;
9204 for gi in 0..gpr {
9205 let mut ga = 0f32;
9206 for k in 0..GROUP_SIZE / 2 {
9211 let e = gi * GROUP_SIZE + k * 2;
9212 ga += buf[e] * x[e] + buf[e + 1] * x[e + 1];
9213 }
9214 acc += ga * gscale(r * gpr + gi);
9215 }
9216 unsafe { *out_addr.at(bi * rows + r) = acc };
9218 }
9219 }
9220 })
9221 };
9222 dispatch_rows(pool, rows, &run);
9223}
9224
9225#[allow(clippy::too_many_arguments)]
9230fn vbitmatmat(
9231 bytes: &[u8],
9232 offsets: &[usize],
9233 xs_all: &[f32],
9234 b: usize,
9235 rows: usize,
9236 cols: usize,
9237 out: &mut [f32],
9238 pool: Option<&Pool>,
9239) {
9240 debug_assert_eq!(xs_all.len(), b * cols);
9241 debug_assert_eq!(out.len(), b * rows);
9242 debug_assert_eq!(offsets.len(), rows + 1);
9243 let ng = cols / GROUP_SIZE;
9244 let bits = &bytes[..rows];
9245 let sc_off = rows;
9246 let gscale = |r: usize, g: usize| {
9247 let so = (r * ng + g) * 2;
9248 f16_to_f32(u16::from_le_bytes([
9249 bytes[sc_off + so],
9250 bytes[sc_off + so + 1],
9251 ]))
9252 };
9253
9254 let decode_f32 = |r: usize, dst: &mut [f32]| {
9256 let bw = bits[r] as usize;
9257 let l = ((1i32 << (bw - 1)) - 1) as f32;
9258 let data = &bytes[offsets[r]..offsets[r + 1]];
9259 let (mut acc, mut nbits, mut idx) = (0u64, 0usize, 0usize);
9260 for d in dst.iter_mut() {
9261 while nbits < bw {
9262 acc = (acc << 8) | data[idx] as u64;
9263 idx += 1;
9264 nbits += 8;
9265 }
9266 let u = ((acc >> (nbits - bw)) & ((1u64 << bw) - 1)) as f32;
9267 nbits -= bw;
9268 *d = u - l;
9269 }
9270 };
9271
9272 if a8w8_enabled() {
9273 let acts: Vec<SplitAct> = (0..b)
9274 .map(|bi| split_act(&xs_all[bi * cols..(bi + 1) * cols]))
9275 .collect();
9276 let acts = &acts;
9277 let out_addr = SendMut(out.as_mut_ptr());
9278 let run = move |start: usize, end: usize| {
9279 for r in start..end {
9280 let bw = bits[r] as usize;
9281 if bw == 8 {
9282 ROW_F32.with(|rb| {
9285 let mut buf = rb.borrow_mut();
9286 buf.resize(cols, 0.0);
9287 decode_f32(r, &mut buf);
9288 for bi in 0..b {
9289 let x = &xs_all[bi * cols..(bi + 1) * cols];
9290 let mut dot = 0f32;
9291 for g in 0..ng {
9292 let mut gd = 0f32;
9293 for k in 0..GROUP_SIZE {
9294 gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9295 }
9296 dot += gd * gscale(r, g);
9297 }
9298 unsafe { *out_addr.at(bi * rows + r) = dot };
9300 }
9301 });
9302 continue;
9303 }
9304 let l = (1i32 << (bw - 1)) - 1;
9305 let data = &bytes[offsets[r]..offsets[r + 1]];
9306 ROW_I8.with(|rb| {
9307 let mut buf = rb.borrow_mut();
9308 buf.resize(cols, 0);
9309 #[inline(always)]
9310 fn fill<const B: usize>(data: &[u8], l: i32, buf: &mut [u8]) {
9311 for (blk, chunk) in buf.chunks_exact_mut(8).enumerate() {
9312 let u = unpack8::<B>(&data[blk * B..]);
9313 for k in 0..8 {
9314 chunk[k] = (u[k] - l) as i8 as u8;
9315 }
9316 }
9317 }
9318 match bw {
9319 3 => fill::<3>(data, l, &mut buf),
9320 4 => vbit_fill4(data, &mut buf),
9321 5 => fill::<5>(data, l, &mut buf),
9322 6 => fill::<6>(data, l, &mut buf),
9323 _ => unreachable!("vbit bit-width {bw} (validated at load)"),
9324 }
9325 let mut bi = 0usize;
9326 #[cfg(target_arch = "x86_64")]
9330 if avx2_enabled() && blocked_enabled() {
9331 while bi + 4 <= acts.len() {
9332 let xs = [
9333 acts[bi].xq.as_slice(),
9334 acts[bi + 1].xq.as_slice(),
9335 acts[bi + 2].xq.as_slice(),
9336 acts[bi + 3].xq.as_slice(),
9337 ];
9338 let sxs = [
9339 acts[bi].sx,
9340 acts[bi + 1].sx,
9341 acts[bi + 2].sx,
9342 acts[bi + 3].sx,
9343 ];
9344 let d = unsafe {
9345 if vnni_tiles_enabled() {
9346 dot_q4b_row_1x4_sx_vnni(
9347 &buf,
9348 &bytes[sc_off..],
9349 r * ng,
9350 ng,
9351 xs,
9352 sxs,
9353 )
9354 } else {
9355 dot_q4b_row_1x4_sx_avx2(
9356 &buf,
9357 &bytes[sc_off..],
9358 r * ng,
9359 ng,
9360 xs,
9361 sxs,
9362 )
9363 }
9364 };
9365 for k in 0..4 {
9366 let act = &acts[bi + k];
9367 let mut dot = d[k];
9368 for &(j, xv) in &act.outliers {
9369 dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9370 }
9371 unsafe { *out_addr.at((bi + k) * rows + r) = dot };
9373 }
9374 bi += 4;
9375 }
9376 }
9377 while bi < acts.len() {
9378 let act = &acts[bi];
9379 let mut dot = 0f32;
9380 for g in 0..ng {
9381 let d = dot_i8_i8(
9382 &buf[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9383 &act.xq[g * GROUP_SIZE..(g + 1) * GROUP_SIZE],
9384 ) as f32
9385 * act.sx;
9386 dot += d * gscale(r, g);
9387 }
9388 for &(j, xv) in &act.outliers {
9389 dot += (buf[j] as i8) as f32 * gscale(r, j / GROUP_SIZE) * xv;
9390 }
9391 unsafe { *out_addr.at(bi * rows + r) = dot };
9393 bi += 1;
9394 }
9395 });
9396 }
9397 };
9398 dispatch_rows(pool, rows, &run);
9399 return;
9400 }
9401
9402 let out_addr = SendMut(out.as_mut_ptr());
9403 let run = move |start: usize, end: usize| {
9404 ROW_F32.with(|rb| {
9405 let mut buf = rb.borrow_mut();
9406 buf.resize(cols, 0.0);
9407 for r in start..end {
9408 decode_f32(r, &mut buf);
9409 for bi in 0..b {
9410 let x = &xs_all[bi * cols..(bi + 1) * cols];
9411 let mut dot = 0f32;
9412 for g in 0..ng {
9413 let mut gd = 0f32;
9414 for k in 0..GROUP_SIZE {
9415 gd += buf[g * GROUP_SIZE + k] * x[g * GROUP_SIZE + k];
9416 }
9417 dot += gd * gscale(r, g);
9418 }
9419 unsafe { *out_addr.at(bi * rows + r) = dot };
9421 }
9422 }
9423 })
9424 };
9425 dispatch_rows(pool, rows, &run);
9426}
9427
9428pub(crate) fn gpu_batch_job<'a>(
9432 t: &'a QTensor,
9433 x: &[f32],
9434) -> Option<(std::sync::Arc<CmfModel>, crate::gpu::BatchJob<'a>)> {
9435 match t {
9436 QTensor::Mapped {
9437 model,
9438 idx,
9439 dtype: dt @ (TensorDtype::Q8Row | TensorDtype::Q8_2f),
9440 rows,
9441 cols,
9442 row_scale,
9443 col_field,
9444 ..
9445 } => Some((
9446 model.clone(),
9447 crate::gpu::BatchJob {
9448 idx: *idx,
9449 rows: *rows,
9450 cols: *cols,
9451 row_scale,
9452 xs: prescale(x, col_field, *dt).into_owned(),
9453 layout: crate::gpu::BatchLayout::Q8,
9454 },
9455 )),
9456 QTensor::Mapped {
9458 model,
9459 idx,
9460 dtype: TensorDtype::Q1,
9461 rows,
9462 cols,
9463 ..
9464 } => Some((
9465 model.clone(),
9466 crate::gpu::BatchJob {
9467 idx: *idx,
9468 rows: *rows,
9469 cols: *cols,
9470 row_scale: &[],
9471 xs: x.to_vec(),
9472 layout: crate::gpu::BatchLayout::Q1,
9473 },
9474 )),
9475 QTensor::Mapped {
9480 model,
9481 idx,
9482 dtype: dt @ (TensorDtype::Q4Tiled | TensorDtype::Q4TiledP),
9483 rows,
9484 cols,
9485 ..
9486 } => Some((
9487 model.clone(),
9488 crate::gpu::BatchJob {
9489 idx: *idx,
9490 rows: *rows,
9491 cols: *cols,
9492 row_scale: &[],
9493 xs: x.to_vec(),
9494 layout: if *dt == TensorDtype::Q4Tiled {
9495 crate::gpu::BatchLayout::Q4t
9496 } else {
9497 crate::gpu::BatchLayout::Q4tp
9498 },
9499 },
9500 )),
9501 _ => None,
9502 }
9503}
9504
9505thread_local! {
9506 static PRESCALE_BUF1: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9507 static PRESCALE_BUF2: std::cell::RefCell<Vec<f32>> = const { std::cell::RefCell::new(Vec::new()) };
9508}
9509
9510pub(crate) fn prescale<'a>(
9511 x: &'a [f32],
9512 col_field: &[f32],
9513 dtype: TensorDtype,
9514) -> std::borrow::Cow<'a, [f32]> {
9515 if dtype == TensorDtype::Q8_2f {
9516 x.iter().zip(col_field).map(|(a, c)| a * c).collect()
9517 } else {
9518 std::borrow::Cow::Borrowed(x)
9519 }
9520}
9521
9522pub(crate) fn prescale_with<R, F: FnOnce(&[f32]) -> R>(
9525 x: &[f32],
9526 col_field: &[f32],
9527 dtype: TensorDtype,
9528 buf_id: u8,
9529 f: F,
9530) -> R {
9531 if dtype == TensorDtype::Q8_2f {
9532 if buf_id == 1 {
9533 PRESCALE_BUF1.with(|b| {
9534 let mut buf = b.borrow_mut();
9535 buf.clear();
9536 buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9537 f(&buf)
9538 })
9539 } else {
9540 PRESCALE_BUF2.with(|b| {
9541 let mut buf = b.borrow_mut();
9542 buf.clear();
9543 buf.extend(x.iter().zip(col_field).map(|(a, c)| a * c));
9544 f(&buf)
9545 })
9546 }
9547 } else {
9548 f(x)
9549 }
9550}
9551
9552#[cfg(target_arch = "x86_64")]
9557pub(crate) fn avx2_enabled() -> bool {
9558 use std::sync::OnceLock;
9559 static ON: OnceLock<bool> = OnceLock::new();
9560 *ON.get_or_init(|| {
9561 std::env::var("CMF_AVX2").map(|v| v != "0").unwrap_or(true)
9562 && std::arch::is_x86_feature_detected!("avx2")
9563 && std::arch::is_x86_feature_detected!("fma")
9564 })
9565}
9566
9567#[cfg(target_arch = "x86_64")]
9572fn avx2_a8w8_enabled() -> bool {
9573 if FLOAT_ACTIVATIONS.get() {
9574 return false;
9575 }
9576 use std::sync::OnceLock;
9577 static ON: OnceLock<bool> = OnceLock::new();
9578 *ON.get_or_init(|| {
9579 avx2_enabled() && std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true)
9580 })
9581}
9582
9583thread_local! {
9584 static FULL_GPU_Q8: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9585}
9586
9587pub(crate) fn enter_full_gpu_q8_scope() -> impl Drop {
9590 struct Restore(bool, std::marker::PhantomData<std::rc::Rc<()>>);
9591 impl Drop for Restore {
9592 fn drop(&mut self) {
9593 FULL_GPU_Q8.set(self.0);
9594 }
9595 }
9596 Restore(FULL_GPU_Q8.replace(true), std::marker::PhantomData)
9597}
9598
9599thread_local! {
9603 static FLOAT_ACTIVATIONS: std::cell::Cell<bool> = const { std::cell::Cell::new(false) };
9604}
9605
9606pub(crate) fn float_activations_scope<R>(f: impl FnOnce() -> R) -> R {
9607 struct Restore(bool);
9608 impl Drop for Restore {
9609 fn drop(&mut self) {
9610 FLOAT_ACTIVATIONS.set(self.0);
9611 }
9612 }
9613 let _restore = Restore(FLOAT_ACTIVATIONS.replace(true));
9614 f()
9615}
9616
9617static ROW_EXACT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
9629
9630pub(crate) fn row_exact() -> bool {
9631 ROW_EXACT.load(std::sync::atomic::Ordering::Acquire) != 0
9632}
9633
9634fn counted_row_exact_scope<R>(active: &std::sync::atomic::AtomicUsize, f: impl FnOnce() -> R) -> R {
9635 struct Restore<'a>(&'a std::sync::atomic::AtomicUsize);
9636 impl Drop for Restore<'_> {
9637 fn drop(&mut self) {
9638 self.0.fetch_sub(1, std::sync::atomic::Ordering::AcqRel);
9639 }
9640 }
9641 active.fetch_add(1, std::sync::atomic::Ordering::AcqRel);
9642 let _restore = Restore(active);
9643 f()
9644}
9645
9646pub(crate) fn row_exact_scope<R>(f: impl FnOnce() -> R) -> R {
9648 counted_row_exact_scope(&ROW_EXACT, f)
9649}
9650
9651#[inline]
9655pub(crate) fn a8w8_enabled() -> bool {
9656 #[cfg(target_arch = "aarch64")]
9657 {
9658 sdot_enabled()
9659 }
9660 #[cfg(target_arch = "x86_64")]
9661 {
9662 avx2_a8w8_enabled()
9663 }
9664 #[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
9665 {
9666 false
9667 }
9668}
9669
9670#[inline]
9673#[allow(unreachable_code)]
9674fn dot_i8_i8(w: &[u8], xq: &[i8]) -> i32 {
9675 #[cfg(target_arch = "aarch64")]
9676 unsafe {
9677 return dot_i8_sdot(w, xq);
9678 }
9679 #[cfg(target_arch = "x86_64")]
9680 unsafe {
9681 if avx512vnni_enabled() {
9682 return dot_i8_i8_vnni(w, xq);
9683 }
9684 return dot_i8_i8_avx2(w, xq);
9685 }
9686 w.iter()
9687 .zip(xq)
9688 .map(|(&a, &b)| (a as i8) as i32 * b as i32)
9689 .sum()
9690}
9691
9692#[cfg(target_arch = "x86_64")]
9696fn avx512vnni_enabled() -> bool {
9697 use std::sync::OnceLock;
9698 static ON: OnceLock<bool> = OnceLock::new();
9699 *ON.get_or_init(|| {
9700 std::env::var("CMF_AVX512")
9701 .map(|v| v != "0")
9702 .unwrap_or(true)
9703 && std::arch::is_x86_feature_detected!("avx512f")
9704 && std::arch::is_x86_feature_detected!("avx512bw")
9705 && std::arch::is_x86_feature_detected!("avx512vl")
9706 && std::arch::is_x86_feature_detected!("avx512vnni")
9707 })
9708}
9709
9710#[cfg(target_arch = "x86_64")]
9718fn vnni_tiles_enabled() -> bool {
9719 use std::sync::OnceLock;
9720 static ON: OnceLock<bool> = OnceLock::new();
9721 *ON.get_or_init(|| {
9722 std::env::var("CMF_VNNI_TILES")
9723 .map(|v| v != "0")
9724 .unwrap_or(true)
9725 && avx512vnni_enabled()
9726 })
9727}
9728
9729#[cfg(target_arch = "x86_64")]
9734#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9735#[inline]
9736unsafe fn dpbusd_hsum(aw: core::arch::x86_64::__m256i, xs: core::arch::x86_64::__m256i) -> i32 {
9737 unsafe {
9739 use core::arch::x86_64::*;
9740 let d = _mm256_dpbusd_epi32(_mm256_setzero_si256(), aw, xs);
9741 let hi128 = _mm256_extracti128_si256::<1>(d);
9742 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9743 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9744 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9745 _mm_cvtsi128_si32(s32)
9746 }
9747}
9748
9749#[cfg(target_arch = "x86_64")]
9754#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
9755unsafe fn dot_i8_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
9756 unsafe {
9758 use core::arch::x86_64::*;
9759 let n = w.len();
9760 let mut j = 0usize;
9761 let mut total: i32;
9762 {
9767 #[inline(always)]
9768 unsafe fn step(
9769 w: *const u8,
9770 x: *const i8,
9771 acc: core::arch::x86_64::__m512i,
9772 ) -> core::arch::x86_64::__m512i {
9773 unsafe {
9774 use core::arch::x86_64::*;
9775 let wv = _mm512_loadu_si512(w as *const _);
9776 let xv = _mm512_loadu_si512(x as *const _);
9777 let aw = _mm512_abs_epi8(wv);
9778 let neg = _mm512_movepi8_mask(wv);
9779 let sx = _mm512_mask_sub_epi8(xv, neg, _mm512_setzero_si512(), xv);
9780 _mm512_dpbusd_epi32(acc, aw, sx)
9781 }
9782 }
9783 let (mut a0, mut a1, mut a2, mut a3) = (
9784 _mm512_setzero_si512(),
9785 _mm512_setzero_si512(),
9786 _mm512_setzero_si512(),
9787 _mm512_setzero_si512(),
9788 );
9789 while j + 256 <= n {
9790 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9791 a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), a1);
9792 a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), a2);
9793 a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), a3);
9794 j += 256;
9795 }
9796 while j + 64 <= n {
9797 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), a0);
9798 j += 64;
9799 }
9800 let s01 = _mm512_add_epi32(a0, a1);
9801 let s23 = _mm512_add_epi32(a2, a3);
9802 total = _mm512_reduce_add_epi32(_mm512_add_epi32(s01, s23));
9803 }
9804 if j + 32 <= n {
9806 let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9807 let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9808 let d = _mm256_dpbusd_epi32(
9809 _mm256_setzero_si256(),
9810 _mm256_abs_epi8(wv),
9811 _mm256_sign_epi8(xv, wv),
9812 );
9813 let hi128 = _mm256_extracti128_si256::<1>(d);
9814 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
9815 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9816 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9817 total += _mm_cvtsi128_si32(s32);
9818 j += 32;
9819 }
9820 while j < n {
9821 total += (w[j] as i8) as i32 * xq[j] as i32;
9822 j += 1;
9823 }
9824 total
9825 }
9826}
9827
9828#[cfg(target_arch = "x86_64")]
9830#[target_feature(enable = "avx2,fma")]
9831unsafe fn dot_i8_f32_avx2(w: &[u8], x: &[f32]) -> f32 {
9832 unsafe {
9834 use core::arch::x86_64::*;
9835 let n = x.len();
9836 let wp = w.as_ptr();
9837 let xp = x.as_ptr();
9838 let (mut a0, mut a1) = (_mm256_setzero_ps(), _mm256_setzero_ps());
9839 let mut j = 0usize;
9840 while j + 16 <= n {
9841 let wb = _mm_loadu_si128(wp.add(j) as *const __m128i);
9842 let lo = _mm256_cvtepi8_epi32(wb);
9843 let hi = _mm256_cvtepi8_epi32(_mm_srli_si128::<8>(wb));
9844 a0 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(lo), _mm256_loadu_ps(xp.add(j)), a0);
9845 a1 = _mm256_fmadd_ps(_mm256_cvtepi32_ps(hi), _mm256_loadu_ps(xp.add(j + 8)), a1);
9846 j += 16;
9847 }
9848 let acc = _mm256_add_ps(a0, a1);
9849 let hi128 = _mm256_extractf128_ps::<1>(acc);
9850 let s128 = _mm_add_ps(_mm256_castps256_ps128(acc), hi128);
9851 let s64 = _mm_add_ps(s128, _mm_movehl_ps(s128, s128));
9852 let s32 = _mm_add_ss(s64, _mm_shuffle_ps::<1>(s64, s64));
9853 let mut sum = _mm_cvtss_f32(s32);
9854 while j < n {
9855 sum += (*wp.add(j) as i8) as f32 * *xp.add(j);
9856 j += 1;
9857 }
9858 sum
9859 }
9860}
9861
9862#[cfg(target_arch = "x86_64")]
9867#[target_feature(enable = "avx2")]
9868unsafe fn dot_i8_i8_avx2(w: &[u8], xq: &[i8]) -> i32 {
9869 unsafe {
9871 use core::arch::x86_64::*;
9872 let n = w.len();
9873 let ones = _mm256_set1_epi16(1);
9874 let mut acc = _mm256_setzero_si256();
9875 let mut j = 0usize;
9876 while j + 32 <= n {
9877 let wv = _mm256_loadu_si256(w.as_ptr().add(j) as *const __m256i);
9878 let xv = _mm256_loadu_si256(xq.as_ptr().add(j) as *const __m256i);
9879 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(wv), _mm256_sign_epi8(xv, wv));
9880 acc = _mm256_add_epi32(acc, _mm256_madd_epi16(p16, ones));
9881 j += 32;
9882 }
9883 let hi128 = _mm256_extracti128_si256::<1>(acc);
9884 let s128 = _mm_add_epi32(_mm256_castsi256_si128(acc), hi128);
9885 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
9886 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
9887 let mut s = _mm_cvtsi128_si32(s32);
9888 while j < n {
9889 s += (w[j] as i8) as i32 * xq[j] as i32;
9890 j += 1;
9891 }
9892 s
9893 }
9894}
9895
9896#[cfg(target_arch = "aarch64")]
9900#[target_feature(enable = "neon,i8mm")]
9901unsafe fn dot_i8_smmla_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9902 unsafe {
9904 use core::arch::aarch64::*;
9905 use core::arch::asm;
9906 let n = w0.len();
9907 let w0p = w0.as_ptr() as *const i8;
9908 let w1p = w1.as_ptr() as *const i8;
9909 let mut acc01 = vdupq_n_s32(0);
9912 let mut acc23 = vdupq_n_s32(0);
9913 let mut i = 0usize;
9914 while i + 8 <= n {
9915 let wa = vcombine_s8(vld1_s8(w0p.add(i)), vld1_s8(w1p.add(i)));
9916 let xb01 = vcombine_s8(
9917 vld1_s8(xs[0].as_ptr().add(i)),
9918 vld1_s8(xs[1].as_ptr().add(i)),
9919 );
9920 let xb23 = vcombine_s8(
9921 vld1_s8(xs[2].as_ptr().add(i)),
9922 vld1_s8(xs[3].as_ptr().add(i)),
9923 );
9924 asm!(
9925 "smmla {a01:v}.4s, {w:v}.16b, {x01:v}.16b",
9926 "smmla {a23:v}.4s, {w:v}.16b, {x23:v}.16b",
9927 a01 = inout(vreg) acc01, a23 = inout(vreg) acc23,
9928 w = in(vreg) wa, x01 = in(vreg) xb01, x23 = in(vreg) xb23,
9929 options(pure, nomem, nostack),
9930 );
9931 i += 8;
9932 }
9933 let mut out = [[0i32; 4]; 2];
9934 let a01: [i32; 4] = core::mem::transmute(acc01);
9935 let a23: [i32; 4] = core::mem::transmute(acc23);
9936 out[0][0] = a01[0];
9937 out[0][1] = a01[1];
9938 out[1][0] = a01[2];
9939 out[1][1] = a01[3];
9940 out[0][2] = a23[0];
9941 out[0][3] = a23[1];
9942 out[1][2] = a23[2];
9943 out[1][3] = a23[3];
9944 if i < n {
9945 for (k, x) in xs.iter().enumerate() {
9946 for j in i..n {
9947 out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9948 out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
9949 }
9950 }
9951 }
9952 out
9953 }
9954}
9955
9956#[cfg(target_arch = "aarch64")]
9960#[target_feature(enable = "neon,dotprod")]
9961unsafe fn dot_i8_sdot_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
9962 unsafe {
9964 use core::arch::aarch64::*;
9965 use core::arch::asm;
9966 let n = w0.len();
9967 let w0p = w0.as_ptr() as *const i8;
9968 let w1p = w1.as_ptr() as *const i8;
9969 let mut acc = [[vdupq_n_s32(0); 4]; 2];
9970 let mut i = 0usize;
9971 while i + 16 <= n {
9972 let wv0 = vld1q_s8(w0p.add(i));
9973 let wv1 = vld1q_s8(w1p.add(i));
9974 for (k, x) in xs.iter().enumerate() {
9975 let xv = vld1q_s8(x.as_ptr().add(i));
9976 let (mut a0, mut a1) = (acc[0][k], acc[1][k]);
9977 asm!(
9978 "sdot {a0:v}.4s, {w0:v}.16b, {x:v}.16b",
9979 "sdot {a1:v}.4s, {w1:v}.16b, {x:v}.16b",
9980 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
9981 w0 = in(vreg) wv0, w1 = in(vreg) wv1, x = in(vreg) xv,
9982 options(pure, nomem, nostack),
9983 );
9984 acc[0][k] = a0;
9985 acc[1][k] = a1;
9986 }
9987 i += 16;
9988 }
9989 let mut out = [[0i32; 4]; 2];
9990 for r in 0..2 {
9991 for k in 0..4 {
9992 out[r][k] = vaddvq_s32(acc[r][k]);
9993 }
9994 }
9995 if i < n {
9996 for (k, x) in xs.iter().enumerate() {
9997 for j in i..n {
9998 out[0][k] += (w0[j] as i8) as i32 * x[j] as i32;
9999 out[1][k] += (w1[j] as i8) as i32 * x[j] as i32;
10000 }
10001 }
10002 }
10003 out
10004 }
10005}
10006
10007#[cfg(target_arch = "x86_64")]
10013#[target_feature(enable = "avx2")]
10014unsafe fn dot_i8_i8_avx2_2x4(w0: &[u8], w1: &[u8], xs: [&[i8]; 4]) -> [[i32; 4]; 2] {
10015 unsafe {
10017 use core::arch::x86_64::*;
10018 let n = w0.len();
10019 let ones = _mm256_set1_epi16(1);
10020 let mut acc = [[_mm256_setzero_si256(); 4]; 2];
10021 let mut j = 0usize;
10022 while j + 32 <= n {
10023 let wv0 = _mm256_loadu_si256(w0.as_ptr().add(j) as *const __m256i);
10024 let wv1 = _mm256_loadu_si256(w1.as_ptr().add(j) as *const __m256i);
10025 let aw0 = _mm256_abs_epi8(wv0);
10026 let aw1 = _mm256_abs_epi8(wv1);
10027 for (k, x) in xs.iter().enumerate() {
10028 let xv = _mm256_loadu_si256(x.as_ptr().add(j) as *const __m256i);
10029 let p0 = _mm256_maddubs_epi16(aw0, _mm256_sign_epi8(xv, wv0));
10030 acc[0][k] = _mm256_add_epi32(acc[0][k], _mm256_madd_epi16(p0, ones));
10031 let p1 = _mm256_maddubs_epi16(aw1, _mm256_sign_epi8(xv, wv1));
10032 acc[1][k] = _mm256_add_epi32(acc[1][k], _mm256_madd_epi16(p1, ones));
10033 }
10034 j += 32;
10035 }
10036 let mut out = [[0i32; 4]; 2];
10037 for r in 0..2 {
10038 for k in 0..4 {
10039 let a = acc[r][k];
10040 let hi128 = _mm256_extracti128_si256::<1>(a);
10041 let s128 = _mm_add_epi32(_mm256_castsi256_si128(a), hi128);
10042 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10043 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10044 out[r][k] = _mm_cvtsi128_si32(s32);
10045 }
10046 }
10047 if j < n {
10048 for (k, x) in xs.iter().enumerate() {
10049 for i in j..n {
10050 out[0][k] += (w0[i] as i8) as i32 * x[i] as i32;
10051 out[1][k] += (w1[i] as i8) as i32 * x[i] as i32;
10052 }
10053 }
10054 }
10055 out
10056 }
10057}
10058
10059#[cfg(target_arch = "x86_64")]
10064#[inline]
10065fn row_dot_avx2(row: &[u8], act: &SplitAct) -> f32 {
10066 let dot = if avx512vnni_enabled() && row.len() >= 64 {
10067 (unsafe { dot_u8p128_i8_vnni(row, &act.xq) }) - 128 * act.xsum
10068 } else {
10069 unsafe { dot_i8_i8_avx2(row, &act.xq) }
10070 };
10071 let mut acc = dot as f32 * act.sx;
10072 for &(j, xv) in &act.outliers {
10073 acc += (row[j] as i8) as f32 * xv;
10074 }
10075 acc
10076}
10077
10078#[cfg(target_arch = "x86_64")]
10082#[target_feature(enable = "avx2,avx512f,avx512bw,avx512vl,avx512vnni")]
10083unsafe fn dot_u8p128_i8_vnni(w: &[u8], xq: &[i8]) -> i32 {
10084 unsafe {
10086 use core::arch::x86_64::*;
10087 let n = w.len();
10088 let flip = _mm512_set1_epi8(-128); #[inline(always)]
10090 unsafe fn step(
10091 w: *const u8,
10092 x: *const i8,
10093 flip: core::arch::x86_64::__m512i,
10094 acc: core::arch::x86_64::__m512i,
10095 ) -> core::arch::x86_64::__m512i {
10096 unsafe {
10097 use core::arch::x86_64::*;
10098 let wv = _mm512_xor_si512(_mm512_loadu_si512(w as *const _), flip);
10099 _mm512_dpbusd_epi32(acc, wv, _mm512_loadu_si512(x as *const _))
10100 }
10101 }
10102 let (mut a0, mut a1, mut a2, mut a3) = (
10103 _mm512_setzero_si512(),
10104 _mm512_setzero_si512(),
10105 _mm512_setzero_si512(),
10106 _mm512_setzero_si512(),
10107 );
10108 let mut j = 0usize;
10109 while j + 256 <= n {
10110 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10111 a1 = step(w.as_ptr().add(j + 64), xq.as_ptr().add(j + 64), flip, a1);
10112 a2 = step(w.as_ptr().add(j + 128), xq.as_ptr().add(j + 128), flip, a2);
10113 a3 = step(w.as_ptr().add(j + 192), xq.as_ptr().add(j + 192), flip, a3);
10114 j += 256;
10115 }
10116 while j + 64 <= n {
10117 a0 = step(w.as_ptr().add(j), xq.as_ptr().add(j), flip, a0);
10118 j += 64;
10119 }
10120 let mut total = _mm512_reduce_add_epi32(_mm512_add_epi32(
10121 _mm512_add_epi32(a0, a1),
10122 _mm512_add_epi32(a2, a3),
10123 ));
10124 while j < n {
10126 total += ((w[j] ^ 0x80) as i32) * xq[j] as i32;
10127 j += 1;
10128 }
10129 total
10130 }
10131}
10132
10133#[cfg(target_arch = "x86_64")]
10139#[target_feature(enable = "avx2")]
10140unsafe fn dot_q4_row_avx2(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10141 unsafe {
10144 use core::arch::x86_64::*;
10145 let lomask = _mm_set1_epi8(0x0F);
10146 let eight = _mm256_set1_epi8(8);
10147 let ones = _mm256_set1_epi16(1);
10148 let mut acc = 0f32;
10149 for gi in 0..gpr {
10150 let g = g0 + gi;
10151 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10152 let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10153 let lo = _mm_and_si128(b, lomask);
10154 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10155 let w = _mm256_sub_epi8(
10156 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10157 eight,
10158 );
10159 let x = _mm256_loadu_si256(xq.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10160 let p16 = _mm256_maddubs_epi16(_mm256_abs_epi8(w), _mm256_sign_epi8(x, w));
10161 let d = _mm256_madd_epi16(p16, ones);
10162 let hi128 = _mm256_extracti128_si256::<1>(d);
10163 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10164 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10165 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10166 acc += _mm_cvtsi128_si32(s32) as f32 * s;
10167 }
10168 acc
10169 }
10170}
10171
10172#[cfg(target_arch = "x86_64")]
10175#[target_feature(enable = "avx2")]
10176unsafe fn dot_q4_row_avx2_2(
10177 packed: &[u8],
10178 scales: &[u8],
10179 g0: usize,
10180 gpr: usize,
10181 xq1: &[i8],
10182 xq2: &[i8],
10183) -> (f32, f32) {
10184 unsafe {
10186 use core::arch::x86_64::*;
10187 let lomask = _mm_set1_epi8(0x0F);
10188 let eight = _mm256_set1_epi8(8);
10189 let ones = _mm256_set1_epi16(1);
10190 let (mut acc1, mut acc2) = (0f32, 0f32);
10191 #[inline(always)]
10192 unsafe fn hsum(d: core::arch::x86_64::__m256i) -> i32 {
10193 unsafe {
10194 use core::arch::x86_64::*;
10195 let hi128 = _mm256_extracti128_si256::<1>(d);
10196 let s128 = _mm_add_epi32(_mm256_castsi256_si128(d), hi128);
10197 let s64 = _mm_add_epi32(s128, _mm_srli_si128::<8>(s128));
10198 let s32 = _mm_add_epi32(s64, _mm_srli_si128::<4>(s64));
10199 _mm_cvtsi128_si32(s32)
10200 }
10201 }
10202 for gi in 0..gpr {
10203 let g = g0 + gi;
10204 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10205 let b = _mm_loadu_si128(packed.as_ptr().add(g * 16) as *const __m128i);
10206 let lo = _mm_and_si128(b, lomask);
10207 let hi = _mm_and_si128(_mm_srli_epi16::<4>(b), lomask);
10208 let w = _mm256_sub_epi8(
10209 _mm256_set_m128i(_mm_unpackhi_epi8(lo, hi), _mm_unpacklo_epi8(lo, hi)),
10210 eight,
10211 );
10212 let aw = _mm256_abs_epi8(w);
10213 let x1 = _mm256_loadu_si256(xq1.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10214 let x2 = _mm256_loadu_si256(xq2.as_ptr().add(gi * GROUP_SIZE) as *const __m256i);
10215 let d1 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x1, w)), ones);
10216 let d2 = _mm256_madd_epi16(_mm256_maddubs_epi16(aw, _mm256_sign_epi8(x2, w)), ones);
10217 acc1 += hsum(d1) as f32 * s;
10218 acc2 += hsum(d2) as f32 * s;
10219 }
10220 (acc1, acc2)
10221 }
10222}
10223
10224#[cfg(target_arch = "x86_64")]
10226fn q8_range_avx2(
10227 q: &[u8],
10228 row_scale: &[f32],
10229 act: &SplitAct,
10230 cols: usize,
10231 out_addr: SendMut,
10232 start: usize,
10233 end: usize,
10234) {
10235 for o in start..end {
10236 let v = row_dot_avx2(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10237 unsafe { *out_addr.at(o) = v };
10239 }
10240}
10241
10242#[cfg(target_arch = "x86_64")]
10244#[allow(clippy::too_many_arguments)]
10245fn q8_range2_avx2(
10246 q: &[u8],
10247 row_scale: &[f32],
10248 a1: &SplitAct,
10249 a2: &SplitAct,
10250 cols: usize,
10251 p1: SendMut,
10252 p2: SendMut,
10253 start: usize,
10254 end: usize,
10255) {
10256 for o in start..end {
10257 let row = &q[o * cols..(o + 1) * cols];
10258 unsafe {
10260 *p1.at(o) = row_dot_avx2(row, a1) * row_scale[o];
10261 *p2.at(o) = row_dot_avx2(row, a2) * row_scale[o];
10262 }
10263 }
10264}
10265
10266#[cfg(target_arch = "aarch64")]
10277fn i8mm_enabled() -> bool {
10278 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
10279 *ON.get_or_init(|| {
10280 std::env::var("CMF_I8MM").map(|v| v == "1").unwrap_or(false)
10281 && std::arch::is_aarch64_feature_detected!("i8mm")
10282 })
10283}
10284
10285#[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
10289fn sdot_enabled() -> bool {
10290 if FLOAT_ACTIVATIONS.get() {
10291 return false;
10292 }
10293 use std::sync::OnceLock;
10294 static ON: OnceLock<bool> = OnceLock::new();
10295 *ON.get_or_init(|| {
10296 let want = std::env::var("CMF_SDOT").map(|v| v != "0").unwrap_or(true);
10297 if !want {
10298 return false;
10299 }
10300
10301 #[cfg(target_arch = "aarch64")]
10302 {
10303 if std::arch::is_aarch64_feature_detected!("dotprod") {
10304 return true;
10305 }
10306 #[cfg(target_os = "android")]
10307 {
10308 if let Ok(cpuinfo) = std::fs::read_to_string("/proc/cpuinfo") {
10309 if cpuinfo.lines().any(|l| {
10310 (l.starts_with("Features") || l.starts_with("features"))
10311 && l.contains("asimddp")
10312 }) {
10313 return true;
10314 }
10315 }
10316 }
10317 false
10318 }
10319 #[cfg(not(target_arch = "aarch64"))]
10320 {
10321 false
10322 }
10323 })
10324}
10325
10326struct SplitAct {
10331 xq: Vec<i8>,
10332 sx: f32,
10333 outliers: Vec<(usize, f32)>,
10334 #[cfg_attr(not(target_arch = "x86_64"), allow(dead_code))]
10337 xsum: i32,
10338}
10339
10340thread_local! {
10341 static XQ_FREE: std::cell::RefCell<Vec<Vec<i8>>> =
10344 const { std::cell::RefCell::new(Vec::new()) };
10345}
10346
10347impl Drop for SplitAct {
10348 fn drop(&mut self) {
10349 let buf = std::mem::take(&mut self.xq);
10350 if buf.capacity() > 0 {
10351 XQ_FREE.with(|f| {
10352 let mut f = f.borrow_mut();
10353 if f.len() < 16 {
10354 f.push(buf);
10355 }
10356 });
10357 }
10358 }
10359}
10360
10361thread_local! {
10362 static KROW: UnsafeCell<Vec<f32>> = const { UnsafeCell::new(Vec::new()) };
10369 static KROWS: UnsafeCell<[Vec<f32>; 2]> =
10374 const { UnsafeCell::new([Vec::new(), Vec::new()]) };
10375}
10376
10377#[inline]
10380fn with_krow<R>(n: usize, f: impl FnOnce(&mut [f32]) -> R) -> R {
10381 KROW.with(|s| {
10382 let b = unsafe { &mut *s.get() };
10385 if b.len() < n {
10386 b.resize(n, 0.0);
10387 }
10388 f(&mut b[..n])
10389 })
10390}
10391
10392#[inline(always)]
10403fn q8_round(t: f32) -> i8 {
10404 let t = t.clamp(-127.0, 127.0);
10405 let i = t as i32;
10406 let f = t - i as f32;
10407 let r = if f >= 0.5 {
10408 i + 1
10409 } else if f <= -0.5 {
10410 i - 1
10411 } else {
10412 i
10413 };
10414 r as i8
10415}
10416
10417#[inline]
10418fn with_krows<R>(n: usize, f: impl FnOnce(&mut [f32], &mut [f32]) -> R) -> R {
10419 KROWS.with(|s| {
10420 let b = unsafe { &mut *s.get() };
10423 for row in b.iter_mut() {
10424 if row.len() < n {
10425 row.resize(n, 0.0);
10426 }
10427 }
10428 let (a, b) = b.split_at_mut(1);
10429 f(&mut a[0][..n], &mut b[0][..n])
10430 })
10431}
10432
10433#[inline]
10434fn silu_mul_limited(mut gate: f32, mut up: f32, limit: f32) -> f32 {
10435 if limit > 0.0 {
10436 up = up.clamp(-limit, limit);
10437 gate = gate.min(limit);
10438 }
10439 gate / (1.0 + (-gate).exp()) * up
10440}
10441
10442fn split_act(x: &[f32]) -> SplitAct {
10443 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10444 let n = x.len();
10445 let rms = (x.iter().map(|&v| (v * v) as f64).sum::<f64>() / n.max(1) as f64).sqrt() as f32;
10446 let thr = 8.0 * rms;
10447 let mut outliers: Vec<(usize, f32)> = Vec::new();
10451 let mut amax = 0f32;
10452 for (j, &v) in x.iter().enumerate() {
10453 let a = v.abs();
10454 if a > thr {
10455 outliers.push((j, v));
10456 } else if a > amax {
10457 amax = a;
10458 }
10459 }
10460 let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10461 let inv = 1.0 / sx;
10462 let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10463 xq.clear();
10464 xq.reserve(n);
10465 if outliers.is_empty() {
10466 xq.extend(
10467 x.iter()
10468 .map(|&v| q8_round(v * inv)),
10469 );
10470 } else {
10471 xq.extend(x.iter().map(|&v| {
10473 if v.abs() > thr {
10474 0
10475 } else {
10476 q8_round(v * inv)
10477 }
10478 }));
10479 }
10480 let xsum = xq.iter().map(|&v| v as i32).sum();
10481 SplitAct {
10482 xq,
10483 sx,
10484 outliers,
10485 xsum,
10486 }
10487}
10488
10489fn split_act_q8_2f(x: &[f32], col: &[f32]) -> SplitAct {
10490 let _prof = crate::cpuprof::time(crate::cpuprof::Slot::SplitAct);
10491 let n = x.len();
10492 let rms = (x
10493 .iter()
10494 .zip(col)
10495 .map(|(&a, &c)| {
10496 let v = a * c;
10497 (v * v) as f64
10498 })
10499 .sum::<f64>()
10500 / n.max(1) as f64)
10501 .sqrt() as f32;
10502 let thr = 8.0 * rms;
10503
10504 let mut outliers = Vec::new();
10505 let mut amax = 0f32;
10506 for (j, (&a, &c)) in x.iter().zip(col).enumerate() {
10507 let v = a * c;
10508 let s = v.abs();
10509 if s > thr {
10510 outliers.push((j, v));
10511 } else if s > amax {
10512 amax = s;
10513 }
10514 }
10515
10516 let sx = if amax > 0.0 { amax / 127.0 } else { 1.0 };
10517 let inv = 1.0 / sx;
10518 let mut xq = XQ_FREE.with(|f| f.borrow_mut().pop()).unwrap_or_default();
10519 xq.clear();
10520 xq.reserve(n);
10521 if outliers.is_empty() {
10522 xq.extend(
10523 x.iter()
10524 .zip(col)
10525 .map(|(&a, &c)| q8_round((a * c) * inv)),
10526 );
10527 } else {
10528 xq.extend(x.iter().zip(col).map(|(&a, &c)| {
10529 let v = a * c;
10530 if v.abs() > thr {
10531 0
10532 } else {
10533 q8_round(v * inv)
10534 }
10535 }));
10536 }
10537 let xsum = xq.iter().map(|&v| v as i32).sum();
10538 SplitAct {
10539 xq,
10540 sx,
10541 outliers,
10542 xsum,
10543 }
10544}
10545
10546#[cfg(target_arch = "aarch64")]
10549#[target_feature(enable = "neon,dotprod")]
10550unsafe fn dot_i8_sdot(w: &[u8], xq: &[i8]) -> i32 {
10551 unsafe {
10553 use core::arch::aarch64::*;
10554 use core::arch::asm;
10555 let wp = w.as_ptr() as *const i8;
10556 let n = w.len();
10557 let (mut a0, mut a1, mut a2, mut a3) = (
10558 vdupq_n_s32(0),
10559 vdupq_n_s32(0),
10560 vdupq_n_s32(0),
10561 vdupq_n_s32(0),
10562 );
10563 let mut i = 0;
10564 while i + 64 <= n {
10565 let (w0, x0) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10566 let (w1, x1) = (vld1q_s8(wp.add(i + 16)), vld1q_s8(xq.as_ptr().add(i + 16)));
10567 let (w2, x2) = (vld1q_s8(wp.add(i + 32)), vld1q_s8(xq.as_ptr().add(i + 32)));
10568 let (w3, x3) = (vld1q_s8(wp.add(i + 48)), vld1q_s8(xq.as_ptr().add(i + 48)));
10569 asm!(
10570 "sdot {a0:v}.4s, {w0:v}.16b, {x0:v}.16b",
10571 "sdot {a1:v}.4s, {w1:v}.16b, {x1:v}.16b",
10572 "sdot {a2:v}.4s, {w2:v}.16b, {x2:v}.16b",
10573 "sdot {a3:v}.4s, {w3:v}.16b, {x3:v}.16b",
10574 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10575 w0 = in(vreg) w0, x0 = in(vreg) x0, w1 = in(vreg) w1, x1 = in(vreg) x1,
10576 w2 = in(vreg) w2, x2 = in(vreg) x2, w3 = in(vreg) w3, x3 = in(vreg) x3,
10577 options(pure, nomem, nostack),
10578 );
10579 i += 64;
10580 }
10581 while i + 16 <= n {
10582 let (wv, xv) = (vld1q_s8(wp.add(i)), vld1q_s8(xq.as_ptr().add(i)));
10583 asm!("sdot {a:v}.4s, {w:v}.16b, {x:v}.16b",
10584 a = inout(vreg) a0, w = in(vreg) wv, x = in(vreg) xv, options(pure, nomem, nostack));
10585 i += 16;
10586 }
10587 let mut s = vaddvq_s32(vaddq_s32(vaddq_s32(a0, a1), vaddq_s32(a2, a3)));
10588 while i < n {
10589 s += (*wp.add(i)) as i32 * xq[i] as i32;
10590 i += 1;
10591 }
10592 s
10593 }
10594}
10595
10596#[cfg(target_arch = "aarch64")]
10600#[target_feature(enable = "neon,dotprod")]
10601unsafe fn dot_i8_sdot_4rows(w0: &[u8], w1: &[u8], w2: &[u8], w3: &[u8], xq: &[i8]) -> [i32; 4] {
10602 unsafe {
10604 use core::arch::aarch64::*;
10605 use core::arch::asm;
10606 let n = xq.len();
10607 let px = xq.as_ptr();
10608 let (p0, p1, p2, p3) = (
10609 w0.as_ptr() as *const i8,
10610 w1.as_ptr() as *const i8,
10611 w2.as_ptr() as *const i8,
10612 w3.as_ptr() as *const i8,
10613 );
10614 let (mut a0, mut a1, mut a2, mut a3) = (
10615 vdupq_n_s32(0),
10616 vdupq_n_s32(0),
10617 vdupq_n_s32(0),
10618 vdupq_n_s32(0),
10619 );
10620 let mut i = 0;
10621 while i + 16 <= n {
10622 let x = vld1q_s8(px.add(i));
10623 let v0 = vld1q_s8(p0.add(i));
10624 let v1 = vld1q_s8(p1.add(i));
10625 let v2 = vld1q_s8(p2.add(i));
10626 let v3 = vld1q_s8(p3.add(i));
10627 asm!(
10628 "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10629 "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10630 "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10631 "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10632 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10633 v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10634 options(pure, nomem, nostack),
10635 );
10636 i += 16;
10637 }
10638 let mut r = [
10639 vaddvq_s32(a0),
10640 vaddvq_s32(a1),
10641 vaddvq_s32(a2),
10642 vaddvq_s32(a3),
10643 ];
10644 while i < n {
10645 let xi = *px.add(i) as i32;
10646 r[0] += (*p0.add(i)) as i32 * xi;
10647 r[1] += (*p1.add(i)) as i32 * xi;
10648 r[2] += (*p2.add(i)) as i32 * xi;
10649 r[3] += (*p3.add(i)) as i32 * xi;
10650 i += 1;
10651 }
10652 r
10653 }
10654}
10655
10656#[cfg(target_arch = "aarch64")]
10663#[target_feature(enable = "neon,dotprod")]
10664unsafe fn dot_i8_sdot_4rows_il(g: &[u8], xq: &[i8]) -> [i32; 4] {
10665 unsafe {
10668 use core::arch::aarch64::*;
10669 use core::arch::asm;
10670 let n = xq.len();
10671 let px = xq.as_ptr();
10672 let pg = g.as_ptr() as *const i8;
10673 let (mut a0, mut a1, mut a2, mut a3) = (
10674 vdupq_n_s32(0),
10675 vdupq_n_s32(0),
10676 vdupq_n_s32(0),
10677 vdupq_n_s32(0),
10678 );
10679 let mut i = 0;
10680 while i + 16 <= n {
10681 let x = vld1q_s8(px.add(i));
10682 let base = pg.add(4 * i);
10683 let v0 = vld1q_s8(base);
10684 let v1 = vld1q_s8(base.add(16));
10685 let v2 = vld1q_s8(base.add(32));
10686 let v3 = vld1q_s8(base.add(48));
10687 asm!(
10688 "sdot {a0:v}.4s, {v0:v}.16b, {x:v}.16b",
10689 "sdot {a1:v}.4s, {v1:v}.16b, {x:v}.16b",
10690 "sdot {a2:v}.4s, {v2:v}.16b, {x:v}.16b",
10691 "sdot {a3:v}.4s, {v3:v}.16b, {x:v}.16b",
10692 a0 = inout(vreg) a0, a1 = inout(vreg) a1, a2 = inout(vreg) a2, a3 = inout(vreg) a3,
10693 v0 = in(vreg) v0, v1 = in(vreg) v1, v2 = in(vreg) v2, v3 = in(vreg) v3, x = in(vreg) x,
10694 options(pure, nomem, nostack),
10695 );
10696 i += 16;
10697 }
10698 [
10699 vaddvq_s32(a0),
10700 vaddvq_s32(a1),
10701 vaddvq_s32(a2),
10702 vaddvq_s32(a3),
10703 ]
10704 }
10705}
10706
10707#[cfg(target_arch = "aarch64")]
10713fn q8_range_sdot(
10714 q: &[u8],
10715 rep: &[u8],
10716 row_scale: &[f32],
10717 act: &SplitAct,
10718 cols: usize,
10719 out_addr: SendMut,
10720 start: usize,
10721 end: usize,
10722) {
10723 let mut o = start;
10724 if !rep.is_empty() {
10727 while o < end && o % 4 != 0 {
10728 let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10729 unsafe { *out_addr.at(o) = v };
10730 o += 1;
10731 }
10732 }
10733 while o + 4 <= end {
10734 let r = if rep.is_empty() {
10735 unsafe {
10736 dot_i8_sdot_4rows(
10737 &q[o * cols..(o + 1) * cols],
10738 &q[(o + 1) * cols..(o + 2) * cols],
10739 &q[(o + 2) * cols..(o + 3) * cols],
10740 &q[(o + 3) * cols..(o + 4) * cols],
10741 &act.xq,
10742 )
10743 }
10744 } else {
10745 unsafe { dot_i8_sdot_4rows_il(&rep[o * cols..(o + 4) * cols], &act.xq) }
10746 };
10747 for k in 0..4 {
10748 let mut acc = r[k] as f32 * act.sx;
10749 for &(j, xv) in &act.outliers {
10750 acc += (q[(o + k) * cols + j] as i8) as f32 * xv;
10751 }
10752 unsafe { *out_addr.at(o + k) = acc * row_scale[o + k] };
10754 }
10755 o += 4;
10756 }
10757 while o < end {
10758 let v = row_dot_sdot(&q[o * cols..(o + 1) * cols], act) * row_scale[o];
10759 unsafe { *out_addr.at(o) = v };
10760 o += 1;
10761 }
10762}
10763
10764#[cfg(target_arch = "aarch64")]
10767#[allow(clippy::too_many_arguments)]
10768fn q8_range2_sdot(
10769 q: &[u8],
10770 row_scale: &[f32],
10771 a1: &SplitAct,
10772 a2: &SplitAct,
10773 cols: usize,
10774 p1: SendMut,
10775 p2: SendMut,
10776 start: usize,
10777 end: usize,
10778) {
10779 for o in start..end {
10780 let row = &q[o * cols..(o + 1) * cols];
10781 unsafe {
10783 *p1.at(o) = row_dot_sdot(row, a1) * row_scale[o];
10784 *p2.at(o) = row_dot_sdot(row, a2) * row_scale[o];
10785 }
10786 }
10787}
10788
10789#[allow(clippy::too_many_arguments)]
10791fn q8_range2_f32(
10792 q: &[u8],
10793 row_scale: &[f32],
10794 x1: &[f32],
10795 x2: &[f32],
10796 cols: usize,
10797 p1: SendMut,
10798 p2: SendMut,
10799 start: usize,
10800 end: usize,
10801) {
10802 for o in start..end {
10803 let row = &q[o * cols..(o + 1) * cols];
10804 unsafe {
10806 *p1.at(o) = dot_i8_f32(row, x1) * row_scale[o];
10807 *p2.at(o) = dot_i8_f32(row, x2) * row_scale[o];
10808 }
10809 }
10810}
10811
10812fn q8_range_f32(
10814 q: &[u8],
10815 row_scale: &[f32],
10816 xs: &[f32],
10817 cols: usize,
10818 out_addr: SendMut,
10819 start: usize,
10820 end: usize,
10821) {
10822 for o in start..end {
10823 let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
10824 unsafe { *out_addr.at(o) = v };
10826 }
10827}
10828
10829#[inline]
10833fn q8_row_dot(row: &[u8], act: &SplitAct) -> f32 {
10834 #[cfg(target_arch = "aarch64")]
10835 return row_dot_sdot(row, act);
10836 #[cfg(target_arch = "x86_64")]
10837 return row_dot_avx2(row, act);
10838 #[allow(unreachable_code)]
10839 q8_row_dot_scalar(row, act)
10840}
10841
10842#[allow(dead_code)]
10843fn q8_row_dot_scalar(row: &[u8], act: &SplitAct) -> f32 {
10844 let mut acc = 0i32;
10845 for (k, &b) in row.iter().enumerate() {
10846 acc += (b as i8) as i32 * act.xq[k] as i32;
10847 }
10848 let mut acc = acc as f32 * act.sx;
10849 for &(j, xv) in &act.outliers {
10850 acc += (row[j] as i8) as f32 * xv;
10851 }
10852 acc
10853}
10854
10855#[cfg(target_arch = "aarch64")]
10858#[inline]
10859fn row_dot_sdot(row: &[u8], act: &SplitAct) -> f32 {
10860 let mut acc = unsafe { dot_i8_sdot(row, &act.xq) } as f32 * act.sx;
10861 for &(j, xv) in &act.outliers {
10862 acc += (row[j] as i8) as f32 * xv;
10863 }
10864 acc
10865}
10866
10867#[cfg(target_arch = "aarch64")]
10875#[target_feature(enable = "neon,dotprod")]
10876unsafe fn dot_q4_row_sdot(packed: &[u8], scales: &[u8], g0: usize, gpr: usize, xq: &[i8]) -> f32 {
10877 unsafe {
10880 use core::arch::aarch64::*;
10881 use core::arch::asm;
10882 let lomask = vdupq_n_u8(0x0F);
10883 let eight = vdupq_n_s8(8);
10884 let mut acc = 0f32;
10885 for gi in 0..gpr {
10886 let g = g0 + gi;
10887 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10888 let b = vld1q_u8(packed.as_ptr().add(g * 16));
10889 let lo = vandq_u8(b, lomask);
10890 let hi = vshrq_n_u8::<4>(b);
10891 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10892 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10893 let x0 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE));
10894 let x1 = vld1q_s8(xq.as_ptr().add(gi * GROUP_SIZE + 16));
10895 let (mut a0, mut a1) = (vdupq_n_s32(0), vdupq_n_s32(0));
10896 asm!(
10897 "sdot {a0:v}.4s, {e0:v}.16b, {x0:v}.16b",
10898 "sdot {a1:v}.4s, {e1:v}.16b, {x1:v}.16b",
10899 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10900 e0 = in(vreg) e0, x0 = in(vreg) x0, e1 = in(vreg) e1, x1 = in(vreg) x1,
10901 options(pure, nomem, nostack),
10902 );
10903 acc += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10904 }
10905 acc
10906 }
10907}
10908
10909#[cfg(target_arch = "aarch64")]
10914#[target_feature(enable = "neon,dotprod")]
10915unsafe fn dot_q4_row_sdot2(
10916 packed: &[u8],
10917 scales: &[u8],
10918 g0: usize,
10919 gpr: usize,
10920 xq1: &[i8],
10921 xq2: &[i8],
10922) -> (f32, f32) {
10923 unsafe {
10926 use core::arch::aarch64::*;
10927 use core::arch::asm;
10928 let lomask = vdupq_n_u8(0x0F);
10929 let eight = vdupq_n_s8(8);
10930 let (mut acc1, mut acc2) = (0f32, 0f32);
10931 for gi in 0..gpr {
10932 let g = g0 + gi;
10933 let s = f16_to_f32(u16::from_le_bytes([scales[g * 2], scales[g * 2 + 1]]));
10934 let b = vld1q_u8(packed.as_ptr().add(g * 16));
10935 let lo = vandq_u8(b, lomask);
10936 let hi = vshrq_n_u8::<4>(b);
10937 let e0 = vsubq_s8(vreinterpretq_s8_u8(vzip1q_u8(lo, hi)), eight);
10938 let e1 = vsubq_s8(vreinterpretq_s8_u8(vzip2q_u8(lo, hi)), eight);
10939 let x10 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE));
10940 let x11 = vld1q_s8(xq1.as_ptr().add(gi * GROUP_SIZE + 16));
10941 let x20 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE));
10942 let x21 = vld1q_s8(xq2.as_ptr().add(gi * GROUP_SIZE + 16));
10943 let (mut a0, mut a1, mut b0, mut b1) = (
10944 vdupq_n_s32(0),
10945 vdupq_n_s32(0),
10946 vdupq_n_s32(0),
10947 vdupq_n_s32(0),
10948 );
10949 asm!(
10950 "sdot {a0:v}.4s, {e0:v}.16b, {x10:v}.16b",
10951 "sdot {a1:v}.4s, {e1:v}.16b, {x11:v}.16b",
10952 "sdot {b0:v}.4s, {e0:v}.16b, {x20:v}.16b",
10953 "sdot {b1:v}.4s, {e1:v}.16b, {x21:v}.16b",
10954 a0 = inout(vreg) a0, a1 = inout(vreg) a1,
10955 b0 = inout(vreg) b0, b1 = inout(vreg) b1,
10956 e0 = in(vreg) e0, e1 = in(vreg) e1,
10957 x10 = in(vreg) x10, x11 = in(vreg) x11,
10958 x20 = in(vreg) x20, x21 = in(vreg) x21,
10959 options(pure, nomem, nostack),
10960 );
10961 acc1 += vaddvq_s32(vaddq_s32(a0, a1)) as f32 * s;
10962 acc2 += vaddvq_s32(vaddq_s32(b0, b1)) as f32 * s;
10963 }
10964 (acc1, acc2)
10965 }
10966}
10967
10968#[inline]
10973pub(crate) fn axpy_i8_f32(acc: &mut [f32], row: &[i8], w: f32) {
10974 #[cfg(target_arch = "aarch64")]
10975 unsafe {
10976 return axpy_i8_f32_neon(acc, row, w);
10977 }
10978 #[cfg(target_arch = "x86_64")]
10979 if avx2_enabled() {
10980 return unsafe { axpy_i8_f32_avx2(acc, row, w) };
10981 }
10982 #[allow(unreachable_code)]
10983 {
10984 for (a, &b) in acc.iter_mut().zip(row) {
10985 *a += w * b as f32;
10986 }
10987 }
10988}
10989
10990#[cfg(target_arch = "x86_64")]
10992#[target_feature(enable = "avx2,fma")]
10993unsafe fn axpy_i8_f32_avx2(acc: &mut [f32], row: &[i8], w: f32) {
10994 unsafe {
10996 use core::arch::x86_64::*;
10997 let n = acc.len().min(row.len());
10998 let ap = acc.as_mut_ptr();
10999 let rp = row.as_ptr();
11000 let wv = _mm256_set1_ps(w);
11001 let mut j = 0usize;
11002 while j + 16 <= n {
11003 let rb = _mm_loadu_si128(rp.add(j) as *const __m128i);
11004 let lo = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(rb));
11005 let hi = _mm256_cvtepi32_ps(_mm256_cvtepi8_epi32(_mm_srli_si128::<8>(rb)));
11006 let v0 = _mm256_fmadd_ps(wv, lo, _mm256_loadu_ps(ap.add(j)));
11007 let v1 = _mm256_fmadd_ps(wv, hi, _mm256_loadu_ps(ap.add(j + 8)));
11008 _mm256_storeu_ps(ap.add(j), v0);
11009 _mm256_storeu_ps(ap.add(j + 8), v1);
11010 j += 16;
11011 }
11012 while j < n {
11013 *ap.add(j) += w * (*rp.add(j)) as f32;
11014 j += 1;
11015 }
11016 }
11017}
11018
11019#[cfg(target_arch = "aarch64")]
11020#[target_feature(enable = "neon")]
11021unsafe fn axpy_i8_f32_neon(acc: &mut [f32], row: &[i8], w: f32) {
11022 unsafe {
11024 use core::arch::aarch64::*;
11025 let n = acc.len().min(row.len());
11026 let ap = acc.as_mut_ptr();
11027 let rp = row.as_ptr();
11028 let wv = vdupq_n_f32(w);
11029 let mut j = 0usize;
11030 while j + 16 <= n {
11031 let rb = vld1q_s8(rp.add(j));
11032 let lo = vmovl_s8(vget_low_s8(rb));
11033 let hi = vmovl_s8(vget_high_s8(rb));
11034 for (off, half) in [(0, lo), (8, hi)] {
11035 let f0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(half)));
11036 let f1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(half)));
11037 let o = j + off;
11038 vst1q_f32(ap.add(o), vfmaq_f32(vld1q_f32(ap.add(o)), wv, f0));
11039 vst1q_f32(ap.add(o + 4), vfmaq_f32(vld1q_f32(ap.add(o + 4)), wv, f1));
11040 }
11041 j += 16;
11042 }
11043 while j < n {
11044 *ap.add(j) += w * (*rp.add(j)) as f32;
11045 j += 1;
11046 }
11047 }
11048}
11049
11050#[inline]
11053pub(crate) fn dot_i8_f32(w: &[u8], x: &[f32]) -> f32 {
11054 #[cfg(target_arch = "aarch64")]
11055 unsafe {
11056 return dot_i8_f32_neon(w, x);
11057 }
11058 #[cfg(target_arch = "x86_64")]
11059 if avx2_enabled() {
11060 return unsafe { dot_i8_f32_avx2(w, x) };
11061 }
11062 #[allow(unreachable_code)]
11063 {
11064 let mut sum = 0.0f32;
11065 for (j, &b) in w.iter().enumerate() {
11066 sum += (b as i8) as f32 * x[j];
11067 }
11068 sum
11069 }
11070}
11071
11072#[inline]
11076fn dot_i8_col_f32(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11077 #[cfg(target_arch = "aarch64")]
11078 unsafe {
11079 return dot_i8_col_f32_neon(w, x, col);
11080 }
11081 #[allow(unreachable_code)]
11082 {
11083 let mut sum = 0.0f32;
11084 for (j, &b) in w.iter().enumerate() {
11085 sum += (b as i8) as f32 * x[j] * col[j];
11086 }
11087 sum
11088 }
11089}
11090
11091#[cfg(target_arch = "aarch64")]
11092#[target_feature(enable = "neon")]
11093unsafe fn dot_i8_col_f32_neon(w: &[u8], x: &[f32], col: &[f32]) -> f32 {
11094 unsafe {
11096 use core::arch::aarch64::*;
11097 let n = x.len();
11098 let wp = w.as_ptr() as *const i8;
11099 let xp = x.as_ptr();
11100 let cp = col.as_ptr();
11101 let (mut a0, mut a1, mut a2, mut a3) = (
11102 vdupq_n_f32(0.0),
11103 vdupq_n_f32(0.0),
11104 vdupq_n_f32(0.0),
11105 vdupq_n_f32(0.0),
11106 );
11107 let mut j = 0usize;
11108 while j + 16 <= n {
11109 let wb = vld1q_s8(wp.add(j));
11110 let lo = vmovl_s8(vget_low_s8(wb));
11111 let hi = vmovl_s8(vget_high_s8(wb));
11112 let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11113 let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11114 let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11115 let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11116 a0 = vfmaq_f32(
11117 a0,
11118 w0,
11119 vmulq_f32(vld1q_f32(xp.add(j)), vld1q_f32(cp.add(j))),
11120 );
11121 a1 = vfmaq_f32(
11122 a1,
11123 w1,
11124 vmulq_f32(vld1q_f32(xp.add(j + 4)), vld1q_f32(cp.add(j + 4))),
11125 );
11126 a2 = vfmaq_f32(
11127 a2,
11128 w2,
11129 vmulq_f32(vld1q_f32(xp.add(j + 8)), vld1q_f32(cp.add(j + 8))),
11130 );
11131 a3 = vfmaq_f32(
11132 a3,
11133 w3,
11134 vmulq_f32(vld1q_f32(xp.add(j + 12)), vld1q_f32(cp.add(j + 12))),
11135 );
11136 j += 16;
11137 }
11138 let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11139 while j < n {
11140 sum += (*wp.add(j)) as f32 * *xp.add(j) * *cp.add(j);
11141 j += 1;
11142 }
11143 sum
11144 }
11145}
11146
11147#[cfg(target_arch = "aarch64")]
11148#[target_feature(enable = "neon")]
11149unsafe fn dot_i8_f32_neon(w: &[u8], x: &[f32]) -> f32 {
11150 unsafe {
11152 use core::arch::aarch64::*;
11153 let n = x.len();
11154 let wp = w.as_ptr() as *const i8;
11155 let xp = x.as_ptr();
11156 let (mut a0, mut a1, mut a2, mut a3) = (
11157 vdupq_n_f32(0.0),
11158 vdupq_n_f32(0.0),
11159 vdupq_n_f32(0.0),
11160 vdupq_n_f32(0.0),
11161 );
11162 let mut j = 0usize;
11163 while j + 16 <= n {
11164 let wb = vld1q_s8(wp.add(j));
11165 let lo = vmovl_s8(vget_low_s8(wb));
11166 let hi = vmovl_s8(vget_high_s8(wb));
11167 let w0 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(lo)));
11168 let w1 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(lo)));
11169 let w2 = vcvtq_f32_s32(vmovl_s16(vget_low_s16(hi)));
11170 let w3 = vcvtq_f32_s32(vmovl_s16(vget_high_s16(hi)));
11171 a0 = vfmaq_f32(a0, w0, vld1q_f32(xp.add(j)));
11172 a1 = vfmaq_f32(a1, w1, vld1q_f32(xp.add(j + 4)));
11173 a2 = vfmaq_f32(a2, w2, vld1q_f32(xp.add(j + 8)));
11174 a3 = vfmaq_f32(a3, w3, vld1q_f32(xp.add(j + 12)));
11175 j += 16;
11176 }
11177 let mut sum = vaddvq_f32(vaddq_f32(vaddq_f32(a0, a1), vaddq_f32(a2, a3)));
11178 while j < n {
11179 sum += (*wp.add(j)) as f32 * *xp.add(j);
11180 j += 1;
11181 }
11182 sum
11183 }
11184}
11185
11186#[allow(clippy::too_many_arguments)]
11187fn qmatvec(
11188 q: &[u8],
11189 rep: &[u8],
11190 row_scale: &[f32],
11191 x: &[f32],
11192 col_field: &[f32],
11193 dtype: TensorDtype,
11194 rows: usize,
11195 cols: usize,
11196 out: &mut [f32],
11197 pool: Option<&Pool>,
11198) {
11199 debug_assert_eq!(out.len(), rows);
11200 #[cfg(not(target_arch = "aarch64"))]
11201 let _ = rep;
11202
11203 #[cfg(target_arch = "aarch64")]
11204 if sdot_enabled() {
11205 let act = if dtype == TensorDtype::Q8_2f {
11206 split_act_q8_2f(x, col_field)
11207 } else {
11208 split_act(x)
11209 };
11210 let out_addr = SendMut(out.as_mut_ptr());
11211 let run_range = |start: usize, end: usize| {
11212 q8_range_sdot(q, rep, row_scale, &act, cols, out_addr, start, end)
11213 };
11214 match pool {
11215 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11216 _ => run_range(0, rows),
11217 }
11218 return;
11219 }
11220 #[cfg(target_arch = "x86_64")]
11223 if avx2_a8w8_enabled() {
11224 let act = if dtype == TensorDtype::Q8_2f {
11225 split_act_q8_2f(x, col_field)
11226 } else {
11227 split_act(x)
11228 };
11229 let out_addr = SendMut(out.as_mut_ptr());
11230 let run_range = |start: usize, end: usize| {
11231 q8_range_avx2(q, row_scale, &act, cols, out_addr, start, end)
11232 };
11233 match pool {
11234 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11235 _ => run_range(0, rows),
11236 }
11237 return;
11238 }
11239
11240 prescale_with(x, col_field, dtype, 1, |xs| {
11241 let out_addr = SendMut(out.as_mut_ptr());
11242 let run_range = move |start: usize, end: usize| {
11243 for o in start..end {
11244 let v = dot_i8_f32(&q[o * cols..(o + 1) * cols], xs) * row_scale[o];
11245 unsafe { *out_addr.at(o) = v };
11247 }
11248 };
11249 match pool {
11250 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11251 _ => run_range(0, rows),
11252 }
11253 });
11254}
11255
11256#[allow(clippy::too_many_arguments)]
11257fn qmatvec2(
11258 q: &[u8],
11259 row_scale: &[f32],
11260 x1: &[f32],
11261 x2: &[f32],
11262 col_field: &[f32],
11263 dtype: TensorDtype,
11264 rows: usize,
11265 cols: usize,
11266 o1: &mut [f32],
11267 o2: &mut [f32],
11268 pool: Option<&Pool>,
11269) {
11270 #[cfg(target_arch = "aarch64")]
11271 if sdot_enabled() {
11272 let a1s = if dtype == TensorDtype::Q8_2f {
11273 split_act_q8_2f(x1, col_field)
11274 } else {
11275 split_act(x1)
11276 };
11277 let a2s = if dtype == TensorDtype::Q8_2f {
11278 split_act_q8_2f(x2, col_field)
11279 } else {
11280 split_act(x2)
11281 };
11282 let p1 = SendMut(o1.as_mut_ptr());
11283 let p2 = SendMut(o2.as_mut_ptr());
11284 let run_range = |start: usize, end: usize| {
11285 q8_range2_sdot(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11286 };
11287 match pool {
11288 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11289 _ => run_range(0, rows),
11290 }
11291 return;
11292 }
11293 #[cfg(target_arch = "x86_64")]
11294 if avx2_a8w8_enabled() {
11295 let a1s = if dtype == TensorDtype::Q8_2f {
11296 split_act_q8_2f(x1, col_field)
11297 } else {
11298 split_act(x1)
11299 };
11300 let a2s = if dtype == TensorDtype::Q8_2f {
11301 split_act_q8_2f(x2, col_field)
11302 } else {
11303 split_act(x2)
11304 };
11305 let p1 = SendMut(o1.as_mut_ptr());
11306 let p2 = SendMut(o2.as_mut_ptr());
11307 let run_range = |start: usize, end: usize| {
11308 q8_range2_avx2(q, row_scale, &a1s, &a2s, cols, p1, p2, start, end)
11309 };
11310 match pool {
11311 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11312 _ => run_range(0, rows),
11313 }
11314 return;
11315 }
11316
11317 prescale_with(x1, col_field, dtype, 1, |x1s| {
11318 prescale_with(x2, col_field, dtype, 2, |x2s| {
11319 let p1 = SendMut(o1.as_mut_ptr());
11320 let p2 = SendMut(o2.as_mut_ptr());
11321 let run_range = move |start: usize, end: usize| {
11322 for o in start..end {
11323 let row = &q[o * cols..(o + 1) * cols];
11324 let s1 = dot_i8_f32(row, x1s) * row_scale[o];
11325 let s2 = dot_i8_f32(row, x2s) * row_scale[o];
11326 unsafe {
11328 *p1.at(o) = s1;
11329 *p2.at(o) = s2;
11330 }
11331 }
11332 };
11333 match pool {
11334 Some(pool) if rows >= 256 => pool.run_rows(rows, &run_range),
11335 _ => run_range(0, rows),
11336 }
11337 });
11338 });
11339}
11340
11341#[derive(Clone, Copy)]
11342struct SendMut(*mut f32);
11343unsafe impl Send for SendMut {}
11344unsafe impl Sync for SendMut {}
11345
11346impl SendMut {
11347 #[inline]
11348 fn at(self, i: usize) -> *mut f32 {
11349 unsafe { self.0.add(i) }
11350 }
11351}
11352
11353#[cfg(test)]
11354mod tests {
11355 #[test]
11359 fn q8_round_is_round_clamp() {
11360 let reference = |t: f32| t.round().clamp(-127.0, 127.0) as i8;
11361 let mut probe = vec![
11362 0.0f32,
11363 -0.0,
11364 f32::NAN,
11365 f32::INFINITY,
11366 f32::NEG_INFINITY,
11367 f32::MAX,
11368 f32::MIN,
11369 1e30,
11370 -1e30,
11371 f32::MIN_POSITIVE,
11372 -f32::MIN_POSITIVE,
11373 ];
11374 for k in -300i32..=300 {
11375 let h = k as f32 * 0.5;
11376 let up = f32::from_bits(h.to_bits() + 1);
11377 let down = f32::from_bits(h.to_bits().wrapping_sub(1));
11378 for t in [h, up, down] {
11379 probe.push(t);
11380 probe.push(-t);
11381 }
11382 }
11383 let mut t = -140.0f32;
11384 while t < 140.0 {
11385 probe.push(t);
11386 t += 0.000_731;
11387 }
11388 for t in probe {
11389 assert_eq!(q8_round(t), reference(t), "t = {t:e} ({:#x})", t.to_bits());
11390 }
11391 }
11392
11393 use super::*;
11394
11395 #[test]
11396 fn q2tp_i8_dot_matches_exact_on_grid() {
11397 let (rows, cols) = (5, 64);
11401 let gpr = cols / GROUP_SIZE;
11402 let chunks: Vec<u8> = (0..rows * gpr * Q2TP_CHUNK)
11406 .map(|i| (i as u32).wrapping_mul(2654435761) as u8)
11407 .collect();
11408 let scales: Vec<f32> = (0..gpr).map(|g| 0.5 + g as f32 * 0.25).collect();
11409 let x: Vec<f32> = (0..cols)
11410 .map(|i| if i % 3 == 0 { -1.0 } else { 1.0 })
11411 .collect();
11412 let act = split_act(&x);
11413 assert!(
11414 act.outliers.is_empty(),
11415 "on-grid input must have no outliers"
11416 );
11417 let gsum = q1_group_sums(&act.xq, gpr);
11418 for r in 0..rows {
11419 let exact = q2tp_row_exact(&chunks, r, gpr, &x, &scales);
11420 let fast = dot_q2tp_row_i8(&chunks, r, gpr, &act.xq, &gsum, &scales) * act.sx;
11421 assert!(
11422 (exact - fast).abs() <= exact.abs() * 1e-5 + 1e-5,
11423 "row {r}: exact {exact} vs i8 {fast}"
11424 );
11425 }
11426 }
11427
11428 #[cfg(target_arch = "x86_64")]
11429 #[test]
11430 fn q2tp_avx2_dot_matches_scalar_for_random_patterns() {
11431 if !std::arch::is_x86_feature_detected!("avx2") {
11436 return;
11437 }
11438 let mut seed = 0x9e3779b9u32;
11439 let mut next = || {
11440 seed = seed.wrapping_mul(1664525).wrapping_add(1013904223);
11441 seed
11442 };
11443 for _ in 0..20_000 {
11444 let mut ch = [0u8; Q2TP_CHUNK];
11445 let mut x = [0i8; GROUP_SIZE];
11446 for b in &mut ch {
11447 *b = next() as u8;
11448 }
11449 for v in &mut x {
11450 *v = (next() >> 24) as i8;
11451 }
11452 let mut reference = 0i32;
11453 for (k, &b) in ch.iter().enumerate() {
11454 reference += (b & 3) as i32 * x[k * 4] as i32;
11455 reference += ((b >> 2) & 3) as i32 * x[k * 4 + 1] as i32;
11456 reference += ((b >> 4) & 3) as i32 * x[k * 4 + 2] as i32;
11457 reference += ((b >> 6) & 3) as i32 * x[k * 4 + 3] as i32;
11458 }
11459 let got = unsafe { q2tp_code_dot_avx2(&ch, &x) };
11462 assert_eq!(got, reference, "packed q2 lane mismatch");
11463 }
11464 }
11465
11466 #[test]
11467 fn q2tp_affine_fuses_half_scale_correction_without_changing_raw_decode() {
11468 let (rows, cols) = (1usize, GROUP_SIZE);
11469 let mut bytes = vec![0u8; Q2TP_CHUNK + 4 + 1];
11470 bytes[..Q2TP_CHUNK].fill(0x24); bytes[Q2TP_CHUNK..Q2TP_CHUNK + 2].copy_from_slice(&0u16.to_le_bytes());
11474 bytes[Q2TP_CHUNK + 2..Q2TP_CHUNK + 4].copy_from_slice(&0u16.to_le_bytes());
11475 bytes[Q2TP_CHUNK + 4] = 1; let x = vec![1.0f32; cols];
11477 let mut raw = vec![0.0f32; rows];
11478 let mut affine = vec![0.0f32; rows];
11479 q2tp_matvec_for_test(&bytes, &x, rows, cols, &mut raw);
11480 q2tp_affine_matvec_for_test(&bytes, &x, rows, cols, &mut affine);
11481 assert_eq!(raw, vec![-24.0]);
11482 assert_eq!(affine, vec![-8.0]);
11483 assert!((affine[0] - (raw[0] + 0.5 * cols as f32)).abs() < 1e-6);
11484 }
11485
11486 #[test]
11487 fn q8_row_dot_fast_matches_scalar() {
11488 let cols = 96;
11491 let row: Vec<u8> = (0..cols)
11492 .map(|i| ((i * 37 % 251) - 125) as i8 as u8)
11493 .collect();
11494 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.13).sin()).collect();
11495 let act = split_act(&x);
11496 let fast = q8_row_dot(&row, &act);
11497 let scalar = q8_row_dot_scalar(&row, &act);
11498 assert!(
11499 (fast - scalar).abs() <= scalar.abs() * 1e-5 + 1e-5,
11500 "fast {fast} vs scalar {scalar}"
11501 );
11502 }
11503
11504 #[test]
11505 fn f32_matvec_matches_matvec_rows_bitexact() {
11506 let (rows, cols) = (300, 40);
11507 let w: Vec<f32> = (0..rows * cols).map(|i| (i as f32 * 0.017).sin()).collect();
11508 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.05).cos()).collect();
11509 let qt = QTensor::from_f32(w.clone(), rows, cols);
11510
11511 let mut a = vec![0.0f32; rows];
11512 matvec_rows(None, &w, &x, &mut a);
11513 let mut b = vec![0.0f32; rows];
11514 qt.matvec(&x, &mut b, None);
11515 assert_eq!(a, b);
11516 }
11517
11518 #[test]
11519 fn sdot_kernel_exact_on_grid() {
11520 eprintln!("sdot_enabled = {}", sdot_enabled());
11525 let (rows, cols) = (9, 80); let w: Vec<u8> = (0..rows * cols)
11527 .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11528 .collect();
11529 let scales: Vec<f32> = (0..rows).map(|o| 0.005 + o as f32 * 0.001).collect();
11530 let x: Vec<f32> = (0..cols)
11531 .map(|i| match i % 3 {
11532 0 => 1.0,
11533 1 => -1.0,
11534 _ => 0.0,
11535 })
11536 .collect();
11537 let mut a = vec![0.0f32; rows];
11538 qmatvec(
11539 &w,
11540 &[],
11541 &scales,
11542 &x,
11543 &[],
11544 TensorDtype::Q8Row,
11545 rows,
11546 cols,
11547 &mut a,
11548 None,
11549 );
11550 for o in 0..rows {
11551 let mut acc = 0.0f32;
11552 for j in 0..cols {
11553 acc += (w[o * cols + j] as i8) as f32 * x[j];
11554 }
11555 let expect = acc * scales[o];
11556 assert!(
11557 (a[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11558 "row {o}: {} vs {expect}",
11559 a[o]
11560 );
11561 }
11562 }
11563
11564 #[test]
11565 fn q1_tbl_fast_path_matches_reference() {
11566 let (rows, cols) = (5, 256);
11571 let gpr = cols / GROUP_SIZE;
11572 let mut bytes = Vec::new();
11573 for t in 0..rows * gpr {
11574 let s = 0.007 + (t % 11) as f32 * 0.004;
11575 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11576 for j in 0..4 {
11577 bytes.push(((t * 53 + j * 89 + 7) % 249) as u8);
11578 }
11579 }
11580 let x: Vec<f32> = (0..cols)
11581 .map(|i| if (i * 5) % 7 < 3 { 1.0 } else { -1.0 })
11582 .collect();
11583 let mut w = vec![0.0f32; rows * cols];
11584 cortiq_core::quant::dequant_q1(&bytes, &mut w);
11585 let mut got = vec![0.0f32; rows];
11586 q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11587 for o in 0..rows {
11588 let expect: f32 = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11589 assert!(
11590 (got[o] - expect).abs() < 1e-3 * expect.abs().max(1e-3),
11591 "row {o}: {} vs {expect}",
11592 got[o]
11593 );
11594 }
11595 let b = 5usize;
11598 let mut xs_all = Vec::new();
11599 for bi in 0..b {
11600 xs_all.extend(x.iter().map(|v| if bi % 2 == 0 { *v } else { -*v }));
11601 }
11602 let mut mm = vec![0.0f32; b * rows];
11603 q1_matmat(&bytes, &xs_all, b, rows, cols, &mut mm, None);
11604 for bi in 0..b {
11605 let mut single = vec![0.0f32; rows];
11606 q1_matvec(
11607 &bytes,
11608 &xs_all[bi * cols..(bi + 1) * cols],
11609 rows,
11610 cols,
11611 &mut single,
11612 None,
11613 );
11614 assert_eq!(&mm[bi * rows..(bi + 1) * rows], &single[..], "stream {bi}");
11615 }
11616 }
11617
11618 #[test]
11619 fn q1_kernels_match_exact_reference() {
11620 let (rows, cols) = (7, 96);
11622 let gpr = cols / GROUP_SIZE;
11623 let mut bytes = Vec::new();
11624 for t in 0..rows * gpr {
11625 let s = 0.01 + (t % 13) as f32 * 0.003;
11626 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11627 for j in 0..4 {
11628 bytes.push(((t * 31 + j * 97) % 251) as u8);
11629 }
11630 }
11631 let x: Vec<f32> = (0..cols)
11633 .map(|i| if i % 3 == 0 { 1.0 } else { -1.0 })
11634 .collect();
11635 let mut w = vec![0.0f32; rows * cols];
11637 cortiq_core::quant::dequant_q1(&bytes, &mut w);
11638 let mut expect = vec![0.0f32; rows];
11639 for o in 0..rows {
11640 expect[o] = (0..cols).map(|j| w[o * cols + j] * x[j]).sum();
11641 }
11642 let mut got = vec![0.0f32; rows];
11643 q1_matvec(&bytes, &x, rows, cols, &mut got, None);
11644 for o in 0..rows {
11645 assert!(
11646 (got[o] - expect[o]).abs() < 1e-3 * expect[o].abs().max(1e-3),
11647 "row {o}: {} vs {}",
11648 got[o],
11649 expect[o]
11650 );
11651 }
11652 let x2: Vec<f32> = x.iter().map(|v| -v).collect();
11654 let (mut a1, mut a2) = (vec![0.0f32; rows], vec![0.0f32; rows]);
11655 q1_matvec2(&bytes, &x, &x2, rows, cols, &mut a1, &mut a2, None);
11656 assert_eq!(a1, got);
11657 let mut xs = x.clone();
11658 xs.extend_from_slice(&x2);
11659 let mut mm = vec![0.0f32; 2 * rows];
11660 q1_matmat(&bytes, &xs, 2, rows, cols, &mut mm, None);
11661 assert_eq!(&mm[..rows], got.as_slice());
11662 assert_eq!(&mm[rows..], a2.as_slice());
11663 }
11664
11665 #[test]
11666 fn repack_is_bit_identical() {
11667 let (rows, cols) = (267, 96); let w: Vec<u8> = (0..rows * cols)
11673 .map(|i| (((i * 89) % 253) as i32 - 126) as i8 as u8)
11674 .collect();
11675 let scales: Vec<f32> = (0..rows).map(|o| 0.003 + o as f32 * 0.0007).collect();
11676 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.37).sin() * 2.0).collect();
11677 let rep = q8_repack_layout(&w, rows, cols);
11678 for g in 0..rows / 4 {
11680 for c in 0..cols / 16 {
11681 for lane in 0..4 {
11682 assert_eq!(
11683 &rep[g * 4 * cols + c * 64 + lane * 16
11684 ..g * 4 * cols + c * 64 + lane * 16 + 16],
11685 &w[(g * 4 + lane) * cols + c * 16..(g * 4 + lane) * cols + c * 16 + 16],
11686 );
11687 }
11688 }
11689 }
11690 let mut a = vec![0.0f32; rows];
11691 qmatvec(
11692 &w,
11693 &[],
11694 &scales,
11695 &x,
11696 &[],
11697 TensorDtype::Q8Row,
11698 rows,
11699 cols,
11700 &mut a,
11701 None,
11702 );
11703 let mut b = vec![0.0f32; rows];
11704 qmatvec(
11705 &w,
11706 &rep,
11707 &scales,
11708 &x,
11709 &[],
11710 TensorDtype::Q8Row,
11711 rows,
11712 cols,
11713 &mut b,
11714 None,
11715 );
11716 assert_eq!(a, b, "full-range repack output diverged");
11717
11718 #[cfg(target_arch = "aarch64")]
11719 if sdot_enabled() {
11720 let act = split_act(&x);
11722 let mut c1 = vec![0.0f32; rows];
11723 let mut c2 = vec![0.0f32; rows];
11724 q8_range_sdot(
11725 &w,
11726 &[],
11727 &scales,
11728 &act,
11729 cols,
11730 SendMut(c1.as_mut_ptr()),
11731 3,
11732 rows - 2,
11733 );
11734 q8_range_sdot(
11735 &w,
11736 &rep,
11737 &scales,
11738 &act,
11739 cols,
11740 SendMut(c2.as_mut_ptr()),
11741 3,
11742 rows - 2,
11743 );
11744 assert_eq!(c1, c2, "unaligned-range repack output diverged");
11745 }
11746 }
11747
11748 #[test]
11749 fn sdot_a8w8_noise_is_bounded() {
11750 let (rows, cols) = (16, 512);
11754 let w: Vec<u8> = (0..rows * cols)
11755 .map(|i| (((i * 37) % 251) as i32 - 125) as i8 as u8)
11756 .collect();
11757 let scales = vec![0.01f32; rows];
11758 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.21).sin()).collect();
11759 let mut a = vec![0.0f32; rows];
11760 qmatvec(
11761 &w,
11762 &[],
11763 &scales,
11764 &x,
11765 &[],
11766 TensorDtype::Q8Row,
11767 rows,
11768 cols,
11769 &mut a,
11770 None,
11771 );
11772 let (mut num, mut den) = (0f64, 0f64);
11773 for o in 0..rows {
11774 let mut acc = 0.0f32;
11775 for j in 0..cols {
11776 acc += (w[o * cols + j] as i8) as f32 * x[j];
11777 }
11778 let expect = acc * scales[o];
11779 num += ((a[o] - expect) as f64).powi(2);
11780 den += (expect as f64).powi(2);
11781 }
11782 let rel = (num / den.max(1e-12)).sqrt();
11783 assert!(rel < 0.05, "A8W8 relative L2 error too high: {rel}");
11784 }
11785
11786 #[test]
11787 fn i8_dot_neon_matches_scalar() {
11788 let n = 100;
11789 let w: Vec<u8> = (0..n).map(|i| ((i * 37 + 11) % 251) as u8).collect();
11790 let x: Vec<f32> = (0..n).map(|i| (i as f32 * 0.13).sin()).collect();
11791 let mut scalar = 0.0f32;
11792 for j in 0..n {
11793 scalar += (w[j] as i8) as f32 * x[j];
11794 }
11795 let fast = dot_i8_f32(&w, &x);
11796 assert!((scalar - fast).abs() < 1e-3 * scalar.abs().max(1.0));
11797 }
11798
11799 #[test]
11801 fn vbitmatvec_matches_full_dequant() {
11802 let (rows, cols) = (6, 64);
11803 let ng = cols / GROUP_SIZE;
11804 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
11806 let mut bytes = bits.clone();
11807 for g in 0..rows * ng {
11808 let s = 0.02 + 0.001 * g as f32;
11809 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11810 }
11811 for r in 0..rows {
11812 let b = bits[r] as usize;
11813 let (mut acc, mut nb) = (0u64, 0usize);
11814 let mut rowbytes = Vec::new();
11815 for i in 0..cols {
11816 let v = ((i * 7 + r * 13) % (1 << b)) as u64;
11817 acc = (acc << b) | v;
11818 nb += b;
11819 while nb >= 8 {
11820 nb -= 8;
11821 rowbytes.push(((acc >> nb) & 0xFF) as u8);
11822 }
11823 }
11824 if nb > 0 {
11825 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11826 }
11827 bytes.extend_from_slice(&rowbytes);
11828 }
11829 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
11830
11831 let mut reference = vec![0f32; rows * cols];
11832 cortiq_core::quant::dequant_vbit(&bytes, rows, cols, &mut reference).unwrap();
11833 let mut expect = vec![0f32; rows];
11834 for r in 0..rows {
11835 expect[r] = reference[r * cols..(r + 1) * cols]
11836 .iter()
11837 .zip(&x)
11838 .map(|(w, xv)| w * xv)
11839 .sum();
11840 }
11841 let mut got = vec![0f32; rows];
11842 let offsets = vbit_row_offsets(&bytes, rows, cols);
11843 vbitmatvec(&bytes, &offsets, &x, rows, cols, &mut got, None);
11844 let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
11848 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
11849 for r in 0..rows {
11850 assert!(
11851 (got[r] - expect[r]).abs() < tol * scale,
11852 "row {r}: {} vs {}",
11853 got[r],
11854 expect[r]
11855 );
11856 }
11857 }
11858
11859 #[test]
11864 #[cfg(target_arch = "x86_64")]
11865 fn vbit_matmat_blocked_matches_per_row() {
11866 let (rows, cols, b) = (64usize, 128usize, 9usize);
11867 let ng = cols / GROUP_SIZE;
11868 let bits: Vec<u8> = (0..rows).map(|r| [3u8, 4, 5, 6][r % 4]).collect();
11869 let mut bytes = bits.clone();
11870 for g in 0..rows * ng {
11871 let sc = 0.02 + 0.0005 * g as f32;
11872 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11873 }
11874 for r in 0..rows {
11875 let bw = bits[r] as usize;
11876 let (mut acc, mut nb) = (0u64, 0usize);
11877 let mut rowbytes = Vec::new();
11878 for i in 0..cols {
11879 let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
11880 acc = (acc << bw) | v;
11881 nb += bw;
11882 while nb >= 8 {
11883 nb -= 8;
11884 rowbytes.push(((acc >> nb) & 0xFF) as u8);
11885 }
11886 }
11887 if nb > 0 {
11888 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
11889 }
11890 bytes.extend_from_slice(&rowbytes);
11891 }
11892 let x: Vec<f32> = (0..b * cols)
11893 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11894 .collect();
11895 let offsets = vbit_row_offsets(&bytes, rows, cols);
11896 let mut y_a = vec![0f32; b * rows];
11897 let mut y_b = vec![0f32; b * rows];
11898 unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11899 vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_a, None);
11900 unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11901 vbitmatmat(&bytes, &offsets, &x, b, rows, cols, &mut y_b, None);
11902 unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11903 let max_d = y_a
11904 .iter()
11905 .zip(&y_b)
11906 .map(|(p, q)| (p - q).abs())
11907 .fold(0.0f32, f32::max);
11908 assert!(max_d < 1e-4, "vbit blocked ≠ per-row: max|Δ| = {max_d}");
11909 }
11910
11911 #[test]
11919 fn q4t_matmat_blocked_matches_per_row() {
11920 let (rows, cols, b) = (16usize, 64usize, 9usize);
11921 let gpr = cols / GROUP_SIZE;
11922 let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
11923 for r in 0..rows {
11924 for g in 0..gpr {
11925 let t = (r * gpr + g) * Q4_TILE;
11926 let sc = 0.02 + 0.001 * (r * gpr + g) as f32;
11927 bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
11928 for k in 0..16 {
11929 bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11930 }
11931 }
11932 }
11933 let x: Vec<f32> = (0..b * cols)
11934 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
11935 .collect();
11936 let mut y_blk = vec![0f32; b * rows];
11937 let mut y_row = vec![0f32; b * rows];
11938 unsafe { std::env::set_var("CMF_X86_BLOCKED", "1") };
11939 q4t_matmat(&bytes, &x, b, rows, cols, &mut y_blk, None);
11940 unsafe { std::env::set_var("CMF_X86_BLOCKED", "0") };
11941 q4t_matmat(&bytes, &x, b, rows, cols, &mut y_row, None);
11942 unsafe { std::env::remove_var("CMF_X86_BLOCKED") };
11943 assert_eq!(y_blk, y_row, "q4t blocked 1x4 ≠ per-row");
11944 }
11945
11946 fn synth_q4tp(rows: usize, cols: usize) -> Vec<u8> {
11953 use cortiq_core::quant::{f32_to_f16, q4tp_code_stride, q4tp_put_code};
11954 let gpr = cols / GROUP_SIZE;
11955 let stride = q4tp_code_stride(gpr);
11956 let (params_off, codes_off, _) = q4tp_sections(rows, cols);
11957 let mut b = vec![0u8; codes_off + rows * stride];
11958 for r in 0..rows {
11959 for g in 0..gpr {
11960 let t = (r * gpr + g) * Q4TP_NIB;
11961 for k in 0..16 {
11962 b[t + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
11963 }
11964 }
11965 let lo = -6.0 - 0.03 * (r % 17) as f32;
11966 let step = 0.01 + 0.004 * (r % 11) as f32;
11967 let p = params_off + r * 4;
11968 b[p..p + 2].copy_from_slice(&f32_to_f16(lo).to_le_bytes());
11969 b[p + 2..p + 4].copy_from_slice(&f32_to_f16(step).to_le_bytes());
11970 let crow = &mut b[codes_off + r * stride..codes_off + (r + 1) * stride];
11971 for g in 0..gpr {
11972 q4tp_put_code(crow, g, (r * 5 + g * 3) % 32);
11973 }
11974 }
11975 b
11976 }
11977
11978 fn q4tp_as_q4t(bytes: &[u8], rows: usize, cols: usize) -> Vec<u8> {
11982 let gpr = cols / GROUP_SIZE;
11983 let v = Q4tpView::new(bytes, rows, cols);
11984 let mut out = vec![0u8; rows * gpr * Q4_TILE];
11985 let mut sc = vec![0f32; gpr];
11986 for r in 0..rows {
11987 v.scales_into(r, gpr, &mut sc);
11988 for g in 0..gpr {
11989 let t = (r * gpr + g) * Q4_TILE;
11990 let s = sc[g];
11991 out[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
11992 let src = (r * gpr + g) * Q4TP_NIB;
11993 out[t + 2..t + Q4_TILE].copy_from_slice(&v.nib[src..src + Q4TP_NIB]);
11994 }
11995 }
11996 out
11997 }
11998
11999 #[test]
12005 fn q4tp_exact_path_matches_dequant_reference() {
12006 let (rows, cols) = (256usize, 512usize);
12007 let gpr = cols / GROUP_SIZE;
12008 let bytes = synth_q4tp(rows, cols);
12009 let mut w = vec![0f32; rows * cols];
12010 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12011
12012 let x: Vec<f32> = (0..cols)
12013 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12014 .collect();
12015 let v = Q4tpView::new(&bytes, rows, cols);
12016 let mut sc = vec![0f32; gpr];
12017 for r in 0..rows {
12018 v.scales_into(r, gpr, &mut sc);
12019 let got = q4tp_row_exact(v.nib, r, gpr, &x, &sc);
12020 let want: f32 = (0..cols).map(|c| w[r * cols + c] * x[c]).sum();
12021 let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12025 assert!(
12026 (got - want).abs() <= 1e-5 * mag,
12027 "row {r}: kernel {got} vs dequant {want}"
12028 );
12029 }
12030 }
12031
12032 #[test]
12038 fn q4tp_matvec_matches_the_q4t_kernel_it_was_ported_from() {
12039 let (rows, cols) = (256usize, 512usize);
12040 let bytes = synth_q4tp(rows, cols);
12041 let twin = q4tp_as_q4t(&bytes, rows, cols);
12042 let x: Vec<f32> = (0..cols)
12043 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12044 .collect();
12045
12046 let mut got = vec![0f32; rows];
12047 q4tp_matvec(&bytes, &x, rows, cols, &mut got, None);
12048 let mut want = vec![0f32; rows];
12049 q4t_matvec(&twin, &x, rows, cols, &mut want, None);
12050
12051 let mut w = vec![0f32; rows * cols];
12054 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12055 for r in 0..rows {
12056 let mag: f32 = (0..cols).map(|c| (w[r * cols + c] * x[c]).abs()).sum();
12057 assert!(
12058 (got[r] - want[r]).abs() <= 1e-3 * mag,
12059 "row {r}: q4tp {} vs q4t {}",
12060 got[r],
12061 want[r]
12062 );
12063 }
12064 }
12065
12066 #[test]
12071 fn q4tp_matmat_matches_the_q4t_kernel_it_was_ported_from() {
12072 let (rows, cols, b) = (256usize, 512usize, 5usize);
12073 let bytes = synth_q4tp(rows, cols);
12074 let twin = q4tp_as_q4t(&bytes, rows, cols);
12075 let xs: Vec<f32> = (0..b * cols)
12076 .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12077 .collect();
12078
12079 let mut got = vec![0f32; b * rows];
12080 q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, None);
12081 let mut want = vec![0f32; b * rows];
12082 q4t_matmat(&twin, &xs, b, rows, cols, &mut want, None);
12083
12084 let mut w = vec![0f32; rows * cols];
12085 cortiq_core::quant::dequant_q4tp(&bytes, rows, cols, &mut w);
12086 for t in 0..b {
12087 for r in 0..rows {
12088 let mag: f32 = (0..cols)
12089 .map(|c| (w[r * cols + c] * xs[t * cols + c]).abs())
12090 .sum();
12091 let (g, wa) = (got[t * rows + r], want[t * rows + r]);
12092 assert!(
12093 (g - wa).abs() <= 1e-3 * mag,
12094 "batch {t} row {r}: q4tp {g} vs q4t {wa}"
12095 );
12096 }
12097 }
12098 }
12099
12100 #[test]
12101 fn q4tp_matvec2_matches_the_single_stream_kernel() {
12102 let (rows, cols) = (128usize, 256usize);
12103 let gpr = cols / GROUP_SIZE;
12104 let bytes = synth_q4tp(rows, cols);
12105 let xs: Vec<f32> = (0..2 * cols)
12106 .map(|i| ((i * 29 + 11) % 89) as f32 / 89.0 - 0.5)
12107 .collect();
12108
12109 let (mut o1, mut o2) = (vec![0f32; rows], vec![0f32; rows]);
12110 q4tp_matvec2(
12111 &bytes,
12112 &xs[..cols],
12113 &xs[cols..],
12114 rows,
12115 cols,
12116 &mut o1,
12117 &mut o2,
12118 None,
12119 );
12120
12121 let v = Q4tpView::new(&bytes, rows, cols);
12124 let mut sc = vec![0f32; gpr];
12125 for r in 0..rows {
12126 v.scales_into(r, gpr, &mut sc);
12127 assert_eq!(o1[r], q4tp_row_exact(v.nib, r, gpr, &xs[..cols], &sc));
12128 assert_eq!(o2[r], q4tp_row_exact(v.nib, r, gpr, &xs[cols..], &sc));
12129 }
12130 }
12131
12132 #[test]
12139 fn q4tp_matvec_keeps_pace_with_q4t() {
12140 let (rows, cols) = (4096usize, 3072usize);
12141 let bytes = synth_q4tp(rows, cols);
12142 let twin = q4tp_as_q4t(&bytes, rows, cols);
12143 let x: Vec<f32> = (0..cols).map(|i| (i % 97) as f32 / 97.0 - 0.5).collect();
12144 let mut o = vec![0f32; rows];
12145 let n = 12;
12146 let mut best = (f64::MAX, f64::MAX);
12147 for _ in 0..3 {
12150 let t0 = std::time::Instant::now();
12151 for _ in 0..n {
12152 q4t_matvec(&twin, &x, rows, cols, &mut o, None);
12153 }
12154 best.0 = best.0.min(t0.elapsed().as_secs_f64());
12155 let t0 = std::time::Instant::now();
12156 for _ in 0..n {
12157 q4tp_matvec(&bytes, &x, rows, cols, &mut o, None);
12158 }
12159 best.1 = best.1.min(t0.elapsed().as_secs_f64());
12160 }
12161 let ratio = best.1 / best.0;
12162 println!(
12163 "q4t {:.3} ms | q4tp {:.3} ms | {ratio:.2}x",
12164 best.0 * 1e3 / n as f64,
12165 best.1 * 1e3 / n as f64
12166 );
12167 assert!(ratio < 2.0, "q4tp matvec {ratio:.2}x slower than q4t");
12168 }
12169
12170 #[cfg(target_os = "macos")]
12171 #[test]
12172 fn q4t_matmat_accel_matches_dequant_reference() {
12173 if !accel_gemm_enabled() {
12174 return; }
12176 let (rows, cols, b) = (512usize, 1024usize, 8usize); let gpr = cols / GROUP_SIZE;
12178 let mut bytes = vec![0u8; rows * gpr * Q4_TILE];
12179 for r in 0..rows {
12180 for g in 0..gpr {
12181 let t = (r * gpr + g) * Q4_TILE;
12182 let sc = 0.02 + 0.0005 * ((r * gpr + g) % 64) as f32;
12183 bytes[t..t + 2].copy_from_slice(&cortiq_core::quant::f32_to_f16(sc).to_le_bytes());
12184 for k in 0..16 {
12185 bytes[t + 2 + k] = ((r * 31 + g * 7 + k * 13) % 251) as u8;
12186 }
12187 }
12188 }
12189 let x: Vec<f32> = (0..b * cols)
12190 .map(|i| ((i * 13 + 7) % 97) as f32 / 97.0 - 0.5)
12191 .collect();
12192 let mut got = vec![0f32; b * rows];
12193 q4t_matmat(&bytes, &x, b, rows, cols, &mut got, None);
12194 let mut w = vec![0f32; rows * cols];
12196 for r in 0..rows {
12197 for g in 0..gpr {
12198 let t = (r * gpr + g) * Q4_TILE;
12199 let s = f16_to_f32(u16::from_le_bytes([bytes[t], bytes[t + 1]]));
12200 for (k, &bb) in bytes[t + 2..t + Q4_TILE].iter().enumerate() {
12201 w[r * cols + g * GROUP_SIZE + k * 2] = ((bb & 0x0F) as f32 - 8.0) * s;
12202 w[r * cols + g * GROUP_SIZE + k * 2 + 1] =
12203 (((bb >> 4) & 0x0F) as f32 - 8.0) * s;
12204 }
12205 }
12206 }
12207 for bi in 0..b {
12208 for r in 0..rows {
12209 let want: f32 = (0..cols).map(|j| x[bi * cols + j] * w[r * cols + j]).sum();
12210 let d = (got[bi * rows + r] - want).abs();
12211 assert!(
12212 d <= want.abs().max(1.0) * 1e-4,
12213 "accel q4t GEMM diverged at ({bi},{r}): {} vs {want}",
12214 got[bi * rows + r]
12215 );
12216 }
12217 }
12218 }
12219
12220 #[test]
12221 fn q4matvec_matches_full_dequant() {
12222 let (rows, cols) = (8, 64);
12223 let groups = rows * cols / GROUP_SIZE;
12224 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12226 for i in 0..groups * 16 {
12227 bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12228 }
12229 for g in 0..groups {
12230 let s = 0.01 + 0.003 * g as f32;
12231 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12232 }
12233 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12234
12235 let mut reference = vec![0.0f32; rows * cols];
12236 cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
12237 let mut expect = vec![0.0f32; rows];
12238 for r in 0..rows {
12239 expect[r] = reference[r * cols..(r + 1) * cols]
12240 .iter()
12241 .zip(&x)
12242 .map(|(w, xv)| w * xv)
12243 .sum();
12244 }
12245
12246 let mut got = vec![0.0f32; rows];
12247 q4matvec(&bytes, &x, rows, cols, &mut got, None);
12248 let tol = if a8w8_enabled() { 6e-2 } else { 1e-4 };
12252 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
12253 for r in 0..rows {
12254 assert!(
12255 (got[r] - expect[r]).abs() < tol * scale,
12256 "row {r}: {} vs {}",
12257 got[r],
12258 expect[r]
12259 );
12260 }
12261 }
12262
12263 #[test]
12266 fn vbitmatvec2_equals_two_singles() {
12267 let (rows, cols) = (6, 64);
12268 let ng = cols / GROUP_SIZE;
12269 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4];
12270 let mut bytes = bits.clone();
12271 for g in 0..rows * ng {
12272 let s = 0.02 + 0.001 * g as f32;
12273 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12274 }
12275 for r in 0..rows {
12276 let b = bits[r] as usize;
12277 let (mut acc, mut nb) = (0u64, 0usize);
12278 let mut rowbytes = Vec::new();
12279 for i in 0..cols {
12280 let v = ((i * 7 + r * 13) % (1 << b)) as u64;
12281 acc = (acc << b) | v;
12282 nb += b;
12283 while nb >= 8 {
12284 nb -= 8;
12285 rowbytes.push(((acc >> nb) & 0xFF) as u8);
12286 }
12287 }
12288 if nb > 0 {
12289 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12290 }
12291 bytes.extend_from_slice(&rowbytes);
12292 }
12293 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.19).sin()).collect();
12294 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).cos()).collect();
12295 let offsets = vbit_row_offsets(&bytes, rows, cols);
12296
12297 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12298 vbitmatvec(&bytes, &offsets, &x1, rows, cols, &mut a1, None);
12299 vbitmatvec(&bytes, &offsets, &x2, rows, cols, &mut a2, None);
12300 let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12301 vbitmatvec2(
12302 &bytes, &offsets, &x1, &x2, rows, cols, &mut b1, &mut b2, None,
12303 );
12304 assert_eq!(a1, b1, "fused vbit lane 1 must be bit-identical");
12305 assert_eq!(a2, b2, "fused vbit lane 2 must be bit-identical");
12306 }
12307
12308 #[test]
12310 fn q4matvec2_equals_two_singles() {
12311 let (rows, cols) = (8, 128);
12312 let groups = rows * cols / GROUP_SIZE;
12313 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
12314 for i in 0..groups * 16 {
12315 bytes.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12316 }
12317 for g in 0..groups {
12318 let s = 0.01 + 0.003 * g as f32;
12319 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
12320 }
12321 let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
12324 x1[9] = 250.0;
12325 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
12326
12327 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
12328 q4matvec(&bytes, &x1, rows, cols, &mut a1, None);
12329 q4matvec(&bytes, &x2, rows, cols, &mut a2, None);
12330 let (mut b1, mut b2) = (vec![0f32; rows], vec![0f32; rows]);
12331 q4matvec2(&bytes, &x1, &x2, rows, cols, &mut b1, &mut b2, None);
12332 assert_eq!(a1, b1, "fused q4 lane 1 must be bit-identical");
12333 assert_eq!(a2, b2, "fused q4 lane 2 must be bit-identical");
12334 }
12335
12336 #[test]
12339 fn matvec_many_equals_separate_matvecs() {
12340 use crate::pool::Pool;
12341 let (r1, r2, cols) = (300, 200, 64);
12342 let mk = |salt: usize, rows: usize| {
12343 QTensor::from_f32(
12344 (0..rows * cols)
12345 .map(|i| ((i * 7 + salt) % 97) as f32 / 97.0 - 0.5)
12346 .collect(),
12347 rows,
12348 cols,
12349 )
12350 };
12351 let (a, b) = (mk(1, r1), mk(5, r2));
12352 let x: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.11).sin()).collect();
12353 let pool = Pool::new(3);
12354
12355 let (mut ea, mut eb) = (vec![0f32; r1], vec![0f32; r2]);
12356 a.matvec(&x, &mut ea, Some(&pool));
12357 b.matvec(&x, &mut eb, Some(&pool));
12358 let (mut ga, mut gb) = (vec![0f32; r1], vec![0f32; r2]);
12359 QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12360 assert_eq!(ea, ga, "fused multi-matrix lane 1 must be bit-identical");
12361 assert_eq!(eb, gb, "fused multi-matrix lane 2 must be bit-identical");
12362 }
12363
12364 #[test]
12369 fn q4tp_matvec_many_equals_separate_matvecs() {
12370 use crate::pool::Pool;
12371 use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12372
12373 let (r1, r2, cols) = (300usize, 200usize, 64usize);
12374 let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12375 "arch_name": "tiny-q4tp",
12376 "hidden_size": cols,
12377 "intermediate_size": cols * 2,
12378 "num_layers": 1,
12379 "num_attention_heads": 2,
12380 "num_kv_heads": 1,
12381 "head_dim": 32,
12382 "vocab_size": r1,
12383 "layer_types": ["FullAttention"],
12384 "rms_norm_eps": 1e-6,
12385 "max_position_embeddings": 8,
12386 "linear_conv_kernel_dim": 0,
12387 "linear_num_key_heads": 0,
12388 "linear_num_value_heads": 0
12389 }))
12390 .unwrap();
12391 let header = CmfHeader {
12392 format: "cmf".into(),
12393 version: CMF_VERSION,
12394 arch,
12395 quant_type: QuantType::Q4Block,
12396 provenance: None,
12397 tokenizer_config: None,
12398 section_hashes: None,
12399 skills: Vec::new(),
12400 shard: None,
12401 calibration: None,
12402 routing: None,
12403 genome: None,
12404 lineage: Vec::new(),
12405 router: None,
12406 segments: Vec::new(),
12407 };
12408 let specs = [
12409 TensorSpec {
12410 name: "q".into(),
12411 dtype: TensorDtype::Q4TiledP,
12412 shape: vec![r1, cols],
12413 data: synth_q4tp(r1, cols),
12414 },
12415 TensorSpec {
12416 name: "kv".into(),
12417 dtype: TensorDtype::Q4TiledP,
12418 shape: vec![r2, cols],
12419 data: synth_q4tp(r2, cols),
12420 },
12421 ];
12422 let dir = std::env::temp_dir().join(format!("cmf-q4tp-many-{}", std::process::id()));
12423 std::fs::create_dir_all(&dir).unwrap();
12424 let path = dir.join("m.cmf");
12425 CmfModel::write(&path, &header, &specs, None, None).unwrap();
12426 let model = Arc::new(CmfModel::open(&path).unwrap());
12427 let (a, b) = (
12428 QTensor::from_model(&model, "q").unwrap(),
12429 QTensor::from_model(&model, "kv").unwrap(),
12430 );
12431 assert_eq!(a.model_dtype(), Some(TensorDtype::Q4TiledP));
12432 assert_eq!(b.model_dtype(), Some(TensorDtype::Q4TiledP));
12433 let x: Vec<f32> = (0..cols)
12434 .map(|i| ((i * 17 + 3) % 97) as f32 / 97.0 - 0.5)
12435 .collect();
12436 let pool = Pool::new(3);
12437 let (mut ea, mut eb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12438 a.matvec(&x, &mut ea, Some(&pool));
12439 b.matvec(&x, &mut eb, Some(&pool));
12440 let (mut ga, mut gb) = (vec![0.0f32; r1], vec![0.0f32; r2]);
12441 QTensor::matvec_many([&a, &b], &x, [&mut ga, &mut gb], Some(&pool));
12442 assert_eq!(ea, ga, "Q4TP fused lane 1 must be bit-identical");
12443 assert_eq!(eb, gb, "Q4TP fused lane 2 must be bit-identical");
12444 let _ = std::fs::remove_dir_all(&dir);
12445 }
12446
12447 #[test]
12453 fn multi_token_moe_rows_equal_single_token_decode() {
12454 use crate::pool::Pool;
12455 use cortiq_core::{CMF_VERSION, CmfHeader, CmfModel, QuantType, TensorSpec};
12456
12457 let (h, inter, ne) = (64usize, 128usize, 3usize);
12458 let arch: cortiq_core::ModelArch = serde_json::from_value(serde_json::json!({
12459 "arch_name": "tiny-q4tp-moe",
12460 "hidden_size": h,
12461 "intermediate_size": inter,
12462 "num_layers": 1,
12463 "num_attention_heads": 2,
12464 "num_kv_heads": 1,
12465 "head_dim": 32,
12466 "vocab_size": 8,
12467 "layer_types": ["FullAttention"],
12468 "rms_norm_eps": 1e-6,
12469 "max_position_embeddings": 8,
12470 "linear_conv_kernel_dim": 0,
12471 "linear_num_key_heads": 0,
12472 "linear_num_value_heads": 0
12473 }))
12474 .unwrap();
12475 let header = CmfHeader {
12476 format: "cmf".into(),
12477 version: CMF_VERSION,
12478 arch,
12479 quant_type: QuantType::Q4Block,
12480 provenance: None,
12481 tokenizer_config: None,
12482 section_hashes: None,
12483 skills: Vec::new(),
12484 shard: None,
12485 calibration: None,
12486 routing: None,
12487 genome: None,
12488 lineage: Vec::new(),
12489 router: None,
12490 segments: Vec::new(),
12491 };
12492 let mut specs = Vec::new();
12493 for e in 0..ne {
12494 for (k, (n, r, c)) in [("g", inter, h), ("u", inter, h), ("d", h, inter)]
12495 .into_iter()
12496 .enumerate()
12497 {
12498 let mut data = synth_q4tp(r, c);
12501 for (i, byte) in data[..r * (c / GROUP_SIZE) * Q4TP_NIB]
12502 .iter_mut()
12503 .enumerate()
12504 {
12505 *byte ^= ((i * (e * 3 + k + 1)) % 251) as u8;
12506 }
12507 specs.push(TensorSpec {
12508 name: format!("{n}{e}"),
12509 dtype: TensorDtype::Q4TiledP,
12510 shape: vec![r, c],
12511 data,
12512 });
12513 }
12514 }
12515 let dir = std::env::temp_dir().join(format!(
12516 "cmf-moe-rows-{}-{}",
12517 std::process::id(),
12518 FLOAT_ACTIVATIONS.get()
12519 ));
12520 std::fs::create_dir_all(&dir).unwrap();
12521 let path = dir.join("m.cmf");
12522 CmfModel::write(&path, &header, &specs, None, None).unwrap();
12523 let model = Arc::new(CmfModel::open(&path).unwrap());
12524 let t = |n: String| QTensor::from_model(&model, &n).unwrap();
12525 let g: Vec<QTensor> = (0..ne).map(|e| t(format!("g{e}"))).collect();
12526 let u: Vec<QTensor> = (0..ne).map(|e| t(format!("u{e}"))).collect();
12527 let d: Vec<QTensor> = (0..ne).map(|e| t(format!("d{e}"))).collect();
12528 let b = 4usize;
12529 let mut xs: Vec<f32> = (0..b * h)
12530 .map(|i| ((i * 31 + 7) % 89) as f32 / 89.0 - 0.5)
12531 .collect();
12532 xs[5] = 9.0; let routes: Vec<(Vec<usize>, Vec<f32>)> = vec![
12535 (vec![2, 0], vec![0.6, 0.4]),
12536 (vec![0, 1, 2], vec![0.2, 0.5, 0.3]),
12537 (vec![1], vec![1.0]),
12538 (vec![2, 1, 0], vec![0.25, 0.25, 0.5]),
12539 ];
12540 let pool = Pool::new(3);
12541 let mut want = vec![0f32; b * h];
12543 for (tk, (idx, w)) in routes.iter().enumerate() {
12544 let x = &xs[tk * h..(tk + 1) * h];
12545 let pairs: Vec<(&QTensor, &QTensor)> = idx.iter().map(|&e| (&g[e], &u[e])).collect();
12546 let mut gs: Vec<Vec<f32>> = idx.iter().map(|_| vec![0f32; inter]).collect();
12547 assert!(QTensor::moe_gate_up_many(&pairs, x, &mut gs, Some(&pool)));
12548 if FLOAT_ACTIVATIONS.get() {
12549 for (slot, &e) in idx.iter().enumerate() {
12550 let (mut gate, mut up) = (vec![0.0; inter], vec![0.0; inter]);
12551 g[e].matvec(x, &mut gate, Some(&pool));
12552 u[e].matvec(x, &mut up, Some(&pool));
12553 for (v, u) in gate.iter_mut().zip(up) {
12554 *v = (*v / (1.0 + (-*v).exp())) * u;
12555 }
12556 assert_eq!(gs[slot], gate, "float gate/up must equal ordinary matvecs");
12557 }
12558 }
12559 let downs: Vec<&QTensor> = idx.iter().map(|&e| &d[e]).collect();
12560 assert!(QTensor::moe_down_many(
12561 &downs,
12562 &gs,
12563 w,
12564 &mut want[tk * h..(tk + 1) * h],
12565 Some(&pool)
12566 ));
12567 }
12568 if FLOAT_ACTIVATIONS.get() {
12569 for (tk, (idx, w)) in routes.iter().enumerate() {
12570 let mut scalar = vec![0.0; h];
12571 for (&e, &weight) in idx.iter().zip(w) {
12572 let (mut gate, mut up, mut down) =
12573 (vec![0.0; inter], vec![0.0; inter], vec![0.0; h]);
12574 g[e].matvec(&xs[tk * h..(tk + 1) * h], &mut gate, Some(&pool));
12575 u[e].matvec(&xs[tk * h..(tk + 1) * h], &mut up, Some(&pool));
12576 for (v, u) in gate.iter_mut().zip(up) {
12577 *v = (*v / (1.0 + (-*v).exp())) * u;
12578 }
12579 d[e].matvec(&gate, &mut down, Some(&pool));
12580 for (v, d) in scalar.iter_mut().zip(down) {
12581 *v += weight * d;
12582 }
12583 }
12584 assert_eq!(
12585 &want[tk * h..(tk + 1) * h],
12586 scalar,
12587 "float many equals scalar experts"
12588 );
12589 }
12590 }
12591 let mut experts: Vec<usize> = Vec::new();
12593 let mut groups: Vec<Vec<usize>> = Vec::new();
12594 for (tk, (idx, _)) in routes.iter().enumerate() {
12595 for &e in idx {
12596 match experts.iter().position(|&x| x == e) {
12597 Some(k) => groups[k].push(tk),
12598 None => {
12599 experts.push(e);
12600 groups.push(vec![tk]);
12601 }
12602 }
12603 }
12604 }
12605 let n_pairs: usize = groups.iter().map(|g| g.len()).sum();
12606 let pairs: Vec<(&QTensor, &QTensor)> = experts.iter().map(|&e| (&g[e], &u[e])).collect();
12607 let mut gs: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; inter]).collect();
12608 assert!(QTensor::moe_gate_up_rows(
12609 &pairs,
12610 &groups,
12611 &xs,
12612 &mut gs,
12613 Some(&pool)
12614 ));
12615 let downs: Vec<&QTensor> = experts.iter().map(|&e| &d[e]).collect();
12616 let lens: Vec<usize> = groups.iter().map(|g| g.len()).collect();
12617 let mut ds: Vec<Vec<f32>> = (0..n_pairs).map(|_| vec![0f32; h]).collect();
12618 assert!(QTensor::moe_down_rows(
12619 &downs,
12620 &lens,
12621 &gs,
12622 &mut ds,
12623 Some(&pool)
12624 ));
12625 let slot = |tk: usize, e: usize| {
12626 let k = experts.iter().position(|&x| x == e).unwrap();
12627 groups[..k].iter().map(|g| g.len()).sum::<usize>()
12628 + groups[k].iter().position(|&x| x == tk).unwrap()
12629 };
12630 let mut got = vec![0f32; b * h];
12631 for (tk, (idx, w)) in routes.iter().enumerate() {
12632 for i in 0..h {
12633 let mut acc = 0f32;
12634 for (&e, &we) in idx.iter().zip(w) {
12635 acc += we * ds[slot(tk, e)][i];
12636 }
12637 got[tk * h + i] = acc;
12638 }
12639 }
12640 assert!(want.iter().any(|v| *v != 0.0));
12641 assert_eq!(
12642 want.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12643 got.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12644 "multi-token MoE must equal decode bit for bit"
12645 );
12646
12647 let b5 = 5usize;
12650 let x5: Vec<f32> = (0..b5 * h)
12651 .map(|i| ((i * 13 + 5) % 71) as f32 / 71.0 - 0.5)
12652 .collect();
12653 let mut mm = vec![0f32; b5 * inter];
12654 row_exact_scope(|| g[1].matmat(&x5, b5, &mut mm, Some(&pool)));
12655 for tk in 0..b5 {
12656 let mut mv = vec![0f32; inter];
12657 g[1].matvec(&x5[tk * h..(tk + 1) * h], &mut mv, Some(&pool));
12658 assert_eq!(
12659 mv.iter().map(|v| v.to_bits()).collect::<Vec<_>>(),
12660 mm[tk * inter..(tk + 1) * inter]
12661 .iter()
12662 .map(|v| v.to_bits())
12663 .collect::<Vec<_>>(),
12664 "row-exact matmat token {tk}"
12665 );
12666 }
12667 let _ = std::fs::remove_dir_all(&dir);
12670 }
12671
12672 #[test]
12680 fn q4tp_matmat_fast_path_unchanged_outside_row_exact() {
12681 use crate::pool::Pool;
12682 use std::sync::atomic::Ordering::Relaxed;
12683 let _alt = Q4TP_ALT_TEST_LOCK.lock().unwrap_or_else(|e| e.into_inner());
12684 Q4TP_ALT.store(2, Relaxed);
12687 let pool = Pool::new(3);
12688 for &(rows, cols, b) in &[(64usize, 256usize, 7usize), (320, 1024, 9)] {
12692 let bytes = synth_q4tp(rows, cols);
12693 let mut xs: Vec<f32> = (0..b * cols)
12694 .map(|i| ((i * 29 + 11) % 83) as f32 / 83.0 - 0.5)
12695 .collect();
12696 xs[3] = 7.5; let run = |exact: bool| {
12698 let mut out = vec![0f32; b * rows];
12699 q4tp_matmat_with(&bytes, &xs, b, rows, cols, &mut out, Some(&pool), exact);
12700 out
12701 };
12702 let (fast, exact) = (run(false), run(true));
12703 #[cfg(not(target_arch = "aarch64"))]
12704 let _ = fast;
12705 let bits = |v: &[f32]| v.iter().map(|x| x.to_bits()).collect::<Vec<_>>();
12706 let mut matvecs = vec![0f32; b * rows];
12707 for (bi, o) in matvecs.chunks_mut(rows).enumerate() {
12708 q4tp_matvec(
12709 &bytes,
12710 &xs[bi * cols..(bi + 1) * cols],
12711 rows,
12712 cols,
12713 o,
12714 Some(&pool),
12715 );
12716 }
12717 assert!(matvecs.iter().any(|v| *v != 0.0));
12718 assert_eq!(
12719 bits(&exact),
12720 bits(&matvecs),
12721 "{rows}x{cols} b={b}: row-exact matmat must equal per-token matvecs"
12722 );
12723 #[cfg(target_arch = "aarch64")]
12724 {
12725 let gpr = cols / GROUP_SIZE;
12727 let v = Q4tpView::new(&bytes, rows, cols);
12728 let mut old = vec![0f32; b * rows];
12729 let mut sc = vec![0f32; gpr];
12730 let a8w8 = a8w8_enabled();
12731 let blocked = sdot_enabled() && blocked_enabled();
12732 let acts: Vec<SplitAct> = (0..b)
12733 .map(|bi| split_act(&xs[bi * cols..(bi + 1) * cols]))
12734 .collect();
12735 for r in 0..rows {
12736 v.scales_into(r, gpr, &mut sc);
12737 if !a8w8 {
12738 for bi in 0..b {
12739 let x = &xs[bi * cols..(bi + 1) * cols];
12740 old[bi * rows + r] = q4tp_row_exact(v.nib, r, gpr, x, &sc);
12741 }
12742 continue;
12743 }
12744 let finish = |d: f32, act: &SplitAct| {
12745 let mut acc = d * act.sx;
12746 for &(j, xv) in &act.outliers {
12747 let (w, s) = q4tp_outlier(v.nib, r, gpr, j, &sc);
12748 acc += w * s * xv;
12749 }
12750 acc
12751 };
12752 let mut bi = 0usize;
12753 while blocked && bi + 4 <= b {
12754 let xs4 = [
12755 acts[bi].xq.as_slice(),
12756 acts[bi + 1].xq.as_slice(),
12757 acts[bi + 2].xq.as_slice(),
12758 acts[bi + 3].xq.as_slice(),
12759 ];
12760 let d = unsafe { dot_q4tp_row_1x4_sdot(v.nib, r, gpr, xs4, &sc) };
12761 for k in 0..4 {
12762 old[(bi + k) * rows + r] = finish(d[k], &acts[bi + k]);
12763 }
12764 bi += 4;
12765 }
12766 for (bi, act) in acts.iter().enumerate().skip(bi) {
12767 let d = dot_q4tp_row_i8(v.nib, r, gpr, &act.xq, &sc);
12768 old[bi * rows + r] = finish(d, act);
12769 }
12770 }
12771 assert_eq!(
12772 bits(&fast),
12773 bits(&old),
12774 "{rows}x{cols} b={b}: the fast path changed outside row_exact"
12775 );
12776 if blocked {
12779 assert_ne!(
12780 bits(&fast),
12781 bits(&matvecs),
12782 "{rows}x{cols} b={b}: the tuned tile no longer runs outside row_exact"
12783 );
12784 }
12785 }
12786 }
12787 Q4TP_ALT.store(0, Relaxed);
12788 }
12789
12790 #[test]
12791 fn row_exact_scopes_survive_overlap_nesting_and_unwind() {
12792 use std::sync::{Barrier, atomic::{AtomicUsize, Ordering}};
12793 let active = AtomicUsize::new(0);
12796 counted_row_exact_scope(&active, || {
12797 assert_eq!(active.load(Ordering::Acquire), 1);
12798 counted_row_exact_scope(&active, || {
12799 assert_eq!(active.load(Ordering::Acquire), 2);
12800 });
12801 assert_eq!(active.load(Ordering::Acquire), 1);
12802 });
12803 assert_eq!(active.load(Ordering::Acquire), 0);
12804
12805 let both_entered = Barrier::new(2);
12806 let release_last = Barrier::new(2);
12807 std::thread::scope(|s| {
12808 let first = s.spawn(|| counted_row_exact_scope(&active, || {
12809 both_entered.wait();
12810 }));
12811 let last = s.spawn(|| counted_row_exact_scope(&active, || {
12812 both_entered.wait();
12813 release_last.wait();
12814 }));
12815 first.join().unwrap();
12816 let after_first = active.load(Ordering::Acquire);
12817 release_last.wait();
12818 last.join().unwrap();
12819 assert_eq!(after_first, 1, "second request must remain exact");
12820 });
12821 assert_eq!(active.load(Ordering::Acquire), 0);
12822 let panic = std::panic::catch_unwind(|| {
12823 counted_row_exact_scope(&active, || panic!("scope unwind"));
12824 });
12825 assert!(panic.is_err());
12826 assert_eq!(active.load(Ordering::Acquire), 0);
12827 }
12828
12829 #[test]
12830 #[cfg(target_arch = "x86_64")]
12831 fn q4tp_float_avx2_is_bitwise_scalar() {
12832 if !avx2_enabled() {
12833 return;
12834 }
12835 for cols in [32, 64, 96, 2048, 4096] {
12836 let rows = 9;
12837 let bytes = synth_q4tp(rows, cols);
12838 let v = Q4tpView::new(&bytes, rows, cols);
12839 let gpr = cols / GROUP_SIZE;
12840 let mut sc = vec![0.0; gpr];
12841 for seed in 1..=5 {
12842 let xs: Vec<f32> = (0..cols)
12843 .map(|i| (((i * 104729 + seed * 8191) % 100003) as f32 - 50001.0) / 7919.0)
12844 .collect();
12845 for r in 0..rows {
12846 v.scales_into(r, gpr, &mut sc);
12847 let scalar = q4tp_row_float_scalar(v.nib, r, gpr, &xs, &sc);
12848 let vector = unsafe { q4tp_row_float_avx2(v.nib, r, gpr, &xs, &sc) };
12849 assert_eq!(
12850 scalar.to_bits(),
12851 vector.to_bits(),
12852 "cols={cols} row={r} seed={seed}"
12853 );
12854 }
12855 }
12856 }
12857 }
12858
12859 #[test]
12860 fn multi_token_moe_rows_float_equal_single_token_decode() {
12861 float_activations_scope(multi_token_moe_rows_equal_single_token_decode);
12862 }
12863
12864 #[test]
12865 fn full_gpu_q8_scope_is_nested_and_thread_local() {
12866 assert!(!FULL_GPU_Q8.get());
12867 let before = gpu_split_frac();
12868 {
12869 let _guard = enter_full_gpu_q8_scope();
12870 assert_eq!(gpu_split_frac(), 1.0);
12871 {
12872 let _nested = enter_full_gpu_q8_scope();
12873 }
12874 assert_eq!(gpu_split_frac(), 1.0);
12875 std::thread::spawn(|| assert!(!FULL_GPU_Q8.get())).join().unwrap();
12876 }
12877 assert!(!FULL_GPU_Q8.get());
12878 assert_eq!(gpu_split_frac(), before);
12879 }
12880
12881 #[test]
12882 fn float_activation_scope_is_nested_thread_local_and_unwind_safe() {
12883 assert!(!FLOAT_ACTIVATIONS.get());
12884 let before = a8w8_enabled();
12885 float_activations_scope(|| {
12886 assert!(!a8w8_enabled());
12887 float_activations_scope(|| assert!(!a8w8_enabled()));
12888 assert!(FLOAT_ACTIVATIONS.get());
12889 std::thread::spawn(|| assert!(!FLOAT_ACTIVATIONS.get()))
12890 .join()
12891 .unwrap();
12892 });
12893 assert!(!FLOAT_ACTIVATIONS.get());
12894 assert_eq!(a8w8_enabled(), before);
12895 let _ = std::panic::catch_unwind(|| float_activations_scope(|| panic!("test unwind")));
12896 assert!(!FLOAT_ACTIVATIONS.get());
12897 }
12898
12899 #[test]
12902 fn batched_matmat_equals_per_position_matvec() {
12903 let (rows, cols, b) = (8, 64, 5);
12904 let groups = rows * cols / GROUP_SIZE;
12906 let mut q4 = Vec::new();
12907 for i in 0..groups * 16 {
12908 q4.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12909 }
12910 for g in 0..groups {
12911 q4.extend_from_slice(
12912 &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
12913 );
12914 }
12915 let ng = cols / GROUP_SIZE;
12917 let bits: Vec<u8> = vec![3, 4, 5, 6, 8, 4, 5, 3];
12918 let mut vb = bits.clone();
12919 for g in 0..rows * ng {
12920 vb.extend_from_slice(
12921 &cortiq_core::quant::f32_to_f16(0.02 + 0.001 * g as f32).to_le_bytes(),
12922 );
12923 }
12924 for r in 0..rows {
12925 let bw = bits[r] as usize;
12926 let (mut acc, mut nb) = (0u64, 0usize);
12927 let mut rowbytes = Vec::new();
12928 for i in 0..cols {
12929 let v = ((i * 7 + r * 13) % (1 << bw)) as u64;
12930 acc = (acc << bw) | v;
12931 nb += bw;
12932 while nb >= 8 {
12933 nb -= 8;
12934 rowbytes.push(((acc >> nb) & 0xFF) as u8);
12935 }
12936 }
12937 if nb > 0 {
12938 rowbytes.push(((acc << (8 - nb)) & 0xFF) as u8);
12939 }
12940 vb.extend_from_slice(&rowbytes);
12941 }
12942 let offsets = vbit_row_offsets(&vb, rows, cols);
12943
12944 let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
12945
12946 let mut got = vec![0f32; b * rows];
12948 q4matmat(&q4, &xs, b, rows, cols, &mut got, None);
12949 for bi in 0..b {
12950 let mut expect = vec![0f32; rows];
12951 q4matvec(
12952 &q4,
12953 &xs[bi * cols..(bi + 1) * cols],
12954 rows,
12955 cols,
12956 &mut expect,
12957 None,
12958 );
12959 assert_eq!(
12960 &got[bi * rows..(bi + 1) * rows],
12961 &expect[..],
12962 "q4 batch pos {bi}"
12963 );
12964 }
12965
12966 let mut got = vec![0f32; b * rows];
12968 vbitmatmat(&vb, &offsets, &xs, b, rows, cols, &mut got, None);
12969 for bi in 0..b {
12970 let mut expect = vec![0f32; rows];
12971 vbitmatvec(
12972 &vb,
12973 &offsets,
12974 &xs[bi * cols..(bi + 1) * cols],
12975 rows,
12976 cols,
12977 &mut expect,
12978 None,
12979 );
12980 assert_eq!(
12981 &got[bi * rows..(bi + 1) * rows],
12982 &expect[..],
12983 "vbit batch pos {bi}"
12984 );
12985 }
12986 }
12987
12988 #[test]
12992 fn q4_tiled_matches_q4_block_bitexact() {
12993 let (rows, cols, b) = (8usize, 128usize, 3usize);
12994 let groups = rows * cols / GROUP_SIZE;
12995 let mut split = Vec::with_capacity(groups * 18);
12996 for i in 0..groups * 16 {
12997 split.push((((i * 7 + 3) % 256) & 0xFF) as u8);
12998 }
12999 for g in 0..groups {
13000 split.extend_from_slice(
13001 &cortiq_core::quant::f32_to_f16(0.01 + 0.003 * g as f32).to_le_bytes(),
13002 );
13003 }
13004 let (packed, scales) = split.split_at(groups * 16);
13006 let mut tiled = Vec::with_capacity(groups * Q4_TILE);
13007 for g in 0..groups {
13008 tiled.extend_from_slice(&scales[g * 2..g * 2 + 2]);
13009 tiled.extend_from_slice(&packed[g * 16..(g + 1) * 16]);
13010 }
13011
13012 let mut x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).sin()).collect();
13013 x1[9] = 250.0; let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.23).cos()).collect();
13015
13016 let (mut a, mut t) = (vec![0f32; rows], vec![0f32; rows]);
13017 q4matvec(&split, &x1, rows, cols, &mut a, None);
13018 q4t_matvec(&tiled, &x1, rows, cols, &mut t, None);
13019 assert_eq!(a, t, "q4t matvec must match q4 bit-for-bit");
13020
13021 let (mut a1, mut a2) = (vec![0f32; rows], vec![0f32; rows]);
13022 let (mut t1, mut t2) = (vec![0f32; rows], vec![0f32; rows]);
13023 q4matvec2(&split, &x1, &x2, rows, cols, &mut a1, &mut a2, None);
13024 q4t_matvec2(&tiled, &x1, &x2, rows, cols, &mut t1, &mut t2, None);
13025 assert_eq!(a1, t1);
13026 assert_eq!(a2, t2);
13027
13028 let xs: Vec<f32> = (0..b * cols).map(|i| (i as f32 * 0.13).sin()).collect();
13029 let (mut am, mut tm) = (vec![0f32; b * rows], vec![0f32; b * rows]);
13030 q4matmat(&split, &xs, b, rows, cols, &mut am, None);
13031 q4t_matmat(&tiled, &xs, b, rows, cols, &mut tm, None);
13032 assert_eq!(am, tm, "q4t matmat must match q4 bit-for-bit");
13033 }
13034
13035 #[test]
13042 fn q4matvec_sdot_outlier_exact() {
13043 let (rows, cols) = (4, 128);
13044 let groups = rows * cols / GROUP_SIZE;
13045 let mut bytes = Vec::with_capacity(groups * 16 + groups * 2);
13046 for i in 0..groups * 16 {
13047 bytes.push(((i * 11 + 5) % 256) as u8);
13048 }
13049 for g in 0..groups {
13050 let s = 0.02 + 0.002 * g as f32;
13051 bytes.extend_from_slice(&cortiq_core::quant::f32_to_f16(s).to_le_bytes());
13052 }
13053 let mut x: Vec<f32> = (0..cols)
13054 .map(|i| match i % 3 {
13055 0 => 1.0,
13056 1 => -1.0,
13057 _ => 0.0,
13058 })
13059 .collect();
13060 x[17] = 300.0; let mut reference = vec![0.0f32; rows * cols];
13063 cortiq_core::quant::dequant_q4_block(&bytes, &mut reference);
13064 let mut expect = vec![0.0f32; rows];
13065 for r in 0..rows {
13066 expect[r] = reference[r * cols..(r + 1) * cols]
13067 .iter()
13068 .zip(&x)
13069 .map(|(w, xv)| w * xv)
13070 .sum();
13071 }
13072 let mut got = vec![0.0f32; rows];
13073 q4matvec(&bytes, &x, rows, cols, &mut got, None);
13074 let scale = expect.iter().fold(0f32, |m, v| m.max(v.abs())).max(1.0);
13075 for r in 0..rows {
13076 assert!(
13077 (got[r] - expect[r]).abs() < 2e-3 * scale,
13078 "row {r}: {} vs {} (outlier term must be exact)",
13079 got[r],
13080 expect[r]
13081 );
13082 }
13083 }
13084
13085 #[test]
13089 fn q1t_matvec_matches_reference() {
13090 use cortiq_core::quant::{dequant_q1t, f32_to_f16};
13091 let (rows, cols) = (3usize, 64usize); let gpr = cols / GROUP_SIZE;
13093 let scales = [0.5f32, 0.3, 0.7, 0.2, 0.6, 0.15];
13094 let outliers: [(u32, f32); 3] = [(5, 9.0), (70, -4.5), (150, 3.25)];
13096 let is_out = |flat: usize| outliers.iter().any(|&(i, _)| i as usize == flat);
13097 let mut bytes = Vec::new();
13098 for r in 0..rows {
13099 for g in 0..gpr {
13100 bytes.extend_from_slice(&f32_to_f16(scales[r * gpr + g]).to_le_bytes());
13101 let mut c = [0u8; 7];
13102 for k in 0..GROUP_SIZE {
13103 let code = if is_out(r * cols + g * GROUP_SIZE + k) {
13105 0
13106 } else {
13107 ((k + r * 3 + g) % 3) as u8 };
13109 cortiq_core::quant::q1t_pack(&mut c, k, code);
13110 }
13111 bytes.extend_from_slice(&c);
13112 }
13113 }
13114 let mut row_ptr = vec![0u32; rows + 1];
13117 for &(idx, _) in &outliers {
13118 row_ptr[idx as usize / cols + 1] += 1;
13119 }
13120 for r in 0..rows {
13121 row_ptr[r + 1] += row_ptr[r];
13122 }
13123 for &p in &row_ptr {
13124 bytes.extend_from_slice(&p.to_le_bytes());
13125 }
13126 for &(idx, v) in &outliers {
13127 bytes.extend_from_slice(&((idx as usize % cols) as u16).to_le_bytes());
13128 bytes.extend_from_slice(&f32_to_f16(v).to_le_bytes());
13129 }
13130
13131 let mut refw = vec![0f32; rows * cols];
13132 dequant_q1t(&bytes, rows, cols, &mut refw);
13133 let x: Vec<f32> = (0..cols)
13136 .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13137 .collect();
13138 let mut expect = vec![0f32; rows];
13139 for r in 0..rows {
13140 let mut a = 0.0f32;
13141 for j in 0..cols {
13142 a += refw[r * cols + j] * x[j];
13143 }
13144 expect[r] = a;
13145 }
13146 let tol = |e: f32| 1e-3 * e.abs().max(1e-3);
13147 let mut got = vec![0f32; rows];
13148 q1t_matvec(&bytes, &x, rows, cols, &mut got, None);
13149 for r in 0..rows {
13150 assert!(
13151 (got[r] - expect[r]).abs() < tol(expect[r]),
13152 "row {r}: {} vs {}",
13153 got[r],
13154 expect[r]
13155 );
13156 }
13157 let x2: Vec<f32> = x.iter().chain(x.iter()).copied().collect();
13159 let mut gm = vec![0f32; 2 * rows];
13160 q1t_matmat(&bytes, &x2, 2, rows, cols, &mut gm, None);
13161 for r in 0..rows {
13162 assert!((gm[r] - expect[r]).abs() < tol(expect[r]));
13163 assert!((gm[rows + r] - expect[r]).abs() < tol(expect[r]));
13164 }
13165 let xb: Vec<f32> = (0..cols)
13169 .map(|j| if j % 5 == 0 { -1.0 } else { 1.0 })
13170 .collect();
13171 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13172 q1t_matvec(&bytes, &x, rows, cols, &mut s1, None);
13173 q1t_matvec(&bytes, &xb, rows, cols, &mut s2, None);
13174 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13175 q1t_matvec2(&bytes, &x, &xb, rows, cols, &mut p1, &mut p2, None);
13176 assert_eq!(p1, s1, "q1t pair lane 1 ≠ single matvec");
13177 assert_eq!(p2, s2, "q1t pair lane 2 ≠ single matvec");
13178 }
13179
13180 #[test]
13183 fn q1t_matvec2_odd_gpr_matches_singles() {
13184 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13185 let (rows, cols) = (5usize, 96usize); let gpr = cols / GROUP_SIZE;
13187 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13188 for r in 0..rows {
13189 for g in 0..gpr {
13190 bytes.extend_from_slice(&f32_to_f16(0.1 + 0.05 * (r + g) as f32).to_le_bytes());
13191 let mut c = [0u8; 7];
13192 for k in 0..GROUP_SIZE {
13193 q1t_pack(&mut c, k, ((k * 7 + r * 5 + g * 3) % 3) as u8);
13194 }
13195 bytes.extend_from_slice(&c);
13196 }
13197 }
13198 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13199 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13200 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13201 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13202 q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13203 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13204 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13205 assert_eq!(p1, s1, "odd-gpr pair lane 1 ≠ single");
13206 assert_eq!(p2, s2, "odd-gpr pair lane 2 ≠ single");
13207 }
13208
13209 #[test]
13213 #[ignore]
13214 fn q1t_matvec2_speed() {
13215 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_pack};
13216 use std::time::Instant;
13217 let (rows, cols) = (8192usize, 4096usize);
13218 let gpr = cols / GROUP_SIZE;
13219 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE);
13220 for r in 0..rows {
13221 for g in 0..gpr {
13222 let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13223 bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13224 let mut c = [0u8; 7];
13225 for k in 0..GROUP_SIZE {
13226 q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13227 }
13228 bytes.extend_from_slice(&c);
13229 }
13230 }
13231 let x1: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.31).sin()).collect();
13232 let x2: Vec<f32> = (0..cols).map(|i| (i as f32 * 0.17).cos()).collect();
13233 let (mut s1, mut s2) = (vec![0f32; rows], vec![0f32; rows]);
13234 let (mut p1, mut p2) = (vec![0f32; rows], vec![0f32; rows]);
13235 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13237 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13238 let (mut t_pair, mut t_two) = (f64::MAX, f64::MAX);
13239 for _ in 0..8 {
13240 let t0 = Instant::now();
13241 q1t_matvec2(&bytes, &x1, &x2, rows, cols, &mut p1, &mut p2, None);
13242 t_pair = t_pair.min(t0.elapsed().as_secs_f64() * 1000.0);
13243 let t1 = Instant::now();
13244 q1t_matvec(&bytes, &x1, rows, cols, &mut s1, None);
13245 q1t_matvec(&bytes, &x2, rows, cols, &mut s2, None);
13246 t_two = t_two.min(t1.elapsed().as_secs_f64() * 1000.0);
13247 }
13248 assert_eq!(p1, s1);
13249 assert_eq!(p2, s2);
13250 println!("q1t pair {rows}x{cols}: fused {t_pair:.2} ms | two singles {t_two:.2} ms");
13251 }
13252
13253 #[test]
13257 #[ignore]
13258 fn q1t_matvec_speed() {
13259 use cortiq_core::quant::{Q1T_TILE, f32_to_f16, q1t_code, q1t_pack};
13260 use std::time::Instant;
13261 let (rows, cols) = (8192usize, 4096usize); let gpr = cols / GROUP_SIZE;
13263 let mut bytes = Vec::with_capacity(rows * gpr * Q1T_TILE + 16);
13264 for r in 0..rows {
13265 for g in 0..gpr {
13266 let s = 0.1 + ((r + g) % 7) as f32 * 0.01;
13267 bytes.extend_from_slice(&f32_to_f16(s).to_le_bytes());
13268 let mut c = [0u8; 7];
13269 for k in 0..GROUP_SIZE {
13270 q1t_pack(&mut c, k, ((k * 7 + r + g) % 3) as u8);
13271 }
13272 bytes.extend_from_slice(&c);
13273 }
13274 }
13275 let (n, stride) = (rows * cols, 40usize); let mut row_ptr = vec![0u32; rows + 1];
13277 let mut idx = 0usize;
13278 while idx < n {
13279 row_ptr[idx / cols + 1] += 1;
13280 idx += stride;
13281 }
13282 for r in 0..rows {
13283 row_ptr[r + 1] += row_ptr[r];
13284 }
13285 for &p in &row_ptr {
13286 bytes.extend_from_slice(&p.to_le_bytes());
13287 }
13288 let mut idx = 0usize;
13289 while idx < n {
13290 bytes.extend_from_slice(&((idx % cols) as u16).to_le_bytes());
13291 bytes.extend_from_slice(&f32_to_f16((idx % 13) as f32 * 0.1 - 0.6).to_le_bytes());
13292 idx += stride;
13293 }
13294 let x: Vec<f32> = (0..cols)
13297 .map(|j| if j % 3 == 0 { 1.0 } else { -1.0 })
13298 .collect();
13299 let (rp_off, ent_off, has_ov) = q1t_overlay(&bytes, rows * gpr * Q1T_TILE, rows);
13300
13301 let slow = |out: &mut [f32]| {
13303 let mut buf = vec![0f32; cols];
13304 for r in 0..rows {
13305 for g in 0..gpr {
13306 let off = (r * gpr + g) * Q1T_TILE;
13307 let s = f16_to_f32(u16::from_le_bytes([bytes[off], bytes[off + 1]]));
13308 let codes = &bytes[off + 2..off + Q1T_TILE];
13309 for k in 0..GROUP_SIZE {
13310 buf[g * GROUP_SIZE + k] = match q1t_code(codes, k) {
13311 1 => s,
13312 2 => -s,
13313 _ => 0.0,
13314 };
13315 }
13316 }
13317 out[r] = q1t_row_outlier_correction(&bytes, r, rp_off, ent_off, has_ov, &x)
13318 + (0..cols).map(|j| buf[j] * x[j]).sum::<f32>();
13319 }
13320 };
13321 let iters = 5;
13322 let mut a = vec![0f32; rows];
13323 slow(&mut a); let t = Instant::now();
13325 for _ in 0..iters {
13326 slow(&mut a);
13327 }
13328 let slow_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13329
13330 let mut b = vec![0f32; rows];
13331 q1t_matvec(&bytes, &x, rows, cols, &mut b, None); let t = Instant::now();
13333 for _ in 0..iters {
13334 q1t_matvec(&bytes, &x, rows, cols, &mut b, None);
13335 }
13336 let fast_ms = t.elapsed().as_secs_f64() * 1e3 / iters as f64;
13337
13338 for r in 0..rows {
13339 assert!((a[r] - b[r]).abs() < 1e-2, "mismatch row {r}");
13340 }
13341 println!(
13342 "q1t matvec {rows}x{cols} (1 thread): div-decode {slow_ms:.2} ms fused-LUT {fast_ms:.2} ms => {:.2}x",
13343 slow_ms / fast_ms
13344 );
13345 }
13346}
13347
13348#[cfg(test)]
13349mod gemm_bench {
13350 #[test]
13361 #[ignore]
13362 fn q4tp_matmat_throughput() {
13363 let _alt = super::Q4TP_ALT_TEST_LOCK
13364 .lock()
13365 .unwrap_or_else(|e| e.into_inner());
13366 let b: usize = std::env::var("CMF_BENCH_B")
13370 .ok()
13371 .and_then(|v| v.parse().ok())
13372 .unwrap_or(296);
13373 let (rows, cols) = (9216usize, 2304usize);
13374 let (_, _, _) = (rows, cols, b);
13375 let total =
13376 cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4TiledP, &[rows, cols])
13377 .unwrap();
13378 let (params_off, codes_off, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13383 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 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 _ = codes_off;
13392 let xs: Vec<f32> = (0..b * cols)
13393 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13394 .collect();
13395 let mut out = vec![0f32; b * rows];
13396 let pool = crate::pool::Pool::from_env();
13397 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13403 let reps: usize = std::env::var("CMF_BENCH_REPS")
13404 .ok()
13405 .and_then(|v| v.parse().ok())
13406 .unwrap_or(10);
13407 let mut best = [f64::MAX; 2];
13408 let mut sums = [0f32; 2];
13409 for _ in 0..reps {
13410 for (k, w) in [(0usize, 1u8), (1usize, 2u8)] {
13411 super::Q4TP_ALT.store(w, std::sync::atomic::Ordering::Relaxed);
13412 let t = std::time::Instant::now();
13413 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13414 best[k] = best[k].min(t.elapsed().as_secs_f64());
13415 sums[k] = out.iter().take(64).sum::<f32>();
13416 }
13417 }
13418 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13419 for (k, name) in ["previous", "tuned "].iter().enumerate() {
13420 println!(
13421 "q4tp matmat {rows}x{cols} b={b} {name}: {:.1} ms {:.1} GFLOP/s (checksum {:.3})",
13422 best[k] * 1e3,
13423 flops / best[k] / 1e9,
13424 sums[k]
13425 );
13426 }
13427 assert!(
13428 (sums[0] - sums[1]).abs() < 1e-2,
13429 "the tuned kernel changed the result: {} vs {}",
13430 sums[0],
13431 sums[1]
13432 );
13433 }
13434
13435 #[test]
13441 fn q4tp_matmat_blocked_matches_scalar() {
13442 use std::sync::atomic::Ordering::Relaxed;
13443 let _alt = super::Q4TP_ALT_TEST_LOCK
13444 .lock()
13445 .unwrap_or_else(|e| e.into_inner());
13446 for &(rows, cols, b) in &[
13454 (64usize, 128usize, 7usize),
13455 (33, 96, 4),
13456 (16, 256, 9),
13457 (192, 2304, 37),
13458 ] {
13459 let total = cortiq_core::quant::expected_nbytes(
13460 cortiq_core::TensorDtype::Q4TiledP,
13461 &[rows, cols],
13462 )
13463 .unwrap();
13464 let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13465 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 61 % 251) as u8).collect();
13466 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13467 let step = cortiq_core::quant::f32_to_f16(0.1);
13468 for r in 0..rows {
13469 let o = params_off + r * 4;
13470 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13471 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13472 }
13473 let xs: Vec<f32> = (0..b * cols)
13474 .map(|i| ((i % 89) as f32 - 44.0) / 44.0)
13475 .collect();
13476 let mut got = vec![0f32; b * rows];
13477 let mut want = vec![0f32; b * rows];
13478 let gpr = cols / 32;
13479 let view = super::Q4tpView::new(&bytes, rows, cols);
13480 let pool = crate::pool::Pool::from_env();
13481 super::Q4TP_ALT.store(2, Relaxed);
13482 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut got, pool.as_deref());
13483 super::Q4TP_ALT.store(1, Relaxed);
13484 super::q4tp_matmat(&bytes, &xs, b, rows, cols, &mut want, pool.as_deref());
13485 super::Q4TP_ALT.store(0, Relaxed);
13486 let scale = want.iter().fold(0f32, |m, v| m.max(v.abs())).max(1e-6);
13493 let (mut worst, mut at) = (0f32, 0usize);
13494 for (i, (g, w)) in got.iter().zip(&want).enumerate() {
13495 if (g - w).abs() > worst {
13496 worst = (g - w).abs();
13497 at = i;
13498 }
13499 }
13500 assert!(
13501 worst <= 1e-4 * scale,
13502 "{rows}x{cols} b={b}: blocked and scalar disagree by {worst:.3e} \
13503 (scale {scale:.3e}) at cell {at}: {} vs {}",
13504 got[at],
13505 want[at]
13506 );
13507
13508 let (mut e_blocked, mut e_scalar) = (0f64, 0f64);
13515 for bi in 0..b {
13516 let act = super::split_act(&xs[bi * cols..(bi + 1) * cols]);
13517 for r in 0..rows {
13518 let mut sc = vec![0f32; gpr];
13519 view.scales_into(r, gpr, &mut sc);
13520 let mut exact = 0f64;
13521 for j in 0..cols {
13522 let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13523 exact += w as f64 * sq as f64 * act.xq[j] as f64;
13524 }
13525 exact *= act.sx as f64;
13526 for &(j, xv) in &act.outliers {
13527 let (w, sq) = super::q4tp_outlier(view.nib, r, gpr, j, &sc);
13528 exact += w as f64 * sq as f64 * xv as f64;
13529 }
13530 let i = bi * rows + r;
13531 e_blocked = e_blocked.max((got[i] as f64 - exact).abs());
13532 e_scalar = e_scalar.max((want[i] as f64 - exact).abs());
13533 }
13534 }
13535 println!(
13536 "{rows}x{cols} b={b}: worst error vs f64 — blocked {e_blocked:.3e}, \
13537 per-column {e_scalar:.3e}"
13538 );
13539 assert!(
13544 e_blocked <= 1e-5 * scale as f64 && e_scalar <= 1e-5 * scale as f64,
13545 "{rows}x{cols} b={b}: error against f64 too large — blocked \
13546 {e_blocked:.3e}, per-column {e_scalar:.3e}, scale {scale:.3e}"
13547 );
13548 }
13549 }
13550
13551 #[test]
13557 #[ignore]
13558 fn q4tp_matmat_row_exact_cost() {
13559 let _alt = super::Q4TP_ALT_TEST_LOCK
13560 .lock()
13561 .unwrap_or_else(|e| e.into_inner());
13562 let pool = crate::pool::Pool::from_env();
13563 let reps: usize = std::env::var("CMF_BENCH_REPS")
13564 .ok()
13565 .and_then(|v| v.parse().ok())
13566 .unwrap_or(30);
13567 for &(rows, cols) in &[(2048usize, 4096usize), (4096, 2048), (4096, 4096)] {
13568 let total = cortiq_core::quant::expected_nbytes(
13569 cortiq_core::TensorDtype::Q4TiledP,
13570 &[rows, cols],
13571 )
13572 .unwrap();
13573 let (params_off, _, _) = cortiq_core::quant::q4tp_sections(rows, cols);
13574 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13575 let lo = cortiq_core::quant::f32_to_f16(-4.0);
13576 let step = cortiq_core::quant::f32_to_f16(0.1);
13577 for r in 0..rows {
13578 let o = params_off + r * 4;
13579 bytes[o..o + 2].copy_from_slice(&lo.to_le_bytes());
13580 bytes[o + 2..o + 4].copy_from_slice(&step.to_le_bytes());
13581 }
13582 for &b in &[2usize, 4, 5, 8] {
13583 let xs: Vec<f32> = (0..b * cols)
13584 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13585 .collect();
13586 let mut out = vec![0f32; b * rows];
13587 let mut best = [f64::MAX; 3];
13588 for _ in 0..reps {
13589 for (k, best_k) in best.iter_mut().enumerate() {
13590 let t = std::time::Instant::now();
13591 match k {
13592 0 | 1 => super::q4tp_matmat_with(
13593 &bytes,
13594 &xs,
13595 b,
13596 rows,
13597 cols,
13598 &mut out,
13599 pool.as_deref(),
13600 k == 1,
13601 ),
13602 _ => {
13603 for (bi, o) in out.chunks_mut(rows).enumerate() {
13604 super::q4tp_matvec(
13605 &bytes,
13606 &xs[bi * cols..(bi + 1) * cols],
13607 rows,
13608 cols,
13609 o,
13610 pool.as_deref(),
13611 );
13612 }
13613 }
13614 }
13615 *best_k = best_k.min(t.elapsed().as_secs_f64());
13616 }
13617 }
13618 println!(
13619 "q4tp {rows}x{cols} b={b}: fast {:.3} ms, row-exact {:.3} ms, \
13620 {b} matvecs {:.3} ms",
13621 best[0] * 1e3,
13622 best[1] * 1e3,
13623 best[2] * 1e3
13624 );
13625 }
13626 }
13627 }
13628
13629 #[test]
13633 #[ignore]
13634 fn q4t_matmat_throughput() {
13635 let (rows, cols, b) = (9216usize, 2304usize, 296usize);
13636 let total =
13637 cortiq_core::quant::expected_nbytes(cortiq_core::TensorDtype::Q4Tiled, &[rows, cols])
13638 .unwrap();
13639 let mut bytes: Vec<u8> = (0..total).map(|i| (i * 37 % 251) as u8).collect();
13642 let sc = cortiq_core::quant::f32_to_f16(0.02);
13643 for t in bytes.chunks_mut(super::Q4_TILE) {
13644 t[..2].copy_from_slice(&sc.to_le_bytes());
13645 }
13646 let xs: Vec<f32> = (0..b * cols)
13647 .map(|i| ((i % 97) as f32 - 48.0) / 48.0)
13648 .collect();
13649 let mut out = vec![0f32; b * rows];
13650 let pool = crate::pool::Pool::from_env();
13651 super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13652 let reps: usize = std::env::var("CMF_BENCH_REPS")
13653 .ok()
13654 .and_then(|v| v.parse().ok())
13655 .unwrap_or(10);
13656 let mut best = f64::MAX;
13657 for _ in 0..reps {
13658 let t = std::time::Instant::now();
13659 super::q4t_matmat(&bytes, &xs, b, rows, cols, &mut out, pool.as_deref());
13660 best = best.min(t.elapsed().as_secs_f64());
13661 }
13662 let flops = 2.0 * b as f64 * rows as f64 * cols as f64;
13663 println!(
13664 "q4t matmat {rows}x{cols} b={b}: {:.1} ms {:.1} GFLOP/s (checksum {:.3})",
13665 best * 1e3,
13666 flops / best / 1e9,
13667 out.iter().take(64).sum::<f32>()
13668 );
13669 }
13670}