#ifndef EXPLICIT_EINSUM_H #define EXPLICIT_EINSUM_H #include "Fastor/tensor_algebra/indicial.h" #include "Fastor/meta/opmin_meta.h" #include "Fastor/tensor_algebra/permutation.h" #include "Fastor/tensor_algebra/permute.h" #include "Fastor/tensor_algebra/einsum.h" #include "Fastor/tensor_algebra/network_einsum.h" #include "Fastor/tensor_algebra/abstract_contraction.h" namespace Fastor { #if FASTOR_CXX_VERSION >= 2017 // Single tensor //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE typename permute_helper>::resulting_index, typename Index_O::parent_type>, typename einsum_helper>::resulting_tensor>::resulting_tensor einsum(const Tensor &a) { using _einsum_helper = einsum_helper>; using resulting_index_einsum = typename _einsum_helper::resulting_index; using resulting_tensor_einsum = typename _einsum_helper::resulting_tensor; auto res = einsum(a); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } //-----------------------------------------------------------------------------------------------------------------------// // Two tensor (by-pair) //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE typename permute_helper,Tensor>::resulting_index, typename Index_O::parent_type>, typename einsum_helper,Tensor>::resulting_tensor>::resulting_tensor einsum(const Tensor &a, const Tensor &b) { using _einsum_helper = einsum_helper,Tensor>; using resulting_index_einsum = typename _einsum_helper::resulting_index; using resulting_tensor_einsum = typename _einsum_helper::resulting_tensor; auto res = einsum(a,b); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } //-----------------------------------------------------------------------------------------------------------------------// // 3 tensor network //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE typename permute_helper< internal::permute_mapped_index_t< typename einsum_helper,Tensor,Tensor>::resulting_index, typename Index_O::parent_type >, typename einsum_helper,Tensor,Tensor>::resulting_tensor>::resulting_tensor einsum(const Tensor &a, const Tensor &b, const Tensor &c) { using _einsum_helper = einsum_helper,Tensor,Tensor>; using resulting_index_einsum = typename _einsum_helper::resulting_index; using resulting_tensor_einsum = typename _einsum_helper::resulting_tensor; auto res = einsum(a,b,c); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } //-----------------------------------------------------------------------------------------------------------------------// // 4 tensor network //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE typename permute_helper< internal::permute_mapped_index_t< typename einsum_helper,Tensor,Tensor,Tensor>::resulting_index, typename Index_O::parent_type >, typename einsum_helper,Tensor,Tensor,Tensor>::resulting_tensor>::resulting_tensor einsum(const Tensor &a, const Tensor &b, const Tensor &c, const Tensor &d) { using _einsum_helper = einsum_helper,Tensor,Tensor,Tensor>; using resulting_index_einsum = typename _einsum_helper::resulting_index; using resulting_tensor_einsum = typename _einsum_helper::resulting_tensor; auto res = einsum(a,b,c,d); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } //-----------------------------------------------------------------------------------------------------------------------// // 5 tensor network //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE typename permute_helper< internal::permute_mapped_index_t< typename einsum_helper,Tensor,Tensor,Tensor,Tensor>::resulting_index, typename Index_O::parent_type >, typename einsum_helper,Tensor,Tensor,Tensor,Tensor>::resulting_tensor >::resulting_tensor einsum(const Tensor &a, const Tensor &b, const Tensor &c, const Tensor &d, const Tensor &e) { using _einsum_helper = einsum_helper,Tensor,Tensor,Tensor,Tensor>; using resulting_index_einsum = typename _einsum_helper::resulting_index; using resulting_tensor_einsum = typename _einsum_helper::resulting_tensor; auto res = einsum(a,b,c,d,e); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } //-----------------------------------------------------------------------------------------------------------------------// // 6 tensor network //-----------------------------------------------------------------------------------------------------------------------// template FASTOR_INLINE typename permute_helper< internal::permute_mapped_index_t< typename einsum_helper,Tensor,Tensor,Tensor,Tensor,Tensor>::resulting_index, typename Index_O::parent_type >, typename einsum_helper,Tensor,Tensor,Tensor,Tensor,Tensor>::resulting_tensor >::resulting_tensor einsum(const Tensor &a, const Tensor &b, const Tensor &c, const Tensor &d, const Tensor &e, const Tensor &f) { using _einsum_helper = einsum_helper,Tensor,Tensor,Tensor,Tensor,Tensor>; using resulting_index_einsum = typename _einsum_helper::resulting_index; using resulting_tensor_einsum = typename _einsum_helper::resulting_tensor; auto res = einsum(a,b,c,d,e,f); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } //-----------------------------------------------------------------------------------------------------------------------// // network einsum for expressions //-----------------------------------------------------------------------------------------------------------------------// // single expression explicit einsum template,bool> = false> FASTOR_INLINE decltype(auto) einsum(const AbstractTensorType0& a) { decltype(auto) tmp = evaluate(a); auto res = einsum(tmp); using resulting_index_einsum = typename einsum_helper::resulting_index; using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } // pair expression explicit einsum template FASTOR_INLINE decltype(auto) einsum(const AbstractTensor& a, const AbstractTensor& b) { auto res = einsum(a.self(),b.self()); using ttype0 = typename Derived0::resulting_type; using ttype1 = typename Derived1::resulting_type; using resulting_index_einsum = typename einsum_helper::resulting_index; using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } // 3 tensor network expression explicit einsum template FASTOR_INLINE decltype(auto) einsum(const AbstractTensor& a, const AbstractTensor& b, const AbstractTensor& c) { auto res = einsum(a.self(),b.self(),c.self()); using ttype0 = typename Derived0::resulting_type; using ttype1 = typename Derived1::resulting_type; using ttype2 = typename Derived2::resulting_type; using resulting_index_einsum = typename einsum_helper::resulting_index; using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } // 4 tensor network expression explicit einsum template FASTOR_INLINE decltype(auto) einsum( const AbstractTensor& a, const AbstractTensor& b, const AbstractTensor& c, const AbstractTensor& d) { auto res = einsum(a.self(),b.self(),c.self(),d.self()); using ttype0 = typename Derived0::resulting_type; using ttype1 = typename Derived1::resulting_type; using ttype2 = typename Derived2::resulting_type; using ttype3 = typename Derived3::resulting_type; using resulting_index_einsum = typename einsum_helper::resulting_index; using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } // 5 tensor network expression explicit einsum template FASTOR_INLINE decltype(auto) einsum( const AbstractTensor& a, const AbstractTensor& b, const AbstractTensor& c, const AbstractTensor& d, const AbstractTensor& e) { auto res = einsum(a.self(),b.self(),c.self(),d.self(),e.self()); using ttype0 = typename Derived0::resulting_type; using ttype1 = typename Derived1::resulting_type; using ttype2 = typename Derived2::resulting_type; using ttype3 = typename Derived3::resulting_type; using ttype4 = typename Derived4::resulting_type; using resulting_index_einsum = typename einsum_helper::resulting_index; using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } // 6 tensor network expression explicit einsum template FASTOR_INLINE decltype(auto) einsum( const AbstractTensor& a, const AbstractTensor& b, const AbstractTensor& c, const AbstractTensor& d, const AbstractTensor& e, const AbstractTensor& f) { auto res = einsum(a.self(),b.self(),c.self(),d.self(),e.self(),f.self()); using ttype0 = typename Derived0::resulting_type; using ttype1 = typename Derived1::resulting_type; using ttype2 = typename Derived2::resulting_type; using ttype3 = typename Derived3::resulting_type; using ttype4 = typename Derived4::resulting_type; using ttype5 = typename Derived5::resulting_type; using resulting_index_einsum = typename einsum_helper::resulting_index; using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } #if !defined(FASTOR_MSVC) template= 6, bool > = false> FASTOR_INLINE decltype(auto) einsum(const AbstractTensorType0& a, const AbstractTensorTypes& ... rest) { auto res = internal::unpack_einsum_tuple::apply(internal::contraction_chain_evaluate(a,rest...)); auto res_idx = internal::unpack_einsum_helper_tuple::apply(internal::contraction_chain_evaluate(a,rest...)); using mapped_index = internal::permute_mapped_index_t; constexpr bool requires_permutation = requires_permute_v; FASTOR_IF_CONSTEXPR(!requires_permutation) return res; return permute(res); } #endif //------------------------------------------------------------------------------------------------- //------------------------------------------------------------------------------------------------- #endif // CXX 2017 } // end of namespace #endif // EXPLICIT_EINSUM_H