#ifndef EINSUM_H #define EINSUM_H #include "Fastor/backend/backend.h" #include "Fastor/tensor/Tensor.h" #include "Fastor/meta/einsum_meta.h" #include "Fastor/tensor_algebra/indicial.h" #include "Fastor/backend/voigt.h" #include "Fastor/tensor_algebra/permutation.h" #include "Fastor/tensor_algebra/permute.h" #include "Fastor/tensor_algebra/innerproduct.h" #include "Fastor/tensor_algebra/outerproduct.h" #include "Fastor/tensor_algebra/contraction.h" #include "Fastor/tensor_algebra/contraction_single.h" #include "Fastor/tensor_algebra/strided_contraction.h" namespace Fastor { // Single tensor //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE auto einsum(const Tensor &a) -> decltype(extractor_contract_1::contract_impl(a)) { static_assert(einsum_index_checker::value, "INDICES FOR EINSUM FUNCTION CANNOT APPEAR MORE THAN TWICE. USE INNER INSTEAD"); return extractor_contract_1::contract_impl(a); } //-----------------------------------------------------------------------------------------------------------------------// // Two tensor (by-pair) //-----------------------------------------------------------------------------------------------------------------------// // Inner product case //-----------------------------------------------------------------------------------------------------------------------// template,bool> = false> FASTOR_INLINE Tensor einsum(const Tensor &a, const Tensor &b) { return inner(a,b); } // General by-pair product cases //-----------------------------------------------------------------------------------------------------------------------// template::value && !internal::is_generalised_matrix_vector::value && !internal::is_generalised_vector_matrix::value && !internal::is_generalised_matrix_matrix::value ,bool>::type=0> FASTOR_INLINE auto einsum(const Tensor &a, const Tensor &b) -> decltype(extractor_contract_2::contract_impl(a,b)) { static_assert(einsum_index_checker::type>::value, "INDICES FOR EINSUM FUNCTION CANNOT APPEAR MORE THAN TWICE. USE INNER INSTEAD"); // // Dispatch to the right routine // using vectorisability = is_vectorisable>; // // constexpr bool is_reducible = vectorisability::last_index_contracted; // constexpr bool is_reducible = vectorisability::is_reducible; // FASTOR_IF_CONSTEXPR (is_reducible) { // return extractor_reducible_contract::contract_impl(a,b); // } // else { // return extractor_contract_2::contract_impl(a,b); // } return extractor_contract_2::contract_impl(a,b); } template::value, bool>::type = 0> FASTOR_INLINE auto einsum(const Tensor &a, const Tensor &b) //{ -> decltype(extractor_contract_2::contract_impl(a,b)) { constexpr size_t which_one_is_vector = internal::is_generalised_matrix_vector::which_one_is_vector; constexpr size_t matches_up_to = internal::is_generalised_matrix_vector::matches_up_to; constexpr size_t rest0[sizeof...(Rest0)] = {Rest0...}; constexpr size_t rest1[sizeof...(Rest1)] = {Rest1...}; constexpr size_t product = which_one_is_vector == 1 ? partial_prod(rest0, matches_up_to) : partial_prod(rest1, matches_up_to); constexpr size_t vec_product = which_one_is_vector == 1 ? pack_prod::value : pack_prod::value; decltype(extractor_contract_2::contract_impl(a,b)) out; which_one_is_vector == 1 ? _matmul(a.data(),b.data(),out.data()) :\ _matmul(b.data(),a.data(),out.data()); return out; } template::value, bool>::type = 0> FASTOR_INLINE auto einsum(const Tensor &a, const Tensor &b) //{ -> decltype(extractor_contract_2::contract_impl(a,b)) { constexpr size_t which_one_is_vector = internal::is_generalised_vector_matrix::which_one_is_vector; constexpr size_t matches_up_to = internal::is_generalised_vector_matrix::matches_up_to; constexpr size_t rest0[sizeof...(Rest0)] = {Rest0...}; constexpr size_t rest1[sizeof...(Rest1)] = {Rest1...}; constexpr size_t product = which_one_is_vector == 1 ? partial_prod_reverse(rest0, matches_up_to) : partial_prod_reverse(rest1, matches_up_to); constexpr size_t vec_product = which_one_is_vector == 1 ? pack_prod::value : pack_prod::value; decltype(extractor_contract_2::contract_impl(a,b)) out; which_one_is_vector == 1 ? _matmul(b.data(),a.data(),out.data()) :\ _matmul(a.data(),b.data(),out.data()); return out; } template::value, bool>::type = 0> FASTOR_INLINE auto einsum(const Tensor &a, const Tensor &b) //{ -> decltype(extractor_contract_2::contract_impl(a,b)) { constexpr size_t matches_up_to = internal::is_generalised_matrix_matrix::ncontracted; constexpr size_t rest0[sizeof...(Rest0)] = {Rest0...}; constexpr size_t rest1[sizeof...(Rest1)] = {Rest1...}; constexpr size_t K_product = partial_prod(rest1, matches_up_to - 1); constexpr size_t M = partial_prod(rest0, sizeof...(Rest0) - matches_up_to - 1); constexpr size_t N = partial_prod(rest1, sizeof...(Rest1) - 1, matches_up_to); decltype(extractor_contract_2::contract_impl(a,b)) out; _matmul(a.data(),b.data(),out.data()); return out; } //-----------------------------------------------------------------------------------------------------------------------// // matmul dispatcher for 2nd order tensors (matrix-matrix) // also includes matrix-vector and vector-matrix when vector is of size // nx1 or 1xn template::type = 0> FASTOR_INLINE Tensor einsum(const Tensor &a, const Tensor &b) { Tensor out; _matmul(a.data(),b.data(),out.data()); return out; } // matmul dispatcher for matrix-vector template::type = 0> FASTOR_INLINE Tensor einsum(const Tensor &a, const Tensor &b) { Tensor out; _matmul(a.data(),b.data(),out.data()); return out; } // matmul dispatcher for matrix-vector template::type = 0> FASTOR_INLINE Tensor einsum(const Tensor &a, const Tensor &b) { Tensor out; _matmul(b.data(),a.data(),out.data()); return out; } // matmul dispatcher for vector-matrix template::type = 0> FASTOR_INLINE Tensor einsum(const Tensor &a, const Tensor &b) { Tensor out; _matmul(a.data(),b.data(),out.data()); return out; } // matmul dispatcher for vector-matrix template::type = 0> FASTOR_INLINE Tensor einsum(const Tensor &a, const Tensor &b) { Tensor out; _matmul(b.data(),a.data(),out.data()); return out; } #ifdef FASTOR_AVX_IMPL // Specific overloads // With Voigt conversion template::value || std::is_same::value) && I==J && J==K && K==L && (I==2 || I==3) && Ind0::NoIndices==2 && Ind1::NoIndices==2 && Convert==FASTOR_Voigt,bool>::type = 0> FASTOR_INLINE typename VoigtType::return_type einsum(const Tensor & a, const Tensor &b) { using OutTensor = typename VoigtType::return_type; OutTensor out; constexpr int i = static_cast(Ind0::values[0]); constexpr int j = static_cast(Ind0::values[1]); constexpr int k = static_cast(Ind1::values[0]); constexpr int l = static_cast(Ind1::values[1]); constexpr bool is_dyadic = ik && j(a.data(),b.data(),out.data()); } if (is_cyclic) { _cyclic(a.data(),b.data(),out.data()); } return out; } template::value && !std::is_same::value) && I==J && J==K && K==L && (I==2 || I==3) && Ind0::Size==2 && Ind1::Size==2 && Convert==FASTOR_Voigt,bool>::type = 0> FASTOR_INLINE typename VoigtType::return_type einsum(const Tensor & a, const Tensor &b) { constexpr int i = static_cast(Ind0::values[0]); constexpr int j = static_cast(Ind0::values[1]); constexpr int k = static_cast(Ind1::values[0]); constexpr int l = static_cast(Ind1::values[1]); constexpr bool is_dyadic = ik && j::type; if (is_dyadic) { auto out = contraction(a,b); return voigt(out); } if (is_cyclic) { auto out = permutation(contraction(a,b)); return voigt(out); } } #endif } // end of namespace #endif // EINSUM_H