11#ifndef EIGEN_GENERAL_BLOCK_PANEL_H
12#define EIGEN_GENERAL_BLOCK_PANEL_H
15#include "../InternalHeaderCheck.h"
21#pragma warning(disable : 4804)
28enum GEBPPacketSizeType { GEBPPacketFull = 0, GEBPPacketHalf, GEBPPacketQuarter };
30template <
typename LhsScalar_,
typename RhsScalar_,
bool ConjLhs_ =
false,
bool ConjRhs_ =
false,
31 int Arch = Architecture::Target,
int PacketSize_ = GEBPPacketFull>
35inline std::ptrdiff_t manage_caching_sizes_helper(std::ptrdiff_t a, std::ptrdiff_t b) {
return a <= 0 ? b : a; }
37#if defined(EIGEN_DEFAULT_L1_CACHE_SIZE)
38#define EIGEN_SET_DEFAULT_L1_CACHE_SIZE(val) EIGEN_DEFAULT_L1_CACHE_SIZE
40#define EIGEN_SET_DEFAULT_L1_CACHE_SIZE(val) val
43#if defined(EIGEN_DEFAULT_L2_CACHE_SIZE)
44#define EIGEN_SET_DEFAULT_L2_CACHE_SIZE(val) EIGEN_DEFAULT_L2_CACHE_SIZE
46#define EIGEN_SET_DEFAULT_L2_CACHE_SIZE(val) val
49#if defined(EIGEN_DEFAULT_L3_CACHE_SIZE)
50#define EIGEN_SET_DEFAULT_L3_CACHE_SIZE(val) EIGEN_DEFAULT_L3_CACHE_SIZE
52#define EIGEN_SET_DEFAULT_L3_CACHE_SIZE(val) val
55#if EIGEN_ARCH_i386_OR_x86_64
56const std::ptrdiff_t defaultL1CacheSize = EIGEN_SET_DEFAULT_L1_CACHE_SIZE(32 * 1024);
57const std::ptrdiff_t defaultL2CacheSize = EIGEN_SET_DEFAULT_L2_CACHE_SIZE(256 * 1024);
58const std::ptrdiff_t defaultL3CacheSize = EIGEN_SET_DEFAULT_L3_CACHE_SIZE(2 * 1024 * 1024);
60const std::ptrdiff_t defaultL1CacheSize = EIGEN_SET_DEFAULT_L1_CACHE_SIZE(64 * 1024);
62const std::ptrdiff_t defaultL2CacheSize = EIGEN_SET_DEFAULT_L2_CACHE_SIZE(2 * 1024 * 1024);
63const std::ptrdiff_t defaultL3CacheSize = EIGEN_SET_DEFAULT_L3_CACHE_SIZE(8 * 1024 * 1024);
65const std::ptrdiff_t defaultL2CacheSize = EIGEN_SET_DEFAULT_L2_CACHE_SIZE(512 * 1024);
66const std::ptrdiff_t defaultL3CacheSize = EIGEN_SET_DEFAULT_L3_CACHE_SIZE(4 * 1024 * 1024);
68#elif EIGEN_ARCH_ARM_OR_ARM64
69const std::ptrdiff_t defaultL1CacheSize = EIGEN_SET_DEFAULT_L1_CACHE_SIZE(64 * 1024);
70const std::ptrdiff_t defaultL2CacheSize = EIGEN_SET_DEFAULT_L2_CACHE_SIZE(1024 * 1024);
71const std::ptrdiff_t defaultL3CacheSize = EIGEN_SET_DEFAULT_L3_CACHE_SIZE(4 * 1024 * 1024);
73const std::ptrdiff_t defaultL1CacheSize = EIGEN_SET_DEFAULT_L1_CACHE_SIZE(16 * 1024);
74const std::ptrdiff_t defaultL2CacheSize = EIGEN_SET_DEFAULT_L2_CACHE_SIZE(512 * 1024);
75const std::ptrdiff_t defaultL3CacheSize = EIGEN_SET_DEFAULT_L3_CACHE_SIZE(512 * 1024);
78#undef EIGEN_SET_DEFAULT_L1_CACHE_SIZE
79#undef EIGEN_SET_DEFAULT_L2_CACHE_SIZE
80#undef EIGEN_SET_DEFAULT_L3_CACHE_SIZE
84 CacheSizes() : m_l1(-1), m_l2(-1), m_l3(-1), m_l3_per_cpu(0) {
85 std::ptrdiff_t l1CacheSize, l2CacheSize, l3CacheSize;
86 queryCacheSizes(l1CacheSize, l2CacheSize, l3CacheSize, m_l3_per_cpu);
87 m_l1 = manage_caching_sizes_helper(l1CacheSize, defaultL1CacheSize);
88 m_l2 = manage_caching_sizes_helper(l2CacheSize, defaultL2CacheSize);
89 m_l3 = manage_caching_sizes_helper(l3CacheSize, defaultL3CacheSize);
97 std::ptrdiff_t m_l3_per_cpu;
101inline void manage_caching_sizes(Action action, std::ptrdiff_t* l1, std::ptrdiff_t* l2, std::ptrdiff_t* l3,
102 std::ptrdiff_t* l3_per_cpu =
nullptr) {
103 static CacheSizes m_cacheSizes;
105 if (action == SetAction) {
107 eigen_internal_assert(l1 != 0 && l2 != 0);
108 m_cacheSizes.m_l1 = *l1;
109 m_cacheSizes.m_l2 = *l2;
110 m_cacheSizes.m_l3 = *l3;
111 m_cacheSizes.m_l3_per_cpu = l3_per_cpu !=
nullptr ? *l3_per_cpu : 0;
112 }
else if (action == GetAction) {
113 eigen_internal_assert(l1 != 0 && l2 != 0);
114 *l1 = m_cacheSizes.m_l1;
115 *l2 = m_cacheSizes.m_l2;
116 if (l3_per_cpu !=
nullptr) *l3_per_cpu = m_cacheSizes.m_l3_per_cpu;
117 *l3 = m_cacheSizes.m_l3;
119 eigen_internal_assert(
false);
135#ifdef EIGEN_VECTORIZE_SME
140template <
typename LhsScalar,
typename RhsScalar>
141struct sme_has_gebp_kernel : std::false_type {};
143struct sme_has_gebp_kernel<float, float> : std::true_type {};
144#ifdef EIGEN_VECTORIZE_SME_F64F64
146struct sme_has_gebp_kernel<double, double> : std::true_type {};
150template <
typename RealScalar>
151struct sme_has_gebp_kernel<std::complex<RealScalar>, std::complex<RealScalar>>
152 : sme_has_gebp_kernel<RealScalar, RealScalar> {};
157#ifndef EIGEN_SME_MAX_KC
158#define EIGEN_SME_MAX_KC 2048
160#ifndef EIGEN_SME_PACKED_RHS_BUDGET_BYTES
161#define EIGEN_SME_PACKED_RHS_BUDGET_BYTES (32 * 1024 * 1024)
163#ifndef EIGEN_SME_LHS_WORKING_SET_BUDGET_BYTES
164#define EIGEN_SME_LHS_WORKING_SET_BUDGET_BYTES (7 * 1024 * 1024)
166#ifndef EIGEN_SME_SINGLE_PASS_RHS_BUDGET_BYTES
167#define EIGEN_SME_SINGLE_PASS_RHS_BUDGET_BYTES (4 * 1024 * 1024)
170template <
typename LhsScalar,
typename RhsScalar,
typename Index>
171void evaluateProductBlockingSizesHeuristicForSme(Index& k, Index& m, Index& n) {
172 using Traits = gebp_traits<LhsScalar, RhsScalar>;
174 const Index mr =
static_cast<Index
>(Traits::mr);
175 const Index nr =
static_cast<Index
>(Traits::nr);
177#ifdef EIGEN_DEBUG_SMALL_PRODUCT_BLOCKS
182 constexpr Index sme_max_kc =
static_cast<Index
>(128);
183 constexpr Index sme_packed_rhs_budget_bytes =
static_cast<Index
>(128 * 1024);
184 constexpr Index sme_lhs_working_set_budget_bytes =
static_cast<Index
>(128 * 1024);
185 constexpr Index sme_single_pass_rhs_budget_bytes =
static_cast<Index
>(128 * 1024);
187 constexpr Index sme_max_kc =
static_cast<Index
>(EIGEN_SME_MAX_KC);
188 constexpr Index sme_packed_rhs_budget_bytes =
static_cast<Index
>(EIGEN_SME_PACKED_RHS_BUDGET_BYTES);
189 constexpr Index sme_lhs_working_set_budget_bytes =
static_cast<Index
>(EIGEN_SME_LHS_WORKING_SET_BUDGET_BYTES);
190 constexpr Index sme_single_pass_rhs_budget_bytes =
static_cast<Index
>(EIGEN_SME_SINGLE_PASS_RHS_BUDGET_BYTES);
197 const Index max_kc = (numext::maxi)(Index(1), sme_max_kc * Index(
sizeof(
float)) / Index(
sizeof(LhsScalar)));
198 k = (numext::mini)(k, max_kc);
200 const Index block_b_hot_bytes = k * nr * Index(
sizeof(RhsScalar));
201 const Index min_lhs_bytes = mr * k * Index(
sizeof(LhsScalar));
202 const Index block_a_bytes = sme_lhs_working_set_budget_bytes > block_b_hot_bytes
203 ? sme_lhs_working_set_budget_bytes - block_b_hot_bytes
205 Index mc = block_a_bytes / (k * Index(
sizeof(LhsScalar)));
210 const Index rhs_budget = m <= mc ? (numext::mini)(sme_packed_rhs_budget_bytes, sme_single_pass_rhs_budget_bytes)
211 : sme_packed_rhs_budget_bytes;
212 Index nc = rhs_budget / (numext::maxi)(Index(1), k * Index(
sizeof(RhsScalar)));
214 n = (numext::mini)(n, (numext::maxi)(nr, nc));
216 m = (numext::mini)(m, (numext::maxi)(mr, mc));
220template <
typename LhsScalar,
typename RhsScalar,
int KcFactor,
typename Index>
221void evaluateProductBlockingSizesHeuristic(Index& k, Index& m, Index& n, Index num_threads = 1) {
222 using Traits = gebp_traits<LhsScalar, RhsScalar>;
229 std::ptrdiff_t l1, l2, l3, l3_per_cpu;
230 manage_caching_sizes(GetAction, &l1, &l2, &l3, &l3_per_cpu);
231#ifdef EIGEN_VECTORIZE_AVX512
232 const std::ptrdiff_t phys_l1 = l1;
243 if (num_threads > 1) {
244 using ResScalar =
typename Traits::ResScalar;
246 kdiv = KcFactor * (Traits::mr *
sizeof(LhsScalar) + Traits::nr *
sizeof(RhsScalar)),
247 ksub = Traits::mr * (Traits::nr *
sizeof(ResScalar)),
257 const Index k_cache = numext::maxi<Index>(kr, (numext::mini<Index>)(
static_cast<Index
>((l1 - ksub) / kdiv), 320));
259 k = k_cache - (k_cache % kr);
260 eigen_internal_assert(k > 0);
263 const Index n_cache =
static_cast<Index
>((l2 - l1) / (nr *
sizeof(RhsScalar) * k));
264 const Index nr_index =
static_cast<Index
>(nr);
267 const Index n_cache_aligned = n_cache >= nr_index ? n_cache - (n_cache % nr_index) : nr_index;
268 const Index n_per_thread = numext::div_ceil(n, num_threads);
269 if (n_cache <= n_per_thread) {
272 n = (numext::mini<Index>)(n, n_cache_aligned);
273 eigen_internal_assert(n > 0);
275 n = (numext::mini<Index>)(n, (n_per_thread + nr - 1) - ((n_per_thread + nr - 1) % nr));
280 const Index m_cache =
static_cast<Index
>((l3 - l2) / (
sizeof(LhsScalar) * k * num_threads));
281 const Index m_per_thread = numext::div_ceil(m, num_threads);
282 if (m_cache < m_per_thread && m_cache >=
static_cast<Index
>(mr)) {
283 m = m_cache - (m_cache % mr);
284 eigen_internal_assert(m > 0);
286 m = (numext::mini<Index>)(m, (m_per_thread + mr - 1) - ((m_per_thread + mr - 1) % mr));
292#ifdef EIGEN_DEBUG_SMALL_PRODUCT_BLOCKS
304 if ((numext::maxi)(k, (numext::maxi)(m, n)) < 48)
return;
306#ifdef EIGEN_VECTORIZE_SME
309 EIGEN_IF_CONSTEXPR ((sme_has_gebp_kernel<LhsScalar, RhsScalar>::value)) {
310 evaluateProductBlockingSizesHeuristicForSme<LhsScalar, RhsScalar>(k, m, n);
315 using ResScalar =
typename Traits::ResScalar;
318 k_div = KcFactor * (Traits::mr *
sizeof(LhsScalar) + Traits::nr *
sizeof(RhsScalar)),
319 k_sub = Traits::mr * (Traits::nr *
sizeof(ResScalar))
329 const Index max_kc = numext::maxi<Index>(
static_cast<Index
>(((l1 - k_sub) / k_div) & (~(k_peeling - 1))), 1);
330 const Index old_k = k;
335 k = (k % max_kc) == 0 ? max_kc
336 : max_kc - k_peeling * ((max_kc - 1 - (k % max_kc)) / (k_peeling * (k / max_kc + 1)));
338 eigen_internal_assert(((old_k / k) == (old_k / max_kc)) &&
"the number of sweeps has to remain the same");
341#ifdef EIGEN_VECTORIZE_AVX512
346 const Index phys_l1_eff = convert_index<Index>(phys_l1 * 85 / 100);
347 const Index max_kc_phys = numext::maxi<Index>(((phys_l1_eff - k_sub) / k_div) & (~(k_peeling - 1)), k_peeling);
348 if (max_kc_phys < k) {
349 k = (old_k % max_kc_phys) == 0 ? max_kc_phys
350 : max_kc_phys - k_peeling * ((max_kc_phys - 1 - (old_k % max_kc_phys)) /
351 (k_peeling * (old_k / max_kc_phys + 1)));
362#ifdef EIGEN_DEBUG_SMALL_PRODUCT_BLOCKS
363 const Index actual_l2 =
static_cast<Index
>(l3);
365 const Index actual_l2 =
static_cast<Index
>(l2 * 3 / 2);
375 const Index rhs_panel_budget = numext::maxi<Index>(actual_l2,
static_cast<Index
>(l3_per_cpu));
384 const Index lhs_bytes = m * k *
sizeof(LhsScalar);
385 const Index remaining_l1 =
static_cast<Index
>(l1 - k_sub - lhs_bytes);
386 if (remaining_l1 >= Index(Traits::nr *
sizeof(RhsScalar)) * k) {
388 max_nc = remaining_l1 / (k *
sizeof(RhsScalar));
392 max_nc = (3 * rhs_panel_budget) / (2 * 2 * k *
sizeof(RhsScalar));
395 Index nc = numext::mini<Index>(rhs_panel_budget / (2 * k *
sizeof(RhsScalar)), max_nc) & (~(Traits::nr - 1));
398 nc = numext::maxi<Index>(nc, Traits::nr);
404 n = (n % nc) == 0 ? nc : (nc - Traits::nr * ((nc - (n % nc)) / (Traits::nr * (n / nc + 1))));
405 }
else if (old_k == k) {
415 Index problem_size = k * n *
sizeof(LhsScalar);
416 Index actual_lm = actual_l2;
418 if (problem_size <= 1024) {
421 actual_lm =
static_cast<Index
>(l1);
422 }
else if (l3 != 0 && problem_size <= l1) {
427 actual_lm =
static_cast<Index
>(l2);
428 max_mc = (numext::mini<Index>)(576, max_mc);
430 Index mc = (numext::mini<Index>)(actual_lm / (3 * k *
sizeof(LhsScalar)), max_mc);
432 mc -= mc % Traits::mr;
435 m = (m % mc) == 0 ? mc : (mc - Traits::mr * ((mc - (m % mc)) / (Traits::mr * (m / mc + 1))));
440template <
typename Index>
441inline bool useSpecificBlockingSizes(Index& k, Index& m, Index& n) {
442#ifdef EIGEN_TEST_SPECIFIC_BLOCKING_SIZES
443 if (EIGEN_TEST_SPECIFIC_BLOCKING_SIZES) {
444 k = numext::mini<Index>(k, EIGEN_TEST_SPECIFIC_BLOCKING_SIZE_K);
445 m = numext::mini<Index>(m, EIGEN_TEST_SPECIFIC_BLOCKING_SIZE_M);
446 n = numext::mini<Index>(n, EIGEN_TEST_SPECIFIC_BLOCKING_SIZE_N);
450 EIGEN_UNUSED_VARIABLE(k);
451 EIGEN_UNUSED_VARIABLE(m);
452 EIGEN_UNUSED_VARIABLE(n);
475template <
typename LhsScalar,
typename RhsScalar,
int KcFactor,
typename Index>
476void computeProductBlockingSizes(Index& k, Index& m, Index& n, Index num_threads = 1) {
477 if (!useSpecificBlockingSizes(k, m, n)) {
478 evaluateProductBlockingSizesHeuristic<LhsScalar, RhsScalar, KcFactor, Index>(k, m, n, num_threads);
482template <
typename LhsScalar,
typename RhsScalar,
typename Index>
483inline void computeProductBlockingSizes(Index& k, Index& m, Index& n, Index num_threads = 1) {
484 computeProductBlockingSizes<LhsScalar, RhsScalar, 1, Index>(k, m, n, num_threads);
487template <
typename RhsPacket,
typename RhsPacketx4,
int registers_taken>
488struct RhsPanelHelper {
490 static constexpr int remaining_registers =
491 (std::max)(
int(EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS) - registers_taken, 0);
494 using type = std::conditional_t<remaining_registers >= 4, RhsPacketx4, RhsPacket>;
497template <
typename Packet>
499 Packet B_0, B1, B2, B3;
500 const Packet& get(
const FixedInt<0>&)
const {
return B_0; }
501 const Packet& get(
const FixedInt<1>&)
const {
return B1; }
502 const Packet& get(
const FixedInt<2>&)
const {
return B2; }
503 const Packet& get(
const FixedInt<3>&)
const {
return B3; }
506template <
int N,
typename T1,
typename T2,
typename T3>
507struct packet_conditional {
511template <
typename T1,
typename T2,
typename T3>
512struct packet_conditional<GEBPPacketFull, T1, T2, T3> {
516template <
typename T1,
typename T2,
typename T3>
517struct packet_conditional<GEBPPacketHalf, T1, T2, T3> {
521#define PACKET_DECL_COND_POSTFIX(postfix, name, packet_size) \
522 typedef typename packet_conditional< \
523 packet_size, typename packet_traits<name##Scalar>::type, typename packet_traits<name##Scalar>::half, \
524 typename unpacket_traits<typename packet_traits<name##Scalar>::half>::half>::type name##Packet##postfix
526#define PACKET_DECL_COND(name, packet_size) \
527 typedef typename packet_conditional< \
528 packet_size, typename packet_traits<name##Scalar>::type, typename packet_traits<name##Scalar>::half, \
529 typename unpacket_traits<typename packet_traits<name##Scalar>::half>::half>::type name##Packet
531#define PACKET_DECL_COND_SCALAR_POSTFIX(postfix, packet_size) \
532 typedef typename packet_conditional< \
533 packet_size, typename packet_traits<Scalar>::type, typename packet_traits<Scalar>::half, \
534 typename unpacket_traits<typename packet_traits<Scalar>::half>::half>::type ScalarPacket##postfix
536#define PACKET_DECL_COND_SCALAR(packet_size) \
537 typedef typename packet_conditional< \
538 packet_size, typename packet_traits<Scalar>::type, typename packet_traits<Scalar>::half, \
539 typename unpacket_traits<typename packet_traits<Scalar>::half>::half>::type ScalarPacket
551template <
typename LhsScalar_,
typename RhsScalar_,
bool ConjLhs_,
bool ConjRhs_,
int Arch,
int PacketSize_>
554 using LhsScalar = LhsScalar_;
555 using RhsScalar = RhsScalar_;
556 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
558 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
559 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
560 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
565 Vectorizable = unpacket_traits<LhsPacket_>::vectorizable && unpacket_traits<RhsPacket_>::vectorizable,
566 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
567 RhsPacketSize = Vectorizable ? unpacket_traits<RhsPacket_>::size : 1,
568 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
570 NumberOfRegisters = EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS,
576 default_mr = (plain_enum_min(16, NumberOfRegisters) / 2 / nr) * LhsPacketSize,
577#
if defined(EIGEN_HAS_SINGLE_INSTRUCTION_MADD) && !defined(EIGEN_VECTORIZE_ALTIVEC) && \
578 !defined(EIGEN_VECTORIZE_VSX) && ((!EIGEN_COMP_MSVC) || (EIGEN_COMP_MSVC >= 1914))
583 mr = Vectorizable ? 3 * LhsPacketSize : default_mr,
588 LhsProgress = LhsPacketSize,
592 using LhsPacket = std::conditional_t<Vectorizable, LhsPacket_, LhsScalar>;
593 using RhsPacket = std::conditional_t<Vectorizable, RhsPacket_, RhsScalar>;
594 using ResPacket = std::conditional_t<Vectorizable, ResPacket_, ResScalar>;
595 using LhsPacket4Packing = LhsPacket;
597 using RhsPacketx4 = QuadPacket<RhsPacket>;
598 using AccPacket = ResPacket;
600 EIGEN_STRONG_INLINE
void initAcc(AccPacket& p)
const { p = pset1<ResPacket>(ResScalar(0)); }
602 template <
typename RhsPacketType>
603 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketType& dest)
const {
604 dest = pset1<RhsPacketType>(*b);
607 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const {
608 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
611 template <
typename RhsPacketType>
612 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, RhsPacketType& dest)
const {
616 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
618 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, RhsPacket& dest)
const { dest = ploadquad<RhsPacket>(b); }
620 template <
typename LhsPacketType>
621 EIGEN_STRONG_INLINE
void loadLhs(
const LhsScalar* a, LhsPacketType& dest)
const {
622 dest = pload<LhsPacketType>(a);
625 template <
typename LhsPacketType>
626 EIGEN_STRONG_INLINE
void loadLhsUnaligned(
const LhsScalar* a, LhsPacketType& dest)
const {
627 dest = ploadu<LhsPacketType>(a);
630 template <
typename LhsPacketType,
typename RhsPacketType,
typename AccPacketType,
typename LaneIdType>
631 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketType& b, AccPacketType& c, RhsPacketType& tmp,
632 const LaneIdType&)
const {
633 conj_helper<LhsPacketType, RhsPacketType, ConjLhs, ConjRhs> cj;
638#ifdef EIGEN_HAS_SINGLE_INSTRUCTION_MADD
639 EIGEN_UNUSED_VARIABLE(tmp);
640 c = cj.pmadd(a, b, c);
643 tmp = cj.pmul(a, tmp);
648 template <
typename LhsPacketType,
typename AccPacketType,
typename LaneIdType>
649 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
650 const LaneIdType& lane)
const {
651 madd(a, b.get(lane), c, tmp, lane);
654 EIGEN_STRONG_INLINE
void acc(
const AccPacket& c,
const ResPacket& alpha, ResPacket& r)
const {
655 r = pmadd(c, alpha, r);
658 template <
typename ResPacketHalf>
659 EIGEN_STRONG_INLINE
void acc(
const ResPacketHalf& c,
const ResPacketHalf& alpha, ResPacketHalf& r)
const {
660 r = pmadd(c, alpha, r);
664template <
typename RealScalar,
bool ConjLhs_,
int Arch,
int PacketSize_>
665class gebp_traits<std::complex<RealScalar>, RealScalar, ConjLhs_, false, Arch, PacketSize_> {
667 using LhsScalar = std::complex<RealScalar>;
668 using RhsScalar = RealScalar;
669 using ResScalar =
typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType;
671 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
672 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
673 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
678 Vectorizable = unpacket_traits<LhsPacket_>::vectorizable && unpacket_traits<RhsPacket_>::vectorizable,
679 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
680 RhsPacketSize = Vectorizable ? unpacket_traits<RhsPacket_>::size : 1,
681 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
683 NumberOfRegisters = EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS,
685#if defined(EIGEN_HAS_SINGLE_INSTRUCTION_MADD) && !defined(EIGEN_VECTORIZE_ALTIVEC) && !defined(EIGEN_VECTORIZE_VSX)
687 mr = 3 * LhsPacketSize,
689 mr = (plain_enum_min(16, NumberOfRegisters) / 2 / nr) * LhsPacketSize,
692 LhsProgress = LhsPacketSize,
696 using LhsPacket = std::conditional_t<Vectorizable, LhsPacket_, LhsScalar>;
697 using RhsPacket = std::conditional_t<Vectorizable, RhsPacket_, RhsScalar>;
698 using ResPacket = std::conditional_t<Vectorizable, ResPacket_, ResScalar>;
699 using LhsPacket4Packing = LhsPacket;
701 using RhsPacketx4 = QuadPacket<RhsPacket>;
703 using AccPacket = ResPacket;
705 EIGEN_STRONG_INLINE
void initAcc(AccPacket& p)
const { p = pset1<ResPacket>(ResScalar(0)); }
707 template <
typename RhsPacketType>
708 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketType& dest)
const {
709 dest = pset1<RhsPacketType>(*b);
712 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const {
713 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
716 template <
typename RhsPacketType>
717 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, RhsPacketType& dest)
const {
721 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
723 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, RhsPacket& dest)
const {
724 loadRhsQuad_impl(b, dest, bool_constant<RhsPacketSize == 16>());
727 EIGEN_STRONG_INLINE
void loadRhsQuad_impl(
const RhsScalar* b, RhsPacket& dest,
const std::true_type&)
const {
729 RhsScalar tmp[4] = {b[0], b[0], b[1], b[1]};
730 dest = ploadquad<RhsPacket>(tmp);
733 EIGEN_STRONG_INLINE
void loadRhsQuad_impl(
const RhsScalar* b, RhsPacket& dest,
const std::false_type&)
const {
734 eigen_internal_assert(RhsPacketSize <= 8);
735 dest = pset1<RhsPacket>(*b);
738 EIGEN_STRONG_INLINE
void loadLhs(
const LhsScalar* a, LhsPacket& dest)
const { dest = pload<LhsPacket>(a); }
740 template <
typename LhsPacketType>
741 EIGEN_STRONG_INLINE
void loadLhsUnaligned(
const LhsScalar* a, LhsPacketType& dest)
const {
742 dest = ploadu<LhsPacketType>(a);
745 template <
typename LhsPacketType,
typename RhsPacketType,
typename AccPacketType,
typename LaneIdType>
746 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketType& b, AccPacketType& c, RhsPacketType& tmp,
747 const LaneIdType&)
const {
748 madd_impl(a, b, c, tmp, bool_constant<Vectorizable>());
751 template <
typename LhsPacketType,
typename RhsPacketType,
typename AccPacketType>
752 EIGEN_STRONG_INLINE
void madd_impl(
const LhsPacketType& a,
const RhsPacketType& b, AccPacketType& c,
753 RhsPacketType& tmp,
const std::true_type&)
const {
754#ifdef EIGEN_HAS_SINGLE_INSTRUCTION_MADD
755 EIGEN_UNUSED_VARIABLE(tmp);
756 c.v = pmadd(a.v, b, c.v);
759 tmp = pmul(a.v, tmp);
760 c.v = padd(c.v, tmp);
764 EIGEN_STRONG_INLINE
void madd_impl(
const LhsScalar& a,
const RhsScalar& b, ResScalar& c, RhsScalar& ,
765 const std::false_type&)
const {
769 template <
typename LhsPacketType,
typename AccPacketType,
typename LaneIdType>
770 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
771 const LaneIdType& lane)
const {
772 madd(a, b.get(lane), c, tmp, lane);
775 template <
typename ResPacketType,
typename AccPacketType>
776 EIGEN_STRONG_INLINE
void acc(
const AccPacketType& c,
const ResPacketType& alpha, ResPacketType& r)
const {
777 conj_helper<ResPacketType, ResPacketType, ConjLhs, false> cj;
778 r = cj.pmadd(c, alpha, r);
782template <
typename Packet>
788template <
typename Packet>
789DoublePacket<Packet> padd(
const DoublePacket<Packet>& a,
const DoublePacket<Packet>& b) {
790 DoublePacket<Packet> res;
791 res.first = padd(a.first, b.first);
792 res.second = padd(a.second, b.second);
796template <typename Packet, std::enable_if_t<unpacket_traits<Packet>::size <= 8,
int> = 0>
797const DoublePacket<Packet>& predux_half(
const DoublePacket<Packet>& a) {
801template <typename Packet, std::enable_if_t<unpacket_traits<Packet>::size >= 16 &&
802 !NumTraits<typename unpacket_traits<Packet>::type>::IsComplex,
804DoublePacket<typename unpacket_traits<Packet>::half> predux_half(
const DoublePacket<Packet>& a) {
806 DoublePacket<typename unpacket_traits<Packet>::half> res;
807 using Cplx = std::complex<typename unpacket_traits<Packet>::type>;
808 using CplxPacket =
typename packet_traits<Cplx>::type;
809 res.first = predux_half(CplxPacket(a.first)).v;
810 res.second = predux_half(CplxPacket(a.second)).v;
815template <typename Scalar, typename RealPacket, std::enable_if_t<unpacket_traits<RealPacket>::size <= 8,
int> = 0>
816void loadQuadToDoublePacket(
const Scalar* b, DoublePacket<RealPacket>& dest) {
817 dest.first = pset1<RealPacket>(numext::real(*b));
818 dest.second = pset1<RealPacket>(numext::imag(*b));
825template <
typename Scalar,
typename RealPacket, std::enable_if_t<(unpacket_traits<RealPacket>::size > 8),
int> = 0>
826void loadQuadToDoublePacket(
const Scalar* b, DoublePacket<RealPacket>& dest) {
827 using RealScalar =
typename NumTraits<Scalar>::Real;
828 constexpr int kQuads = unpacket_traits<RealPacket>::size / 4;
829 RealScalar r[kQuads], i[kQuads];
830 for (
int j = 0; j < kQuads; ++j) {
831 r[j] = numext::real(b[j / 2]);
832 i[j] = numext::imag(b[j / 2]);
834 dest.first = ploadquad<RealPacket>(r);
835 dest.second = ploadquad<RealPacket>(i);
838template <
typename Packet>
839struct unpacket_traits<DoublePacket<Packet>> {
840 using half = DoublePacket<typename unpacket_traits<Packet>::half>;
841 enum { size = 2 * unpacket_traits<Packet>::size };
844template <
typename RealScalar,
bool ConjLhs_,
bool ConjRhs_,
int Arch,
int PacketSize_>
845class gebp_traits<std::complex<RealScalar>, std::complex<RealScalar>, ConjLhs_, ConjRhs_, Arch, PacketSize_> {
847 using Scalar = std::complex<RealScalar>;
848 using LhsScalar = std::complex<RealScalar>;
849 using RhsScalar = std::complex<RealScalar>;
850 using ResScalar = std::complex<RealScalar>;
852 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
853 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
854 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
855 PACKET_DECL_COND(Real, PacketSize_);
856 PACKET_DECL_COND_SCALAR(PacketSize_);
861 Vectorizable = unpacket_traits<RealPacket>::vectorizable && unpacket_traits<ScalarPacket>::vectorizable,
862 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
863 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
864 RhsPacketSize = Vectorizable ? unpacket_traits<RhsScalar>::size : 1,
865 RealPacketSize = Vectorizable ? unpacket_traits<RealPacket>::size : 1,
866 NumberOfRegisters = EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS,
869 mr = (plain_enum_min(16, NumberOfRegisters) / 2 / nr) * ResPacketSize,
871 LhsProgress = ResPacketSize,
875 using DoublePacketType = DoublePacket<RealPacket>;
877 using LhsPacket4Packing = std::conditional_t<Vectorizable, ScalarPacket, Scalar>;
878 using LhsPacket = std::conditional_t<Vectorizable, RealPacket, Scalar>;
879 using RhsPacket = std::conditional_t<Vectorizable, DoublePacketType, Scalar>;
880 using ResPacket = std::conditional_t<Vectorizable, ScalarPacket, Scalar>;
881 using AccPacket = std::conditional_t<Vectorizable, DoublePacketType, Scalar>;
884 using RhsPacketx4 = QuadPacket<RhsPacket>;
886 EIGEN_STRONG_INLINE
void initAcc(Scalar& p)
const { p = Scalar(0); }
888 EIGEN_STRONG_INLINE
void initAcc(DoublePacketType& p)
const {
889 p.first = pset1<RealPacket>(RealScalar(0));
890 p.second = pset1<RealPacket>(RealScalar(0));
894 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, ScalarPacket& dest)
const { dest = pset1<ScalarPacket>(*b); }
897 template <
typename RealPacketType>
898 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, DoublePacket<RealPacketType>& dest)
const {
899 dest.first = pset1<RealPacketType>(numext::real(*b));
900 dest.second = pset1<RealPacketType>(numext::imag(*b));
903 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const {
904 loadRhs(b, dest.B_0);
905 loadRhs(b + 1, dest.B1);
906 loadRhs(b + 2, dest.B2);
907 loadRhs(b + 3, dest.B3);
911 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, ScalarPacket& dest)
const { loadRhs(b, dest); }
914 template <
typename RealPacketType>
915 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, DoublePacket<RealPacketType>& dest)
const {
919 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
921 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, ResPacket& dest)
const { loadRhs(b, dest); }
922 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, DoublePacketType& dest)
const {
923 loadQuadToDoublePacket(b, dest);
927 EIGEN_STRONG_INLINE
void loadLhs(
const LhsScalar* a, LhsPacket& dest)
const {
928 dest = pload<LhsPacket>((
const typename unpacket_traits<LhsPacket>::type*)(a));
931 template <
typename LhsPacketType>
932 EIGEN_STRONG_INLINE
void loadLhsUnaligned(
const LhsScalar* a, LhsPacketType& dest)
const {
933 dest = ploadu<LhsPacketType>((
const typename unpacket_traits<LhsPacketType>::type*)(a));
936 template <
typename LhsPacketType,
typename RhsPacketType,
typename ResPacketType,
typename TmpType,
938 EIGEN_STRONG_INLINE std::enable_if_t<!std::is_same<RhsPacketType, RhsPacketx4>::value> madd(
939 const LhsPacketType& a,
const RhsPacketType& b, DoublePacket<ResPacketType>& c, TmpType& ,
940 const LaneIdType&)
const {
941 c.first = pmadd(a, b.first, c.first);
942 c.second = pmadd(a, b.second, c.second);
945 template <
typename LaneIdType>
946 EIGEN_STRONG_INLINE
void madd(
const LhsPacket& a,
const RhsPacket& b, ResPacket& c, RhsPacket& ,
947 const LaneIdType&)
const {
948 c = cj.pmadd(a, b, c);
951 template <
typename LhsPacketType,
typename AccPacketType,
typename LaneIdType>
952 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
953 const LaneIdType& lane)
const {
954 madd(a, b.get(lane), c, tmp, lane);
957 EIGEN_STRONG_INLINE
void acc(
const Scalar& c,
const Scalar& alpha, Scalar& r)
const { r += alpha * c; }
959 template <
typename RealPacketType,
typename ResPacketType>
960 EIGEN_STRONG_INLINE
void acc(
const DoublePacket<RealPacketType>& c,
const ResPacketType& alpha,
961 ResPacketType& r)
const {
964 EIGEN_IF_CONSTEXPR ((!ConjLhs) && (!ConjRhs)) {
965 tmp = pcplxflip(pconj(ResPacketType(c.second)));
966 tmp = padd(ResPacketType(c.first), tmp);
967 }
else EIGEN_IF_CONSTEXPR ((!ConjLhs) && (ConjRhs)) {
968 tmp = pconj(pcplxflip(ResPacketType(c.second)));
969 tmp = padd(ResPacketType(c.first), tmp);
970 }
else EIGEN_IF_CONSTEXPR ((ConjLhs) && (!ConjRhs)) {
971 tmp = pcplxflip(ResPacketType(c.second));
972 tmp = padd(pconj(ResPacketType(c.first)), tmp);
974 tmp = pcplxflip(ResPacketType(c.second));
975 tmp = psub(pconj(ResPacketType(c.first)), tmp);
978 r = pmadd(tmp, alpha, r);
982 conj_helper<LhsScalar, RhsScalar, ConjLhs, ConjRhs> cj;
985template <
typename RealScalar,
bool ConjRhs_,
int Arch,
int PacketSize_>
986class gebp_traits<RealScalar, std::complex<RealScalar>, false, ConjRhs_, Arch, PacketSize_> {
988 using Scalar = std::complex<RealScalar>;
989 using LhsScalar = RealScalar;
990 using RhsScalar = Scalar;
991 using ResScalar = Scalar;
993 PACKET_DECL_COND_POSTFIX(_, Lhs, PacketSize_);
994 PACKET_DECL_COND_POSTFIX(_, Rhs, PacketSize_);
995 PACKET_DECL_COND_POSTFIX(_, Res, PacketSize_);
996 PACKET_DECL_COND_POSTFIX(_, Real, PacketSize_);
997 PACKET_DECL_COND_SCALAR_POSTFIX(_, PacketSize_);
999#undef PACKET_DECL_COND_SCALAR_POSTFIX
1000#undef PACKET_DECL_COND_POSTFIX
1001#undef PACKET_DECL_COND_SCALAR
1002#undef PACKET_DECL_COND
1007 Vectorizable = unpacket_traits<RealPacket_>::vectorizable && unpacket_traits<ScalarPacket_>::vectorizable,
1008 LhsPacketSize = Vectorizable ? unpacket_traits<LhsPacket_>::size : 1,
1009 RhsPacketSize = Vectorizable ? unpacket_traits<RhsPacket_>::size : 1,
1010 ResPacketSize = Vectorizable ? unpacket_traits<ResPacket_>::size : 1,
1012 NumberOfRegisters = EIGEN_ARCH_DEFAULT_NUMBER_OF_REGISTERS,
1015 mr = (plain_enum_min(16, NumberOfRegisters) / 2 / nr) * ResPacketSize,
1017 LhsProgress = ResPacketSize,
1021 using LhsPacket = std::conditional_t<Vectorizable, LhsPacket_, LhsScalar>;
1022 using RhsPacket = std::conditional_t<Vectorizable, RhsPacket_, RhsScalar>;
1023 using ResPacket = std::conditional_t<Vectorizable, ResPacket_, ResScalar>;
1024 using LhsPacket4Packing = LhsPacket;
1025 using RhsPacketx4 = QuadPacket<RhsPacket>;
1026 using AccPacket = ResPacket;
1028 EIGEN_STRONG_INLINE
void initAcc(AccPacket& p)
const { p = pset1<ResPacket>(ResScalar(0)); }
1030 template <
typename RhsPacketType>
1031 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketType& dest)
const {
1032 dest = pset1<RhsPacketType>(*b);
1035 EIGEN_STRONG_INLINE
void loadRhs(
const RhsScalar* b, RhsPacketx4& dest)
const {
1036 pbroadcast4(b, dest.B_0, dest.B1, dest.B2, dest.B3);
1039 template <
typename RhsPacketType>
1040 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar* b, RhsPacketType& dest)
const {
1044 EIGEN_STRONG_INLINE
void updateRhs(
const RhsScalar*, RhsPacketx4&)
const {}
1046 EIGEN_STRONG_INLINE
void loadLhs(
const LhsScalar* a, LhsPacket& dest)
const { dest = ploaddup<LhsPacket>(a); }
1048 EIGEN_STRONG_INLINE
void loadRhsQuad(
const RhsScalar* b, RhsPacket& dest)
const { dest = ploadquad<RhsPacket>(b); }
1050 template <
typename LhsPacketType>
1051 EIGEN_STRONG_INLINE
void loadLhsUnaligned(
const LhsScalar* a, LhsPacketType& dest)
const {
1052 dest = ploaddup<LhsPacketType>(a);
1055 template <
typename LhsPacketType,
typename RhsPacketType,
typename AccPacketType,
typename LaneIdType>
1056 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketType& b, AccPacketType& c, RhsPacketType& tmp,
1057 const LaneIdType&)
const {
1058 madd_impl(a, b, c, tmp, bool_constant<Vectorizable>());
1061 template <
typename LhsPacketType,
typename RhsPacketType,
typename AccPacketType>
1062 EIGEN_STRONG_INLINE
void madd_impl(
const LhsPacketType& a,
const RhsPacketType& b, AccPacketType& c,
1063 RhsPacketType& tmp,
const std::true_type&)
const {
1064#ifdef EIGEN_HAS_SINGLE_INSTRUCTION_MADD
1065 EIGEN_UNUSED_VARIABLE(tmp);
1066 c.v = pmadd(a, b.v, c.v);
1069 tmp.v = pmul(a, tmp.v);
1074 EIGEN_STRONG_INLINE
void madd_impl(
const LhsScalar& a,
const RhsScalar& b, ResScalar& c, RhsScalar& ,
1075 const std::false_type&)
const {
1079 template <
typename LhsPacketType,
typename AccPacketType,
typename LaneIdType>
1080 EIGEN_STRONG_INLINE
void madd(
const LhsPacketType& a,
const RhsPacketx4& b, AccPacketType& c, RhsPacket& tmp,
1081 const LaneIdType& lane)
const {
1082 madd(a, b.get(lane), c, tmp, lane);
1085 template <
typename ResPacketType,
typename AccPacketType>
1086 EIGEN_STRONG_INLINE
void acc(
const AccPacketType& c,
const ResPacketType& alpha, ResPacketType& r)
const {
1087 conj_helper<ResPacketType, ResPacketType, false, ConjRhs> cj;
1088 r = cj.pmadd(alpha, c, r);
1099template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename DataMapper,
int mr,
int nr,
1100 bool ConjugateLhs,
bool ConjugateRhs>
1102 using Traits = gebp_traits<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs, Architecture::Target>;
1104 gebp_traits<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs, Architecture::Target, GEBPPacketHalf>;
1105 using QuarterTraits =
1106 gebp_traits<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs, Architecture::Target, GEBPPacketQuarter>;
1108 using ResScalar =
typename Traits::ResScalar;
1109 using LhsPacket =
typename Traits::LhsPacket;
1110 using RhsPacket =
typename Traits::RhsPacket;
1111 using ResPacket =
typename Traits::ResPacket;
1112 using AccPacket =
typename Traits::AccPacket;
1113 using RhsPacketx4 =
typename Traits::RhsPacketx4;
1115 using SwappedTraits = gebp_traits<RhsScalar, LhsScalar, ConjugateRhs, ConjugateLhs, Architecture::Target>;
1117 using SLhsPacket =
typename SwappedTraits::LhsPacket;
1118 using SRhsPacket =
typename SwappedTraits::RhsPacket;
1119 using SResPacket =
typename SwappedTraits::ResPacket;
1120 using SAccPacket =
typename SwappedTraits::AccPacket;
1122 using LhsPacketHalf =
typename HalfTraits::LhsPacket;
1123 using RhsPacketHalf =
typename HalfTraits::RhsPacket;
1124 using ResPacketHalf =
typename HalfTraits::ResPacket;
1125 using AccPacketHalf =
typename HalfTraits::AccPacket;
1127 using LhsPacketQuarter =
typename QuarterTraits::LhsPacket;
1128 using RhsPacketQuarter =
typename QuarterTraits::RhsPacket;
1129 using ResPacketQuarter =
typename QuarterTraits::ResPacket;
1130 using AccPacketQuarter =
typename QuarterTraits::AccPacket;
1132 using LinearMapper =
typename DataMapper::LinearMapper;
1135 Vectorizable = Traits::Vectorizable,
1136 LhsProgress = Traits::LhsProgress,
1137 LhsProgressHalf = HalfTraits::LhsProgress,
1138 LhsProgressQuarter = QuarterTraits::LhsProgress,
1139 RhsProgress = Traits::RhsProgress,
1140 RhsProgressHalf = HalfTraits::RhsProgress,
1141 RhsProgressQuarter = QuarterTraits::RhsProgress,
1142 ResPacketSize = Traits::ResPacketSize
1145 EIGEN_DONT_INLINE
void operator()(
const DataMapper& res,
const LhsScalar* blockA,
const RhsScalar* blockB, Index rows,
1146 Index depth, Index cols, ResScalar alpha, Index strideA = -1, Index strideB = -1,
1147 Index offsetA = 0, Index offsetB = 0)
const;
1150template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename DataMapper,
int mr,
int nr,
1151 bool ConjugateLhs,
bool ConjugateRhs,
1152 int SwappedLhsProgress =
1153 gebp_traits<RhsScalar, LhsScalar, ConjugateRhs, ConjugateLhs, Architecture::Target>::LhsProgress>
1154struct last_row_process_16_packets {
1155 using Traits = gebp_traits<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs, Architecture::Target>;
1156 using SwappedTraits = gebp_traits<RhsScalar, LhsScalar, ConjugateRhs, ConjugateLhs, Architecture::Target>;
1158 using ResScalar =
typename Traits::ResScalar;
1159 using SLhsPacket =
typename SwappedTraits::LhsPacket;
1160 using SRhsPacket =
typename SwappedTraits::RhsPacket;
1161 using SResPacket =
typename SwappedTraits::ResPacket;
1162 using SAccPacket =
typename SwappedTraits::AccPacket;
1164 EIGEN_STRONG_INLINE
void operator()(
const DataMapper& res, SwappedTraits& straits,
const LhsScalar* blA,
1165 const RhsScalar* blB, Index depth,
const Index endk, Index i, Index j2,
1166 ResScalar alpha, SAccPacket& C0)
const {
1167 EIGEN_UNUSED_VARIABLE(res);
1168 EIGEN_UNUSED_VARIABLE(straits);
1169 EIGEN_UNUSED_VARIABLE(blA);
1170 EIGEN_UNUSED_VARIABLE(blB);
1171 EIGEN_UNUSED_VARIABLE(depth);
1172 EIGEN_UNUSED_VARIABLE(endk);
1173 EIGEN_UNUSED_VARIABLE(i);
1174 EIGEN_UNUSED_VARIABLE(j2);
1175 EIGEN_UNUSED_VARIABLE(alpha);
1176 EIGEN_UNUSED_VARIABLE(C0);
1180template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename DataMapper,
int mr,
int nr,
1181 bool ConjugateLhs,
bool ConjugateRhs>
1182struct last_row_process_16_packets<LhsScalar, RhsScalar, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs, 16> {
1183 using Traits = gebp_traits<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs, Architecture::Target>;
1184 using SwappedTraits = gebp_traits<RhsScalar, LhsScalar, ConjugateRhs, ConjugateLhs, Architecture::Target>;
1186 using ResScalar =
typename Traits::ResScalar;
1187 using SLhsPacket =
typename SwappedTraits::LhsPacket;
1188 using SRhsPacket =
typename SwappedTraits::RhsPacket;
1189 using SResPacket =
typename SwappedTraits::ResPacket;
1190 using SAccPacket =
typename SwappedTraits::AccPacket;
1192 EIGEN_STRONG_INLINE
void operator()(
const DataMapper& res, SwappedTraits& straits,
const LhsScalar* blA,
1193 const RhsScalar* blB, Index depth,
const Index endk, Index i, Index j2,
1194 ResScalar alpha, SAccPacket& C0)
const {
1195 using SResPacketQuarter =
typename unpacket_traits<typename unpacket_traits<SResPacket>::half>::half;
1196 using SLhsPacketQuarter =
typename unpacket_traits<typename unpacket_traits<SLhsPacket>::half>::half;
1197 using SRhsPacketQuarter =
typename unpacket_traits<typename unpacket_traits<SRhsPacket>::half>::half;
1198 using SAccPacketQuarter =
typename unpacket_traits<typename unpacket_traits<SAccPacket>::half>::half;
1200 SResPacketQuarter R = res.template gatherPacket<SResPacketQuarter>(i, j2);
1201 SResPacketQuarter alphav = pset1<SResPacketQuarter>(alpha);
1203 if (depth - endk > 0) {
1206 SAccPacketQuarter c0 = predux_half(predux_half(C0));
1208 for (Index kk = endk; kk < depth; kk++) {
1209 SLhsPacketQuarter a0;
1210 SRhsPacketQuarter b0;
1211 straits.loadLhsUnaligned(blB, a0);
1212 straits.loadRhs(blA, b0);
1213 straits.madd(a0, b0, c0, b0, fix<0>);
1214 blB += SwappedTraits::LhsProgress / 4;
1217 straits.acc(c0, alphav, R);
1219 straits.acc(predux_half(predux_half(C0)), alphav, R);
1221 res.scatterPacket(i, j2, R);
1228template <
int J,
int MrPackets,
int NrCols,
bool Continue = (J < NrCols)>
1229struct gebp_rhs_cols;
1232template <
int J,
int MrPackets,
int NrCols>
1233struct gebp_rhs_cols<J, MrPackets, NrCols, false> {
1234 template <
typename GEBPTraits,
typename LhsArray,
typename RhsPanelType,
typename RhsPacketType,
typename AccArray,
1236 static EIGEN_ALWAYS_INLINE
void run(GEBPTraits&,
const RhsScalar*, Index, LhsArray&, RhsPanelType&, RhsPacketType&,
1241template <
int J,
int MrPackets,
int NrCols>
1242struct gebp_rhs_cols<J, MrPackets, NrCols, true> {
1243 template <
typename GEBPTraits,
typename LhsArray,
typename RhsPanelType,
typename RhsPacketType,
typename AccArray,
1245 static EIGEN_ALWAYS_INLINE
void run(GEBPTraits& traits,
const RhsScalar* blB, Index rhs_offset, LhsArray& A,
1246 RhsPanelType& rhs_panel, RhsPacketType& T0, AccArray& C) {
1247 constexpr int lane = J % 4;
1248 EIGEN_IF_CONSTEXPR (lane == 0)
1249 traits.loadRhs(blB + (J + rhs_offset) * GEBPTraits::RhsProgress, rhs_panel);
1251 traits.updateRhs(blB + (J + rhs_offset) * GEBPTraits::RhsProgress, rhs_panel);
1253 EIGEN_IF_CONSTEXPR (MrPackets >= 1) traits.madd(A[0], rhs_panel, C[J + 0 * NrCols], T0, fix<lane>);
1254 EIGEN_IF_CONSTEXPR (MrPackets >= 2) traits.madd(A[1], rhs_panel, C[J + 1 * NrCols], T0, fix<lane>);
1255 EIGEN_IF_CONSTEXPR (MrPackets >= 3) traits.madd(A[2], rhs_panel, C[J + 2 * NrCols], T0, fix<lane>);
1257 gebp_rhs_cols<J + 1, MrPackets, NrCols>::run(traits, blB, rhs_offset, A, rhs_panel, T0, C);
1263template <
int K,
int MrPackets,
int NrCols>
1264struct gebp_micro_step {
1265 template <
typename GEBPTraits,
typename LhsScalar_,
typename RhsScalar_,
typename LhsArray,
typename RhsPanelType,
1266 typename RhsPacketType,
typename AccArray>
1267 static EIGEN_ALWAYS_INLINE
void run(GEBPTraits& traits,
const LhsScalar_* blA,
const RhsScalar_* blB, LhsArray& A,
1268 RhsPanelType& rhs_panel, RhsPacketType& T0, AccArray& C) {
1269 constexpr int LhsProg = GEBPTraits::LhsProgress;
1271 EIGEN_IF_CONSTEXPR (MrPackets >= 1) traits.loadLhs(&blA[(0 + MrPackets * K) * LhsProg], A[0]);
1272 EIGEN_IF_CONSTEXPR (MrPackets >= 2) traits.loadLhs(&blA[(1 + MrPackets * K) * LhsProg], A[1]);
1273 EIGEN_IF_CONSTEXPR (MrPackets >= 3) traits.loadLhs(&blA[(2 + MrPackets * K) * LhsProg], A[2]);
1275 gebp_rhs_cols<0, MrPackets, NrCols>::run(traits, blB, Index(NrCols * K), A, rhs_panel, T0, C);
1286template <
int MrPackets, typename GEBPTraits_, typename FullLhsPacket_, typename LhsArray_>
1287EIGEN_ALWAYS_INLINE
void gebp_neon_3p_workaround(LhsArray_& A) {
1288#if EIGEN_ARCH_ARM64 && defined(EIGEN_VECTORIZE_NEON) && EIGEN_GNUC_STRICT_LESS_THAN(9, 0, 0)
1289 using LhsElement = std::remove_all_extents_t<std::remove_reference_t<LhsArray_>>;
1290 constexpr bool apply = GEBPTraits_::Vectorizable && MrPackets == 3 && std::is_same<LhsElement, FullLhsPacket_>::value;
1291 EIGEN_IF_CONSTEXPR (apply) {
1292 __asm__(
"" :
"+w,m"(A[0]),
"+w,m"(A[1]),
"+w,m"(A[2]));
1295 EIGEN_UNUSED_VARIABLE(A);
1302template <
int MrPackets,
int NrCols,
typename GEBPTraits_,
typename FullLhsPacket_,
typename LhsArray_,
1304EIGEN_ALWAYS_INLINE
void gebp_sse_spilling_workaround(LhsArray_& A, AccArray_& ACC) {
1305 EIGEN_UNUSED_VARIABLE(A);
1306 EIGEN_UNUSED_VARIABLE(ACC);
1307#if EIGEN_GNUC_STRICT_AT_LEAST(6, 0, 0) && defined(EIGEN_VECTORIZE_SSE)
1308 using LhsElement = std::remove_all_extents_t<std::remove_reference_t<LhsArray_>>;
1309 constexpr bool apply =
1310 GEBPTraits_::Vectorizable && MrPackets <= 2 && NrCols >= 4 && std::is_same<LhsElement, FullLhsPacket_>::value;
1311 EIGEN_IF_CONSTEXPR (apply) {
1312#ifdef EIGEN_HAS_CXX17_IFCONSTEXPR
1313 using AccElement = std::decay_t<
decltype(ACC[0])>;
1314 constexpr bool pin_acc = std::is_same<AccElement, FullLhsPacket_>::value && MrPackets == 2 && NrCols == 4;
1315 if constexpr (pin_acc) {
1317 :
"+x"(ACC[0]),
"+x"(ACC[1]),
"+x"(ACC[2]),
"+x"(ACC[3]),
"+x"(ACC[4]),
"+x"(ACC[5]),
"+x"(ACC[6]),
1321 EIGEN_IF_CONSTEXPR (MrPackets == 2) {
1322 __asm__(
"" :
"+x,m"(A[0]),
"+x,m"(A[1]));
1331template <
int MrPackets,
int NrCols>
1332struct gebp_peeled_loop {
1333 template <
typename GEBPTraits,
typename LhsScalar_,
typename RhsScalar_,
typename LhsArray,
typename RhsPanelType,
1334 typename RhsPacketType,
typename AccArray,
typename AccArrayD,
typename FullLhsPacket>
1335 static EIGEN_ALWAYS_INLINE
void run(GEBPTraits& traits,
const LhsScalar_* blA,
const RhsScalar_* blB, LhsArray& A,
1336 RhsPanelType& rhs_panel, RhsPacketType& T0, AccArray& C, AccArrayD& D) {
1337 constexpr bool use_double_accum = (MrPackets == 1 && NrCols == 4);
1340 EIGEN_IF_CONSTEXPR (NrCols == 4) {
1341 internal::prefetch(blB + (48 + 0));
1345#define EIGEN_GEBP_DO_STEP(KVAL, ACC) \
1347 gebp_micro_step<KVAL, MrPackets, NrCols>::run(traits, blA, blB, A, rhs_panel, T0, ACC); \
1348 gebp_neon_3p_workaround<MrPackets, GEBPTraits, FullLhsPacket>(A); \
1349 gebp_sse_spilling_workaround<MrPackets, NrCols, GEBPTraits, FullLhsPacket>(A, ACC); \
1351 EIGEN_IF_CONSTEXPR ((MrPackets == 2 || MrPackets == 3) && NrCols == 4) { \
1352 internal::prefetch(blA + (MrPackets * KVAL + 16) * GEBPTraits::LhsProgress); \
1353 if (EIGEN_ARCH_ARM || EIGEN_ARCH_MIPS) { \
1354 internal::prefetch(blB + (NrCols * KVAL + 16) * GEBPTraits::RhsProgress); \
1359 EIGEN_IF_CONSTEXPR (use_double_accum) {
1360 EIGEN_GEBP_DO_STEP(0, C);
1361 EIGEN_GEBP_DO_STEP(1, D);
1362 EIGEN_GEBP_DO_STEP(2, C);
1363 EIGEN_GEBP_DO_STEP(3, D);
1364 EIGEN_IF_CONSTEXPR (NrCols == 4) {
1365 internal::prefetch(blB + (48 + 16));
1367 EIGEN_GEBP_DO_STEP(4, C);
1368 EIGEN_GEBP_DO_STEP(5, D);
1369 EIGEN_GEBP_DO_STEP(6, C);
1370 EIGEN_GEBP_DO_STEP(7, D);
1372 EIGEN_GEBP_DO_STEP(0, C);
1373 EIGEN_GEBP_DO_STEP(1, C);
1374 EIGEN_GEBP_DO_STEP(2, C);
1375 EIGEN_GEBP_DO_STEP(3, C);
1376 EIGEN_IF_CONSTEXPR (NrCols == 4 && MrPackets == 2) {
1377 internal::prefetch(blB + (48 + 16));
1379 EIGEN_GEBP_DO_STEP(4, C);
1380 EIGEN_GEBP_DO_STEP(5, C);
1381 EIGEN_GEBP_DO_STEP(6, C);
1382 EIGEN_GEBP_DO_STEP(7, C);
1385#undef EIGEN_GEBP_DO_STEP
1392template <
int MrPackets,
int NrCols,
typename GEBPTraits,
typename LhsScalar_,
typename RhsScalar_,
typename ResScalar_,
1393 typename Index_,
typename DataMapper_,
typename LinearMapper_,
typename FullLhsPacket>
1394EIGEN_ALWAYS_INLINE
void gebp_micro_panel_impl(GEBPTraits& traits,
const DataMapper_& res,
const LhsScalar_* blockA,
1395 const RhsScalar_* blockB, ResScalar_ alpha, Index_ i, Index_ j2,
1396 Index_ depth, Index_ strideA, Index_ strideB, Index_ offsetA,
1397 Index_ offsetB,
int prefetch_res_offset, Index_ peeled_kc,
int pk) {
1398 using LhsPacketLocal =
typename GEBPTraits::LhsPacket;
1399 using RhsPacketLocal =
typename GEBPTraits::RhsPacket;
1400 using ResPacketLocal =
typename GEBPTraits::ResPacket;
1401 using AccPacketLocal =
typename GEBPTraits::AccPacket;
1402 using RhsPacketx4Local =
typename GEBPTraits::RhsPacketx4;
1403 constexpr int LhsProg = GEBPTraits::LhsProgress;
1404 constexpr int RhsProg = GEBPTraits::RhsProgress;
1405 constexpr int ResPacketSz = GEBPTraits::ResPacketSize;
1408 using RhsPanelType = std::conditional_t<
1409 NrCols == 1, RhsPacketLocal,
1410 typename RhsPanelHelper<RhsPacketLocal, RhsPacketx4Local, MrPackets * NrCols + MrPackets>::type>;
1412 const LhsScalar_* blA = &blockA[i * strideA + offsetA * (MrPackets * LhsProg)];
1418#ifdef EIGEN_HAS_CXX17_IFCONSTEXPR
1419 constexpr int CSize = MrPackets * NrCols;
1421 constexpr int CSize = 3 * NrCols > MrPackets * NrCols ? 3 * NrCols : MrPackets * NrCols;
1423 alignas(AccPacketLocal) AccPacketLocal C[CSize];
1424 for (
int n = 0; n < MrPackets * NrCols; ++n) traits.initAcc(C[n]);
1427 constexpr bool use_double_accum = (MrPackets == 1 && NrCols == 4);
1428#ifdef EIGEN_HAS_CXX17_IFCONSTEXPR
1429 alignas(AccPacketLocal) AccPacketLocal D[use_double_accum ? NrCols : 1];
1433 alignas(AccPacketLocal) AccPacketLocal D[CSize];
1435 EIGEN_IF_CONSTEXPR (use_double_accum) {
1436 for (
int n = 0; n < NrCols; ++n) traits.initAcc(D[n]);
1440 for (
int j = 0; j < NrCols; ++j) res.getLinearMapper(i, j2 + j).prefetch(NrCols > 1 ? prefetch_res_offset : 0);
1443 const RhsScalar_* blB = &blockB[j2 * strideB + offsetB * NrCols];
1447#ifdef EIGEN_HAS_CXX17_IFCONSTEXPR
1448 alignas(LhsPacketLocal) LhsPacketLocal A[MrPackets];
1450 alignas(LhsPacketLocal) LhsPacketLocal A[3];
1454#if defined(EIGEN_VECTORIZE_RVV10) && EIGEN_GNUC_STRICT_AT_LEAST(15, 0, 0) && EIGEN_GNUC_STRICT_LESS_THAN(17, 0, 0)
1458 for (Index_ k = 0; k < peeled_kc; k += pk) {
1459 alignas(RhsPanelType) RhsPanelType rhs_panel;
1460 alignas(RhsPacketLocal) RhsPacketLocal T0;
1462 gebp_peeled_loop<MrPackets, NrCols>::template run<GEBPTraits, LhsScalar_, RhsScalar_,
decltype(A), RhsPanelType,
1463 RhsPacketLocal,
decltype(C),
decltype(D), FullLhsPacket>(
1464 traits, blA, blB, A, rhs_panel, T0, C, D);
1466 blB += pk * NrCols * RhsProg;
1467 blA += pk * MrPackets * LhsProg;
1471 EIGEN_IF_CONSTEXPR (use_double_accum) {
1472 for (
int n = 0; n < NrCols; ++n) C[n] = padd(C[n], D[n]);
1476 for (Index_ k = peeled_kc; k < depth; k++) {
1477 alignas(RhsPanelType) RhsPanelType rhs_panel;
1478 alignas(RhsPacketLocal) RhsPacketLocal T0;
1480 gebp_micro_step<0, MrPackets, NrCols>::run(traits, blA, blB, A, rhs_panel, T0, C);
1482 blB += NrCols * RhsProg;
1483 blA += MrPackets * LhsProg;
1487 alignas(ResPacketLocal) ResPacketLocal alphav = pset1<ResPacketLocal>(alpha);
1488 for (
int j = 0; j < NrCols; ++j) {
1489 LinearMapper_ r = res.getLinearMapper(i, j2 + j);
1490 for (
int p = 0; p < MrPackets; ++p) {
1491 alignas(ResPacketLocal) ResPacketLocal R = r.template loadPacket<ResPacketLocal>(p * ResPacketSz);
1492 traits.acc(C[j + p * NrCols], alphav, R);
1493 r.storePacket(p * ResPacketSz, R);
1505#if EIGEN_COMP_GNUC_STRICT && EIGEN_ARCH_ARM64 && !defined(EIGEN_DONT_DISABLE_GEBP_INSN_SCHEDULING)
1506#pragma GCC push_options
1507#pragma GCC optimize("no-schedule-insns")
1508#define EIGEN_GEBP_DISABLED_INSN_SCHEDULING
1510template <
typename LhsScalar,
typename RhsScalar,
typename Index,
typename DataMapper,
int mr,
int nr,
1511 bool ConjugateLhs,
bool ConjugateRhs>
1512EIGEN_DONT_INLINE
void gebp_kernel<LhsScalar, RhsScalar, Index, DataMapper, mr, nr, ConjugateLhs,
1513 ConjugateRhs>::operator()(
const DataMapper& res,
const LhsScalar* blockA,
1514 const RhsScalar* blockB, Index rows, Index depth,
1515 Index cols, ResScalar alpha, Index strideA, Index strideB,
1516 Index offsetA, Index offsetB)
const {
1518 SwappedTraits straits;
1520 if (strideA == -1) strideA = depth;
1521 if (strideB == -1) strideB = depth;
1522 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
1523 Index packet_cols4 = nr >= 4 ? (cols / 4) * 4 : 0;
1524 Index packet_cols8 = nr >= 8 ? (cols / 8) * 8 : 0;
1525 const Index peeled_mc3 = mr >= 3 * Traits::LhsProgress ? (rows / (3 * LhsProgress)) * (3 * LhsProgress) : 0;
1526 const Index peeled_mc2 =
1527 mr >= 2 * Traits::LhsProgress ? peeled_mc3 + ((rows - peeled_mc3) / (2 * LhsProgress)) * (2 * LhsProgress) : 0;
1528 const Index peeled_mc1 =
1529 mr >= 1 * Traits::LhsProgress ? peeled_mc2 + ((rows - peeled_mc2) / (1 * LhsProgress)) * (1 * LhsProgress) : 0;
1530 const Index peeled_mc_half =
1531 mr >= LhsProgressHalf ? peeled_mc1 + ((rows - peeled_mc1) / (LhsProgressHalf)) * (LhsProgressHalf) : 0;
1532 const Index peeled_mc_quarter =
1533 mr >= LhsProgressQuarter
1534 ? peeled_mc_half + ((rows - peeled_mc_half) / (LhsProgressQuarter)) * (LhsProgressQuarter)
1537 const Index peeled_kc = depth & ~(pk - 1);
1538 const int prefetch_res_offset = 32 /
sizeof(ResScalar);
1545 auto micro_panel = [&](
auto mrp_tag,
auto nrc_tag,
auto& local_traits, Index i, Index j2) EIGEN_LAMBDA_ALWAYS_INLINE {
1546 constexpr int MrP =
decltype(mrp_tag)::value;
1547 constexpr int NrC =
decltype(nrc_tag)::value;
1548 using LTraits = std::remove_reference_t<
decltype(local_traits)>;
1549 gebp_micro_panel_impl<MrP, NrC, LTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1550 LhsPacket>(local_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB, offsetA,
1551 offsetB, prefetch_res_offset, peeled_kc, pk);
1564 std::ptrdiff_t l1, l2, l3;
1565 manage_caching_sizes(GetAction, &l1, &l2, &l3);
1566#if EIGEN_ARCH_i386_OR_x86_64
1567 lhs_budget =
static_cast<Index
>(l2 / 2);
1569 lhs_budget =
static_cast<Index
>(l1);
1574 EIGEN_IF_CONSTEXPR (mr >= 3 * Traits::LhsProgress) {
1575 const Index rhs_block =
sizeof(ResScalar) * mr * nr + depth * nr *
sizeof(RhsScalar);
1576 const Index lhs_strip = depth *
sizeof(LhsScalar) * 3 * LhsProgress;
1577 const Index lhs_avail = (lhs_budget > rhs_block) ? (lhs_budget - rhs_block) : 0;
1578 const Index actual_panel_rows = (lhs_avail >= peeled_mc3 * depth *
static_cast<Index
>(
sizeof(LhsScalar)))
1580 : (3 * LhsProgress) * std::max<Index>(1, lhs_avail / lhs_strip);
1581 for (Index i1 = 0; i1 < peeled_mc3; i1 += actual_panel_rows) {
1582 const Index actual_panel_end = (std::min)(i1 + actual_panel_rows, peeled_mc3);
1583 EIGEN_IF_CONSTEXPR (nr >= 8) {
1584 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1585 for (Index i = i1; i < actual_panel_end; i += 3 * LhsProgress) {
1586 micro_panel(fix<3>, fix<8>, traits, i, j2);
1590 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1591 for (Index i = i1; i < actual_panel_end; i += 3 * LhsProgress) {
1592 micro_panel(fix<3>, fix<4>, traits, i, j2);
1595 for (Index j2 = packet_cols4; j2 < cols; j2++) {
1596 for (Index i = i1; i < actual_panel_end; i += 3 * LhsProgress) {
1597 micro_panel(fix<3>, fix<1>, traits, i, j2);
1604 EIGEN_IF_CONSTEXPR (mr >= 2 * Traits::LhsProgress) {
1605 const Index rhs_block2 =
sizeof(ResScalar) * mr * nr + depth * nr *
sizeof(RhsScalar);
1606 const Index lhs_strip2 = depth *
sizeof(LhsScalar) * 2 * LhsProgress;
1607 const Index lhs_avail2 = (lhs_budget > rhs_block2) ? (lhs_budget - rhs_block2) : 0;
1608 const Index mc2_range = peeled_mc2 - peeled_mc3;
1609 Index actual_panel_rows = (lhs_avail2 >= mc2_range * depth *
static_cast<Index
>(
sizeof(LhsScalar)))
1611 : (2 * LhsProgress) * std::max<Index>(1, lhs_avail2 / lhs_strip2);
1612 for (Index i1 = peeled_mc3; i1 < peeled_mc2; i1 += actual_panel_rows) {
1613 Index actual_panel_end = (std::min)(i1 + actual_panel_rows, peeled_mc2);
1614 EIGEN_IF_CONSTEXPR (nr >= 8) {
1615 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1616 for (Index i = i1; i < actual_panel_end; i += 2 * LhsProgress) {
1617 micro_panel(fix<2>, fix<8>, traits, i, j2);
1621 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1622 for (Index i = i1; i < actual_panel_end; i += 2 * LhsProgress) {
1623 micro_panel(fix<2>, fix<4>, traits, i, j2);
1626 for (Index j2 = packet_cols4; j2 < cols; j2++) {
1627 for (Index i = i1; i < actual_panel_end; i += 2 * LhsProgress) {
1628 micro_panel(fix<2>, fix<1>, traits, i, j2);
1635 EIGEN_IF_CONSTEXPR (mr >= 1 * Traits::LhsProgress) {
1636 for (Index i = peeled_mc2; i < peeled_mc1; i += LhsProgress) {
1637 EIGEN_IF_CONSTEXPR (nr >= 8) {
1638 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1639 micro_panel(fix<1>, fix<8>, traits, i, j2);
1642 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1643 micro_panel(fix<1>, fix<4>, traits, i, j2);
1645 for (Index j2 = packet_cols4; j2 < cols; j2++) {
1646 micro_panel(fix<1>, fix<1>, traits, i, j2);
1652 EIGEN_IF_CONSTEXPR ((LhsProgressHalf < LhsProgress) && mr >= LhsProgressHalf) {
1653 HalfTraits half_traits;
1654 for (Index i = peeled_mc1; i < peeled_mc_half; i += LhsProgressHalf) {
1655 EIGEN_IF_CONSTEXPR (nr >= 8) {
1656 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1657 gebp_micro_panel_impl<1, 8, HalfTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1658 LhsPacket>(half_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB,
1659 offsetA, offsetB, prefetch_res_offset, peeled_kc, pk);
1662 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1663 gebp_micro_panel_impl<1, 4, HalfTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1664 LhsPacket>(half_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB,
1665 offsetA, offsetB, prefetch_res_offset, peeled_kc, pk);
1667 for (Index j2 = packet_cols4; j2 < cols; j2++) {
1668 gebp_micro_panel_impl<1, 1, HalfTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1669 LhsPacket>(half_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB,
1670 offsetA, offsetB, prefetch_res_offset, peeled_kc, pk);
1676 EIGEN_IF_CONSTEXPR ((LhsProgressQuarter < LhsProgressHalf) && mr >= LhsProgressQuarter) {
1677 QuarterTraits quarter_traits;
1678 for (Index i = peeled_mc_half; i < peeled_mc_quarter; i += LhsProgressQuarter) {
1679 EIGEN_IF_CONSTEXPR (nr >= 8) {
1680 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1681 gebp_micro_panel_impl<1, 8, QuarterTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1682 LhsPacket>(quarter_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB,
1683 offsetA, offsetB, prefetch_res_offset, peeled_kc, pk);
1686 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1687 gebp_micro_panel_impl<1, 4, QuarterTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1688 LhsPacket>(quarter_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB,
1689 offsetA, offsetB, prefetch_res_offset, peeled_kc, pk);
1691 for (Index j2 = packet_cols4; j2 < cols; j2++) {
1692 gebp_micro_panel_impl<1, 1, QuarterTraits, LhsScalar, RhsScalar, ResScalar, Index, DataMapper, LinearMapper,
1693 LhsPacket>(quarter_traits, res, blockA, blockB, alpha, i, j2, depth, strideA, strideB,
1694 offsetA, offsetB, prefetch_res_offset, peeled_kc, pk);
1700 if (peeled_mc_quarter < rows) {
1701 EIGEN_IF_CONSTEXPR (nr >= 8) {
1703 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
1705 for (Index i = peeled_mc_quarter; i < rows; i += 1) {
1706 const LhsScalar* blA = &blockA[i * strideA + offsetA];
1709 ResScalar C0(0), C1(0), C2(0), C3(0), C4(0), C5(0), C6(0), C7(0);
1710 const RhsScalar* blB = &blockB[j2 * strideB + offsetB * 8];
1711 for (Index k = 0; k < depth; k++) {
1712 LhsScalar A0 = blA[k];
1716 C0 = cj.pmadd(A0, B_0, C0);
1719 C1 = cj.pmadd(A0, B_0, C1);
1722 C2 = cj.pmadd(A0, B_0, C2);
1725 C3 = cj.pmadd(A0, B_0, C3);
1728 C4 = cj.pmadd(A0, B_0, C4);
1731 C5 = cj.pmadd(A0, B_0, C5);
1734 C6 = cj.pmadd(A0, B_0, C6);
1737 C7 = cj.pmadd(A0, B_0, C7);
1741 res(i, j2 + 0) += alpha * C0;
1742 res(i, j2 + 1) += alpha * C1;
1743 res(i, j2 + 2) += alpha * C2;
1744 res(i, j2 + 3) += alpha * C3;
1745 res(i, j2 + 4) += alpha * C4;
1746 res(i, j2 + 5) += alpha * C5;
1747 res(i, j2 + 6) += alpha * C6;
1748 res(i, j2 + 7) += alpha * C7;
1753 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
1755 for (Index i = peeled_mc_quarter; i < rows; i += 1) {
1756 const LhsScalar* blA = &blockA[i * strideA + offsetA];
1758 const RhsScalar* blB = &blockB[j2 * strideB + offsetB * 4];
1763 constexpr int SResPacketHalfSize = unpacket_traits<typename unpacket_traits<SResPacket>::half>::size;
1764 constexpr int SResPacketQuarterSize =
1765 unpacket_traits<typename unpacket_traits<typename unpacket_traits<SResPacket>::half>::half>::size;
1770 constexpr bool kCanLoadSRhsQuad =
1771 (unpacket_traits<SLhsPacket>::size < 4) ||
1772 (unpacket_traits<SRhsPacket>::size % ((std::max<int>)(unpacket_traits<SLhsPacket>::size, 4) / 4)) == 0;
1773 EIGEN_IF_CONSTEXPR (kCanLoadSRhsQuad && (SwappedTraits::LhsProgress % 4) == 0 &&
1774 (SwappedTraits::LhsProgress <= 16) &&
1775 (SwappedTraits::LhsProgress != 8 || SResPacketHalfSize == 4) &&
1776 (SwappedTraits::LhsProgress != 16 || SResPacketQuarterSize == 4)) {
1777 SAccPacket C0, C1, C2, C3;
1778 straits.initAcc(C0);
1779 straits.initAcc(C1);
1780 straits.initAcc(C2);
1781 straits.initAcc(C3);
1783 const Index spk = (std::max)(1, SwappedTraits::LhsProgress / 4);
1784 const Index endk = (depth / spk) * spk;
1785 const Index endk4 = (depth / (spk * 4)) * (spk * 4);
1788 for (; k < endk4; k += 4 * spk) {
1790 SRhsPacket B_0, B_1;
1792 straits.loadLhsUnaligned(blB + 0 * SwappedTraits::LhsProgress, A0);
1793 straits.loadLhsUnaligned(blB + 1 * SwappedTraits::LhsProgress, A1);
1795 straits.loadRhsQuad(blA + 0 * spk, B_0);
1796 straits.loadRhsQuad(blA + 1 * spk, B_1);
1797 straits.madd(A0, B_0, C0, B_0, fix<0>);
1798 straits.madd(A1, B_1, C1, B_1, fix<0>);
1800 straits.loadLhsUnaligned(blB + 2 * SwappedTraits::LhsProgress, A0);
1801 straits.loadLhsUnaligned(blB + 3 * SwappedTraits::LhsProgress, A1);
1802 straits.loadRhsQuad(blA + 2 * spk, B_0);
1803 straits.loadRhsQuad(blA + 3 * spk, B_1);
1804 straits.madd(A0, B_0, C2, B_0, fix<0>);
1805 straits.madd(A1, B_1, C3, B_1, fix<0>);
1807 blB += 4 * SwappedTraits::LhsProgress;
1810 C0 = padd(padd(C0, C1), padd(C2, C3));
1811 for (; k < endk; k += spk) {
1815 straits.loadLhsUnaligned(blB, A0);
1816 straits.loadRhsQuad(blA, B_0);
1817 straits.madd(A0, B_0, C0, B_0, fix<0>);
1819 blB += SwappedTraits::LhsProgress;
1822 if (SwappedTraits::LhsProgress == 8) {
1824 typedef std::conditional_t<SwappedTraits::LhsProgress >= 8,
typename unpacket_traits<SResPacket>::half,
1827 typedef std::conditional_t<SwappedTraits::LhsProgress >= 8,
typename unpacket_traits<SLhsPacket>::half,
1830 typedef std::conditional_t<SwappedTraits::LhsProgress >= 8,
typename unpacket_traits<SRhsPacket>::half,
1833 typedef std::conditional_t<SwappedTraits::LhsProgress >= 8,
typename unpacket_traits<SAccPacket>::half,
1837 SResPacketHalf R = res.template gatherPacket<SResPacketHalf>(i, j2);
1838 SResPacketHalf alphav = pset1<SResPacketHalf>(alpha);
1840 if (depth - endk > 0) {
1844 straits.loadLhsUnaligned(blB, a0);
1845 straits.loadRhs(blA, b0);
1846 SAccPacketHalf c0 = predux_half(C0);
1847 straits.madd(a0, b0, c0, b0, fix<0>);
1848 straits.acc(c0, alphav, R);
1850 straits.acc(predux_half(C0), alphav, R);
1852 res.scatterPacket(i, j2, R);
1853 }
else if (SwappedTraits::LhsProgress == 16) {
1858 last_row_process_16_packets<LhsScalar, RhsScalar, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> p;
1859 p(res, straits, blA, blB, depth, endk, i, j2, alpha, C0);
1861 SResPacket R = res.template gatherPacket<SResPacket>(i, j2);
1862 SResPacket alphav = pset1<SResPacket>(alpha);
1863 straits.acc(C0, alphav, R);
1864 res.scatterPacket(i, j2, R);
1869 ResScalar C0(0), C1(0), C2(0), C3(0);
1871 for (Index k = 0; k < depth; k++) {
1879 C0 = cj.pmadd(A0, B_0, C0);
1880 C1 = cj.pmadd(A0, B_1, C1);
1884 C2 = cj.pmadd(A0, B_0, C2);
1885 C3 = cj.pmadd(A0, B_1, C3);
1889 res(i, j2 + 0) += alpha * C0;
1890 res(i, j2 + 1) += alpha * C1;
1891 res(i, j2 + 2) += alpha * C2;
1892 res(i, j2 + 3) += alpha * C3;
1897 for (Index j2 = packet_cols4; j2 < cols; j2++) {
1899 for (Index i = peeled_mc_quarter; i < rows; i += 1) {
1900 const LhsScalar* blA = &blockA[i * strideA + offsetA];
1904 const RhsScalar* blB = &blockB[j2 * strideB + offsetB];
1905 for (Index k = 0; k < depth; k++) {
1906 LhsScalar A0 = blA[k];
1907 RhsScalar B_0 = blB[k];
1908 C0 = cj.pmadd(A0, B_0, C0);
1910 res(i, j2) += alpha * C0;
1915#ifdef EIGEN_GEBP_DISABLED_INSN_SCHEDULING
1916#pragma GCC pop_options
1917#undef EIGEN_GEBP_DISABLED_INSN_SCHEDULING
1934template <
typename Scalar,
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
1936struct gemm_pack_lhs<Scalar, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate, PanelMode> {
1937 using LinearMapper =
typename DataMapper::LinearMapper;
1938 EIGEN_DONT_INLINE
void operator()(Scalar* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
1939 Index offset = 0)
const;
1942template <
typename Scalar,
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
1944EIGEN_DONT_INLINE
void gemm_pack_lhs<Scalar, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate,
1945 PanelMode>::operator()(Scalar* blockA,
const DataMapper& lhs, Index depth,
1946 Index rows, Index stride, Index offset)
const {
1947 using HalfPacket =
typename unpacket_traits<Packet>::half;
1948 using QuarterPacket =
typename unpacket_traits<typename unpacket_traits<Packet>::half>::half;
1950 PacketSize = unpacket_traits<Packet>::size,
1951 HalfPacketSize = unpacket_traits<HalfPacket>::size,
1952 QuarterPacketSize = unpacket_traits<QuarterPacket>::size,
1953 HasHalf = (int)HalfPacketSize < (
int)PacketSize,
1954 HasQuarter = (int)QuarterPacketSize < (
int)HalfPacketSize
1957 EIGEN_ASM_COMMENT(
"EIGEN PRODUCT PACK LHS");
1958 EIGEN_UNUSED_VARIABLE(stride);
1959 EIGEN_UNUSED_VARIABLE(offset);
1960 eigen_assert(((!PanelMode) && stride == 0 && offset == 0) || (PanelMode && stride >= depth && offset <= stride));
1961 eigen_assert(((Pack1 % PacketSize) == 0 && Pack1 <= 4 * PacketSize) || (Pack1 <= 4) || (Pack1 < PacketSize));
1962 conj_if<NumTraits<Scalar>::IsComplex && Conjugate> cj;
1965 const Index peeled_mc3 = Pack1 >= 3 * PacketSize ? (rows / (3 * PacketSize)) * (3 * PacketSize) : 0;
1966 const Index peeled_mc2 =
1967 Pack1 >= 2 * PacketSize ? peeled_mc3 + ((rows - peeled_mc3) / (2 * PacketSize)) * (2 * PacketSize) : 0;
1968 const Index peeled_mc1 =
1969 Pack1 >= 1 * PacketSize ? peeled_mc2 + ((rows - peeled_mc2) / (1 * PacketSize)) * (1 * PacketSize) : 0;
1970 const Index peeled_mc_half =
1971 Pack1 >= HalfPacketSize ? peeled_mc1 + ((rows - peeled_mc1) / (HalfPacketSize)) * (HalfPacketSize) : 0;
1972 const Index peeled_mc_quarter = Pack1 >= QuarterPacketSize ? (rows / (QuarterPacketSize)) * (QuarterPacketSize) : 0;
1973 const Index last_lhs_progress = rows > peeled_mc_quarter ? (rows - peeled_mc_quarter) & ~1 : 0;
1974 const Index peeled_mc0 = Pack2 >= PacketSize ? peeled_mc_quarter
1975 : Pack2 > 1 && last_lhs_progress ? (rows / last_lhs_progress) * last_lhs_progress
1981 EIGEN_IF_CONSTEXPR (Pack1 >= 3 * PacketSize) {
1982 for (; i < peeled_mc3; i += 3 * PacketSize) {
1983 EIGEN_IF_CONSTEXPR (PanelMode) count += (3 * PacketSize) * offset;
1985 for (Index k = 0; k < depth; k++) {
1987 A = lhs.template loadPacket<Packet>(i + 0 * PacketSize, k);
1988 B = lhs.template loadPacket<Packet>(i + 1 * PacketSize, k);
1989 C = lhs.template loadPacket<Packet>(i + 2 * PacketSize, k);
1990 pstore(blockA + count, cj.pconj(A));
1991 count += PacketSize;
1992 pstore(blockA + count, cj.pconj(B));
1993 count += PacketSize;
1994 pstore(blockA + count, cj.pconj(C));
1995 count += PacketSize;
1997 EIGEN_IF_CONSTEXPR (PanelMode) count += (3 * PacketSize) * (stride - offset - depth);
2001 EIGEN_IF_CONSTEXPR (Pack1 >= 2 * PacketSize) {
2002 for (; i < peeled_mc2; i += 2 * PacketSize) {
2003 EIGEN_IF_CONSTEXPR (PanelMode) count += (2 * PacketSize) * offset;
2005 for (Index k = 0; k < depth; k++) {
2007 A = lhs.template loadPacket<Packet>(i + 0 * PacketSize, k);
2008 B = lhs.template loadPacket<Packet>(i + 1 * PacketSize, k);
2009 pstore(blockA + count, cj.pconj(A));
2010 count += PacketSize;
2011 pstore(blockA + count, cj.pconj(B));
2012 count += PacketSize;
2014 EIGEN_IF_CONSTEXPR (PanelMode) count += (2 * PacketSize) * (stride - offset - depth);
2018 EIGEN_IF_CONSTEXPR (Pack1 >= 1 * PacketSize) {
2019 for (; i < peeled_mc1; i += 1 * PacketSize) {
2020 EIGEN_IF_CONSTEXPR (PanelMode) count += (1 * PacketSize) * offset;
2022 for (Index k = 0; k < depth; k++) {
2024 A = lhs.template loadPacket<Packet>(i + 0 * PacketSize, k);
2025 pstore(blockA + count, cj.pconj(A));
2026 count += PacketSize;
2028 EIGEN_IF_CONSTEXPR (PanelMode) count += (1 * PacketSize) * (stride - offset - depth);
2032 EIGEN_IF_CONSTEXPR (HasHalf && Pack1 >= HalfPacketSize) {
2033 for (; i < peeled_mc_half; i += HalfPacketSize) {
2034 EIGEN_IF_CONSTEXPR (PanelMode) count += (HalfPacketSize)*offset;
2036 for (Index k = 0; k < depth; k++) {
2038 A = lhs.template loadPacket<HalfPacket>(i + 0 * (HalfPacketSize), k);
2039 pstoreu(blockA + count, cj.pconj(A));
2040 count += HalfPacketSize;
2042 EIGEN_IF_CONSTEXPR (PanelMode) count += (HalfPacketSize) * (stride - offset - depth);
2046 EIGEN_IF_CONSTEXPR (HasQuarter && Pack1 >= QuarterPacketSize) {
2047 for (; i < peeled_mc_quarter; i += QuarterPacketSize) {
2048 EIGEN_IF_CONSTEXPR (PanelMode) count += (QuarterPacketSize)*offset;
2050 for (Index k = 0; k < depth; k++) {
2052 A = lhs.template loadPacket<QuarterPacket>(i + 0 * (QuarterPacketSize), k);
2053 pstoreu(blockA + count, cj.pconj(A));
2054 count += QuarterPacketSize;
2056 EIGEN_IF_CONSTEXPR (PanelMode) count += (QuarterPacketSize) * (stride - offset - depth);
2072 EIGEN_IF_CONSTEXPR (Pack2 < PacketSize && Pack2 > 1) {
2073 const Index pack2_progress = (HasHalf || HasQuarter) ? last_lhs_progress : Pack2;
2074 const Index peeled = (HasHalf || HasQuarter) ? peeled_mc0 : (rows / Pack2) * Pack2;
2075 for (; i < peeled; i += pack2_progress) {
2076 EIGEN_IF_CONSTEXPR (PanelMode) count += pack2_progress * offset;
2078 for (Index k = 0; k < depth; k++)
2079 for (Index w = 0; w < pack2_progress; w++) blockA[count++] = cj(lhs(i + w, k));
2081 EIGEN_IF_CONSTEXPR (PanelMode) count += pack2_progress * (stride - offset - depth);
2085 for (; i < rows; i++) {
2086 EIGEN_IF_CONSTEXPR (PanelMode) count += offset;
2087 for (Index k = 0; k < depth; k++) blockA[count++] = cj(lhs(i, k));
2088 EIGEN_IF_CONSTEXPR (PanelMode) count += (stride - offset - depth);
2092template <
typename Scalar,
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
2094struct gemm_pack_lhs<Scalar, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate, PanelMode> {
2095 using LinearMapper =
typename DataMapper::LinearMapper;
2096 EIGEN_DONT_INLINE
void operator()(Scalar* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
2097 Index offset = 0)
const;
2100template <
typename Scalar,
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
2102EIGEN_DONT_INLINE
void gemm_pack_lhs<Scalar, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate,
2103 PanelMode>::operator()(Scalar* blockA,
const DataMapper& lhs, Index depth,
2104 Index rows, Index stride, Index offset)
const {
2105 using HalfPacket =
typename unpacket_traits<Packet>::half;
2106 using QuarterPacket =
typename unpacket_traits<typename unpacket_traits<Packet>::half>::half;
2108 PacketSize = unpacket_traits<Packet>::size,
2109 HalfPacketSize = unpacket_traits<HalfPacket>::size,
2110 QuarterPacketSize = unpacket_traits<QuarterPacket>::size,
2111 HasHalf = (int)HalfPacketSize < (
int)PacketSize,
2112 HasQuarter = (int)QuarterPacketSize < (
int)HalfPacketSize
2115 EIGEN_ASM_COMMENT(
"EIGEN PRODUCT PACK LHS");
2116 EIGEN_UNUSED_VARIABLE(stride);
2117 EIGEN_UNUSED_VARIABLE(offset);
2118 eigen_assert(((!PanelMode) && stride == 0 && offset == 0) || (PanelMode && stride >= depth && offset <= stride));
2119 conj_if<NumTraits<Scalar>::IsComplex && Conjugate> cj;
2121 bool gone_half =
false, gone_quarter =
false, gone_last =
false;
2125 Index psize = PacketSize;
2127 Index remaining_rows = rows - i;
2128 Index peeled_mc = gone_last ? Pack2 > 1 ? (rows / pack) * pack : 0 : i + (remaining_rows / pack) * pack;
2129 Index starting_pos = i;
2130 for (; i < peeled_mc; i += pack) {
2131 EIGEN_IF_CONSTEXPR (PanelMode) count += pack * offset;
2134 if (pack >= psize && psize >= QuarterPacketSize) {
2135 const Index peeled_k = (depth / psize) * psize;
2136 for (; k < peeled_k; k += psize) {
2137 for (Index m = 0; m < pack; m += psize) {
2138 if (psize == PacketSize) {
2139 PacketBlock<Packet> kernel;
2140 for (Index p = 0; p < psize; ++p) kernel.packet[p] = lhs.template loadPacket<Packet>(i + p + m, k);
2142 for (Index p = 0; p < psize; ++p) pstore(blockA + count + m + (pack)*p, cj.pconj(kernel.packet[p]));
2143 }
else if (HasHalf && psize == HalfPacketSize) {
2145 PacketBlock<HalfPacket> kernel_half;
2146 for (Index p = 0; p < psize; ++p)
2147 kernel_half.packet[p] = lhs.template loadPacket<HalfPacket>(i + p + m, k);
2148 ptranspose(kernel_half);
2149 for (Index p = 0; p < psize; ++p) pstore(blockA + count + m + (pack)*p, cj.pconj(kernel_half.packet[p]));
2150 }
else if (HasQuarter && psize == QuarterPacketSize) {
2151 gone_quarter =
true;
2152 PacketBlock<QuarterPacket> kernel_quarter;
2153 for (Index p = 0; p < psize; ++p)
2154 kernel_quarter.packet[p] = lhs.template loadPacket<QuarterPacket>(i + p + m, k);
2155 ptranspose(kernel_quarter);
2156 for (Index p = 0; p < psize; ++p)
2157 pstore(blockA + count + m + (pack)*p, cj.pconj(kernel_quarter.packet[p]));
2160 count += psize * pack;
2164 for (; k < depth; k++) {
2166 for (; w < pack - 3; w += 4) {
2167 Scalar a(cj(lhs(i + w + 0, k))), b(cj(lhs(i + w + 1, k))), c(cj(lhs(i + w + 2, k))), d(cj(lhs(i + w + 3, k)));
2168 blockA[count++] = a;
2169 blockA[count++] = b;
2170 blockA[count++] = c;
2171 blockA[count++] = d;
2174 for (; w < pack; ++w) blockA[count++] = cj(lhs(i + w, k));
2177 EIGEN_IF_CONSTEXPR (PanelMode) count += pack * (stride - offset - depth);
2181 Index left = rows - i;
2183 if (!gone_last && (starting_pos == i || left >= psize / 2 || left >= psize / 4) &&
2184 ((psize / 2 == HalfPacketSize && HasHalf && !gone_half) ||
2185 (psize / 2 == QuarterPacketSize && HasQuarter && !gone_quarter))) {
2203 EIGEN_IF_CONSTEXPR (Pack2 < PacketSize) {
2206 psize = pack = (HasHalf || HasQuarter) ? (left & ~1) : Pack2;
2212 for (; i < rows; i++) {
2213 EIGEN_IF_CONSTEXPR (PanelMode) count += offset;
2214 for (Index k = 0; k < depth; k++) blockA[count++] = cj(lhs(i, k));
2215 EIGEN_IF_CONSTEXPR (PanelMode) count += (stride - offset - depth);
2226template <
typename Scalar,
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
2227struct gemm_pack_rhs<Scalar, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode> {
2228 using Packet =
typename packet_traits<Scalar>::type;
2229 using LinearMapper =
typename DataMapper::LinearMapper;
2230 enum { PacketSize = packet_traits<Scalar>::size };
2231 EIGEN_DONT_INLINE
void operator()(Scalar* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
2232 Index offset = 0)
const;
2235template <
typename Scalar,
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
2236EIGEN_DONT_INLINE
void gemm_pack_rhs<Scalar, Index, DataMapper, nr, ColMajor, Conjugate, PanelMode>::operator()(
2237 Scalar* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset)
const {
2238 EIGEN_ASM_COMMENT(
"EIGEN PRODUCT PACK RHS COLMAJOR");
2239 EIGEN_UNUSED_VARIABLE(stride);
2240 EIGEN_UNUSED_VARIABLE(offset);
2241 eigen_assert(((!PanelMode) && stride == 0 && offset == 0) || (PanelMode && stride >= depth && offset <= stride));
2242 conj_if<NumTraits<Scalar>::IsComplex && Conjugate> cj;
2243 Index packet_cols8 = nr >= 8 ? (cols / 8) * 8 : 0;
2244 Index packet_cols4 = nr >= 4 ? (cols / 4) * 4 : 0;
2246 const Index peeled_k = (depth / PacketSize) * PacketSize;
2248 EIGEN_IF_CONSTEXPR (nr >= 8) {
2249 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
2251 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * offset;
2252 const LinearMapper dm0 = rhs.getLinearMapper(0, j2 + 0);
2253 const LinearMapper dm1 = rhs.getLinearMapper(0, j2 + 1);
2254 const LinearMapper dm2 = rhs.getLinearMapper(0, j2 + 2);
2255 const LinearMapper dm3 = rhs.getLinearMapper(0, j2 + 3);
2256 const LinearMapper dm4 = rhs.getLinearMapper(0, j2 + 4);
2257 const LinearMapper dm5 = rhs.getLinearMapper(0, j2 + 5);
2258 const LinearMapper dm6 = rhs.getLinearMapper(0, j2 + 6);
2259 const LinearMapper dm7 = rhs.getLinearMapper(0, j2 + 7);
2261 EIGEN_IF_CONSTEXPR (PacketSize % 2 == 0 && PacketSize <= 8)
2263 for (; k < peeled_k; k += PacketSize) {
2264 EIGEN_IF_CONSTEXPR (PacketSize == 2) {
2265 PacketBlock<Packet, PacketSize == 2 ? 2 : PacketSize> kernel0, kernel1, kernel2, kernel3;
2266 kernel0.packet[0 % PacketSize] = dm0.template loadPacket<Packet>(k);
2267 kernel0.packet[1 % PacketSize] = dm1.template loadPacket<Packet>(k);
2268 kernel1.packet[0 % PacketSize] = dm2.template loadPacket<Packet>(k);
2269 kernel1.packet[1 % PacketSize] = dm3.template loadPacket<Packet>(k);
2270 kernel2.packet[0 % PacketSize] = dm4.template loadPacket<Packet>(k);
2271 kernel2.packet[1 % PacketSize] = dm5.template loadPacket<Packet>(k);
2272 kernel3.packet[0 % PacketSize] = dm6.template loadPacket<Packet>(k);
2273 kernel3.packet[1 % PacketSize] = dm7.template loadPacket<Packet>(k);
2274 ptranspose(kernel0);
2275 ptranspose(kernel1);
2276 ptranspose(kernel2);
2277 ptranspose(kernel3);
2279 pstoreu(blockB + count + 0 * PacketSize, cj.pconj(kernel0.packet[0 % PacketSize]));
2280 pstoreu(blockB + count + 1 * PacketSize, cj.pconj(kernel1.packet[0 % PacketSize]));
2281 pstoreu(blockB + count + 2 * PacketSize, cj.pconj(kernel2.packet[0 % PacketSize]));
2282 pstoreu(blockB + count + 3 * PacketSize, cj.pconj(kernel3.packet[0 % PacketSize]));
2284 pstoreu(blockB + count + 4 * PacketSize, cj.pconj(kernel0.packet[1 % PacketSize]));
2285 pstoreu(blockB + count + 5 * PacketSize, cj.pconj(kernel1.packet[1 % PacketSize]));
2286 pstoreu(blockB + count + 6 * PacketSize, cj.pconj(kernel2.packet[1 % PacketSize]));
2287 pstoreu(blockB + count + 7 * PacketSize, cj.pconj(kernel3.packet[1 % PacketSize]));
2288 count += 8 * PacketSize;
2289 }
else EIGEN_IF_CONSTEXPR (PacketSize == 4) {
2290 PacketBlock<Packet, PacketSize == 4 ? 4 : PacketSize> kernel0, kernel1;
2292 kernel0.packet[0 % PacketSize] = dm0.template loadPacket<Packet>(k);
2293 kernel0.packet[1 % PacketSize] = dm1.template loadPacket<Packet>(k);
2294 kernel0.packet[2 % PacketSize] = dm2.template loadPacket<Packet>(k);
2295 kernel0.packet[3 % PacketSize] = dm3.template loadPacket<Packet>(k);
2296 kernel1.packet[0 % PacketSize] = dm4.template loadPacket<Packet>(k);
2297 kernel1.packet[1 % PacketSize] = dm5.template loadPacket<Packet>(k);
2298 kernel1.packet[2 % PacketSize] = dm6.template loadPacket<Packet>(k);
2299 kernel1.packet[3 % PacketSize] = dm7.template loadPacket<Packet>(k);
2300 ptranspose(kernel0);
2301 ptranspose(kernel1);
2303 pstoreu(blockB + count + 0 * PacketSize, cj.pconj(kernel0.packet[0 % PacketSize]));
2304 pstoreu(blockB + count + 1 * PacketSize, cj.pconj(kernel1.packet[0 % PacketSize]));
2305 pstoreu(blockB + count + 2 * PacketSize, cj.pconj(kernel0.packet[1 % PacketSize]));
2306 pstoreu(blockB + count + 3 * PacketSize, cj.pconj(kernel1.packet[1 % PacketSize]));
2307 pstoreu(blockB + count + 4 * PacketSize, cj.pconj(kernel0.packet[2 % PacketSize]));
2308 pstoreu(blockB + count + 5 * PacketSize, cj.pconj(kernel1.packet[2 % PacketSize]));
2309 pstoreu(blockB + count + 6 * PacketSize, cj.pconj(kernel0.packet[3 % PacketSize]));
2310 pstoreu(blockB + count + 7 * PacketSize, cj.pconj(kernel1.packet[3 % PacketSize]));
2311 count += 8 * PacketSize;
2312 }
else EIGEN_IF_CONSTEXPR (PacketSize == 8) {
2313 PacketBlock<Packet, PacketSize == 8 ? 8 : PacketSize> kernel0;
2315 kernel0.packet[0 % PacketSize] = dm0.template loadPacket<Packet>(k);
2316 kernel0.packet[1 % PacketSize] = dm1.template loadPacket<Packet>(k);
2317 kernel0.packet[2 % PacketSize] = dm2.template loadPacket<Packet>(k);
2318 kernel0.packet[3 % PacketSize] = dm3.template loadPacket<Packet>(k);
2319 kernel0.packet[4 % PacketSize] = dm4.template loadPacket<Packet>(k);
2320 kernel0.packet[5 % PacketSize] = dm5.template loadPacket<Packet>(k);
2321 kernel0.packet[6 % PacketSize] = dm6.template loadPacket<Packet>(k);
2322 kernel0.packet[7 % PacketSize] = dm7.template loadPacket<Packet>(k);
2323 ptranspose(kernel0);
2325 pstoreu(blockB + count + 0 * PacketSize, cj.pconj(kernel0.packet[0 % PacketSize]));
2326 pstoreu(blockB + count + 1 * PacketSize, cj.pconj(kernel0.packet[1 % PacketSize]));
2327 pstoreu(blockB + count + 2 * PacketSize, cj.pconj(kernel0.packet[2 % PacketSize]));
2328 pstoreu(blockB + count + 3 * PacketSize, cj.pconj(kernel0.packet[3 % PacketSize]));
2329 pstoreu(blockB + count + 4 * PacketSize, cj.pconj(kernel0.packet[4 % PacketSize]));
2330 pstoreu(blockB + count + 5 * PacketSize, cj.pconj(kernel0.packet[5 % PacketSize]));
2331 pstoreu(blockB + count + 6 * PacketSize, cj.pconj(kernel0.packet[6 % PacketSize]));
2332 pstoreu(blockB + count + 7 * PacketSize, cj.pconj(kernel0.packet[7 % PacketSize]));
2333 count += 8 * PacketSize;
2338 for (; k < depth; k++) {
2339 blockB[count + 0] = cj(dm0(k));
2340 blockB[count + 1] = cj(dm1(k));
2341 blockB[count + 2] = cj(dm2(k));
2342 blockB[count + 3] = cj(dm3(k));
2343 blockB[count + 4] = cj(dm4(k));
2344 blockB[count + 5] = cj(dm5(k));
2345 blockB[count + 6] = cj(dm6(k));
2346 blockB[count + 7] = cj(dm7(k));
2350 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * (stride - offset - depth);
2354 EIGEN_IF_CONSTEXPR (nr >= 4) {
2355 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
2357 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * offset;
2358 const LinearMapper dm0 = rhs.getLinearMapper(0, j2 + 0);
2359 const LinearMapper dm1 = rhs.getLinearMapper(0, j2 + 1);
2360 const LinearMapper dm2 = rhs.getLinearMapper(0, j2 + 2);
2361 const LinearMapper dm3 = rhs.getLinearMapper(0, j2 + 3);
2364 EIGEN_IF_CONSTEXPR ((PacketSize % 4) == 0 || PacketSize == 2) {
2365 for (; k < peeled_k; k += PacketSize) {
2366 PacketBlock<Packet, 4> kernel;
2367 kernel.packet[0] = dm0.template loadPacket<Packet>(k);
2368 kernel.packet[1] = dm1.template loadPacket<Packet>(k);
2369 kernel.packet[2] = dm2.template loadPacket<Packet>(k);
2370 kernel.packet[3] = dm3.template loadPacket<Packet>(k);
2371 EIGEN_IF_CONSTEXPR (PacketSize == 2) {
2375 PacketBlock<Packet, 2> tmp01;
2376 tmp01.packet[0] = kernel.packet[0];
2377 tmp01.packet[1] = kernel.packet[1];
2379 PacketBlock<Packet, 2> tmp23;
2380 tmp23.packet[0] = kernel.packet[2];
2381 tmp23.packet[1] = kernel.packet[3];
2383 kernel.packet[0] = tmp01.packet[0];
2384 kernel.packet[1] = tmp23.packet[0];
2385 kernel.packet[2] = tmp01.packet[1];
2386 kernel.packet[3] = tmp23.packet[1];
2390 pstoreu(blockB + count + 0 * PacketSize, cj.pconj(kernel.packet[0]));
2391 pstoreu(blockB + count + 1 * PacketSize, cj.pconj(kernel.packet[1]));
2392 pstoreu(blockB + count + 2 * PacketSize, cj.pconj(kernel.packet[2]));
2393 pstoreu(blockB + count + 3 * PacketSize, cj.pconj(kernel.packet[3]));
2394 count += 4 * PacketSize;
2397 for (; k < depth; k++) {
2398 blockB[count + 0] = cj(dm0(k));
2399 blockB[count + 1] = cj(dm1(k));
2400 blockB[count + 2] = cj(dm2(k));
2401 blockB[count + 3] = cj(dm3(k));
2405 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * (stride - offset - depth);
2410 for (Index j2 = packet_cols4; j2 < cols; ++j2) {
2411 EIGEN_IF_CONSTEXPR (PanelMode) count += offset;
2412 const LinearMapper dm0 = rhs.getLinearMapper(0, j2);
2413 for (Index k = 0; k < depth; k++) {
2414 blockB[count] = cj(dm0(k));
2417 EIGEN_IF_CONSTEXPR (PanelMode) count += (stride - offset - depth);
2422template <
typename Scalar,
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
2423struct gemm_pack_rhs<Scalar, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode> {
2424 using Packet =
typename packet_traits<Scalar>::type;
2425 using HalfPacket =
typename unpacket_traits<Packet>::half;
2426 using QuarterPacket =
typename unpacket_traits<typename unpacket_traits<Packet>::half>::half;
2427 using LinearMapper =
typename DataMapper::LinearMapper;
2429 PacketSize = packet_traits<Scalar>::size,
2430 HalfPacketSize = unpacket_traits<HalfPacket>::size,
2431 QuarterPacketSize = unpacket_traits<QuarterPacket>::size
2433 EIGEN_DONT_INLINE
void operator()(Scalar* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
2434 Index offset = 0)
const {
2435 EIGEN_ASM_COMMENT(
"EIGEN PRODUCT PACK RHS ROWMAJOR");
2436 EIGEN_UNUSED_VARIABLE(stride);
2437 EIGEN_UNUSED_VARIABLE(offset);
2438 eigen_assert(((!PanelMode) && stride == 0 && offset == 0) || (PanelMode && stride >= depth && offset <= stride));
2439 constexpr bool HasHalf = (int)HalfPacketSize < (
int)PacketSize;
2440 constexpr bool HasQuarter = (int)QuarterPacketSize < (
int)HalfPacketSize;
2441 conj_if<NumTraits<Scalar>::IsComplex && Conjugate> cj;
2442 Index packet_cols8 = nr >= 8 ? (cols / 8) * 8 : 0;
2443 Index packet_cols4 = nr >= 4 ? (cols / 4) * 4 : 0;
2446 EIGEN_IF_CONSTEXPR (nr >= 8) {
2447 for (Index j2 = 0; j2 < packet_cols8; j2 += 8) {
2449 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * offset;
2450 for (Index k = 0; k < depth; k++) {
2451 EIGEN_IF_CONSTEXPR (PacketSize == 8) {
2452 Packet A = rhs.template loadPacket<Packet>(k, j2);
2453 pstoreu(blockB + count, cj.pconj(A));
2454 count += PacketSize;
2455 }
else EIGEN_IF_CONSTEXPR (PacketSize == 4) {
2456 Packet A = rhs.template loadPacket<Packet>(k, j2);
2457 Packet B = rhs.template loadPacket<Packet>(k, j2 + 4);
2458 pstoreu(blockB + count, cj.pconj(A));
2459 pstoreu(blockB + count + PacketSize, cj.pconj(B));
2460 count += 2 * PacketSize;
2462 const LinearMapper dm0 = rhs.getLinearMapper(k, j2);
2463 blockB[count + 0] = cj(dm0(0));
2464 blockB[count + 1] = cj(dm0(1));
2465 blockB[count + 2] = cj(dm0(2));
2466 blockB[count + 3] = cj(dm0(3));
2467 blockB[count + 4] = cj(dm0(4));
2468 blockB[count + 5] = cj(dm0(5));
2469 blockB[count + 6] = cj(dm0(6));
2470 blockB[count + 7] = cj(dm0(7));
2475 EIGEN_IF_CONSTEXPR (PanelMode) count += 8 * (stride - offset - depth);
2479 EIGEN_IF_CONSTEXPR (nr >= 4) {
2480 for (Index j2 = packet_cols8; j2 < packet_cols4; j2 += 4) {
2482 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * offset;
2483 for (Index k = 0; k < depth; k++) {
2484 EIGEN_IF_CONSTEXPR (PacketSize == 4) {
2485 Packet A = rhs.template loadPacket<Packet>(k, j2);
2486 pstoreu(blockB + count, cj.pconj(A));
2487 count += PacketSize;
2488 }
else EIGEN_IF_CONSTEXPR (HasHalf && HalfPacketSize == 4) {
2489 HalfPacket A = rhs.template loadPacket<HalfPacket>(k, j2);
2490 pstoreu(blockB + count, cj.pconj(A));
2491 count += HalfPacketSize;
2492 }
else EIGEN_IF_CONSTEXPR (HasQuarter && QuarterPacketSize == 4) {
2493 QuarterPacket A = rhs.template loadPacket<QuarterPacket>(k, j2);
2494 pstoreu(blockB + count, cj.pconj(A));
2495 count += QuarterPacketSize;
2497 const LinearMapper dm0 = rhs.getLinearMapper(k, j2);
2498 blockB[count + 0] = cj(dm0(0));
2499 blockB[count + 1] = cj(dm0(1));
2500 blockB[count + 2] = cj(dm0(2));
2501 blockB[count + 3] = cj(dm0(3));
2506 EIGEN_IF_CONSTEXPR (PanelMode) count += 4 * (stride - offset - depth);
2510 for (Index j2 = packet_cols4; j2 < cols; ++j2) {
2511 EIGEN_IF_CONSTEXPR (PanelMode) count += offset;
2512 for (Index k = 0; k < depth; k++) {
2513 blockB[count] = cj(rhs(k, j2));
2516 EIGEN_IF_CONSTEXPR (PanelMode) count += stride - offset - depth;
2525inline std::ptrdiff_t l1CacheSize() {
2526 std::ptrdiff_t l1, l2, l3;
2527 internal::manage_caching_sizes(GetAction, &l1, &l2, &l3);
2533inline std::ptrdiff_t l2CacheSize() {
2534 std::ptrdiff_t l1, l2, l3;
2535 internal::manage_caching_sizes(GetAction, &l1, &l2, &l3);
2541inline std::ptrdiff_t l3CacheSize() {
2542 std::ptrdiff_t l1, l2, l3;
2543 internal::manage_caching_sizes(GetAction, &l1, &l2, &l3);
2551inline void setCpuCacheSizes(std::ptrdiff_t l1, std::ptrdiff_t l2, std::ptrdiff_t l3) {
2552 internal::manage_caching_sizes(SetAction, &l1, &l2, &l3);
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321