matrixmultiply 0.3.11

General matrix multiplication for f32 and f64 matrices. Operates on matrices with general layout (they can use arbitrary row and column stride). Detects and uses SIMD features on x86/x86-64 and AArch64 transparently for higher performance. Uses a microkernel strategy, so that the implementation is easy to parallelize and optimize. Supports multithreading.
Documentation
// Copyright 2016 - 2021 Ulrik Sverdrup "bluss"
//
// Licensed under the Apache License, Version 2.0 <LICENSE-APACHE or
// http://www.apache.org/licenses/LICENSE-2.0> or the MIT license
// <LICENSE-MIT or http://opensource.org/licenses/MIT>, at your
// option. This file may not be copied, modified, or distributed
// except according to those terms.

use crate::kernel::GemmKernel;
use crate::kernel::GemmSelect;
use crate::kernel::{U2, U4, c32, Element, c32_mul as mul};
#[cfg(has_avx512)]
use crate::kernel::U8;
use crate::kernel_util::at;
use crate::archparam;
use crate::cgemm_common::pack_complex;
use crate::packing::PackSlice;

#[cfg(has_avx512)]
struct KernelAvx512;
#[cfg(any(target_arch="x86", target_arch="x86_64"))]
struct KernelAvx2;

#[cfg(target_arch = "aarch64")]
struct KernelNeon;

struct KernelFallback;

type T = c32;
type TReal = f32;

/// Detect which implementation to use and select it using the selector's
/// .select(Kernel) method.
///
/// This function is called one or more times during a whole program's
/// execution, it may be called for each gemm kernel invocation or fewer times.
#[inline]
pub(crate) fn detect<G>(selector: G) where G: GemmSelect<T> {
    // dispatch to specific compiled versions
    #[cfg(any(target_arch="x86", target_arch="x86_64"))]
    {
        #[cfg(has_avx512)]
        {
            if is_x86_feature_detected_!("avx512f") {
                return selector.select(KernelAvx512);
            }
        }
        if is_x86_feature_detected_!("fma") && is_x86_feature_detected_!("avx2") {
            return selector.select(KernelAvx2);
        }
    }
    #[cfg(target_arch = "aarch64")]
    {
        if is_aarch64_feature_detected_!("neon") {
            return selector.select(KernelNeon);
        }
    }
    return selector.select(KernelFallback);
}

#[cfg(has_avx512)]
impl GemmKernel for KernelAvx512 {
    type Elem = T;

    type MRTy = U8;
    type NRTy = U4;

    const ALIGNMENT: usize = 32;
    const ALWAYS_MASKED: bool = KernelFallback::ALWAYS_MASKED;

    #[inline(always)]
    fn nc() -> usize { archparam::C_NC }
    #[inline(always)]
    fn kc() -> usize { archparam::C_KC }
    #[inline(always)]
    fn mc() -> usize { archparam::C_MC }

    pack_methods!{}

    #[inline(always)]
    unsafe fn kernel(
        k: usize,
        alpha: T,
        a: *const T,
        b: *const T,
        beta: T,
        c: *mut T, rsc: isize, csc: isize) {
        kernel_target_avx512(k, alpha, a, b, beta, c, rsc, csc)
    }
}

#[cfg(any(target_arch="x86", target_arch="x86_64"))]
impl GemmKernel for KernelAvx2 {
    type Elem = T;

    type MRTy = U4;
    type NRTy = U4;

    const ALIGNMENT: usize = 32;
    const ALWAYS_MASKED: bool = KernelFallback::ALWAYS_MASKED;

    #[inline(always)]
    fn nc() -> usize { archparam::C_NC }
    #[inline(always)]
    fn kc() -> usize { archparam::C_KC }
    #[inline(always)]
    fn mc() -> usize { archparam::C_MC }

    pack_methods!{}

    #[inline(always)]
    unsafe fn kernel(
        k: usize,
        alpha: T,
        a: *const T,
        b: *const T,
        beta: T,
        c: *mut T, rsc: isize, csc: isize) {
        kernel_target_avx2(k, alpha, a, b, beta, c, rsc, csc)
    }
}

#[cfg(target_arch = "aarch64")]
impl GemmKernel for KernelNeon {
    type Elem = T;

    type MRTy = U4;
    type NRTy = U2;

    const ALIGNMENT: usize = 16;
    const ALWAYS_MASKED: bool = KernelFallback::ALWAYS_MASKED;

    #[inline(always)]
    fn nc() -> usize { archparam::C_NC }
    #[inline(always)]
    fn kc() -> usize { archparam::C_KC }
    #[inline(always)]
    fn mc() -> usize { archparam::C_MC }

    pack_methods!{}

    #[inline(always)]
    unsafe fn kernel(
        k: usize,
        alpha: T,
        a: *const T,
        b: *const T,
        beta: T,
        c: *mut T, rsc: isize, csc: isize) {
        kernel_target_neon(k, alpha, a, b, beta, c, rsc, csc)
    }
}

