1use super::detect::{simd_level, SimdLevel};
2use super::kernels::*;
3use crate::core::{tropical_gemm_inner_with_workspace, GemmWorkspace, TilingParams, Transpose};
4use crate::types::{
5 TropicalAndOr, TropicalBitwise, TropicalMaxMul, TropicalMaxPlus, TropicalMinPlus,
6 TropicalSemiring,
7};
8
9pub unsafe fn tropical_gemm_dispatch<T: TropicalSemiring + KernelDispatch>(
14 m: usize,
15 n: usize,
16 k: usize,
17 a: *const T::Scalar,
18 lda: usize,
19 trans_a: Transpose,
20 b: *const T::Scalar,
21 ldb: usize,
22 trans_b: Transpose,
23 c: *mut T,
24 ldc: usize,
25) {
26 T::dispatch_gemm(m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc);
27}
28
29#[allow(clippy::too_many_arguments)]
34pub unsafe fn tropical_gemm_with_argmax_dispatch<
35 T: KernelDispatch + crate::TropicalWithArgmax<Index = u32>,
36>(
37 m: usize,
38 n: usize,
39 k: usize,
40 a: *const T::Scalar,
41 lda: usize,
42 trans_a: Transpose,
43 b: *const T::Scalar,
44 ldb: usize,
45 trans_b: Transpose,
46 result: &mut crate::core::GemmWithArgmax<T>,
47) {
48 T::dispatch_gemm_with_argmax(m, n, k, a, lda, trans_a, b, ldb, trans_b, result);
49}
50
51pub trait KernelDispatch: TropicalSemiring {
53 unsafe fn dispatch_gemm(
55 m: usize,
56 n: usize,
57 k: usize,
58 a: *const Self::Scalar,
59 lda: usize,
60 trans_a: Transpose,
61 b: *const Self::Scalar,
62 ldb: usize,
63 trans_b: Transpose,
64 c: *mut Self,
65 ldc: usize,
66 );
67 #[allow(clippy::too_many_arguments)]
72 unsafe fn dispatch_gemm_with_argmax(
73 m: usize,
74 n: usize,
75 k: usize,
76 a: *const Self::Scalar,
77 lda: usize,
78 trans_a: Transpose,
79 b: *const Self::Scalar,
80 ldb: usize,
81 trans_b: Transpose,
82 result: &mut crate::core::GemmWithArgmax<Self>,
83 ) where
84 Self: crate::TropicalWithArgmax<Index = u32>,
85 {
86 crate::core::tropical_gemm_with_argmax_portable(
87 m, n, k, a, lda, trans_a, b, ldb, trans_b, result,
88 );
89 }
90 #[allow(clippy::too_many_arguments)]
96 unsafe fn dispatch_gemm_with_workspace(
97 m: usize,
98 n: usize,
99 k: usize,
100 a: *const Self::Scalar,
101 lda: usize,
102 trans_a: Transpose,
103 b: *const Self::Scalar,
104 ldb: usize,
105 trans_b: Transpose,
106 c: *mut Self,
107 ldc: usize,
108 _workspace: &mut GemmWorkspace<Self::Scalar>,
109 ) {
110 Self::dispatch_gemm(m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc);
111 }
112
113 #[allow(clippy::too_many_arguments)]
118 unsafe fn dispatch_gemm_with_argmax_with_workspace(
119 m: usize,
120 n: usize,
121 k: usize,
122 a: *const Self::Scalar,
123 lda: usize,
124 trans_a: Transpose,
125 b: *const Self::Scalar,
126 ldb: usize,
127 trans_b: Transpose,
128 result: &mut crate::core::GemmWithArgmax<Self>,
129 workspace: &mut GemmWorkspace<Self::Scalar>,
130 ) where
131 Self: crate::TropicalWithArgmax<Index = u32>,
132 {
133 crate::core::tropical_gemm_with_argmax_inner_with_workspace(
134 m,
135 n,
136 k,
137 a,
138 lda,
139 trans_a,
140 b,
141 ldb,
142 trans_b,
143 result,
144 &TilingParams::PORTABLE,
145 &crate::core::PortableMicrokernel,
146 workspace,
147 );
148 }
149}
150
151macro_rules! dispatch_with_local_workspace {
154 () => {
155 unsafe fn dispatch_gemm(
156 m: usize,
157 n: usize,
158 k: usize,
159 a: *const Self::Scalar,
160 lda: usize,
161 trans_a: Transpose,
162 b: *const Self::Scalar,
163 ldb: usize,
164 trans_b: Transpose,
165 c: *mut Self,
166 ldc: usize,
167 ) {
168 Self::dispatch_gemm_with_workspace(
169 m,
170 n,
171 k,
172 a,
173 lda,
174 trans_a,
175 b,
176 ldb,
177 trans_b,
178 c,
179 ldc,
180 &mut GemmWorkspace::new(),
181 );
182 }
183 };
184}
185
186macro_rules! argmax_dispatch {
188 ($scalar:ty, $dispatch:ident, $min:expr, $mul:expr) => {
189 unsafe fn dispatch_gemm_with_argmax(
190 m: usize,
191 n: usize,
192 k: usize,
193 a: *const $scalar,
194 lda: usize,
195 trans_a: Transpose,
196 b: *const $scalar,
197 ldb: usize,
198 trans_b: Transpose,
199 result: &mut crate::core::GemmWithArgmax<Self>,
200 ) {
201 Self::dispatch_gemm_with_argmax_with_workspace(
202 m,
203 n,
204 k,
205 a,
206 lda,
207 trans_a,
208 b,
209 ldb,
210 trans_b,
211 result,
212 &mut GemmWorkspace::new(),
213 );
214 }
215 unsafe fn dispatch_gemm_with_argmax_with_workspace(
216 m: usize,
217 n: usize,
218 k: usize,
219 a: *const $scalar,
220 lda: usize,
221 trans_a: Transpose,
222 b: *const $scalar,
223 ldb: usize,
224 trans_b: Transpose,
225 result: &mut crate::core::GemmWithArgmax<Self>,
226 workspace: &mut GemmWorkspace<Self::Scalar>,
227 ) {
228 super::argmax::$dispatch::<Self, $min, $mul>(
229 m, n, k, a, lda, trans_a, b, ldb, trans_b, result, workspace,
230 );
231 }
232 };
233}
234
235impl KernelDispatch for TropicalMaxPlus<f32> {
236 argmax_dispatch!(f32, dispatch_f32, false, false);
237 dispatch_with_local_workspace!();
238 unsafe fn dispatch_gemm_with_workspace(
239 m: usize,
240 n: usize,
241 k: usize,
242 a: *const f32,
243 lda: usize,
244 trans_a: Transpose,
245 b: *const f32,
246 ldb: usize,
247 trans_b: Transpose,
248 c: *mut Self,
249 ldc: usize,
250 workspace: &mut GemmWorkspace<Self::Scalar>,
251 ) {
252 match simd_level() {
253 #[cfg(target_arch = "x86_64")]
254 SimdLevel::Avx2 | SimdLevel::Avx512 => {
255 let kernel = Avx2MaxPlusF32Kernel;
256 let params = TilingParams::F32_AVX2;
257 tropical_gemm_inner_with_workspace::<Self, _>(
258 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
259 );
260 }
261 #[cfg(target_arch = "aarch64")]
262 SimdLevel::Neon => {
263 let kernel = NeonMaxPlusF32Kernel;
264 let params = TilingParams::new(128, 128, 256, 4, 4);
265 tropical_gemm_inner_with_workspace::<Self, _>(
266 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
267 );
268 }
269 _ => {
270 let kernel = PortableKernel;
271 let params = TilingParams::PORTABLE;
272 tropical_gemm_inner_with_workspace::<Self, _>(
273 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
274 );
275 }
276 }
277 }
278}
279
280impl KernelDispatch for TropicalMaxPlus<f64> {
281 argmax_dispatch!(f64, dispatch_f64, false, false);
282 dispatch_with_local_workspace!();
283 unsafe fn dispatch_gemm_with_workspace(
284 m: usize,
285 n: usize,
286 k: usize,
287 a: *const f64,
288 lda: usize,
289 trans_a: Transpose,
290 b: *const f64,
291 ldb: usize,
292 trans_b: Transpose,
293 c: *mut Self,
294 ldc: usize,
295 workspace: &mut GemmWorkspace<Self::Scalar>,
296 ) {
297 match simd_level() {
298 #[cfg(target_arch = "x86_64")]
299 SimdLevel::Avx2 | SimdLevel::Avx512 => {
300 let kernel = Avx2MaxPlusF64Kernel;
301 let params = TilingParams::F64_AVX2;
302 tropical_gemm_inner_with_workspace::<Self, _>(
303 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
304 );
305 }
306 #[cfg(target_arch = "aarch64")]
307 SimdLevel::Neon => {
308 let kernel = NeonMaxPlusF64Kernel;
309 let params = TilingParams::new(64, 64, 128, 2, 2);
310 tropical_gemm_inner_with_workspace::<Self, _>(
311 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
312 );
313 }
314 _ => {
315 let kernel = PortableKernel;
316 let params = TilingParams::PORTABLE;
317 tropical_gemm_inner_with_workspace::<Self, _>(
318 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
319 );
320 }
321 }
322 }
323}
324
325impl KernelDispatch for TropicalMinPlus<f32> {
326 argmax_dispatch!(f32, dispatch_f32, true, false);
327 dispatch_with_local_workspace!();
328 unsafe fn dispatch_gemm_with_workspace(
329 m: usize,
330 n: usize,
331 k: usize,
332 a: *const f32,
333 lda: usize,
334 trans_a: Transpose,
335 b: *const f32,
336 ldb: usize,
337 trans_b: Transpose,
338 c: *mut Self,
339 ldc: usize,
340 workspace: &mut GemmWorkspace<Self::Scalar>,
341 ) {
342 match simd_level() {
343 #[cfg(target_arch = "x86_64")]
344 SimdLevel::Avx2 | SimdLevel::Avx512 => {
345 let kernel = Avx2MinPlusF32Kernel;
346 let params = TilingParams::F32_AVX2;
347 tropical_gemm_inner_with_workspace::<Self, _>(
348 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
349 );
350 }
351 #[cfg(target_arch = "aarch64")]
352 SimdLevel::Neon => {
353 let kernel = NeonMinPlusF32Kernel;
354 let params = TilingParams::new(128, 128, 256, 4, 4);
355 tropical_gemm_inner_with_workspace::<Self, _>(
356 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
357 );
358 }
359 _ => {
360 let kernel = PortableKernel;
361 let params = TilingParams::PORTABLE;
362 tropical_gemm_inner_with_workspace::<Self, _>(
363 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
364 );
365 }
366 }
367 }
368}
369
370impl KernelDispatch for TropicalMaxMul<f32> {
371 argmax_dispatch!(f32, dispatch_f32, false, true);
372 dispatch_with_local_workspace!();
373 unsafe fn dispatch_gemm_with_workspace(
374 m: usize,
375 n: usize,
376 k: usize,
377 a: *const f32,
378 lda: usize,
379 trans_a: Transpose,
380 b: *const f32,
381 ldb: usize,
382 trans_b: Transpose,
383 c: *mut Self,
384 ldc: usize,
385 workspace: &mut GemmWorkspace<Self::Scalar>,
386 ) {
387 match simd_level() {
388 #[cfg(target_arch = "x86_64")]
389 SimdLevel::Avx2 | SimdLevel::Avx512 => {
390 let kernel = Avx2MaxMulF32Kernel;
391 let params = TilingParams::F32_AVX2;
392 tropical_gemm_inner_with_workspace::<Self, _>(
393 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
394 );
395 }
396 _ => {
397 let kernel = PortableKernel;
398 let params = TilingParams::PORTABLE;
399 tropical_gemm_inner_with_workspace::<Self, _>(
400 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
401 );
402 }
403 }
404 }
405}
406
407impl KernelDispatch for TropicalMinPlus<f64> {
408 argmax_dispatch!(f64, dispatch_f64, true, false);
409 dispatch_with_local_workspace!();
410 unsafe fn dispatch_gemm_with_workspace(
411 m: usize,
412 n: usize,
413 k: usize,
414 a: *const f64,
415 lda: usize,
416 trans_a: Transpose,
417 b: *const f64,
418 ldb: usize,
419 trans_b: Transpose,
420 c: *mut Self,
421 ldc: usize,
422 workspace: &mut GemmWorkspace<Self::Scalar>,
423 ) {
424 crate::core::tropical_gemm_inner_with_workspace(
425 m,
426 n,
427 k,
428 a,
429 lda,
430 trans_a,
431 b,
432 ldb,
433 trans_b,
434 c,
435 ldc,
436 &TilingParams::PORTABLE,
437 &crate::core::PortableMicrokernel,
438 workspace,
439 );
440 }
441}
442
443impl KernelDispatch for TropicalMaxMul<f64> {
444 argmax_dispatch!(f64, dispatch_f64, false, true);
445 dispatch_with_local_workspace!();
446 unsafe fn dispatch_gemm_with_workspace(
447 m: usize,
448 n: usize,
449 k: usize,
450 a: *const f64,
451 lda: usize,
452 trans_a: Transpose,
453 b: *const f64,
454 ldb: usize,
455 trans_b: Transpose,
456 c: *mut Self,
457 ldc: usize,
458 workspace: &mut GemmWorkspace<Self::Scalar>,
459 ) {
460 crate::core::tropical_gemm_inner_with_workspace(
461 m,
462 n,
463 k,
464 a,
465 lda,
466 trans_a,
467 b,
468 ldb,
469 trans_b,
470 c,
471 ldc,
472 &TilingParams::PORTABLE,
473 &crate::core::PortableMicrokernel,
474 workspace,
475 );
476 }
477}
478
479macro_rules! impl_kernel_dispatch_portable {
481 ($($t:ty),*) => {
482 $(
483 impl KernelDispatch for $t {
484 dispatch_with_local_workspace!();
485 unsafe fn dispatch_gemm_with_workspace(
486 m: usize,
487 n: usize,
488 k: usize,
489 a: *const Self::Scalar,
490 lda: usize,
491 trans_a: Transpose,
492 b: *const Self::Scalar,
493 ldb: usize,
494 trans_b: Transpose,
495 c: *mut Self,
496 ldc: usize,
497 workspace: &mut GemmWorkspace<Self::Scalar>,
498 ) {
499 let kernel = PortableKernel;
500 let params = TilingParams::PORTABLE;
501 tropical_gemm_inner_with_workspace::<Self, _>(
502 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &kernel, workspace,
503 );
504 }
505 }
506 )*
507 };
508}
509
510impl_kernel_dispatch_portable!(
511 TropicalAndOr,
512 TropicalMaxPlus<i32>,
513 TropicalMaxPlus<i64>,
514 TropicalMinPlus<i32>,
515 TropicalMinPlus<i64>,
516 TropicalMaxMul<i32>,
517 TropicalMaxMul<i64>
518);
519
520impl KernelDispatch for TropicalBitwise<u32> {
521 dispatch_with_local_workspace!();
522 unsafe fn dispatch_gemm_with_workspace(
523 m: usize,
524 n: usize,
525 k: usize,
526 a: *const u32,
527 lda: usize,
528 trans_a: Transpose,
529 b: *const u32,
530 ldb: usize,
531 trans_b: Transpose,
532 c: *mut Self,
533 ldc: usize,
534 workspace: &mut GemmWorkspace<u32>,
535 ) {
536 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
537 {
538 #[cfg(target_arch = "x86_64")]
539 let supported = is_x86_feature_detected!("avx2");
540 #[cfg(target_arch = "aarch64")]
541 let supported = true;
542 if supported {
543 use super::kernels::bitwise_native::Bitwise32;
544 use crate::core::Microkernel;
545 let params = TilingParams::new(128, 128, 256, 4, Bitwise32::NR);
546 tropical_gemm_inner_with_workspace(
547 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &Bitwise32,
548 workspace,
549 );
550 return;
551 }
552 }
553 tropical_gemm_inner_with_workspace(
554 m,
555 n,
556 k,
557 a,
558 lda,
559 trans_a,
560 b,
561 ldb,
562 trans_b,
563 c,
564 ldc,
565 &TilingParams::PORTABLE,
566 &PortableKernel,
567 workspace,
568 );
569 }
570}
571
572impl KernelDispatch for TropicalBitwise<u64> {
573 dispatch_with_local_workspace!();
574 unsafe fn dispatch_gemm_with_workspace(
575 m: usize,
576 n: usize,
577 k: usize,
578 a: *const u64,
579 lda: usize,
580 trans_a: Transpose,
581 b: *const u64,
582 ldb: usize,
583 trans_b: Transpose,
584 c: *mut Self,
585 ldc: usize,
586 workspace: &mut GemmWorkspace<u64>,
587 ) {
588 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
589 {
590 #[cfg(target_arch = "x86_64")]
591 let supported = is_x86_feature_detected!("avx2");
592 #[cfg(target_arch = "aarch64")]
593 let supported = true;
594 if supported {
595 use super::kernels::bitwise_native::Bitwise64;
596 use crate::core::Microkernel;
597 let params = TilingParams::new(128, 128, 256, 4, Bitwise64::NR);
598 tropical_gemm_inner_with_workspace(
599 m, n, k, a, lda, trans_a, b, ldb, trans_b, c, ldc, ¶ms, &Bitwise64,
600 workspace,
601 );
602 return;
603 }
604 }
605 tropical_gemm_inner_with_workspace(
606 m,
607 n,
608 k,
609 a,
610 lda,
611 trans_a,
612 b,
613 ldb,
614 trans_b,
615 c,
616 ldc,
617 &TilingParams::PORTABLE,
618 &PortableKernel,
619 workspace,
620 );
621 }
622}
623
624#[cfg(test)]
625mod tests {
626 use super::*;
627
628 #[test]
630 fn test_dispatch_maxplus_f32() {
631 let a = vec![1.0f32, 2.0, 3.0, 4.0];
632 let b = vec![1.0f32, 2.0, 3.0, 4.0];
633 let mut c = vec![TropicalMaxPlus::tropical_zero(); 4];
634
635 unsafe {
636 tropical_gemm_dispatch::<TropicalMaxPlus<f32>>(
637 2,
638 2,
639 2,
640 a.as_ptr(),
641 2,
642 Transpose::NoTrans,
643 b.as_ptr(),
644 2,
645 Transpose::NoTrans,
646 c.as_mut_ptr(),
647 2,
648 );
649 }
650
651 assert_eq!(c[0].0, 5.0);
653 }
654
655 #[test]
656 fn test_dispatch_maxplus_f64() {
657 let a = vec![1.0f64, 2.0, 3.0, 4.0];
658 let b = vec![1.0f64, 2.0, 3.0, 4.0];
659 let mut c = vec![TropicalMaxPlus::tropical_zero(); 4];
660
661 unsafe {
662 tropical_gemm_dispatch::<TropicalMaxPlus<f64>>(
663 2,
664 2,
665 2,
666 a.as_ptr(),
667 2,
668 Transpose::NoTrans,
669 b.as_ptr(),
670 2,
671 Transpose::NoTrans,
672 c.as_mut_ptr(),
673 2,
674 );
675 }
676
677 assert_eq!(c[0].0, 5.0);
678 }
679
680 #[test]
681 fn test_dispatch_minplus_f32() {
682 let a = vec![1.0f32, 2.0, 3.0, 4.0];
683 let b = vec![1.0f32, 2.0, 3.0, 4.0];
684 let mut c = vec![TropicalMinPlus::tropical_zero(); 4];
685
686 unsafe {
687 tropical_gemm_dispatch::<TropicalMinPlus<f32>>(
688 2,
689 2,
690 2,
691 a.as_ptr(),
692 2,
693 Transpose::NoTrans,
694 b.as_ptr(),
695 2,
696 Transpose::NoTrans,
697 c.as_mut_ptr(),
698 2,
699 );
700 }
701
702 assert_eq!(c[0].0, 2.0);
704 }
705
706 #[test]
707 fn test_dispatch_minplus_f64() {
708 let a = vec![1.0f64, 2.0, 3.0, 4.0];
709 let b = vec![1.0f64, 2.0, 3.0, 4.0];
710 let mut c = vec![TropicalMinPlus::tropical_zero(); 4];
711
712 unsafe {
713 tropical_gemm_dispatch::<TropicalMinPlus<f64>>(
714 2,
715 2,
716 2,
717 a.as_ptr(),
718 2,
719 Transpose::NoTrans,
720 b.as_ptr(),
721 2,
722 Transpose::NoTrans,
723 c.as_mut_ptr(),
724 2,
725 );
726 }
727
728 assert_eq!(c[0].0, 2.0);
729 }
730
731 #[test]
732 fn test_dispatch_maxmul_f32() {
733 let a = vec![2.0f32, 3.0, 4.0, 5.0];
734 let b = vec![1.0f32, 2.0, 3.0, 4.0];
735 let mut c = vec![TropicalMaxMul::tropical_zero(); 4];
736
737 unsafe {
738 tropical_gemm_dispatch::<TropicalMaxMul<f32>>(
739 2,
740 2,
741 2,
742 a.as_ptr(),
743 2,
744 Transpose::NoTrans,
745 b.as_ptr(),
746 2,
747 Transpose::NoTrans,
748 c.as_mut_ptr(),
749 2,
750 );
751 }
752
753 assert_eq!(c[0].0, 9.0);
755 }
756
757 #[test]
758 fn test_dispatch_maxmul_f64() {
759 let a = vec![2.0f64, 3.0, 4.0, 5.0];
760 let b = vec![1.0f64, 2.0, 3.0, 4.0];
761 let mut c = vec![TropicalMaxMul::tropical_zero(); 4];
762
763 unsafe {
764 tropical_gemm_dispatch::<TropicalMaxMul<f64>>(
765 2,
766 2,
767 2,
768 a.as_ptr(),
769 2,
770 Transpose::NoTrans,
771 b.as_ptr(),
772 2,
773 Transpose::NoTrans,
774 c.as_mut_ptr(),
775 2,
776 );
777 }
778
779 assert_eq!(c[0].0, 9.0);
780 }
781
782 #[test]
783 fn test_dispatch_maxplus_i32() {
784 let a = vec![1i32, 2, 3, 4];
785 let b = vec![1i32, 2, 3, 4];
786 let mut c = vec![TropicalMaxPlus::tropical_zero(); 4];
787
788 unsafe {
789 tropical_gemm_dispatch::<TropicalMaxPlus<i32>>(
790 2,
791 2,
792 2,
793 a.as_ptr(),
794 2,
795 Transpose::NoTrans,
796 b.as_ptr(),
797 2,
798 Transpose::NoTrans,
799 c.as_mut_ptr(),
800 2,
801 );
802 }
803
804 assert_eq!(c[0].0, 5);
805 }
806
807 #[test]
808 fn test_dispatch_maxplus_i64() {
809 let a = vec![1i64, 2, 3, 4];
810 let b = vec![1i64, 2, 3, 4];
811 let mut c = vec![TropicalMaxPlus::tropical_zero(); 4];
812
813 unsafe {
814 tropical_gemm_dispatch::<TropicalMaxPlus<i64>>(
815 2,
816 2,
817 2,
818 a.as_ptr(),
819 2,
820 Transpose::NoTrans,
821 b.as_ptr(),
822 2,
823 Transpose::NoTrans,
824 c.as_mut_ptr(),
825 2,
826 );
827 }
828
829 assert_eq!(c[0].0, 5);
830 }
831
832 #[test]
833 fn test_dispatch_minplus_i32() {
834 let a = vec![1i32, 2, 3, 4];
835 let b = vec![1i32, 2, 3, 4];
836 let mut c = vec![TropicalMinPlus::tropical_zero(); 4];
837
838 unsafe {
839 tropical_gemm_dispatch::<TropicalMinPlus<i32>>(
840 2,
841 2,
842 2,
843 a.as_ptr(),
844 2,
845 Transpose::NoTrans,
846 b.as_ptr(),
847 2,
848 Transpose::NoTrans,
849 c.as_mut_ptr(),
850 2,
851 );
852 }
853
854 assert_eq!(c[0].0, 2);
855 }
856
857 #[test]
858 fn test_dispatch_minplus_i64() {
859 let a = vec![1i64, 2, 3, 4];
860 let b = vec![1i64, 2, 3, 4];
861 let mut c = vec![TropicalMinPlus::tropical_zero(); 4];
862
863 unsafe {
864 tropical_gemm_dispatch::<TropicalMinPlus<i64>>(
865 2,
866 2,
867 2,
868 a.as_ptr(),
869 2,
870 Transpose::NoTrans,
871 b.as_ptr(),
872 2,
873 Transpose::NoTrans,
874 c.as_mut_ptr(),
875 2,
876 );
877 }
878
879 assert_eq!(c[0].0, 2);
880 }
881
882 #[test]
883 fn test_dispatch_maxmul_i32() {
884 let a = vec![2i32, 3, 4, 5];
885 let b = vec![1i32, 2, 3, 4];
886 let mut c = vec![TropicalMaxMul::tropical_zero(); 4];
887
888 unsafe {
889 tropical_gemm_dispatch::<TropicalMaxMul<i32>>(
890 2,
891 2,
892 2,
893 a.as_ptr(),
894 2,
895 Transpose::NoTrans,
896 b.as_ptr(),
897 2,
898 Transpose::NoTrans,
899 c.as_mut_ptr(),
900 2,
901 );
902 }
903
904 assert_eq!(c[0].0, 9);
905 }
906
907 #[test]
908 fn test_dispatch_maxmul_i64() {
909 let a = vec![2i64, 3, 4, 5];
910 let b = vec![1i64, 2, 3, 4];
911 let mut c = vec![TropicalMaxMul::tropical_zero(); 4];
912
913 unsafe {
914 tropical_gemm_dispatch::<TropicalMaxMul<i64>>(
915 2,
916 2,
917 2,
918 a.as_ptr(),
919 2,
920 Transpose::NoTrans,
921 b.as_ptr(),
922 2,
923 Transpose::NoTrans,
924 c.as_mut_ptr(),
925 2,
926 );
927 }
928
929 assert_eq!(c[0].0, 9);
930 }
931
932 #[test]
933 fn test_dispatch_larger_matrix() {
934 let m = 16;
936 let n = 16;
937 let k = 16;
938
939 let a: Vec<f32> = (0..m * k).map(|i| (i % 10) as f32).collect();
940 let b: Vec<f32> = (0..k * n).map(|i| (i % 10) as f32).collect();
941 let mut c = vec![TropicalMaxPlus::tropical_zero(); m * n];
942
943 unsafe {
944 tropical_gemm_dispatch::<TropicalMaxPlus<f32>>(
945 m,
946 n,
947 k,
948 a.as_ptr(),
949 k,
950 Transpose::NoTrans,
951 b.as_ptr(),
952 n,
953 Transpose::NoTrans,
954 c.as_mut_ptr(),
955 n,
956 );
957 }
958
959 let has_non_zero = c.iter().any(|x| x.0 > f32::NEG_INFINITY);
961 assert!(has_non_zero);
962 }
963}