Skip to main content

tropical_gemm/simd/
dispatch.rs

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
9/// Runtime-dispatched GEMM that selects the best kernel for the current CPU.
10///
11/// # Safety
12/// Same requirements as `tropical_gemm_inner`
13pub 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/// Runtime-dispatched GEMM with first-winner indices.
30///
31/// # Safety
32/// Same pointer and output storage requirements as `tropical_gemm_with_argmax_portable`.
33#[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
51/// Trait for types that support kernel dispatch.
52pub trait KernelDispatch: TropicalSemiring {
53    /// Dispatch to the appropriate kernel based on CPU features.
54    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    /// Dispatch argmax, defaulting to the portable kernel for custom types.
68    ///
69    /// # Safety
70    /// Inputs and result must be valid for the requested dimensions and strides.
71    #[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    /// Dispatch using reusable packing storage. Custom implementations retain
91    /// their existing dispatch unless they override this method to use workspace.
92    ///
93    /// # Safety
94    /// Same requirements as `dispatch_gemm`.
95    #[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    /// Dispatch argmax using reusable packing storage.
114    ///
115    /// # Safety
116    /// Same requirements as `dispatch_gemm_with_argmax`.
117    #[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
151// Preserve the existing dispatch entry point while allowing callers to opt in
152// to packing reuse through the workspace-aware variant.
153macro_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
186// Each float semiring supplies its ordered comparison and product operation.
187macro_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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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, &params, &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
479// Fallback implementations for other types
480macro_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, &params, &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, &params, &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, &params, &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 that the dispatch function exists and doesn't panic for small inputs
629    #[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        // C[0,0] = max(A[0,0]+B[0,0], A[0,1]+B[1,0]) = max(1+1, 2+3) = 5
652        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        // C[0,0] = min(A[0,0]+B[0,0], A[0,1]+B[1,0]) = min(1+1, 2+3) = 2
703        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        // C[0,0] = max(A[0,0]*B[0,0], A[0,1]*B[1,0]) = max(2*1, 3*3) = 9
754        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        // Test a larger matrix to exercise blocking
935        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        // Just verify no panic and result is not all zeros
960        let has_non_zero = c.iter().any(|x| x.0 > f32::NEG_INFINITY);
961        assert!(has_non_zero);
962    }
963}