#ifndef MATMUL_H #define MATMUL_H #include "Fastor/meta/meta.h" #include "Fastor/backend/matmul/matmul_kernels.h" #ifdef FASTOR_USE_LIBXSMM #include "Fastor/backend/matmul/libxsmm_backend.h" #endif #ifdef FASTOR_USE_MKL #include "Fastor/backend/matmul/mkl_backend.h" #endif namespace Fastor { // Forward declare //----------------------------------------------------------------------------------------------------------- namespace internal { template FASTOR_INLINE void _matvecmul(const T * FASTOR_RESTRICT a, const T * FASTOR_RESTRICT b, T * FASTOR_RESTRICT out); } // internal //----------------------------------------------------------------------------------------------------------- //----------------------------------------------------------------------------------------------------------- //----------------------------------------------------------------------------------------------------------- #if !defined(FASTOR_USE_LIBXSMM) && !defined(FASTOR_USE_MKL) template || is_same_v_) ),bool> = 0> #else template || is_same_v_) ) && is_less_equal::value,1>::value, bool> = 0> #endif FASTOR_INLINE void _matmul(const T * FASTOR_RESTRICT a, const T * FASTOR_RESTRICT b, T * FASTOR_RESTRICT out) { // Non-primitive types FASTOR_IF_CONSTEXPR (!is_primitive_v_) { internal::_matmul_base_non_primitive(a,b,out); return; } // Matrix-vector specialisation FASTOR_IF_CONSTEXPR (N==1UL) { internal::_matvecmul(a,b,out); return; } using nativeV = SIMDVector; using V = choose_best_simd_t; // Use specialised kernels FASTOR_IF_CONSTEXPR((N==V::Size || N==2*V::Size || N==3*V::Size || N==4*V::Size || N==5*V::Size) && V::Size!=1UL) { internal::_matmul_mk_smalln(a,b,out); return; } #if defined(FASTOR_AVX2_IMPL) || defined(FASTOR_HAS_AVX512_MASKS) FASTOR_IF_CONSTEXPR((N<5*V::Size && N!=1UL)) { internal::_matmul_mk_smalln(a,b,out); return; } #endif #if defined(FASTOR_AVX2_IMPL) || defined(FASTOR_HAS_AVX512_MASKS) FASTOR_IF_CONSTEXPR( M*N*K > 27UL && N % V::Size <= 1UL) { internal::_matmul_base(a,b,out); return; } else FASTOR_IF_CONSTEXPR( M*N*K > 27UL && N % V::Size > 1UL) { internal::_matmul_base_masked(a,b,out); return; } #else FASTOR_IF_CONSTEXPR( M*N*K > 27UL ) { internal::_matmul_base(a,b,out); return; } #endif else { // For all other cases where M,N,K is too small // this simple version is sufficient constexpr int ROUND_ = ROUND_DOWN(N,V::Size); for (size_t j=0; j || is_same_v_) ) && is_greater::value,1>::value, bool> = 0> FASTOR_INLINE void _matmul(const T * FASTOR_RESTRICT a, const T * FASTOR_RESTRICT b, T * FASTOR_RESTRICT c) { blas::matmul_libxsmm(a,b,c); } #endif #if !defined(FASTOR_USE_LIBXSMM) && defined(FASTOR_USE_MKL) template || is_same_v_) ) && is_greater::value,1>::value, bool> = 0> FASTOR_INLINE void _matmul(const T * FASTOR_RESTRICT a, const T * FASTOR_RESTRICT b, T * FASTOR_RESTRICT c) { blas::matmul_mkl(a,b,c); } #endif //----------------------------------------------------------------------------------------------------------- //----------------------------------------------------------------------------------------------------------- //----------------------------------------------------------------------------------------------------------- //----------------------------------------------------------------------------------------------------------- #include "Fastor/backend/matmul/matmul_specialisations_kernels.h" //----------------------------------------------------------------------------------------------------------- //----------------------------------------------------------------------------------------------------------- } // end of namespace #endif // MATMUL_H