1use std::os::raw::c_int;
13use std::os::raw::c_void;
14use std::ptr::NonNull;
15use std::sync::Once;
16
17use rayon::prelude::*;
18
19unsafe extern "C" {
20 fn mlas_sgemm(
23 trans_a: c_int,
24 trans_b: c_int,
25 m: usize,
26 n: usize,
27 k: usize,
28 alpha: f32,
29 a: *const f32,
30 lda: usize,
31 b: *const f32,
32 ldb: usize,
33 beta: f32,
34 c: *mut f32,
35 ldc: usize,
36 );
37
38 fn mlas_sgemm_pack_b_size(trans_a: c_int, trans_b: c_int, n: usize, k: usize) -> usize;
39 fn mlas_sgemm_pack_b(
40 trans_a: c_int,
41 trans_b: c_int,
42 n: usize,
43 k: usize,
44 b: *const f32,
45 ldb: usize,
46 packed_b: *mut u8,
47 );
48 fn mlas_sgemm_packed(
49 trans_a: c_int,
50 trans_b: c_int,
51 m: usize,
52 n: usize,
53 k: usize,
54 alpha: f32,
55 a: *const f32,
56 lda: usize,
57 packed_b: *const u8,
58 beta: f32,
59 c: *mut f32,
60 ldc: usize,
61 );
62
63 fn mlas_float_kernel_id() -> c_int;
64
65 fn mlas_compute_logistic(input: *const f32, output: *mut f32, n: usize);
68 fn mlas_compute_silu(input: *const f32, output: *mut f32, n: usize);
71 fn mlas_eltwise_add(left: *const f32, right: *const f32, output: *mut f32, n: usize);
72 fn mlas_compute_activation(
73 kind: c_int,
74 minimum: f32,
75 maximum: f32,
76 input: *const f32,
77 output: *mut f32,
78 n: usize,
79 );
80
81 fn mlas_conv_prepare(
82 dimensions: usize,
83 batch_count: usize,
84 group_count: usize,
85 input_channels_per_group: usize,
86 input_shape: *const i64,
87 kernel_shape: *const i64,
88 dilation_shape: *const i64,
89 padding: *const i64,
90 stride_shape: *const i64,
91 output_shape: *const i64,
92 filter_count_per_group: usize,
93 working_buffer_elements: *mut usize,
94 ) -> *mut c_void;
95 fn mlas_conv_run(
96 plan: *const c_void,
97 input: *const f32,
98 filter: *const f32,
99 bias: *const f32,
100 working_buffer: *mut f32,
101 output: *mut f32,
102 );
103 fn mlas_conv_plan_destroy(plan: *mut c_void);
104
105 fn mlas_nchwc_block_size() -> usize;
107 fn mlas_nchwc_reorder_input_nchw(
108 source: *const f32,
109 dest: *mut f32,
110 channels: usize,
111 input_size: usize,
112 );
113 fn mlas_nchwc_reorder_output_nchw(output_shape: *const i64, source: *const f32, dest: *mut f32);
114 fn mlas_nchwc_reorder_filter_bibo(filter_shape: *const i64, source: *const f32, dest: *mut f32);
115 fn mlas_nchwc_reorder_filter_bo(filter_shape: *const i64, source: *const f32, dest: *mut f32);
116 #[allow(clippy::too_many_arguments)]
117 fn mlas_nchwc_conv(
118 input_shape: *const i64,
119 kernel_shape: *const i64,
120 dilation_shape: *const i64,
121 padding: *const i64,
122 stride_shape: *const i64,
123 output_shape: *const i64,
124 group_count: usize,
125 input: *const f32,
126 filter: *const f32,
127 bias: *const f32,
128 output: *mut f32,
129 activation_kind: c_int,
130 activation_value0: f32,
131 activation_value1: f32,
132 zero_mode: c_int,
133 );
134 fn mlas_pool(
135 kind: c_int,
136 dimensions: usize,
137 input_shape: *const i64,
138 kernel_shape: *const i64,
139 padding: *const i64,
140 stride_shape: *const i64,
141 output_shape: *const i64,
142 input: *const f32,
143 output: *mut f32,
144 );
145 #[allow(clippy::too_many_arguments)]
146 fn mlas_nchwc_pool(
147 kind: c_int,
148 input_shape: *const i64,
149 kernel_shape: *const i64,
150 dilation_shape: *const i64,
151 padding: *const i64,
152 stride_shape: *const i64,
153 output_shape: *const i64,
154 input: *const f32,
155 output: *mut f32,
156 );
157
158 fn mlas_qnbit_gemm_available(bits: usize, blk_len: usize, comp_type: c_int) -> c_int;
160 fn mlas_qnbit_gemm_pack_b_size(
161 n: usize,
162 k: usize,
163 bits: usize,
164 blk_len: usize,
165 has_zp: c_int,
166 comp_type: c_int,
167 ) -> usize;
168 fn mlas_qnbit_gemm_pack_b(
169 n: usize,
170 k: usize,
171 bits: usize,
172 blk_len: usize,
173 comp_type: c_int,
174 quant_b_data: *const c_void,
175 packed_b: *mut u8,
176 quant_b_scale: *const f32,
177 has_zp: c_int,
178 quant_b_zero_point: *const c_void,
179 );
180 fn mlas_qnbit_gemm_workspace_size(
181 m: usize,
182 n: usize,
183 k: usize,
184 bits: usize,
185 blk_len: usize,
186 has_zp: c_int,
187 comp_type: c_int,
188 ) -> usize;
189 #[allow(clippy::too_many_arguments)]
190 fn mlas_qnbit_gemm(
191 m: usize,
192 n: usize,
193 k: usize,
194 bits: usize,
195 blk_len: usize,
196 comp_type: c_int,
197 a: *const f32,
198 lda: usize,
199 packed_b: *const u8,
200 quant_b_scale: *const f32,
201 has_zp: c_int,
202 quant_b_zero_point: *const c_void,
203 bias: *const f32,
204 c: *mut f32,
205 ldc: usize,
206 workspace: *mut u8,
207 multithread: c_int,
208 );
209
210 fn mlas_set_threading(
214 parallel_for: MlasParallelForFn,
215 max_threads: MlasMaxThreadsFn,
216 rust_ctx: *mut c_void,
217 );
218}
219
220type MlasTaskFn = unsafe extern "C" fn(task_ctx: *mut c_void, tid: isize);
222type MlasParallelForFn = unsafe extern "C" fn(
224 rust_ctx: *mut c_void,
225 iterations: isize,
226 task: MlasTaskFn,
227 task_ctx: *mut c_void,
228);
229type MlasMaxThreadsFn = unsafe extern "C" fn(rust_ctx: *mut c_void) -> c_int;
231
232unsafe extern "C" fn rayon_parallel_for(
236 _rust_ctx: *mut c_void,
237 iterations: isize,
238 task: MlasTaskFn,
239 task_ctx: *mut c_void,
240) {
241 if iterations <= 0 {
242 return;
243 }
244 let task_ctx = task_ctx as usize;
249 (0..iterations).into_par_iter().for_each(|tid| {
250 unsafe { task(task_ctx as *mut c_void, tid) };
253 });
254}
255
256unsafe extern "C" fn rayon_max_threads(_rust_ctx: *mut c_void) -> c_int {
259 rayon::current_num_threads().max(1) as c_int
260}
261
262static THREADING_INIT: Once = Once::new();
263
264fn ensure_threading() {
269 THREADING_INIT.call_once(|| unsafe {
270 mlas_set_threading(rayon_parallel_for, rayon_max_threads, std::ptr::null_mut());
271 });
272}
273
274pub fn selected_float_kernel() -> i32 {
277 unsafe { mlas_float_kernel_id() as i32 }
278}
279
280pub fn compute_logistic(input: &[f32], output: &mut [f32]) {
287 assert_eq!(
288 input.len(),
289 output.len(),
290 "compute_logistic input and output must have equal length"
291 );
292 if input.is_empty() {
293 return;
294 }
295 unsafe { mlas_compute_logistic(input.as_ptr(), output.as_mut_ptr(), input.len()) };
298}
299
300pub fn compute_silu(input: &[f32], output: &mut [f32]) {
305 assert_eq!(
306 input.len(),
307 output.len(),
308 "compute_silu input and output must have equal length"
309 );
310 if input.is_empty() {
311 return;
312 }
313 unsafe { mlas_compute_silu(input.as_ptr(), output.as_mut_ptr(), input.len()) };
316}
317
318pub fn eltwise_add(left: &[f32], right: &[f32], output: &mut [f32]) {
320 assert_eq!(left.len(), right.len());
321 assert_eq!(left.len(), output.len());
322 unsafe {
323 mlas_eltwise_add(
324 left.as_ptr(),
325 right.as_ptr(),
326 output.as_mut_ptr(),
327 output.len(),
328 );
329 }
330}
331
332pub fn compute_relu(input: &[f32], output: &mut [f32]) {
334 assert_eq!(input.len(), output.len());
335 unsafe {
336 mlas_compute_activation(
337 1,
338 0.0,
339 0.0,
340 input.as_ptr(),
341 output.as_mut_ptr(),
342 output.len(),
343 );
344 }
345}
346
347pub fn compute_clip(input: &[f32], output: &mut [f32], minimum: f32, maximum: f32) {
349 assert_eq!(input.len(), output.len());
350 unsafe {
351 mlas_compute_activation(
352 5,
353 minimum,
354 maximum,
355 input.as_ptr(),
356 output.as_mut_ptr(),
357 output.len(),
358 );
359 }
360}
361
362pub struct ConvPlan {
364 ptr: NonNull<c_void>,
365 working_buffer_elements: usize,
366}
367
368unsafe impl Send for ConvPlan {}
371unsafe impl Sync for ConvPlan {}
372
373impl ConvPlan {
374 #[allow(clippy::too_many_arguments)]
376 pub fn new(
377 batch_count: usize,
378 group_count: usize,
379 input_channels_per_group: usize,
380 input_shape: &[i64],
381 kernel_shape: &[i64],
382 dilation_shape: &[i64],
383 padding: &[i64],
384 stride_shape: &[i64],
385 output_shape: &[i64],
386 filter_count_per_group: usize,
387 ) -> Option<Self> {
388 let dimensions = input_shape.len();
389 assert!((1..=3).contains(&dimensions));
390 assert_eq!(kernel_shape.len(), dimensions);
391 assert_eq!(dilation_shape.len(), dimensions);
392 assert_eq!(padding.len(), dimensions * 2);
393 assert_eq!(stride_shape.len(), dimensions);
394 assert_eq!(output_shape.len(), dimensions);
395 ensure_threading();
396 let mut working_buffer_elements = 0;
397 let ptr = unsafe {
398 mlas_conv_prepare(
399 dimensions,
400 batch_count,
401 group_count,
402 input_channels_per_group,
403 input_shape.as_ptr(),
404 kernel_shape.as_ptr(),
405 dilation_shape.as_ptr(),
406 padding.as_ptr(),
407 stride_shape.as_ptr(),
408 output_shape.as_ptr(),
409 filter_count_per_group,
410 &mut working_buffer_elements,
411 )
412 };
413 Some(Self {
414 ptr: NonNull::new(ptr)?,
415 working_buffer_elements,
416 })
417 }
418
419 pub fn working_buffer_elements(&self) -> usize {
421 self.working_buffer_elements
422 }
423
424 pub fn run(
426 &self,
427 input: &[f32],
428 filter: &[f32],
429 bias: Option<&[f32]>,
430 working_buffer: &mut [f32],
431 output: &mut [f32],
432 ) {
433 assert!(working_buffer.len() >= self.working_buffer_elements);
434 ensure_threading();
435 unsafe {
436 mlas_conv_run(
437 self.ptr.as_ptr(),
438 input.as_ptr(),
439 filter.as_ptr(),
440 bias.map_or(std::ptr::null(), <[f32]>::as_ptr),
441 if self.working_buffer_elements == 0 {
442 std::ptr::null_mut()
443 } else {
444 working_buffer.as_mut_ptr()
445 },
446 output.as_mut_ptr(),
447 );
448 }
449 }
450}
451
452impl Drop for ConvPlan {
453 fn drop(&mut self) {
454 unsafe { mlas_conv_plan_destroy(self.ptr.as_ptr()) };
455 }
456}
457
458#[derive(Clone, Copy, Debug)]
463pub struct NchwcActivation {
464 pub kind: i32,
466 pub values: [f32; 2],
468}
469
470impl NchwcActivation {
471 pub const IDENTITY: Self = Self {
473 kind: 0,
474 values: [0.0, 0.0],
475 };
476 pub const RELU: Self = Self {
478 kind: 1,
479 values: [0.0, 0.0],
480 };
481
482 pub fn clip(minimum: f32, maximum: f32) -> Self {
484 Self {
485 kind: 5,
486 values: [minimum, maximum],
487 }
488 }
489}
490
491pub fn nchwc_block_size() -> usize {
495 unsafe { mlas_nchwc_block_size() }
496}
497
498pub fn nchwc_reorder_filter_bibo(filter_shape: &[i64; 4], source: &[f32], dest: &mut [f32]) {
502 unsafe {
503 mlas_nchwc_reorder_filter_bibo(filter_shape.as_ptr(), source.as_ptr(), dest.as_mut_ptr())
504 };
505}
506
507pub fn nchwc_reorder_filter_bo(filter_shape: &[i64; 4], source: &[f32], dest: &mut [f32]) {
512 unsafe {
513 mlas_nchwc_reorder_filter_bo(filter_shape.as_ptr(), source.as_ptr(), dest.as_mut_ptr())
514 };
515}
516
517pub fn nchwc_reorder_input_nchw(
521 source: &[f32],
522 dest: &mut [f32],
523 channels: usize,
524 input_size: usize,
525) {
526 ensure_threading();
527 unsafe {
528 mlas_nchwc_reorder_input_nchw(source.as_ptr(), dest.as_mut_ptr(), channels, input_size)
529 };
530}
531
532pub fn nchwc_reorder_output_nchw(output_shape: &[i64; 4], source: &[f32], dest: &mut [f32]) {
535 ensure_threading();
536 unsafe {
537 mlas_nchwc_reorder_output_nchw(output_shape.as_ptr(), source.as_ptr(), dest.as_mut_ptr())
538 };
539}
540
541#[allow(clippy::too_many_arguments)]
550pub fn nchwc_conv(
551 input_shape: &[i64; 4],
552 kernel_shape: &[i64; 2],
553 dilation_shape: &[i64; 2],
554 padding: &[i64; 4],
555 stride_shape: &[i64; 2],
556 output_shape: &[i64; 4],
557 group_count: usize,
558 input: &[f32],
559 filter: &[f32],
560 bias: Option<&[f32]>,
561 output: &mut [f32],
562 activation: NchwcActivation,
563 zero_mode: bool,
564) {
565 ensure_threading();
566 unsafe {
567 mlas_nchwc_conv(
568 input_shape.as_ptr(),
569 kernel_shape.as_ptr(),
570 dilation_shape.as_ptr(),
571 padding.as_ptr(),
572 stride_shape.as_ptr(),
573 output_shape.as_ptr(),
574 group_count,
575 input.as_ptr(),
576 filter.as_ptr(),
577 bias.map_or(std::ptr::null(), <[f32]>::as_ptr),
578 output.as_mut_ptr(),
579 activation.kind,
580 activation.values[0],
581 activation.values[1],
582 i32::from(zero_mode),
583 );
584 }
585}
586
587#[derive(Clone, Copy, Debug)]
589#[repr(i32)]
590pub enum PoolKind {
591 Maximum = 0,
592 AverageExcludePad = 1,
593 AverageIncludePad = 2,
594}
595
596#[allow(clippy::too_many_arguments)]
598pub fn pool(
599 kind: PoolKind,
600 input_shape: &[i64],
601 kernel_shape: &[i64],
602 padding: &[i64],
603 stride_shape: &[i64],
604 output_shape: &[i64],
605 input: &[f32],
606 output: &mut [f32],
607) {
608 let dimensions = input_shape.len().saturating_sub(2);
609 assert!((1..=3).contains(&dimensions));
610 assert_eq!(kernel_shape.len(), dimensions);
611 assert_eq!(padding.len(), dimensions * 2);
612 assert_eq!(stride_shape.len(), dimensions);
613 assert_eq!(output_shape.len(), dimensions + 2);
614 ensure_threading();
615 unsafe {
616 mlas_pool(
617 kind as c_int,
618 dimensions,
619 input_shape.as_ptr(),
620 kernel_shape.as_ptr(),
621 padding.as_ptr(),
622 stride_shape.as_ptr(),
623 output_shape.as_ptr(),
624 input.as_ptr(),
625 output.as_mut_ptr(),
626 );
627 }
628}
629
630#[allow(clippy::too_many_arguments)]
638pub fn nchwc_pool(
639 kind: PoolKind,
640 input_shape: &[i64; 4],
641 kernel_shape: &[i64; 2],
642 dilation_shape: &[i64; 2],
643 padding: &[i64; 4],
644 stride_shape: &[i64; 2],
645 output_shape: &[i64; 4],
646 input: &[f32],
647 output: &mut [f32],
648) {
649 ensure_threading();
650 unsafe {
651 mlas_nchwc_pool(
652 kind as c_int,
653 input_shape.as_ptr(),
654 kernel_shape.as_ptr(),
655 dilation_shape.as_ptr(),
656 padding.as_ptr(),
657 stride_shape.as_ptr(),
658 output_shape.as_ptr(),
659 input.as_ptr(),
660 output.as_mut_ptr(),
661 );
662 }
663}
664pub struct PackedB {
668 ptr: *mut u8,
669 layout: std::alloc::Layout,
670 n: usize,
671 k: usize,
672}
673
674unsafe impl Send for PackedB {}
677unsafe impl Sync for PackedB {}
678
679impl PackedB {
680 pub fn new(n: usize, k: usize, b: &[f32]) -> Self {
682 assert_eq!(b.len(), k * n);
683 let size = unsafe { mlas_sgemm_pack_b_size(0, 0, n, k) }.max(1);
684 let layout = std::alloc::Layout::from_size_align(size, 64).unwrap();
685 let ptr = unsafe { std::alloc::alloc_zeroed(layout) };
686 assert!(!ptr.is_null(), "packed-B allocation failed");
687 unsafe { mlas_sgemm_pack_b(0, 0, n, k, b.as_ptr(), n, ptr) };
688 Self { ptr, layout, n, k }
689 }
690
691 pub fn dimensions(&self) -> (usize, usize) {
693 (self.k, self.n)
694 }
695}
696
697impl Drop for PackedB {
698 fn drop(&mut self) {
699 unsafe { std::alloc::dealloc(self.ptr, self.layout) };
700 }
701}
702
703pub fn sgemm_nn_packed(m: usize, a: &[f32], packed: &PackedB, c: &mut [f32]) {
705 let (n, k) = (packed.n, packed.k);
706 assert_eq!(a.len(), m * k);
707 assert_eq!(c.len(), m * n);
708 ensure_threading();
709 unsafe {
710 mlas_sgemm_packed(
711 0,
712 0,
713 m,
714 n,
715 k,
716 1.0,
717 a.as_ptr(),
718 k,
719 packed.ptr,
720 0.0,
721 c.as_mut_ptr(),
722 n,
723 );
724 }
725}
726
727pub fn sgemm_nn(m: usize, n: usize, k: usize, a: &[f32], b: &[f32], c: &mut [f32]) {
732 assert_eq!(a.len(), m * k, "A must be m*k");
733 assert_eq!(b.len(), k * n, "B must be k*n");
734 assert_eq!(c.len(), m * n, "C must be m*n");
735 ensure_threading();
736 unsafe {
737 mlas_sgemm(
738 0,
739 0,
740 m,
741 n,
742 k,
743 1.0,
744 a.as_ptr(),
745 k,
746 b.as_ptr(),
747 n,
748 0.0,
749 c.as_mut_ptr(),
750 n,
751 );
752 }
753}
754
755#[allow(clippy::too_many_arguments)]
758pub fn sgemm(
759 trans_a: bool,
760 trans_b: bool,
761 m: usize,
762 n: usize,
763 k: usize,
764 alpha: f32,
765 a: &[f32],
766 lda: usize,
767 b: &[f32],
768 ldb: usize,
769 beta: f32,
770 c: &mut [f32],
771 ldc: usize,
772) {
773 ensure_threading();
774 unsafe {
775 mlas_sgemm(
776 trans_a as c_int,
777 trans_b as c_int,
778 m,
779 n,
780 k,
781 alpha,
782 a.as_ptr(),
783 lda,
784 b.as_ptr(),
785 ldb,
786 beta,
787 c.as_mut_ptr(),
788 ldc,
789 );
790 }
791}
792
793#[derive(Debug, Clone, Copy, PartialEq, Eq)]
797pub enum SQNBitComputeType {
798 Fp32,
800 Int8,
803}
804
805impl SQNBitComputeType {
806 #[inline]
807 fn raw(self) -> c_int {
808 match self {
811 SQNBitComputeType::Fp32 => 0, SQNBitComputeType::Int8 => 3, }
814 }
815}
816
817pub fn sqnbit_gemm_available(bits: usize, blk_len: usize, comp: SQNBitComputeType) -> bool {
821 unsafe { mlas_qnbit_gemm_available(bits, blk_len, comp.raw()) != 0 }
822}
823
824pub struct SQNBitPackedB {
837 ptr: *mut u8,
838 layout: std::alloc::Layout,
839 n: usize,
840 k: usize,
841 bits: usize,
842 blk_len: usize,
843 comp: SQNBitComputeType,
844 has_zp: bool,
845 scale: Vec<f32>,
846 zp: Option<Vec<u8>>,
847}
848
849unsafe impl Send for SQNBitPackedB {}
854unsafe impl Sync for SQNBitPackedB {}
855
856impl SQNBitPackedB {
857 #[allow(clippy::too_many_arguments)]
861 pub fn new(
862 n: usize,
863 k: usize,
864 bits: usize,
865 blk_len: usize,
866 comp: SQNBitComputeType,
867 quant_b_data: &[u8],
868 scale: &[f32],
869 zp: Option<&[u8]>,
870 ) -> Option<Self> {
871 if !sqnbit_gemm_available(bits, blk_len, comp) {
872 return None;
873 }
874 let has_zp = zp.is_some();
875 let size = unsafe {
876 mlas_qnbit_gemm_pack_b_size(n, k, bits, blk_len, has_zp as c_int, comp.raw())
877 };
878 if size == 0 {
879 return None;
880 }
881 let layout = std::alloc::Layout::from_size_align(size, 64).unwrap();
882 let ptr = unsafe { std::alloc::alloc_zeroed(layout) };
883 assert!(!ptr.is_null(), "SQNBit packed-B allocation failed");
884 let zp_ptr = zp.map_or(std::ptr::null(), |z| z.as_ptr()) as *const c_void;
885 unsafe {
886 mlas_qnbit_gemm_pack_b(
887 n,
888 k,
889 bits,
890 blk_len,
891 comp.raw(),
892 quant_b_data.as_ptr() as *const c_void,
893 ptr,
894 scale.as_ptr(),
895 has_zp as c_int,
896 zp_ptr,
897 );
898 }
899 Some(Self {
900 ptr,
901 layout,
902 n,
903 k,
904 bits,
905 blk_len,
906 comp,
907 has_zp,
908 scale: scale.to_vec(),
909 zp: zp.map(<[u8]>::to_vec),
910 })
911 }
912
913 pub fn dimensions(&self) -> (usize, usize) {
915 (self.k, self.n)
916 }
917}
918
919impl Drop for SQNBitPackedB {
920 fn drop(&mut self) {
921 unsafe { std::alloc::dealloc(self.ptr, self.layout) };
922 }
923}
924
925pub fn sqnbit_gemm(
932 packed: &SQNBitPackedB,
933 m: usize,
934 a: &[f32],
935 bias: Option<&[f32]>,
936 c: &mut [f32],
937 multithread: bool,
938) {
939 let n = packed.n;
940 assert_eq!(c.len(), m * n, "C must be m*n");
941 unsafe { sqnbit_gemm_into(packed, m, a, bias, c.as_mut_ptr(), n, multithread) };
945}
946
947pub unsafe fn sqnbit_gemm_into(
959 packed: &SQNBitPackedB,
960 m: usize,
961 a: &[f32],
962 bias: Option<&[f32]>,
963 c: *mut f32,
964 ldc: usize,
965 multithread: bool,
966) {
967 let (k, n) = (packed.k, packed.n);
968 assert_eq!(a.len(), m * k, "A must be m*k");
969 assert!(ldc >= n, "ldc must be >= packed N");
970 if let Some(bias) = bias {
971 assert_eq!(bias.len(), n, "bias must be length n");
972 }
973 ensure_threading();
974
975 let ws_size = unsafe {
976 mlas_qnbit_gemm_workspace_size(
977 m,
978 n,
979 k,
980 packed.bits,
981 packed.blk_len,
982 packed.has_zp as c_int,
983 packed.comp.raw(),
984 )
985 };
986 let mut workspace: Vec<u8> = if ws_size == 0 {
989 Vec::new()
990 } else {
991 vec![0u8; ws_size + 64]
992 };
993 let ws_ptr = if ws_size == 0 {
994 std::ptr::null_mut()
995 } else {
996 workspace.as_mut_ptr()
997 };
998
999 let zp_ptr = packed.zp.as_ref().map_or(std::ptr::null(), |z| z.as_ptr()) as *const c_void;
1000 let bias_ptr = bias.map_or(std::ptr::null(), <[f32]>::as_ptr);
1001
1002 unsafe {
1003 mlas_qnbit_gemm(
1004 m,
1005 n,
1006 k,
1007 packed.bits,
1008 packed.blk_len,
1009 packed.comp.raw(),
1010 a.as_ptr(),
1011 k,
1012 packed.ptr,
1013 packed.scale.as_ptr(),
1014 packed.has_zp as c_int,
1015 zp_ptr,
1016 bias_ptr,
1017 c,
1018 ldc,
1019 ws_ptr,
1020 multithread as c_int,
1021 );
1022 }
1023}
1024
1025#[cfg(test)]
1026mod tests {
1027 use super::*;
1028
1029 fn assert_send_sync<T: Send + Sync>() {}
1030
1031 #[test]
1032 fn packed_b_is_send_sync() {
1033 assert_send_sync::<PackedB>();
1034 }
1035
1036 #[allow(clippy::too_many_arguments)]
1038 fn ref_sgemm(
1039 trans_a: bool,
1040 trans_b: bool,
1041 m: usize,
1042 n: usize,
1043 k: usize,
1044 alpha: f32,
1045 a: &[f32],
1046 lda: usize,
1047 b: &[f32],
1048 ldb: usize,
1049 beta: f32,
1050 c: &mut [f32],
1051 ldc: usize,
1052 ) {
1053 for i in 0..m {
1054 for j in 0..n {
1055 let mut acc = 0.0f32;
1056 for p in 0..k {
1057 let av = if trans_a {
1058 a[p * lda + i]
1059 } else {
1060 a[i * lda + p]
1061 };
1062 let bv = if trans_b {
1063 b[j * ldb + p]
1064 } else {
1065 b[p * ldb + j]
1066 };
1067 acc += av * bv;
1068 }
1069 let cell = &mut c[i * ldc + j];
1070 *cell = alpha * acc + beta * *cell;
1071 }
1072 }
1073 }
1074
1075 fn seq(n: usize, seed: f32) -> Vec<f32> {
1076 (0..n)
1078 .map(|i| ((i as f32 * 0.013 + seed).sin()) * 2.0)
1079 .collect()
1080 }
1081
1082 fn assert_close(a: &[f32], b: &[f32], tol: f32, ctx: &str) {
1083 assert_eq!(a.len(), b.len());
1084 for (idx, (x, y)) in a.iter().zip(b.iter()).enumerate() {
1085 let diff = (x - y).abs();
1086 let rel = diff / (y.abs().max(1.0));
1087 assert!(
1088 diff <= tol || rel <= tol,
1089 "{ctx}: mismatch at {idx}: mlas={x} ref={y} diff={diff}"
1090 );
1091 }
1092 }
1093
1094 fn check_nn(m: usize, n: usize, k: usize) {
1095 let a = seq(m * k, 0.5);
1096 let b = seq(k * n, 1.5);
1097 let mut c_mlas = vec![0.0f32; m * n];
1098 let mut c_ref = vec![0.0f32; m * n];
1099 sgemm_nn(m, n, k, &a, &b, &mut c_mlas);
1100 ref_sgemm(false, false, m, n, k, 1.0, &a, k, &b, n, 0.0, &mut c_ref, n);
1101 assert_close(&c_mlas, &c_ref, 1e-3, &format!("nn {m}x{n}x{k}"));
1102 }
1103
1104 #[test]
1105 fn correctness_square() {
1106 check_nn(64, 64, 64);
1107 }
1108
1109 #[test]
1110 fn correctness_non_square_and_non_tile_multiples() {
1111 check_nn(1, 1, 1);
1113 check_nn(3, 5, 7);
1114 check_nn(17, 31, 13);
1115 check_nn(32, 512, 512);
1116 check_nn(33, 65, 129);
1117 check_nn(100, 1, 100);
1118 check_nn(1, 100, 100);
1119 }
1120
1121 #[test]
1122 fn correctness_alpha_beta() {
1123 let (m, n, k) = (23, 19, 41);
1124 let a = seq(m * k, 0.2);
1125 let b = seq(k * n, 0.7);
1126 let base = seq(m * n, 2.0);
1127 let mut c_mlas = base.clone();
1128 let mut c_ref = base.clone();
1129 sgemm(
1130 false,
1131 false,
1132 m,
1133 n,
1134 k,
1135 0.5,
1136 &a,
1137 k,
1138 &b,
1139 n,
1140 2.0,
1141 &mut c_mlas,
1142 n,
1143 );
1144 ref_sgemm(false, false, m, n, k, 0.5, &a, k, &b, n, 2.0, &mut c_ref, n);
1145 assert_close(&c_mlas, &c_ref, 1e-3, "alpha_beta");
1146 }
1147
1148 #[test]
1149 fn correctness_transpose_b() {
1150 let (m, n, k) = (12, 20, 28);
1152 let a = seq(m * k, 0.3);
1153 let b_t = seq(n * k, 0.9); let mut c_mlas = vec![0.0f32; m * n];
1155 let mut c_ref = vec![0.0f32; m * n];
1156 sgemm(
1157 false,
1158 true,
1159 m,
1160 n,
1161 k,
1162 1.0,
1163 &a,
1164 k,
1165 &b_t,
1166 k,
1167 0.0,
1168 &mut c_mlas,
1169 n,
1170 );
1171 ref_sgemm(
1172 false, true, m, n, k, 1.0, &a, k, &b_t, k, 0.0, &mut c_ref, n,
1173 );
1174 assert_close(&c_mlas, &c_ref, 1e-3, "transpose_b");
1175 }
1176
1177 #[test]
1178 fn correctness_transpose_a() {
1179 let (m, n, k) = (14, 22, 18);
1181 let a_t = seq(k * m, 0.4); let b = seq(k * n, 0.6);
1183 let mut c_mlas = vec![0.0f32; m * n];
1184 let mut c_ref = vec![0.0f32; m * n];
1185 sgemm(
1186 true,
1187 false,
1188 m,
1189 n,
1190 k,
1191 1.0,
1192 &a_t,
1193 m,
1194 &b,
1195 n,
1196 0.0,
1197 &mut c_mlas,
1198 n,
1199 );
1200 ref_sgemm(
1201 true, false, m, n, k, 1.0, &a_t, m, &b, n, 0.0, &mut c_ref, n,
1202 );
1203 assert_close(&c_mlas, &c_ref, 1e-3, "transpose_a");
1204 }
1205
1206 #[test]
1207 fn correctness_packed_b() {
1208 for (m, n, k) in [(32usize, 512usize, 512usize), (7, 13, 19), (1, 64, 64)] {
1209 let a = seq(m * k, 0.5);
1210 let b = seq(k * n, 1.5);
1211 let mut c_mlas = vec![0.0f32; m * n];
1212 let mut c_ref = vec![0.0f32; m * n];
1213 let packed = PackedB::new(n, k, &b);
1214 sgemm_nn_packed(m, &a, &packed, &mut c_mlas);
1215 ref_sgemm(false, false, m, n, k, 1.0, &a, k, &b, n, 0.0, &mut c_ref, n);
1216 assert_close(&c_mlas, &c_ref, 1e-3, &format!("packed {m}x{n}x{k}"));
1217 }
1218 }
1219
1220 #[test]
1221 fn float_kernel_matches_detected_isa() {
1222 let id = selected_float_kernel();
1223 let expected = if std::arch::is_x86_feature_detected!("avx512f") {
1224 512
1225 } else if std::arch::is_x86_feature_detected!("avx2")
1226 && std::arch::is_x86_feature_detected!("fma")
1227 {
1228 3
1229 } else if std::arch::is_x86_feature_detected!("avx") {
1230 1
1231 } else {
1232 -1
1233 };
1234 eprintln!("selected f32 GEMM kernel id = {id}; expected {expected} for host ISA");
1235 assert_eq!(id, expected, "MLAS f32 GEMM dispatch did not match host ISA");
1236 }
1237
1238 #[test]
1243 #[ignore = "perf probe; run explicitly with --ignored --nocapture"]
1244 fn perf_sgemm_medium() {
1245 use std::time::Instant;
1246
1247 let (m, n, k) = (32usize, 512usize, 512usize);
1248 let a = seq(m * k, 0.5);
1249 let b = seq(k * n, 1.5);
1250 let mut c = vec![0.0f32; m * n];
1251
1252 for _ in 0..50 {
1254 sgemm_nn(m, n, k, &a, &b, &mut c);
1255 }
1256
1257 let iters = 5000u32;
1258 let start = Instant::now();
1259 for _ in 0..iters {
1260 sgemm_nn(m, n, k, &a, &b, &mut c);
1261 }
1262 let elapsed = start.elapsed();
1263 let checksum: f32 = c.iter().copied().sum();
1265
1266 let per_us = elapsed.as_secs_f64() * 1e6 / iters as f64;
1267 let flops = 2.0 * m as f64 * n as f64 * k as f64;
1268 let gflops = flops / (per_us * 1e3);
1269 eprintln!(
1270 "vendored-MLAS SGEMM 32x512x512 single-thread (repack B/call): {per_us:.1} us/iter \
1271 ({gflops:.1} GFLOP/s), checksum={checksum:.3}"
1272 );
1273
1274 let packed = PackedB::new(n, k, &b);
1276 for _ in 0..50 {
1277 sgemm_nn_packed(m, &a, &packed, &mut c);
1278 }
1279 let start = Instant::now();
1280 for _ in 0..iters {
1281 sgemm_nn_packed(m, &a, &packed, &mut c);
1282 }
1283 let elapsed_p = start.elapsed();
1284 let checksum_p: f32 = c.iter().copied().sum();
1285 let per_us_p = elapsed_p.as_secs_f64() * 1e6 / iters as f64;
1286 let gflops_p = flops / (per_us_p * 1e3);
1287 eprintln!(
1288 "vendored-MLAS SGEMM 32x512x512 single-thread (pre-packed B): {per_us_p:.1} us/iter \
1289 ({gflops_p:.1} GFLOP/s), checksum={checksum_p:.3}"
1290 );
1291 eprintln!(
1292 "recorded baselines (docs/KERNEL_PERF.md): ORT 1-thread ~131 us, SimdX86 ~285 us"
1293 );
1294 }
1295
1296 #[test]
1301 #[ignore = "perf probe; run explicitly with --ignored --nocapture"]
1302 fn perf_sgemm_multithread() {
1303 use std::time::Instant;
1304
1305 let (m, n, k) = (32usize, 512usize, 512usize);
1306 let a = seq(m * k, 0.5);
1307 let b = seq(k * n, 1.5);
1308 let flops = 2.0 * m as f64 * n as f64 * k as f64;
1309
1310 for threads in [1usize, 8] {
1311 let pool = rayon::ThreadPoolBuilder::new()
1312 .num_threads(threads)
1313 .build()
1314 .unwrap();
1315 let (per_us, checksum) = pool.install(|| {
1316 let mut c = vec![0.0f32; m * n];
1317 for _ in 0..100 {
1318 sgemm_nn(m, n, k, &a, &b, &mut c);
1319 }
1320 let iters = 5000u32;
1321 let start = Instant::now();
1322 for _ in 0..iters {
1323 sgemm_nn(m, n, k, &a, &b, &mut c);
1324 }
1325 let per_us = start.elapsed().as_secs_f64() * 1e6 / iters as f64;
1326 (per_us, c.iter().copied().sum::<f32>())
1327 });
1328 let gflops = flops / (per_us * 1e3);
1329 eprintln!(
1330 "vendored-MLAS SGEMM 32x512x512 repack-B, {threads} thread(s): {per_us:.1} us/iter \
1331 ({gflops:.1} GFLOP/s), checksum={checksum:.3}"
1332 );
1333 }
1334 eprintln!(
1335 "recorded ORT baselines (docs/KERNEL_PERF.md): 1-thread ~131 us, 8-thread ~28-30 us"
1336 );
1337 }
1338
1339 fn quantize_int4(
1348 weights_nk: &[f32],
1349 n: usize,
1350 k: usize,
1351 block_size: usize,
1352 asymmetric: bool,
1353 ) -> (Vec<u8>, Vec<f32>, Option<Vec<u8>>, Vec<f32>) {
1354 let blocks = k.div_ceil(block_size);
1355 let blob = block_size / 2;
1356 let zp_row = blocks.div_ceil(2);
1357 let mut packed = vec![0u8; n * blocks * blob];
1358 let mut scales = vec![0.0f32; n * blocks];
1359 let mut zps = vec![0u8; n * zp_row];
1360 let mut dequant = vec![0.0f32; n * k];
1361 for row in 0..n {
1362 for block in 0..blocks {
1363 let start = block * block_size;
1364 let end = (start + block_size).min(k);
1365 let values = &weights_nk[row * k + start..row * k + end];
1366 let (scale, zp) = if asymmetric {
1367 let min = values.iter().copied().fold(f32::INFINITY, f32::min);
1368 let max = values.iter().copied().fold(f32::NEG_INFINITY, f32::max);
1369 let scale = ((max - min) / 15.0).max(1e-6);
1370 (scale, (-min / scale).round().clamp(0.0, 15.0) as u8)
1371 } else {
1372 let max_abs = values.iter().map(|v| v.abs()).fold(0.0, f32::max);
1373 ((max_abs / 7.0).max(1e-6), 8u8)
1374 };
1375 scales[row * blocks + block] = scale;
1376 if asymmetric {
1377 zps[row * zp_row + block / 2] |= zp << (4 * (block % 2));
1378 }
1379 for (offset, &value) in values.iter().enumerate() {
1380 let q = (value / scale + zp as f32).round().clamp(0.0, 15.0) as u8;
1381 packed[(row * blocks + block) * blob + offset / 2] |= q << (4 * (offset % 2));
1382 dequant[row * k + start + offset] = (q as f32 - zp as f32) * scale;
1383 }
1384 }
1385 }
1386 (packed, scales, asymmetric.then_some(zps), dequant)
1387 }
1388
1389 fn ref_gemm_nk(
1390 a: &[f32],
1391 w_nk: &[f32],
1392 m: usize,
1393 k: usize,
1394 n: usize,
1395 bias: Option<&[f32]>,
1396 ) -> Vec<f32> {
1397 let mut c = vec![0.0f32; m * n];
1398 for row in 0..m {
1399 for col in 0..n {
1400 let mut acc = bias.map_or(0.0, |b| b[col]);
1401 for depth in 0..k {
1402 acc += a[row * k + depth] * w_nk[col * k + depth];
1403 }
1404 c[row * n + col] = acc;
1405 }
1406 }
1407 c
1408 }
1409
1410 fn check_sqnbit(
1411 comp: SQNBitComputeType,
1412 m: usize,
1413 n: usize,
1414 k: usize,
1415 block_size: usize,
1416 asymmetric: bool,
1417 with_bias: bool,
1418 ) {
1419 let weights: Vec<f32> = (0..n * k).map(|i| (i as f32 * 0.017 + 0.3).sin()).collect();
1420 let (packed_b, scales, zps, dequant) =
1421 quantize_int4(&weights, n, k, block_size, asymmetric);
1422 let a: Vec<f32> = (0..m * k)
1423 .map(|i| ((i as f32 * 0.011 + 0.7).cos()) * 0.5)
1424 .collect();
1425 let bias: Option<Vec<f32>> =
1426 with_bias.then(|| (0..n).map(|i| (i as f32 * 0.03).sin()).collect());
1427
1428 let packed = match SQNBitPackedB::new(
1429 n,
1430 k,
1431 4,
1432 block_size,
1433 comp,
1434 &packed_b,
1435 &scales,
1436 zps.as_deref(),
1437 ) {
1438 Some(p) => p,
1439 None => {
1440 eprintln!(
1441 "SQNBit int4 blk={block_size} comp={comp:?} unavailable on host; skipping"
1442 );
1443 return;
1444 }
1445 };
1446 let mut c = vec![0.0f32; m * n];
1447 sqnbit_gemm(&packed, m, &a, bias.as_deref(), &mut c, true);
1448 let expected = ref_gemm_nk(&a, &dequant, m, k, n, bias.as_deref());
1449 assert_close(
1450 &c,
1451 &expected,
1452 2e-2,
1453 &format!(
1454 "sqnbit {comp:?} m{m} n{n} k{k} blk{block_size} asym{asymmetric} bias{with_bias}"
1455 ),
1456 );
1457 }
1458
1459 #[test]
1460 fn sqnbit_packed_b_is_send_sync() {
1461 assert_send_sync::<SQNBitPackedB>();
1462 }
1463
1464 #[test]
1465 fn sqnbit_int4_compfp32_matches_reference() {
1466 for &blk in &[32usize, 64, 128] {
1467 for &m in &[1usize, 5] {
1468 for &asym in &[false, true] {
1469 check_sqnbit(SQNBitComputeType::Fp32, m, 96, 256, blk, asym, false);
1470 }
1471 }
1472 }
1473 check_sqnbit(SQNBitComputeType::Fp32, 4, 128, 512, 32, false, true);
1474 }
1475
1476 #[test]
1489 fn sqnbit_int4_n_shards_match_full() {
1490 let n = 96usize;
1491 for &(k, block_size) in &[(256usize, 32usize), (256, 64), (256, 128), (384, 64)] {
1494 for &m in &[1usize, 5] {
1495 for &asym in &[false, true] {
1496 for &with_bias in &[false, true] {
1497 let weights: Vec<f32> =
1498 (0..n * k).map(|i| (i as f32 * 0.017 + 0.3).sin()).collect();
1499 let (packed_b, scales, zps, _) =
1500 quantize_int4(&weights, n, k, block_size, asym);
1501 let a: Vec<f32> = (0..m * k)
1502 .map(|i| ((i as f32 * 0.011 + 0.7).cos()) * 0.5)
1503 .collect();
1504 let bias: Option<Vec<f32>> =
1505 with_bias.then(|| (0..n).map(|i| (i as f32 * 0.03).sin()).collect());
1506
1507 let full = match SQNBitPackedB::new(
1508 n,
1509 k,
1510 4,
1511 block_size,
1512 SQNBitComputeType::Fp32,
1513 &packed_b,
1514 &scales,
1515 zps.as_deref(),
1516 ) {
1517 Some(p) => p,
1518 None => {
1519 eprintln!("SQNBit blk={block_size} unavailable; skipping");
1520 return;
1521 }
1522 };
1523 let mut c_full = vec![0.0f32; m * n];
1524 sqnbit_gemm(&full, m, &a, bias.as_deref(), &mut c_full, true);
1525
1526 let blocks = k.div_ceil(block_size);
1527 let blob = block_size / 2;
1528 let zp_row = blocks.div_ceil(2);
1529 let shards: &[(usize, usize)] = &[(0, 17), (17, 30), (47, 1), (48, 48)];
1532 for &mt in &[false, true] {
1535 let mut c_shard = vec![0.0f32; m * n];
1536 for &(start, len) in shards {
1537 let pb =
1538 &packed_b[start * blocks * blob..(start + len) * blocks * blob];
1539 let sc = &scales[start * blocks..(start + len) * blocks];
1540 let zp = zps
1541 .as_deref()
1542 .map(|z| &z[start * zp_row..(start + len) * zp_row]);
1543 let packed = SQNBitPackedB::new(
1544 len,
1545 k,
1546 4,
1547 block_size,
1548 SQNBitComputeType::Fp32,
1549 pb,
1550 sc,
1551 zp,
1552 )
1553 .expect("shard packs when the full weight packs");
1554 let bias_shard = bias.as_deref().map(|b| &b[start..start + len]);
1555 unsafe {
1558 sqnbit_gemm_into(
1559 &packed,
1560 m,
1561 &a,
1562 bias_shard,
1563 c_shard.as_mut_ptr().add(start),
1564 n,
1565 mt,
1566 );
1567 }
1568 }
1569 assert_close(
1573 &c_shard,
1574 &c_full,
1575 1e-3,
1576 &format!(
1577 "N-sharded (multithread={mt}) vs full: \
1578 k{k} blk{block_size} m{m} asym{asym} bias{with_bias}"
1579 ),
1580 );
1581 }
1582 }
1583 }
1584 }
1585 }
1586 }
1587
1588 #[test]
1601 fn sqnbit_int4_tile_aligned_shards_are_bit_exact() {
1602 let n = 176usize;
1604 let mut any_mid_tile_drift = false;
1605 for &(k, block_size) in &[(256usize, 128usize), (512, 32), (256, 64)] {
1606 for &asym in &[false, true] {
1607 let weights: Vec<f32> =
1608 (0..n * k).map(|i| (i as f32 * 0.017 + 0.3).sin()).collect();
1609 let (packed_b, scales, zps, _) = quantize_int4(&weights, n, k, block_size, asym);
1610 let a: Vec<f32> = (0..k)
1611 .map(|i| ((i as f32 * 0.011 + 0.7).cos()) * 0.5)
1612 .collect();
1613
1614 let full = match SQNBitPackedB::new(
1615 n,
1616 k,
1617 4,
1618 block_size,
1619 SQNBitComputeType::Fp32,
1620 &packed_b,
1621 &scales,
1622 zps.as_deref(),
1623 ) {
1624 Some(p) => p,
1625 None => {
1626 eprintln!("SQNBit blk={block_size} unavailable; skipping");
1627 return;
1628 }
1629 };
1630 let mut c_full = vec![0.0f32; n];
1631 sqnbit_gemm(&full, 1, &a, None, &mut c_full, false);
1632
1633 let blocks = k.div_ceil(block_size);
1634 let blob = block_size / 2;
1635 let zp_row = blocks.div_ceil(2);
1636 let run_shards = |shards: &[(usize, usize)]| -> Vec<f32> {
1637 let mut c = vec![0.0f32; n];
1638 for &(start, len) in shards {
1639 let pb = &packed_b[start * blocks * blob..(start + len) * blocks * blob];
1640 let sc = &scales[start * blocks..(start + len) * blocks];
1641 let zp = zps
1642 .as_deref()
1643 .map(|z| &z[start * zp_row..(start + len) * zp_row]);
1644 let packed = SQNBitPackedB::new(
1645 len,
1646 k,
1647 4,
1648 block_size,
1649 SQNBitComputeType::Fp32,
1650 pb,
1651 sc,
1652 zp,
1653 )
1654 .expect("shard packs when the full weight packs");
1655 unsafe {
1657 sqnbit_gemm_into(
1658 &packed,
1659 1,
1660 &a,
1661 None,
1662 c.as_mut_ptr().add(start),
1663 n,
1664 false,
1665 );
1666 }
1667 }
1668 c
1669 };
1670
1671 let aligned: &[(usize, usize)] = &[(0, 16), (16, 48), (64, 64), (128, 48)];
1673 assert_eq!(aligned.iter().map(|&(_, l)| l).sum::<usize>(), n);
1674 let c_aligned = run_shards(aligned);
1675 let aligned_bits_match = c_aligned
1676 .iter()
1677 .zip(&c_full)
1678 .all(|(a, b)| a.to_bits() == b.to_bits());
1679 assert!(
1680 aligned_bits_match,
1681 "k{k} blk{block_size} asym{asym}: 16-aligned N shards must be \
1682 bit-identical to full-width, but differ"
1683 );
1684
1685 let mid_tile: &[(usize, usize)] = &[(0, 17), (17, 30), (47, 1), (48, 128)];
1691 assert_eq!(mid_tile.iter().map(|&(_, l)| l).sum::<usize>(), n);
1692 let c_mid = run_shards(mid_tile);
1693 any_mid_tile_drift |= c_mid
1694 .iter()
1695 .zip(&c_full)
1696 .any(|(a, b)| a.to_bits() != b.to_bits());
1697 }
1698 }
1699 assert!(
1700 any_mid_tile_drift,
1701 "expected at least one mid-tile-split N shard layout to drift from full-width \
1702 (non-vacuous guard); none did, so the 16-alignment fix is untested on this host"
1703 );
1704 }
1705
1706 #[test]
1707 fn sqnbit_int4_compint8_matches_reference() {
1708 if !std::arch::is_x86_feature_detected!("avx512f") {
1712 eprintln!(
1713 "skipping SQNBit int4 CompInt8 reference check: AVX-512F is unavailable; \
1714 AVX2 CompInt8 SQNBit M=1 asymmetric-weight bug: microsoft/onnxruntime#29853"
1715 );
1716 return;
1717 }
1718 #[cfg(target_arch = "x86_64")]
1741 let host_has_avx512 = std::arch::is_x86_feature_detected!("avx512f")
1742 && std::arch::is_x86_feature_detected!("avx512bw")
1743 && std::arch::is_x86_feature_detected!("avx512dq")
1744 && std::arch::is_x86_feature_detected!("avx512vl");
1745 #[cfg(not(target_arch = "x86_64"))]
1746 let host_has_avx512 = false;
1747 for &blk in &[32usize, 64, 128] {
1748 for &m in &[1usize, 8] {
1749 for &asym in &[false, true] {
1750 if m == 1 && asym && !host_has_avx512 {
1751 eprintln!(
1752 "skipping MLAS-broken AVX2 M=1 asymmetric CompInt8 blk{blk} \
1753 (production uses the hand int8 kernel here)"
1754 );
1755 continue;
1756 }
1757 check_sqnbit(SQNBitComputeType::Int8, m, 96, 256, blk, asym, false);
1758 }
1759 }
1760 }
1761 check_sqnbit(SQNBitComputeType::Int8, 4, 128, 512, 32, false, true);
1762 }
1763
1764 #[test]
1768 #[ignore = "perf probe; run explicitly with --ignored --nocapture"]
1769 fn perf_sqnbit() {
1770 use std::time::Instant;
1771 for &(k, n) in &[(2048usize, 2048usize), (4096, 11008)] {
1772 let weights: Vec<f32> = (0..n * k).map(|i| (i as f32 * 0.017).sin()).collect();
1773 let (packed_b, scales, _zps, _d) = quantize_int4(&weights, n, k, 32, false);
1774 for comp in [SQNBitComputeType::Fp32, SQNBitComputeType::Int8] {
1775 let packed = match SQNBitPackedB::new(n, k, 4, 32, comp, &packed_b, &scales, None) {
1776 Some(p) => p,
1777 None => continue,
1778 };
1779 for &m in &[1usize, 32] {
1780 let a: Vec<f32> = (0..m * k).map(|i| (i as f32 * 0.011).cos()).collect();
1781 for threads in [1usize, 8] {
1782 let pool = rayon::ThreadPoolBuilder::new()
1783 .num_threads(threads)
1784 .build()
1785 .unwrap();
1786 let per_us = pool.install(|| {
1787 let mut c = vec![0.0f32; m * n];
1788 for _ in 0..20 {
1789 sqnbit_gemm(&packed, m, &a, None, &mut c, true);
1790 }
1791 let iters = 200u32;
1792 let start = Instant::now();
1793 for _ in 0..iters {
1794 sqnbit_gemm(&packed, m, &a, None, &mut c, true);
1795 }
1796 start.elapsed().as_secs_f64() * 1e6 / iters as f64
1797 });
1798 eprintln!(
1799 "SQNBit int4 {comp:?} K={k} N={n} M={m} {threads}t: {per_us:.1} us/iter"
1800 );
1801 }
1802 }
1803 }
1804 }
1805 }
1806
1807 fn round_up(value: usize, multiple: usize) -> usize {
1808 value.div_ceil(multiple) * multiple
1809 }
1810
1811 #[allow(clippy::too_many_arguments)]
1813 fn ref_conv_nchw(
1814 input: &[f32],
1815 filter: &[f32],
1816 bias: Option<&[f32]>,
1817 n: usize,
1818 cin: usize,
1819 hin: usize,
1820 win: usize,
1821 cout: usize,
1822 kh: usize,
1823 kw: usize,
1824 pad: [usize; 4],
1825 stride: [usize; 2],
1826 group: usize,
1827 ) -> (Vec<f32>, usize, usize) {
1828 let hout = (hin + pad[0] + pad[2] - kh) / stride[0] + 1;
1829 let wout = (win + pad[1] + pad[3] - kw) / stride[1] + 1;
1830 let cin_g = cin / group;
1831 let cout_g = cout / group;
1832 let mut out = vec![0.0f32; n * cout * hout * wout];
1833 for ni in 0..n {
1834 for oc in 0..cout {
1835 let g = oc / cout_g;
1836 for oy in 0..hout {
1837 for ox in 0..wout {
1838 let mut acc = bias.map_or(0.0, |b| b[oc]);
1839 for icg in 0..cin_g {
1840 let ic = g * cin_g + icg;
1841 for ky in 0..kh {
1842 let iy = oy * stride[0] + ky;
1843 if iy < pad[0] || iy - pad[0] >= hin {
1844 continue;
1845 }
1846 let iy = iy - pad[0];
1847 for kx in 0..kw {
1848 let ix = ox * stride[1] + kx;
1849 if ix < pad[1] || ix - pad[1] >= win {
1850 continue;
1851 }
1852 let ix = ix - pad[1];
1853 let iv = input[((ni * cin + ic) * hin + iy) * win + ix];
1854 let fv = filter[(((oc * cin_g) + icg) * kh + ky) * kw + kx];
1855 acc += iv * fv;
1856 }
1857 }
1858 }
1859 out[((ni * cout + oc) * hout + oy) * wout + ox] = acc;
1860 }
1861 }
1862 }
1863 }
1864 (out, hout, wout)
1865 }
1866
1867 #[allow(clippy::too_many_arguments)]
1868 fn run_nchwc_group1(
1869 input: &[f32],
1870 filter: &[f32],
1871 bias: Option<&[f32]>,
1872 n: usize,
1873 cin: usize,
1874 hin: usize,
1875 win: usize,
1876 cout: usize,
1877 kh: usize,
1878 kw: usize,
1879 pad: [usize; 4],
1880 stride: [usize; 2],
1881 ) -> Vec<f32> {
1882 let block = nchwc_block_size();
1883 let hout = (hin + pad[0] + pad[2] - kh) / stride[0] + 1;
1884 let wout = (win + pad[1] + pad[3] - kw) / stride[1] + 1;
1885 let nchwc_cout = round_up(cout, block);
1886 let filter_shape = [cout as i64, cin as i64, kh as i64, kw as i64];
1887
1888 let reorder_input = cin >= block;
1889 let (packed_filter, conv_input, in_channels_for_shape) = if reorder_input {
1890 let nchwc_cin = round_up(cin, block);
1891 let mut pf = vec![0.0f32; nchwc_cout * nchwc_cin * kh * kw];
1892 nchwc_reorder_filter_bibo(&filter_shape, filter, &mut pf);
1893 let mut blocked = vec![0.0f32; n * nchwc_cin * hin * win];
1894 for ni in 0..n {
1895 nchwc_reorder_input_nchw(
1896 &input[ni * cin * hin * win..(ni + 1) * cin * hin * win],
1897 &mut blocked[ni * nchwc_cin * hin * win..(ni + 1) * nchwc_cin * hin * win],
1898 cin,
1899 hin * win,
1900 );
1901 }
1902 (pf, blocked, nchwc_cin)
1903 } else {
1904 let mut pf = vec![0.0f32; nchwc_cout * cin * kh * kw];
1905 nchwc_reorder_filter_bo(&filter_shape, filter, &mut pf);
1906 (pf, input.to_vec(), cin)
1907 };
1908
1909 let padded_bias = bias.map(|b| {
1910 let mut pb = vec![0.0f32; nchwc_cout];
1911 pb[..cout].copy_from_slice(b);
1912 pb
1913 });
1914
1915 let mut blocked_out = vec![0.0f32; n * nchwc_cout * hout * wout];
1916 nchwc_conv(
1917 &[
1918 n as i64,
1919 in_channels_for_shape as i64,
1920 hin as i64,
1921 win as i64,
1922 ],
1923 &[kh as i64, kw as i64],
1924 &[1, 1],
1925 &[pad[0] as i64, pad[1] as i64, pad[2] as i64, pad[3] as i64],
1926 &[stride[0] as i64, stride[1] as i64],
1927 &[n as i64, nchwc_cout as i64, hout as i64, wout as i64],
1928 1,
1929 &conv_input,
1930 &packed_filter,
1931 padded_bias.as_deref(),
1932 &mut blocked_out,
1933 NchwcActivation::IDENTITY,
1934 true,
1935 );
1936
1937 let mut out = vec![0.0f32; n * cout * hout * wout];
1938 nchwc_reorder_output_nchw(
1939 &[n as i64, cout as i64, hout as i64, wout as i64],
1940 &blocked_out,
1941 &mut out,
1942 );
1943 out
1944 }
1945
1946 fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 {
1947 a.iter()
1948 .zip(b)
1949 .fold(0.0f32, |m, (x, y)| m.max((x - y).abs()))
1950 }
1951
1952 #[test]
1953 fn nchwc_block_size_is_supported() {
1954 assert!(nchwc_block_size() >= 8, "block size {}", nchwc_block_size());
1956 }
1957
1958 #[test]
1959 fn nchwc_conv_pointwise_matches_reference() {
1960 let block = nchwc_block_size();
1961 let (n, cin, hin, win, cout) = (1, 2 * block, 7, 7, 3 * block);
1962 let input: Vec<f32> = (0..n * cin * hin * win)
1963 .map(|i| ((i % 13) as f32 - 6.0) * 0.1)
1964 .collect();
1965 let filter: Vec<f32> = (0..cout * cin)
1966 .map(|i| ((i % 7) as f32 - 3.0) * 0.05)
1967 .collect();
1968 let bias: Vec<f32> = (0..cout).map(|i| (i as f32) * 0.01).collect();
1969 let (want, _, _) = ref_conv_nchw(
1970 &input,
1971 &filter,
1972 Some(&bias),
1973 n,
1974 cin,
1975 hin,
1976 win,
1977 cout,
1978 1,
1979 1,
1980 [0; 4],
1981 [1, 1],
1982 1,
1983 );
1984 let got = run_nchwc_group1(
1985 &input,
1986 &filter,
1987 Some(&bias),
1988 n,
1989 cin,
1990 hin,
1991 win,
1992 cout,
1993 1,
1994 1,
1995 [0; 4],
1996 [1, 1],
1997 );
1998 assert!(
1999 max_abs_diff(&want, &got) < 1e-4,
2000 "diff {}",
2001 max_abs_diff(&want, &got)
2002 );
2003 }
2004
2005 #[test]
2006 fn nchwc_conv_3x3_blocked_matches_reference() {
2007 let block = nchwc_block_size();
2008 let (n, cin, hin, win, cout) = (1, block, 9, 9, block);
2009 let input: Vec<f32> = (0..n * cin * hin * win)
2010 .map(|i| ((i % 17) as f32 - 8.0) * 0.05)
2011 .collect();
2012 let filter: Vec<f32> = (0..cout * cin * 9)
2013 .map(|i| ((i % 11) as f32 - 5.0) * 0.03)
2014 .collect();
2015 let (want, _, _) = ref_conv_nchw(
2016 &input,
2017 &filter,
2018 None,
2019 n,
2020 cin,
2021 hin,
2022 win,
2023 cout,
2024 3,
2025 3,
2026 [1, 1, 1, 1],
2027 [1, 1],
2028 1,
2029 );
2030 let got = run_nchwc_group1(
2031 &input,
2032 &filter,
2033 None,
2034 n,
2035 cin,
2036 hin,
2037 win,
2038 cout,
2039 3,
2040 3,
2041 [1, 1, 1, 1],
2042 [1, 1],
2043 );
2044 assert!(
2045 max_abs_diff(&want, &got) < 1e-4,
2046 "diff {}",
2047 max_abs_diff(&want, &got)
2048 );
2049 }
2050
2051 #[test]
2052 fn nchwc_conv_first_layer_nchw_input_matches_reference() {
2053 let block = nchwc_block_size();
2055 let (n, cin, hin, win, cout) = (1, 3, 16, 16, block + block / 2);
2056 let input: Vec<f32> = (0..n * cin * hin * win)
2057 .map(|i| ((i % 19) as f32 - 9.0) * 0.04)
2058 .collect();
2059 let filter: Vec<f32> = (0..cout * cin * 9)
2060 .map(|i| ((i % 13) as f32 - 6.0) * 0.02)
2061 .collect();
2062 let bias: Vec<f32> = (0..cout).map(|i| (i as f32) * 0.02 - 0.3).collect();
2063 let (want, _, _) = ref_conv_nchw(
2064 &input,
2065 &filter,
2066 Some(&bias),
2067 n,
2068 cin,
2069 hin,
2070 win,
2071 cout,
2072 3,
2073 3,
2074 [1, 1, 1, 1],
2075 [2, 2],
2076 1,
2077 );
2078 let got = run_nchwc_group1(
2079 &input,
2080 &filter,
2081 Some(&bias),
2082 n,
2083 cin,
2084 hin,
2085 win,
2086 cout,
2087 3,
2088 3,
2089 [1, 1, 1, 1],
2090 [2, 2],
2091 );
2092 assert!(
2093 max_abs_diff(&want, &got) < 1e-4,
2094 "diff {}",
2095 max_abs_diff(&want, &got)
2096 );
2097 }
2098
2099 #[test]
2100 fn nchwc_conv_depthwise_matches_reference() {
2101 let block = nchwc_block_size();
2103 let channels = 2 * block; let (n, hin, win) = (1, 8, 8);
2105 let input: Vec<f32> = (0..n * channels * hin * win)
2106 .map(|i| ((i % 15) as f32 - 7.0) * 0.06)
2107 .collect();
2108 let filter: Vec<f32> = (0..channels * 9)
2110 .map(|i| ((i % 7) as f32 - 3.0) * 0.05)
2111 .collect();
2112 let bias: Vec<f32> = (0..channels).map(|i| (i as f32) * 0.01).collect();
2113 let (want, hout, wout) = ref_conv_nchw(
2114 &input,
2115 &filter,
2116 Some(&bias),
2117 n,
2118 channels,
2119 hin,
2120 win,
2121 channels,
2122 3,
2123 3,
2124 [1, 1, 1, 1],
2125 [1, 1],
2126 channels,
2127 );
2128
2129 let nchwc_ch = round_up(channels, block);
2130 let mut pf = vec![0.0f32; nchwc_ch * 9];
2131 nchwc_reorder_filter_bo(&[channels as i64, 1, 3, 3], &filter, &mut pf);
2132 let mut blocked_in = vec![0.0f32; n * nchwc_ch * hin * win];
2133 nchwc_reorder_input_nchw(&input, &mut blocked_in, channels, hin * win);
2134 let mut padded_bias = vec![0.0f32; nchwc_ch];
2135 padded_bias[..channels].copy_from_slice(&bias);
2136 let mut blocked_out = vec![0.0f32; n * nchwc_ch * hout * wout];
2137 nchwc_conv(
2138 &[n as i64, nchwc_ch as i64, hin as i64, win as i64],
2139 &[3, 3],
2140 &[1, 1],
2141 &[1, 1, 1, 1],
2142 &[1, 1],
2143 &[n as i64, nchwc_ch as i64, hout as i64, wout as i64],
2144 nchwc_ch, &blocked_in,
2146 &pf,
2147 Some(&padded_bias),
2148 &mut blocked_out,
2149 NchwcActivation::IDENTITY,
2150 true,
2151 );
2152 let mut got = vec![0.0f32; n * channels * hout * wout];
2153 nchwc_reorder_output_nchw(
2154 &[n as i64, channels as i64, hout as i64, wout as i64],
2155 &blocked_out,
2156 &mut got,
2157 );
2158 assert!(
2159 max_abs_diff(&want, &got) < 1e-4,
2160 "diff {}",
2161 max_abs_diff(&want, &got)
2162 );
2163 }
2164
2165 #[test]
2166 fn nchwc_conv_relu_activation_matches_reference() {
2167 let block = nchwc_block_size();
2168 let (n, cin, hin, win, cout) = (1, block, 5, 5, block);
2169 let input: Vec<f32> = (0..n * cin * hin * win)
2170 .map(|i| ((i % 9) as f32 - 4.0) * 0.2)
2171 .collect();
2172 let filter: Vec<f32> = (0..cout * cin)
2173 .map(|i| ((i % 5) as f32 - 2.0) * 0.1)
2174 .collect();
2175 let (mut want, _, _) = ref_conv_nchw(
2176 &input,
2177 &filter,
2178 None,
2179 n,
2180 cin,
2181 hin,
2182 win,
2183 cout,
2184 1,
2185 1,
2186 [0; 4],
2187 [1, 1],
2188 1,
2189 );
2190 for v in &mut want {
2191 *v = v.max(0.0);
2192 }
2193 let nchwc_cout = round_up(cout, block);
2195 let nchwc_cin = round_up(cin, block);
2196 let mut pf = vec![0.0f32; nchwc_cout * nchwc_cin];
2197 nchwc_reorder_filter_bibo(&[cout as i64, cin as i64, 1, 1], &filter, &mut pf);
2198 let mut blocked_in = vec![0.0f32; n * nchwc_cin * hin * win];
2199 nchwc_reorder_input_nchw(&input, &mut blocked_in, cin, hin * win);
2200 let mut blocked_out = vec![0.0f32; n * nchwc_cout * hin * win];
2201 nchwc_conv(
2202 &[n as i64, nchwc_cin as i64, hin as i64, win as i64],
2203 &[1, 1],
2204 &[1, 1],
2205 &[0; 4],
2206 &[1, 1],
2207 &[n as i64, nchwc_cout as i64, hin as i64, win as i64],
2208 1,
2209 &blocked_in,
2210 &pf,
2211 None,
2212 &mut blocked_out,
2213 NchwcActivation::RELU,
2214 true,
2215 );
2216 let mut got = vec![0.0f32; n * cout * hin * win];
2217 nchwc_reorder_output_nchw(
2218 &[n as i64, cout as i64, hin as i64, win as i64],
2219 &blocked_out,
2220 &mut got,
2221 );
2222 assert!(
2223 max_abs_diff(&want, &got) < 1e-4,
2224 "diff {}",
2225 max_abs_diff(&want, &got)
2226 );
2227 }
2228
2229 #[test]
2230 fn nchwc_pool_max_and_average_match_reference() {
2231 let block = nchwc_block_size();
2232 let channels = block + block / 2; let (n, hin, win) = (1, 8, 8);
2234 let (kh, kw) = (2usize, 2usize);
2235 let (sh, sw) = (2usize, 2usize);
2236 let hout = (hin - kh) / sh + 1;
2237 let wout = (win - kw) / sw + 1;
2238 let input: Vec<f32> = (0..n * channels * hin * win)
2239 .map(|i| ((i % 23) as f32 - 11.0) * 0.13)
2240 .collect();
2241
2242 let nchwc_ch = round_up(channels, block);
2243 let mut blocked_in = vec![0.0f32; n * nchwc_ch * hin * win];
2244 nchwc_reorder_input_nchw(&input, &mut blocked_in, channels, hin * win);
2245
2246 for kind in [PoolKind::Maximum, PoolKind::AverageIncludePad] {
2247 let mut blocked_out = vec![0.0f32; n * nchwc_ch * hout * wout];
2248 nchwc_pool(
2249 kind,
2250 &[n as i64, nchwc_ch as i64, hin as i64, win as i64],
2251 &[kh as i64, kw as i64],
2252 &[1, 1],
2253 &[0, 0, 0, 0],
2254 &[sh as i64, sw as i64],
2255 &[n as i64, nchwc_ch as i64, hout as i64, wout as i64],
2256 &blocked_in,
2257 &mut blocked_out,
2258 );
2259 let mut got = vec![0.0f32; n * channels * hout * wout];
2260 nchwc_reorder_output_nchw(
2261 &[n as i64, channels as i64, hout as i64, wout as i64],
2262 &blocked_out,
2263 &mut got,
2264 );
2265
2266 let mut want = vec![0.0f32; n * channels * hout * wout];
2267 for c in 0..channels {
2268 for oh in 0..hout {
2269 for ow in 0..wout {
2270 let mut acc = if matches!(kind, PoolKind::Maximum) {
2271 f32::NEG_INFINITY
2272 } else {
2273 0.0
2274 };
2275 for ky in 0..kh {
2276 for kx in 0..kw {
2277 let ih = oh * sh + ky;
2278 let iw = ow * sw + kx;
2279 let v = input[((c * hin) + ih) * win + iw];
2280 if matches!(kind, PoolKind::Maximum) {
2281 acc = acc.max(v);
2282 } else {
2283 acc += v;
2284 }
2285 }
2286 }
2287 if !matches!(kind, PoolKind::Maximum) {
2288 acc /= (kh * kw) as f32;
2289 }
2290 want[((c * hout) + oh) * wout + ow] = acc;
2291 }
2292 }
2293 }
2294 assert!(
2295 max_abs_diff(&want, &got) < 1e-4,
2296 "kind {kind:?} diff {}",
2297 max_abs_diff(&want, &got)
2298 );
2299 }
2300 }
2301
2302 #[test]
2308 fn nchwc_reorder_round_trip_is_identity() {
2309 let block = nchwc_block_size();
2310 for &channels in &[block, block + 4] {
2313 let (n, h, w) = (1usize, 5usize, 7usize);
2314 let plane = h * w;
2315 let input: Vec<f32> = (0..n * channels * plane)
2316 .map(|i| ((i % 17) as f32 - 8.0) * 0.07)
2317 .collect();
2318
2319 let nchwc_ch = round_up(channels, block);
2320 let mut blocked = vec![7.0f32; n * nchwc_ch * plane]; nchwc_reorder_input_nchw(&input, &mut blocked, channels, plane);
2322
2323 let mut back = vec![0.0f32; n * channels * plane];
2324 nchwc_reorder_output_nchw(
2325 &[n as i64, channels as i64, h as i64, w as i64],
2326 &blocked,
2327 &mut back,
2328 );
2329
2330 assert_eq!(back, input, "round-trip mismatch for channels={channels}");
2331 }
2332 }
2333}