impl GemmKernel for KernelFallback {
    type Elem = T;

    type MRTy = U4;
    type NRTy = U2;

    const ALWAYS_MASKED: bool = true;

    #[inline(always)]
    fn nc() -> usize { archparam::C_NC }
    #[inline(always)]
    fn kc() -> usize { archparam::C_KC }
    #[inline(always)]
    fn mc() -> usize { archparam::C_MC }

    pack_methods!{}

    #[inline(always)]
    unsafe fn kernel(
        k: usize,
        alpha: T,
        a: *const T,
        b: *const T,
        beta: T,
        c: *mut T, rsc: isize, csc: isize) {
        kernel_fallback_impl(k, alpha, a, b, beta, c, rsc, csc)
    }
}

// Kernel AVX-512
#[cfg(has_avx512)]
macro_rules! loop_m { ($i:ident, $e:expr) => { loop8!($i, $e) }; }
#[cfg(has_avx512)]
macro_rules! loop_n { ($j:ident, $e:expr) => { loop4!($j, $e) }; }

#[cfg(has_avx512)]
kernel_fallback_impl_complex! {
    // instantiate separately
    [inline target_feature(enable="avx512f")] [fma_yes]
    kernel_target_avx512, T, TReal, KernelAvx512::MR, KernelAvx512::NR, 4
}

// Kernel AVX2
#[cfg(any(target_arch="x86", target_arch="x86_64"))]
macro_rules! loop_m { ($i:ident, $e:expr) => { loop4!($i, $e) }; }
#[cfg(any(target_arch="x86", target_arch="x86_64"))]
macro_rules! loop_n { ($j:ident, $e:expr) => { loop4!($j, $e) }; }

#[cfg(any(target_arch="x86", target_arch="x86_64"))]
kernel_fallback_impl_complex! {
    // instantiate separately
    [inline target_feature(enable="fma,avx2")] [fma_yes]
    kernel_target_avx2, T, TReal, KernelAvx2::MR, KernelAvx2::NR, 4
}

// Kernel neon

#[cfg(target_arch = "aarch64")]
macro_rules! loop_m { ($i:ident, $e:expr) => { loop4!($i, $e) }; }
#[cfg(target_arch = "aarch64")]
macro_rules! loop_n { ($j:ident, $e:expr) => { loop2!($j, $e) }; }

#[cfg(target_arch = "aarch64")]
kernel_fallback_impl_complex! {
    [inline target_feature(enable="neon")] [fma_yes]
    kernel_target_neon, T, TReal, KernelNeon::MR, KernelNeon::NR, 1
}

// Kernel fallback

macro_rules! loop_m { ($i:ident, $e:expr) => { loop4!($i, $e) }; }
macro_rules! loop_n { ($j:ident, $e:expr) => { loop2!($j, $e) }; }

kernel_fallback_impl_complex! {
    [inline(always)] [fma_no]
    kernel_fallback_impl, T, TReal, KernelFallback::MR, KernelFallback::NR, 1
}

#[cfg(test)]
mod tests {
    use super::*;
    use crate::kernel::test::test_complex_packed_kernel;

    #[test]
    fn test_kernel_fallback_impl() {
        test_complex_packed_kernel::<KernelFallback, _, TReal>("kernel");
    }

    #[cfg(target_arch = "aarch64")]
    mod test_kernel_aarch64 {
        use super::test_complex_packed_kernel;
        use super::super::*;
        #[cfg(feature = "std")]
        use std::println;
        macro_rules! test_arch_kernels {
            ($($feature_name:tt, $name:ident, $kernel_ty:ty),*) => {
                $(
                #[test]
                fn $name() {
                    if is_aarch64_feature_detected_!($feature_name) {
                        test_complex_packed_kernel::<$kernel_ty, _, TReal>(stringify!($name));
                    } else {
                        #[cfg(feature = "std")]
                        println!("Skipping, host does not have feature: {:?}", $feature_name);
                    }
                }
                )*
            }
        }

        test_arch_kernels! {
            "neon", neon, KernelNeon
        }
    }

    #[cfg(any(target_arch="x86", target_arch="x86_64"))]
    mod test_arch_kernels {
        use super::test_complex_packed_kernel;
        use super::super::*;
        #[cfg(feature = "std")]
        use std::println;
        macro_rules! test_arch_kernels_x86 {
            ($($feature_name:tt, $name:ident, $kernel_ty:ty),*) => {
                $(
                #[test]
                fn $name() {
                    if is_x86_feature_detected_!($feature_name) {
                        test_complex_packed_kernel::<$kernel_ty, _, TReal>(stringify!($name));
                    } else {
                        #[cfg(feature = "std")]
                        println!("Skipping, host does not have feature: {:?}", $feature_name);
                    }
                }
                )*
            }
        }

        test_arch_kernels_x86! {
            "avx2", avx2, KernelAvx2
        }

        #[cfg(has_avx512)]
        test_arch_kernels_x86! {
            "avx512f", avx512f, KernelAvx512
        }
    }
}