12#ifndef EIGEN_MATRIX_PRODUCT_ALTIVEC_H
13#define EIGEN_MATRIX_PRODUCT_ALTIVEC_H
15#ifndef EIGEN_ALTIVEC_USE_CUSTOM_PACK
16#define EIGEN_ALTIVEC_USE_CUSTOM_PACK 1
19#if !defined(EIGEN_ALTIVEC_DISABLE_MMA)
20#define EIGEN_ALTIVEC_DISABLE_MMA 0
24#if !EIGEN_ALTIVEC_DISABLE_MMA && defined(__has_builtin)
25#if __has_builtin(__builtin_mma_assemble_acc)
26#define EIGEN_ALTIVEC_MMA_SUPPORT
31#if defined(EIGEN_ALTIVEC_MMA_SUPPORT)
33#if !defined(EIGEN_ALTIVEC_ENABLE_MMA_DYNAMIC_DISPATCH)
34#define EIGEN_ALTIVEC_ENABLE_MMA_DYNAMIC_DISPATCH 0
38#if EIGEN_ALTIVEC_ENABLE_MMA_DYNAMIC_DISPATCH && !EIGEN_COMP_LLVM
39#define EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH 1
42#define EIGEN_ALTIVEC_MMA_ONLY 1
47#include "MatrixProductCommon.h"
49#if defined(EIGEN_ALTIVEC_MMA_ONLY) || defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH)
50#include "MatrixProductMMA.h"
54#include "../../InternalHeaderCheck.h"
63template <
typename Scalar>
65 typedef typename packet_traits<Scalar>::type vectortype;
66 typedef PacketBlock<vectortype, 4> type;
67 typedef vectortype rhstype;
68 enum { vectorsize = packet_traits<Scalar>::size, size = 4, rows = 4 };
72struct quad_traits<double> {
73 typedef Packet2d vectortype;
74 typedef PacketBlock<vectortype, 4> type;
75 typedef PacketBlock<Packet2d, 2> rhstype;
76 enum { vectorsize = packet_traits<double>::size, size = 2, rows = 4 };
80struct quad_traits<bfloat16> {
81 typedef Packet8bf vectortype;
82 typedef PacketBlock<vectortype, 4> type;
83 typedef vectortype rhstype;
84 enum { vectorsize = packet_traits<bfloat16>::size, size = 8, rows = 4 };
91const static Packet16uc p16uc_GETREAL32 = {0, 1, 2, 3, 8, 9, 10, 11, 16, 17, 18, 19, 24, 25, 26, 27};
93const static Packet16uc p16uc_GETIMAG32 = {4, 5, 6, 7, 12, 13, 14, 15, 20, 21, 22, 23, 28, 29, 30, 31};
95const static Packet16uc p16uc_GETREAL32b = {0, 1, 2, 3, 16, 17, 18, 19, 8, 9, 10, 11, 24, 25, 26, 27};
97const static Packet16uc p16uc_GETIMAG32b = {4, 5, 6, 7, 20, 21, 22, 23, 12, 13, 14, 15, 28, 29, 30, 31};
117template <
typename Scalar,
int StorageOrder>
118EIGEN_ALWAYS_INLINE std::complex<Scalar> getAdjointVal(
119 Index i, Index j, const_blas_data_mapper<std::complex<Scalar>, Index, StorageOrder>& dt) {
120 std::complex<Scalar> v;
122 v.real(dt(j, i).real());
123 v.imag(-dt(j, i).imag());
125 v.real(dt(i, j).real());
126 v.imag(dt(i, j).imag());
128 v.real(dt(i, j).real());
134template <
typename Scalar,
int StorageOrder,
int N>
135EIGEN_STRONG_INLINE
void symm_pack_complex_rhs_helper(std::complex<Scalar>* blockB,
const std::complex<Scalar>* _rhs,
136 Index rhsStride, Index rows, Index cols, Index k2) {
137 const Index depth = k2 + rows;
138 const_blas_data_mapper<std::complex<Scalar>, Index, StorageOrder> rhs(_rhs, rhsStride);
139 const Index vectorSize = N * quad_traits<Scalar>::vectorsize;
140 const Index vectorDelta = vectorSize * rows;
141 Scalar* blockBf =
reinterpret_cast<Scalar*
>(blockB);
143 Index rir = 0, rii, j = 0;
144 for (; j + vectorSize <= cols; j += vectorSize) {
145 rii = rir + vectorDelta;
147 for (Index i = k2; i < depth; i++) {
148 for (Index k = 0; k < vectorSize; k++) {
149 std::complex<Scalar> v = getAdjointVal<Scalar, StorageOrder>(i, j + k, rhs);
151 blockBf[rir + k] = v.real();
152 blockBf[rii + k] = v.imag();
161 for (; j < cols; j++) {
164 for (Index i = k2; i < depth; i++) {
165 std::complex<Scalar> v = getAdjointVal<Scalar, StorageOrder>(i, j, rhs);
167 blockBf[rir] = v.real();
168 blockBf[rii] = v.imag();
178template <
typename Scalar,
int StorageOrder>
179EIGEN_STRONG_INLINE
void symm_pack_complex_lhs_helper(std::complex<Scalar>* blockA,
const std::complex<Scalar>* _lhs,
180 Index lhsStride, Index cols, Index rows) {
181 const Index depth = cols;
182 const_blas_data_mapper<std::complex<Scalar>, Index, StorageOrder> lhs(_lhs, lhsStride);
183 const Index vectorSize = quad_traits<Scalar>::vectorsize;
184 const Index vectorDelta = vectorSize * depth;
185 Scalar* blockAf =
reinterpret_cast<Scalar*
>(blockA);
187 Index rir = 0, rii, j = 0;
188 for (; j + vectorSize <= rows; j += vectorSize) {
189 rii = rir + vectorDelta;
191 for (Index i = 0; i < depth; i++) {
192 for (Index k = 0; k < vectorSize; k++) {
193 std::complex<Scalar> v = getAdjointVal<Scalar, StorageOrder>(j + k, i, lhs);
195 blockAf[rir + k] = v.real();
196 blockAf[rii + k] = v.imag();
206 rii = rir + ((rows - j) * depth);
208 for (Index i = 0; i < depth; i++) {
210 for (; k < rows; k++) {
211 std::complex<Scalar> v = getAdjointVal<Scalar, StorageOrder>(k, i, lhs);
213 blockAf[rir] = v.real();
214 blockAf[rii] = v.imag();
223template <
typename Scalar,
int StorageOrder,
int N>
224EIGEN_STRONG_INLINE
void symm_pack_rhs_helper(Scalar* blockB,
const Scalar* _rhs, Index rhsStride, Index rows,
225 Index cols, Index k2) {
226 const Index depth = k2 + rows;
227 const_blas_data_mapper<Scalar, Index, StorageOrder> rhs(_rhs, rhsStride);
228 const Index vectorSize = quad_traits<Scalar>::vectorsize;
231 for (; j + N * vectorSize <= cols; j += N * vectorSize) {
233 for (; i < depth; i++) {
234 for (Index k = 0; k < N * vectorSize; k++) {
236 blockB[ri + k] = rhs(j + k, i);
238 blockB[ri + k] = rhs(i, j + k);
240 ri += N * vectorSize;
244 for (; j < cols; j++) {
245 for (Index i = k2; i < depth; i++) {
247 blockB[ri] = rhs(i, j);
249 blockB[ri] = rhs(j, i);
255template <
typename Scalar,
int StorageOrder>
256EIGEN_STRONG_INLINE
void symm_pack_lhs_helper(Scalar* blockA,
const Scalar* _lhs, Index lhsStride, Index cols,
258 const Index depth = cols;
259 const_blas_data_mapper<Scalar, Index, StorageOrder> lhs(_lhs, lhsStride);
260 const Index vectorSize = quad_traits<Scalar>::vectorsize;
263 for (; j + vectorSize <= rows; j += vectorSize) {
266 for (; i < depth; i++) {
267 for (Index k = 0; k < vectorSize; k++) {
269 blockA[ri + k] = lhs(j + k, i);
271 blockA[ri + k] = lhs(i, j + k);
278 for (Index i = 0; i < depth; i++) {
280 for (; k < rows; k++) {
282 blockA[ri] = lhs(k, i);
284 blockA[ri] = lhs(i, k);
291template <
typename Index,
int nr,
int StorageOrder>
292struct symm_pack_rhs<std::complex<float>, Index, nr, StorageOrder> {
293 void operator()(std::complex<float>* blockB,
const std::complex<float>* _rhs, Index rhsStride, Index rows, Index cols,
295 symm_pack_complex_rhs_helper<float, StorageOrder, 1>(blockB, _rhs, rhsStride, rows, cols, k2);
299template <
typename Index,
int Pack1,
int Pack2_dummy,
int StorageOrder>
300struct symm_pack_lhs<std::complex<float>, Index, Pack1, Pack2_dummy, StorageOrder> {
301 void operator()(std::complex<float>* blockA,
const std::complex<float>* _lhs, Index lhsStride, Index cols,
303 symm_pack_complex_lhs_helper<float, StorageOrder>(blockA, _lhs, lhsStride, cols, rows);
309template <
typename Index,
int nr,
int StorageOrder>
310struct symm_pack_rhs<std::complex<double>, Index, nr, StorageOrder> {
311 void operator()(std::complex<double>* blockB,
const std::complex<double>* _rhs, Index rhsStride, Index rows,
312 Index cols, Index k2) {
313 symm_pack_complex_rhs_helper<double, StorageOrder, 2>(blockB, _rhs, rhsStride, rows, cols, k2);
317template <
typename Index,
int Pack1,
int Pack2_dummy,
int StorageOrder>
318struct symm_pack_lhs<std::complex<double>, Index, Pack1, Pack2_dummy, StorageOrder> {
319 void operator()(std::complex<double>* blockA,
const std::complex<double>* _lhs, Index lhsStride, Index cols,
321 symm_pack_complex_lhs_helper<double, StorageOrder>(blockA, _lhs, lhsStride, cols, rows);
326template <
typename Index,
int nr,
int StorageOrder>
327struct symm_pack_rhs<float, Index, nr, StorageOrder> {
328 void operator()(
float* blockB,
const float* _rhs, Index rhsStride, Index rows, Index cols, Index k2) {
329 symm_pack_rhs_helper<float, StorageOrder, 1>(blockB, _rhs, rhsStride, rows, cols, k2);
333template <
typename Index,
int Pack1,
int Pack2_dummy,
int StorageOrder>
334struct symm_pack_lhs<float, Index, Pack1, Pack2_dummy, StorageOrder> {
335 void operator()(
float* blockA,
const float* _lhs, Index lhsStride, Index cols, Index rows) {
336 symm_pack_lhs_helper<float, StorageOrder>(blockA, _lhs, lhsStride, cols, rows);
341template <
typename Index,
int nr,
int StorageOrder>
342struct symm_pack_rhs<double, Index, nr, StorageOrder> {
343 void operator()(
double* blockB,
const double* _rhs, Index rhsStride, Index rows, Index cols, Index k2) {
344 symm_pack_rhs_helper<double, StorageOrder, 2>(blockB, _rhs, rhsStride, rows, cols, k2);
348template <
typename Index,
int Pack1,
int Pack2_dummy,
int StorageOrder>
349struct symm_pack_lhs<double, Index, Pack1, Pack2_dummy, StorageOrder> {
350 void operator()(
double* blockA,
const double* _lhs, Index lhsStride, Index cols, Index rows) {
351 symm_pack_lhs_helper<double, StorageOrder>(blockA, _lhs, lhsStride, cols, rows);
366template <
typename Scalar,
typename Packet,
int N>
367EIGEN_ALWAYS_INLINE
void storeBlock(Scalar* to, PacketBlock<Packet, N>& block) {
368 const Index size = 16 /
sizeof(Scalar);
369 pstore<Scalar>(to + (0 * size), block.packet[0]);
370 pstore<Scalar>(to + (1 * size), block.packet[1]);
371 EIGEN_IF_CONSTEXPR (N > 2) {
372 pstore<Scalar>(to + (2 * size), block.packet[2]);
374 EIGEN_IF_CONSTEXPR (N > 3) {
375 pstore<Scalar>(to + (3 * size), block.packet[3]);
380template <
typename Scalar,
typename DataMapper,
typename Packet,
typename PacketC,
int StorageOrder,
bool Conjugate,
381 bool PanelMode,
bool UseLhs>
383 template <
bool transpose>
384 EIGEN_ALWAYS_INLINE
void dhs_cblock(PacketBlock<PacketC, 8>& cblock, PacketBlock<Packet, 4>& block,
385 const Packet16uc& permute) {
387 block.packet[0] = vec_perm(cblock.packet[0].v, cblock.packet[1].v, permute);
388 block.packet[1] = vec_perm(cblock.packet[2].v, cblock.packet[3].v, permute);
389 block.packet[2] = vec_perm(cblock.packet[4].v, cblock.packet[5].v, permute);
390 block.packet[3] = vec_perm(cblock.packet[6].v, cblock.packet[7].v, permute);
392 Packet4f t0, t1, t2, t3;
393#ifdef EIGEN_VECTORIZE_VSX
394 t0 =
reinterpret_cast<Packet
>(
395 vec_mergeh(
reinterpret_cast<Packet2ul
>(block.packet[0]),
reinterpret_cast<Packet2ul
>(block.packet[1])));
396 t1 =
reinterpret_cast<Packet
>(
397 vec_mergel(
reinterpret_cast<Packet2ul
>(block.packet[0]),
reinterpret_cast<Packet2ul
>(block.packet[1])));
398 t2 =
reinterpret_cast<Packet
>(
399 vec_mergeh(
reinterpret_cast<Packet2ul
>(block.packet[2]),
reinterpret_cast<Packet2ul
>(block.packet[3])));
400 t3 =
reinterpret_cast<Packet
>(
401 vec_mergel(
reinterpret_cast<Packet2ul
>(block.packet[2]),
reinterpret_cast<Packet2ul
>(block.packet[3])));
403 t0 =
reinterpret_cast<Packet
>(vec_perm(block.packet[0], block.packet[1], p16uc_TRANSPOSE64_HI));
404 t1 =
reinterpret_cast<Packet
>(vec_perm(block.packet[0], block.packet[1], p16uc_TRANSPOSE64_LO));
405 t2 =
reinterpret_cast<Packet
>(vec_perm(block.packet[2], block.packet[3], p16uc_TRANSPOSE64_HI));
406 t3 =
reinterpret_cast<Packet
>(vec_perm(block.packet[2], block.packet[3], p16uc_TRANSPOSE64_LO));
409 block.packet[0] = t0;
410 block.packet[1] = t1;
411 block.packet[2] = t2;
412 block.packet[3] = t3;
414 block.packet[0] = vec_perm(cblock.packet[0].v, cblock.packet[4].v, permute);
415 block.packet[1] = vec_perm(cblock.packet[1].v, cblock.packet[5].v, permute);
416 block.packet[2] = vec_perm(cblock.packet[2].v, cblock.packet[6].v, permute);
417 block.packet[3] = vec_perm(cblock.packet[3].v, cblock.packet[7].v, permute);
421 EIGEN_ALWAYS_INLINE
void dhs_ccopy(Scalar* blockAt,
const DataMapper& lhs2, Index& i, Index& rir, Index& rii,
422 Index depth,
const Index vectorSize) {
423 PacketBlock<Packet, 4> blockr, blocki;
424 PacketBlock<PacketC, 8> cblock;
426 for (; i + vectorSize <= depth; i += vectorSize) {
427 EIGEN_IF_CONSTEXPR (UseLhs) {
428 bload<DataMapper, PacketC, 2, StorageOrder, true, 4>(cblock, lhs2, 0, i);
430 bload<DataMapper, PacketC, 2, StorageOrder, true, 4>(cblock, lhs2, i, 0);
433 EIGEN_IF_CONSTEXPR (((StorageOrder ==
RowMajor) && UseLhs) || (((StorageOrder ==
ColMajor) && !UseLhs))) {
434 dhs_cblock<true>(cblock, blockr, p16uc_GETREAL32b);
435 dhs_cblock<true>(cblock, blocki, p16uc_GETIMAG32b);
437 dhs_cblock<false>(cblock, blockr, p16uc_GETREAL32);
438 dhs_cblock<false>(cblock, blocki, p16uc_GETIMAG32);
441 EIGEN_IF_CONSTEXPR (Conjugate) {
442 blocki.packet[0] = -blocki.packet[0];
443 blocki.packet[1] = -blocki.packet[1];
444 blocki.packet[2] = -blocki.packet[2];
445 blocki.packet[3] = -blocki.packet[3];
448 storeBlock<Scalar, Packet, 4>(blockAt + rir, blockr);
449 storeBlock<Scalar, Packet, 4>(blockAt + rii, blocki);
451 rir += 4 * vectorSize;
452 rii += 4 * vectorSize;
456 EIGEN_STRONG_INLINE
void operator()(std::complex<Scalar>* blockA,
const DataMapper& lhs, Index depth, Index rows,
457 Index stride, Index offset) {
458 const Index vectorSize = quad_traits<Scalar>::vectorsize;
459 const Index vectorDelta = vectorSize * ((PanelMode) ? stride : depth);
460 Index rir = ((PanelMode) ? (vectorSize * offset) : 0), rii;
461 Scalar* blockAt =
reinterpret_cast<Scalar*
>(blockA);
464 for (; j + vectorSize <= rows; j += vectorSize) {
465 const DataMapper lhs2 = UseLhs ? lhs.getSubMapper(j, 0) : lhs.getSubMapper(0, j);
468 rii = rir + vectorDelta;
470 dhs_ccopy(blockAt, lhs2, i, rir, rii, depth, vectorSize);
472 for (; i < depth; i++) {
473 PacketBlock<Packet, 1> blockr, blocki;
474 PacketBlock<PacketC, 2> cblock;
476 EIGEN_IF_CONSTEXPR (((StorageOrder ==
ColMajor) && UseLhs) || (((StorageOrder ==
RowMajor) && !UseLhs))) {
477 EIGEN_IF_CONSTEXPR (UseLhs) {
478 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i);
479 cblock.packet[1] = lhs2.template loadPacket<PacketC>(2, i);
481 cblock.packet[0] = lhs2.template loadPacket<PacketC>(i, 0);
482 cblock.packet[1] = lhs2.template loadPacket<PacketC>(i, 2);
485 EIGEN_IF_CONSTEXPR (UseLhs) {
486 cblock.packet[0] = pload2(lhs2(0, i), lhs2(1, i));
487 cblock.packet[1] = pload2(lhs2(2, i), lhs2(3, i));
489 cblock.packet[0] = pload2(lhs2(i, 0), lhs2(i, 1));
490 cblock.packet[1] = pload2(lhs2(i, 2), lhs2(i, 3));
494 blockr.packet[0] = vec_perm(cblock.packet[0].v, cblock.packet[1].v, p16uc_GETREAL32);
495 blocki.packet[0] = vec_perm(cblock.packet[0].v, cblock.packet[1].v, p16uc_GETIMAG32);
497 EIGEN_IF_CONSTEXPR (Conjugate) {
498 blocki.packet[0] = -blocki.packet[0];
501 pstore<Scalar>(blockAt + rir, blockr.packet[0]);
502 pstore<Scalar>(blockAt + rii, blocki.packet[0]);
508 rir += ((PanelMode) ? (vectorSize * (2 * stride - depth)) : vectorDelta);
511 EIGEN_IF_CONSTEXPR (!UseLhs) {
512 EIGEN_IF_CONSTEXPR (PanelMode) rir -= (offset * (vectorSize - 1));
514 for (; j < rows; j++) {
515 const DataMapper lhs2 = lhs.getSubMapper(0, j);
516 rii = rir + ((PanelMode) ? stride : depth);
518 for (Index i = 0; i < depth; i++) {
519 blockAt[rir] = lhs2(i, 0).real();
521 EIGEN_IF_CONSTEXPR (Conjugate)
522 blockAt[rii] = -lhs2(i, 0).imag();
524 blockAt[rii] = lhs2(i, 0).imag();
530 rir += ((PanelMode) ? (2 * stride - depth) : depth);
534 EIGEN_IF_CONSTEXPR (PanelMode) rir += (offset * (rows - j - vectorSize));
535 rii = rir + (((PanelMode) ? stride : depth) * (rows - j));
537 for (Index i = 0; i < depth; i++) {
539 for (; k < rows; k++) {
540 blockAt[rir] = lhs(k, i).real();
542 EIGEN_IF_CONSTEXPR (Conjugate)
543 blockAt[rii] = -lhs(k, i).imag();
545 blockAt[rii] = lhs(k, i).imag();
557template <
typename Scalar,
typename DataMapper,
typename Packet,
int StorageOrder,
bool PanelMode,
bool UseLhs>
560 EIGEN_ALWAYS_INLINE
void dhs_copy(Scalar* blockA,
const DataMapper& lhs2, Index& i, Index& ri, Index depth,
561 const Index vectorSize) {
562 PacketBlock<Packet, 4> block[n];
564 for (; i + n * vectorSize <= depth; i += n * vectorSize) {
565 for (Index k = 0; k < n; k++) {
566 EIGEN_IF_CONSTEXPR (UseLhs) {
567 bload<DataMapper, Packet, 4, StorageOrder, false, 4>(block[k], lhs2, 0, i + k * vectorSize);
569 bload<DataMapper, Packet, 4, StorageOrder, false, 4>(block[k], lhs2, i + k * vectorSize, 0);
573 EIGEN_IF_CONSTEXPR (((StorageOrder ==
RowMajor) && UseLhs) || ((StorageOrder ==
ColMajor) && !UseLhs)) {
574 for (Index k = 0; k < n; k++) {
575 ptranspose(block[k]);
579 for (Index k = 0; k < n; k++) {
580 storeBlock<Scalar, Packet, 4>(blockA + ri + k * 4 * vectorSize, block[k]);
583 ri += n * 4 * vectorSize;
587 EIGEN_STRONG_INLINE
void operator()(Scalar* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride,
589 const Index vectorSize = quad_traits<Scalar>::vectorsize;
592 for (; j + vectorSize <= rows; j += vectorSize) {
593 const DataMapper lhs2 = UseLhs ? lhs.getSubMapper(j, 0) : lhs.getSubMapper(0, j);
596 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * offset;
598 dhs_copy<4>(blockA, lhs2, i, ri, depth, vectorSize);
599 dhs_copy<2>(blockA, lhs2, i, ri, depth, vectorSize);
600 dhs_copy<1>(blockA, lhs2, i, ri, depth, vectorSize);
602 for (; i < depth; i++) {
603 EIGEN_IF_CONSTEXPR (((StorageOrder ==
RowMajor) && UseLhs) || ((StorageOrder ==
ColMajor) && !UseLhs)) {
604 EIGEN_IF_CONSTEXPR (UseLhs) {
605 blockA[ri + 0] = lhs2(0, i);
606 blockA[ri + 1] = lhs2(1, i);
607 blockA[ri + 2] = lhs2(2, i);
608 blockA[ri + 3] = lhs2(3, i);
610 blockA[ri + 0] = lhs2(i, 0);
611 blockA[ri + 1] = lhs2(i, 1);
612 blockA[ri + 2] = lhs2(i, 2);
613 blockA[ri + 3] = lhs2(i, 3);
617 EIGEN_IF_CONSTEXPR (UseLhs) {
618 lhsV = lhs2.template loadPacket<Packet>(0, i);
620 lhsV = lhs2.template loadPacket<Packet>(i, 0);
622 pstore<Scalar>(blockA + ri, lhsV);
628 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * (stride - offset - depth);
631 EIGEN_IF_CONSTEXPR (!UseLhs) {
632 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset;
634 for (; j < rows; j++) {
635 const DataMapper lhs2 = lhs.getSubMapper(0, j);
636 for (Index i = 0; i < depth; i++) {
637 blockA[ri] = lhs2(i, 0);
641 EIGEN_IF_CONSTEXPR (PanelMode) ri += stride - depth;
645 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (rows - j);
647 for (Index i = 0; i < depth; i++) {
649 for (; k < rows; k++) {
650 blockA[ri] = lhs(k, i);
660template <
typename DataMapper,
int StorageOrder,
bool PanelMode>
661struct dhs_pack<double, DataMapper, Packet2d, StorageOrder, PanelMode, true> {
663 EIGEN_ALWAYS_INLINE
void dhs_copy(
double* blockA,
const DataMapper& lhs2, Index& i, Index& ri, Index depth,
664 const Index vectorSize) {
665 PacketBlock<Packet2d, 2> block[n];
667 for (; i + n * vectorSize <= depth; i += n * vectorSize) {
668 for (Index k = 0; k < n; k++) {
669 EIGEN_IF_CONSTEXPR (StorageOrder ==
RowMajor) {
670 block[k].packet[0] = lhs2.template loadPacket<Packet2d>(0, i + k * vectorSize);
671 block[k].packet[1] = lhs2.template loadPacket<Packet2d>(1, i + k * vectorSize);
673 block[k].packet[0] = lhs2.template loadPacket<Packet2d>(0, i + k * vectorSize + 0);
674 block[k].packet[1] = lhs2.template loadPacket<Packet2d>(0, i + k * vectorSize + 1);
678 EIGEN_IF_CONSTEXPR (StorageOrder ==
RowMajor) {
679 for (Index k = 0; k < n; k++) {
680 ptranspose(block[k]);
684 for (Index k = 0; k < n; k++) {
685 storeBlock<double, Packet2d, 2>(blockA + ri + k * 2 * vectorSize, block[k]);
688 ri += n * 2 * vectorSize;
692 EIGEN_STRONG_INLINE
void operator()(
double* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride,
694 const Index vectorSize = quad_traits<double>::vectorsize;
697 for (; j + vectorSize <= rows; j += vectorSize) {
698 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
701 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * offset;
703 dhs_copy<4>(blockA, lhs2, i, ri, depth, vectorSize);
704 dhs_copy<2>(blockA, lhs2, i, ri, depth, vectorSize);
705 dhs_copy<1>(blockA, lhs2, i, ri, depth, vectorSize);
707 for (; i < depth; i++) {
708 EIGEN_IF_CONSTEXPR (StorageOrder ==
RowMajor) {
709 blockA[ri + 0] = lhs2(0, i);
710 blockA[ri + 1] = lhs2(1, i);
712 Packet2d lhsV = lhs2.template loadPacket<Packet2d>(0, i);
713 pstore<double>(blockA + ri, lhsV);
719 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * (stride - offset - depth);
723 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (rows - j);
725 for (Index i = 0; i < depth; i++) {
727 for (; k < rows; k++) {
728 blockA[ri] = lhs(k, i);
737template <
typename DataMapper,
int StorageOrder,
bool PanelMode>
738struct dhs_pack<double, DataMapper, Packet2d, StorageOrder, PanelMode, false> {
740 EIGEN_ALWAYS_INLINE
void dhs_copy(
double* blockB,
const DataMapper& rhs2, Index& i, Index& ri, Index depth,
741 const Index vectorSize) {
742 PacketBlock<Packet2d, 2> block1[n], block2[n];
743 PacketBlock<Packet2d, 4> block3[n];
745 for (; i + n * vectorSize <= depth; i += n * vectorSize) {
746 for (Index k = 0; k < n; k++) {
747 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
748 block1[k].packet[0] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize, 0);
749 block1[k].packet[1] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize, 1);
750 block2[k].packet[0] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize, 2);
751 block2[k].packet[1] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize, 3);
753 block3[k].packet[0] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 0, 0);
754 block3[k].packet[1] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 0, 2);
755 block3[k].packet[2] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 1, 0);
756 block3[k].packet[3] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 1, 2);
760 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
761 for (Index k = 0; k < n; k++) {
762 ptranspose(block1[k]);
763 ptranspose(block2[k]);
767 for (Index k = 0; k < n; k++) {
768 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
769 pstore<double>(blockB + ri + k * 4 * vectorSize, block1[k].packet[0]);
770 pstore<double>(blockB + ri + k * 4 * vectorSize + 2, block2[k].packet[0]);
771 pstore<double>(blockB + ri + k * 4 * vectorSize + 4, block1[k].packet[1]);
772 pstore<double>(blockB + ri + k * 4 * vectorSize + 6, block2[k].packet[1]);
774 storeBlock<double, Packet2d, 4>(blockB + ri + k * 4 * vectorSize, block3[k]);
778 ri += n * 4 * vectorSize;
782 EIGEN_STRONG_INLINE
void operator()(
double* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride,
784 const Index vectorSize = quad_traits<double>::vectorsize;
787 for (; j + 2 * vectorSize <= cols; j += 2 * vectorSize) {
788 const DataMapper rhs2 = rhs.getSubMapper(0, j);
791 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (2 * vectorSize);
793 dhs_copy<4>(blockB, rhs2, i, ri, depth, vectorSize);
794 dhs_copy<2>(blockB, rhs2, i, ri, depth, vectorSize);
795 dhs_copy<1>(blockB, rhs2, i, ri, depth, vectorSize);
797 for (; i < depth; i++) {
798 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
799 blockB[ri + 0] = rhs2(i, 0);
800 blockB[ri + 1] = rhs2(i, 1);
804 blockB[ri + 0] = rhs2(i, 2);
805 blockB[ri + 1] = rhs2(i, 3);
807 Packet2d rhsV = rhs2.template loadPacket<Packet2d>(i, 0);
808 pstore<double>(blockB + ri, rhsV);
812 rhsV = rhs2.template loadPacket<Packet2d>(i, 2);
813 pstore<double>(blockB + ri, rhsV);
818 EIGEN_IF_CONSTEXPR (PanelMode) ri += (2 * vectorSize) * (stride - offset - depth);
821 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset;
823 for (; j < cols; j++) {
824 const DataMapper rhs2 = rhs.getSubMapper(0, j);
825 for (Index i = 0; i < depth; i++) {
826 blockB[ri] = rhs2(i, 0);
830 EIGEN_IF_CONSTEXPR (PanelMode) ri += stride - depth;
836template <
typename DataMapper,
int StorageOrder,
bool PanelMode>
837struct dhs_pack<bfloat16, DataMapper, Packet8bf, StorageOrder, PanelMode, true> {
838 EIGEN_STRONG_INLINE
void operator()(bfloat16* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride,
840 const Index vectorSize = quad_traits<bfloat16>::vectorsize;
843 for (; j + 2 * vectorSize <= rows; j += 2 * vectorSize) {
844 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
847 EIGEN_IF_CONSTEXPR (PanelMode) ri += 2 * vectorSize * offset;
849 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
850 for (; i + 2 <= depth; i += 2) {
851 PacketBlock<Packet8bf, 4> block;
853 block.packet[0] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
854 block.packet[1] = lhs2.template loadPacket<Packet8bf>(1 * vectorSize, i + 0);
855 block.packet[2] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 1);
856 block.packet[3] = lhs2.template loadPacket<Packet8bf>(1 * vectorSize, i + 1);
859 t0 = vec_mergeh(block.packet[0].m_val, block.packet[2].m_val);
860 t1 = vec_mergel(block.packet[0].m_val, block.packet[2].m_val);
861 block.packet[2] = vec_mergeh(block.packet[1].m_val, block.packet[3].m_val);
862 block.packet[3] = vec_mergel(block.packet[1].m_val, block.packet[3].m_val);
863 block.packet[0] = t0;
864 block.packet[1] = t1;
866 storeBlock<bfloat16, Packet8bf, 4>(blockA + ri, block);
868 ri += 2 * 2 * vectorSize;
871 PacketBlock<Packet8bf, 2> block;
873 block.packet[0] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
874 block.packet[1] = lhs2.template loadPacket<Packet8bf>(1 * vectorSize, i + 0);
876 storeBlock<bfloat16, Packet8bf, 2>(blockA + ri, block);
878 ri += 2 * vectorSize;
881 for (; i + vectorSize <= depth; i += vectorSize) {
882 PacketBlock<Packet8bf, 8> block1, block2;
884 bload<DataMapper, Packet8bf, 8, StorageOrder, false, 8>(block1, lhs2, 0 * vectorSize, i);
885 bload<DataMapper, Packet8bf, 8, StorageOrder, false, 8>(block2, lhs2, 1 * vectorSize, i);
887 Packet4ui v1[8], v2[8];
889 v1[0] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[0].m_val),
890 reinterpret_cast<Packet4ui
>(block1.packet[1].m_val));
891 v1[1] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[0].m_val),
892 reinterpret_cast<Packet4ui
>(block1.packet[1].m_val));
893 v1[2] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[2].m_val),
894 reinterpret_cast<Packet4ui
>(block1.packet[3].m_val));
895 v1[3] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[2].m_val),
896 reinterpret_cast<Packet4ui
>(block1.packet[3].m_val));
897 v1[4] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[4].m_val),
898 reinterpret_cast<Packet4ui
>(block1.packet[5].m_val));
899 v1[5] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[4].m_val),
900 reinterpret_cast<Packet4ui
>(block1.packet[5].m_val));
901 v1[6] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[6].m_val),
902 reinterpret_cast<Packet4ui
>(block1.packet[7].m_val));
903 v1[7] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[6].m_val),
904 reinterpret_cast<Packet4ui
>(block1.packet[7].m_val));
905 v2[0] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block2.packet[0].m_val),
906 reinterpret_cast<Packet4ui
>(block2.packet[1].m_val));
907 v2[1] = vec_mergel(
reinterpret_cast<Packet4ui
>(block2.packet[0].m_val),
908 reinterpret_cast<Packet4ui
>(block2.packet[1].m_val));
909 v2[2] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block2.packet[2].m_val),
910 reinterpret_cast<Packet4ui
>(block2.packet[3].m_val));
911 v2[3] = vec_mergel(
reinterpret_cast<Packet4ui
>(block2.packet[2].m_val),
912 reinterpret_cast<Packet4ui
>(block2.packet[3].m_val));
913 v2[4] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block2.packet[4].m_val),
914 reinterpret_cast<Packet4ui
>(block2.packet[5].m_val));
915 v2[5] = vec_mergel(
reinterpret_cast<Packet4ui
>(block2.packet[4].m_val),
916 reinterpret_cast<Packet4ui
>(block2.packet[5].m_val));
917 v2[6] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block2.packet[6].m_val),
918 reinterpret_cast<Packet4ui
>(block2.packet[7].m_val));
919 v2[7] = vec_mergel(
reinterpret_cast<Packet4ui
>(block2.packet[6].m_val),
920 reinterpret_cast<Packet4ui
>(block2.packet[7].m_val));
922#ifdef EIGEN_VECTORIZE_VSX
923 block1.packet[0] =
reinterpret_cast<Packet8us
>(
924 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[0]),
reinterpret_cast<Packet2ul
>(v1[2])));
925 block1.packet[2] =
reinterpret_cast<Packet8us
>(
926 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[0]),
reinterpret_cast<Packet2ul
>(v1[2])));
927 block1.packet[4] =
reinterpret_cast<Packet8us
>(
928 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[1]),
reinterpret_cast<Packet2ul
>(v1[3])));
929 block1.packet[6] =
reinterpret_cast<Packet8us
>(
930 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[1]),
reinterpret_cast<Packet2ul
>(v1[3])));
931 block1.packet[1] =
reinterpret_cast<Packet8us
>(
932 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[4]),
reinterpret_cast<Packet2ul
>(v1[6])));
933 block1.packet[3] =
reinterpret_cast<Packet8us
>(
934 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[4]),
reinterpret_cast<Packet2ul
>(v1[6])));
935 block1.packet[5] =
reinterpret_cast<Packet8us
>(
936 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[5]),
reinterpret_cast<Packet2ul
>(v1[7])));
937 block1.packet[7] =
reinterpret_cast<Packet8us
>(
938 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[5]),
reinterpret_cast<Packet2ul
>(v1[7])));
939 block2.packet[0] =
reinterpret_cast<Packet8us
>(
940 vec_mergeh(
reinterpret_cast<Packet2ul
>(v2[0]),
reinterpret_cast<Packet2ul
>(v2[2])));
941 block2.packet[2] =
reinterpret_cast<Packet8us
>(
942 vec_mergel(
reinterpret_cast<Packet2ul
>(v2[0]),
reinterpret_cast<Packet2ul
>(v2[2])));
943 block2.packet[4] =
reinterpret_cast<Packet8us
>(
944 vec_mergeh(
reinterpret_cast<Packet2ul
>(v2[1]),
reinterpret_cast<Packet2ul
>(v2[3])));
945 block2.packet[6] =
reinterpret_cast<Packet8us
>(
946 vec_mergel(
reinterpret_cast<Packet2ul
>(v2[1]),
reinterpret_cast<Packet2ul
>(v2[3])));
947 block2.packet[1] =
reinterpret_cast<Packet8us
>(
948 vec_mergeh(
reinterpret_cast<Packet2ul
>(v2[4]),
reinterpret_cast<Packet2ul
>(v2[6])));
949 block2.packet[3] =
reinterpret_cast<Packet8us
>(
950 vec_mergel(
reinterpret_cast<Packet2ul
>(v2[4]),
reinterpret_cast<Packet2ul
>(v2[6])));
951 block2.packet[5] =
reinterpret_cast<Packet8us
>(
952 vec_mergeh(
reinterpret_cast<Packet2ul
>(v2[5]),
reinterpret_cast<Packet2ul
>(v2[7])));
953 block2.packet[7] =
reinterpret_cast<Packet8us
>(
954 vec_mergel(
reinterpret_cast<Packet2ul
>(v2[5]),
reinterpret_cast<Packet2ul
>(v2[7])));
956 block1.packet[0] =
reinterpret_cast<Packet8us
>(vec_perm(v1[0], v1[2], p16uc_TRANSPOSE64_HI));
957 block1.packet[2] =
reinterpret_cast<Packet8us
>(vec_perm(v1[0], v1[2], p16uc_TRANSPOSE64_LO));
958 block1.packet[4] =
reinterpret_cast<Packet8us
>(vec_perm(v1[1], v1[3], p16uc_TRANSPOSE64_HI));
959 block1.packet[6] =
reinterpret_cast<Packet8us
>(vec_perm(v1[1], v1[3], p16uc_TRANSPOSE64_LO));
960 block1.packet[1] =
reinterpret_cast<Packet8us
>(vec_perm(v1[4], v1[6], p16uc_TRANSPOSE64_HI));
961 block1.packet[3] =
reinterpret_cast<Packet8us
>(vec_perm(v1[4], v1[6], p16uc_TRANSPOSE64_LO));
962 block1.packet[5] =
reinterpret_cast<Packet8us
>(vec_perm(v1[5], v1[7], p16uc_TRANSPOSE64_HI));
963 block1.packet[7] =
reinterpret_cast<Packet8us
>(vec_perm(v1[5], v1[7], p16uc_TRANSPOSE64_LO));
964 block2.packet[0] =
reinterpret_cast<Packet8us
>(vec_perm(v2[0], v2[2], p16uc_TRANSPOSE64_HI));
965 block2.packet[2] =
reinterpret_cast<Packet8us
>(vec_perm(v2[0], v2[2], p16uc_TRANSPOSE64_LO));
966 block2.packet[4] =
reinterpret_cast<Packet8us
>(vec_perm(v2[1], v2[3], p16uc_TRANSPOSE64_HI));
967 block2.packet[6] =
reinterpret_cast<Packet8us
>(vec_perm(v2[1], v2[3], p16uc_TRANSPOSE64_LO));
968 block2.packet[1] =
reinterpret_cast<Packet8us
>(vec_perm(v2[4], v2[6], p16uc_TRANSPOSE64_HI));
969 block2.packet[3] =
reinterpret_cast<Packet8us
>(vec_perm(v2[4], v2[6], p16uc_TRANSPOSE64_LO));
970 block2.packet[5] =
reinterpret_cast<Packet8us
>(vec_perm(v2[5], v2[7], p16uc_TRANSPOSE64_HI));
971 block2.packet[7] =
reinterpret_cast<Packet8us
>(vec_perm(v2[5], v2[7], p16uc_TRANSPOSE64_LO));
974 for (Index M = 0; M < 8; M += 2) {
975 pstore<bfloat16>(blockA + ri + (0 * vectorSize) + (2 * vectorSize * M), block1.packet[M + 0]);
976 pstore<bfloat16>(blockA + ri + (1 * vectorSize) + (2 * vectorSize * M), block1.packet[M + 1]);
977 pstore<bfloat16>(blockA + ri + (2 * vectorSize) + (2 * vectorSize * M), block2.packet[M + 0]);
978 pstore<bfloat16>(blockA + ri + (3 * vectorSize) + (2 * vectorSize * M), block2.packet[M + 1]);
981 ri += 2 * vectorSize * vectorSize;
983 for (; i + 2 <= depth; i += 2) {
984 for (Index M = 0; M < 2 * vectorSize; M++) {
985 blockA[ri + (M * 2) + 0] = lhs2(M, i + 0);
986 blockA[ri + (M * 2) + 1] = lhs2(M, i + 1);
989 ri += 2 * 2 * vectorSize;
992 for (Index M = 0; M < 2 * vectorSize; M++) {
993 blockA[ri + M] = lhs2(M, i);
995 ri += 2 * vectorSize;
999 EIGEN_IF_CONSTEXPR (PanelMode) ri += 2 * vectorSize * (stride - offset - depth);
1001 for (; j + vectorSize <= rows; j += vectorSize) {
1002 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
1005 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * offset;
1007 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
1008 for (; i + 2 <= depth; i += 2) {
1009 PacketBlock<Packet8bf, 2> block;
1011 block.packet[0] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
1012 block.packet[1] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 1);
1015 t0 = vec_mergeh(block.packet[0].m_val, block.packet[1].m_val);
1016 block.packet[1] = vec_mergel(block.packet[0].m_val, block.packet[1].m_val);
1017 block.packet[0] = t0;
1019 storeBlock<bfloat16, Packet8bf, 2>(blockA + ri, block);
1021 ri += 2 * vectorSize;
1024 Packet8bf lhsV = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
1025 pstore<bfloat16>(blockA + ri, lhsV);
1030 for (; i + vectorSize <= depth; i += vectorSize) {
1031 PacketBlock<Packet8bf, 8> block1;
1033 bload<DataMapper, Packet8bf, 8, StorageOrder, false, 8>(block1, lhs2, 0 * vectorSize, i);
1038 v1[0] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[0].m_val),
1039 reinterpret_cast<Packet4ui
>(block1.packet[1].m_val));
1040 v1[1] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[0].m_val),
1041 reinterpret_cast<Packet4ui
>(block1.packet[1].m_val));
1042 v1[2] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[2].m_val),
1043 reinterpret_cast<Packet4ui
>(block1.packet[3].m_val));
1044 v1[3] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[2].m_val),
1045 reinterpret_cast<Packet4ui
>(block1.packet[3].m_val));
1046 v1[4] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[4].m_val),
1047 reinterpret_cast<Packet4ui
>(block1.packet[5].m_val));
1048 v1[5] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[4].m_val),
1049 reinterpret_cast<Packet4ui
>(block1.packet[5].m_val));
1050 v1[6] = vec_mergeh(
reinterpret_cast<Packet4ui
>(block1.packet[6].m_val),
1051 reinterpret_cast<Packet4ui
>(block1.packet[7].m_val));
1052 v1[7] = vec_mergel(
reinterpret_cast<Packet4ui
>(block1.packet[6].m_val),
1053 reinterpret_cast<Packet4ui
>(block1.packet[7].m_val));
1055#ifdef EIGEN_VECTORIZE_VSX
1056 block1.packet[0] =
reinterpret_cast<Packet8us
>(
1057 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[0]),
reinterpret_cast<Packet2ul
>(v1[2])));
1058 block1.packet[2] =
reinterpret_cast<Packet8us
>(
1059 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[0]),
reinterpret_cast<Packet2ul
>(v1[2])));
1060 block1.packet[4] =
reinterpret_cast<Packet8us
>(
1061 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[1]),
reinterpret_cast<Packet2ul
>(v1[3])));
1062 block1.packet[6] =
reinterpret_cast<Packet8us
>(
1063 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[1]),
reinterpret_cast<Packet2ul
>(v1[3])));
1064 block1.packet[1] =
reinterpret_cast<Packet8us
>(
1065 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[4]),
reinterpret_cast<Packet2ul
>(v1[6])));
1066 block1.packet[3] =
reinterpret_cast<Packet8us
>(
1067 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[4]),
reinterpret_cast<Packet2ul
>(v1[6])));
1068 block1.packet[5] =
reinterpret_cast<Packet8us
>(
1069 vec_mergeh(
reinterpret_cast<Packet2ul
>(v1[5]),
reinterpret_cast<Packet2ul
>(v1[7])));
1070 block1.packet[7] =
reinterpret_cast<Packet8us
>(
1071 vec_mergel(
reinterpret_cast<Packet2ul
>(v1[5]),
reinterpret_cast<Packet2ul
>(v1[7])));
1073 block1.packet[0] =
reinterpret_cast<Packet8us
>(vec_perm(v1[0], v1[2], p16uc_TRANSPOSE64_HI));
1074 block1.packet[2] =
reinterpret_cast<Packet8us
>(vec_perm(v1[0], v1[2], p16uc_TRANSPOSE64_LO));
1075 block1.packet[4] =
reinterpret_cast<Packet8us
>(vec_perm(v1[1], v1[3], p16uc_TRANSPOSE64_HI));
1076 block1.packet[6] =
reinterpret_cast<Packet8us
>(vec_perm(v1[1], v1[3], p16uc_TRANSPOSE64_LO));
1077 block1.packet[1] =
reinterpret_cast<Packet8us
>(vec_perm(v1[4], v1[6], p16uc_TRANSPOSE64_HI));
1078 block1.packet[3] =
reinterpret_cast<Packet8us
>(vec_perm(v1[4], v1[6], p16uc_TRANSPOSE64_LO));
1079 block1.packet[5] =
reinterpret_cast<Packet8us
>(vec_perm(v1[5], v1[7], p16uc_TRANSPOSE64_HI));
1080 block1.packet[7] =
reinterpret_cast<Packet8us
>(vec_perm(v1[5], v1[7], p16uc_TRANSPOSE64_LO));
1083 for (Index M = 0; M < 8; M++) {
1084 pstore<bfloat16>(blockA + ri + (vectorSize * M), block1.packet[M]);
1087 ri += vectorSize * vectorSize;
1089 for (; i + 2 <= depth; i += 2) {
1090 for (Index M = 0; M < vectorSize; M++) {
1091 blockA[ri + (M * 2) + 0] = lhs2(M, i + 0);
1092 blockA[ri + (M * 2) + 1] = lhs2(M, i + 1);
1095 ri += 2 * vectorSize;
1098 for (Index M = 0; M < vectorSize; M++) {
1099 blockA[ri + M] = lhs2(M, i);
1106 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * (stride - offset - depth);
1108 if (j + 4 <= rows) {
1109 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
1112 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * offset;
1114 for (; i + 2 <= depth; i += 2) {
1115 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
1116 PacketBlock<Packet8bf, 2> block;
1118 block.packet[0] = lhs2.template loadPacketPartial<Packet8bf>(0, i + 0, 4);
1119 block.packet[1] = lhs2.template loadPacketPartial<Packet8bf>(0, i + 1, 4);
1121 block.packet[0] = vec_mergeh(block.packet[0].m_val, block.packet[1].m_val);
1123 pstore<bfloat16>(blockA + ri, block.packet[0]);
1125 blockA[ri + 0] = lhs2(0, i + 0);
1126 blockA[ri + 1] = lhs2(0, i + 1);
1127 blockA[ri + 2] = lhs2(1, i + 0);
1128 blockA[ri + 3] = lhs2(1, i + 1);
1129 blockA[ri + 4] = lhs2(2, i + 0);
1130 blockA[ri + 5] = lhs2(2, i + 1);
1131 blockA[ri + 6] = lhs2(3, i + 0);
1132 blockA[ri + 7] = lhs2(3, i + 1);
1138 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
1139 Packet8bf lhsV = lhs2.template loadPacketPartial<Packet8bf>(0, i + 0, 4);
1141 pstore_partial<bfloat16>(blockA + ri, lhsV, 4);
1143 blockA[ri + 0] = lhs2(0, i);
1144 blockA[ri + 1] = lhs2(1, i);
1145 blockA[ri + 2] = lhs2(2, i);
1146 blockA[ri + 3] = lhs2(3, i);
1152 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * (stride - offset - depth);
1157 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (rows - j);
1160 for (; i + 2 <= depth; i += 2) {
1162 for (; k < rows; k++) {
1163 blockA[ri + 0] = lhs(k, i + 0);
1164 blockA[ri + 1] = lhs(k, i + 1);
1169 for (; j < rows; j++) {
1170 blockA[ri] = lhs(j, i);
1179template <
typename DataMapper,
int StorageOrder,
bool PanelMode>
1180struct dhs_pack<bfloat16, DataMapper, Packet8bf, StorageOrder, PanelMode, false> {
1181 EIGEN_STRONG_INLINE
void operator()(bfloat16* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride,
1183 const Index vectorSize = quad_traits<bfloat16>::vectorsize;
1184 Index ri = 0, j = 0;
1186 for (; j + 4 <= cols; j += 4) {
1187 const DataMapper rhs2 = rhs.getSubMapper(0, j);
1190 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * offset;
1192 for (; i + vectorSize <= depth; i += vectorSize) {
1193 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
1194 PacketBlock<Packet8bf, 4> block;
1196 bload<DataMapper, Packet8bf, 4, StorageOrder, false, 4>(block, rhs2, i, 0);
1198 Packet4ui t0, t1, t2, t3;
1200 t0 = vec_mergeh(
reinterpret_cast<Packet4ui
>(block.packet[0].m_val),
1201 reinterpret_cast<Packet4ui
>(block.packet[1].m_val));
1202 t1 = vec_mergel(
reinterpret_cast<Packet4ui
>(block.packet[0].m_val),
1203 reinterpret_cast<Packet4ui
>(block.packet[1].m_val));
1204 t2 = vec_mergeh(
reinterpret_cast<Packet4ui
>(block.packet[2].m_val),
1205 reinterpret_cast<Packet4ui
>(block.packet[3].m_val));
1206 t3 = vec_mergel(
reinterpret_cast<Packet4ui
>(block.packet[2].m_val),
1207 reinterpret_cast<Packet4ui
>(block.packet[3].m_val));
1209#ifdef EIGEN_VECTORIZE_VSX
1211 reinterpret_cast<Packet8us
>(vec_mergeh(
reinterpret_cast<Packet2ul
>(t0),
reinterpret_cast<Packet2ul
>(t2)));
1213 reinterpret_cast<Packet8us
>(vec_mergel(
reinterpret_cast<Packet2ul
>(t0),
reinterpret_cast<Packet2ul
>(t2)));
1215 reinterpret_cast<Packet8us
>(vec_mergeh(
reinterpret_cast<Packet2ul
>(t1),
reinterpret_cast<Packet2ul
>(t3)));
1217 reinterpret_cast<Packet8us
>(vec_mergel(
reinterpret_cast<Packet2ul
>(t1),
reinterpret_cast<Packet2ul
>(t3)));
1219 block.packet[0] =
reinterpret_cast<Packet8us
>(vec_perm(t0, t2, p16uc_TRANSPOSE64_HI));
1220 block.packet[1] =
reinterpret_cast<Packet8us
>(vec_perm(t0, t2, p16uc_TRANSPOSE64_LO));
1221 block.packet[2] =
reinterpret_cast<Packet8us
>(vec_perm(t1, t3, p16uc_TRANSPOSE64_HI));
1222 block.packet[3] =
reinterpret_cast<Packet8us
>(vec_perm(t1, t3, p16uc_TRANSPOSE64_LO));
1225 storeBlock<bfloat16, Packet8bf, 4>(blockB + ri, block);
1227 PacketBlock<Packet8bf, 8> block;
1229 for (
int M = 0; M < 8; M++) {
1230 block.packet[M] = rhs2.template loadPacketPartial<Packet8bf>(i + M, 0, 4);
1233 block.packet[0] = vec_mergeh(block.packet[0].m_val, block.packet[1].m_val);
1234 block.packet[1] = vec_mergeh(block.packet[2].m_val, block.packet[3].m_val);
1235 block.packet[2] = vec_mergeh(block.packet[4].m_val, block.packet[5].m_val);
1236 block.packet[3] = vec_mergeh(block.packet[6].m_val, block.packet[7].m_val);
1238 const Index size = 16 /
sizeof(bfloat16);
1240 for (
int M = 0; M < 4; M++) {
1241 pstore<bfloat16>(blockB + ri + (M * size), block.packet[M]);
1245 ri += 4 * vectorSize;
1247 for (; i + 2 <= depth; i += 2) {
1248 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
1249 blockB[ri + 0] = rhs2(i + 0, 0);
1250 blockB[ri + 1] = rhs2(i + 1, 0);
1251 blockB[ri + 2] = rhs2(i + 0, 1);
1252 blockB[ri + 3] = rhs2(i + 1, 1);
1253 blockB[ri + 4] = rhs2(i + 0, 2);
1254 blockB[ri + 5] = rhs2(i + 1, 2);
1255 blockB[ri + 6] = rhs2(i + 0, 3);
1256 blockB[ri + 7] = rhs2(i + 1, 3);
1258 PacketBlock<Packet8bf, 2> block;
1260 for (
int M = 0; M < 2; M++) {
1261 block.packet[M] = rhs2.template loadPacketPartial<Packet8bf>(i + M, 0, 4);
1264 block.packet[0] = vec_mergeh(block.packet[0].m_val, block.packet[1].m_val);
1266 pstore<bfloat16>(blockB + ri, block.packet[0]);
1272 blockB[ri + 0] = rhs2(i, 0);
1273 blockB[ri + 1] = rhs2(i, 1);
1274 blockB[ri + 2] = rhs2(i, 2);
1275 blockB[ri + 3] = rhs2(i, 3);
1280 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * (stride - offset - depth);
1284 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (cols - j);
1287 for (; i + 2 <= depth; i += 2) {
1289 for (; k < cols; k++) {
1290 blockB[ri + 0] = rhs(i + 0, k);
1291 blockB[ri + 1] = rhs(i + 1, k);
1296 for (; j < cols; j++) {
1297 blockB[ri] = rhs(i, j);
1306template <
typename DataMapper,
typename Packet,
typename PacketC,
int StorageOrder,
bool Conjugate,
bool PanelMode>
1307struct dhs_cpack<double, DataMapper, Packet, PacketC, StorageOrder, Conjugate, PanelMode, true> {
1308 EIGEN_ALWAYS_INLINE
void dhs_ccopy(
double* blockAt,
const DataMapper& lhs2, Index& i, Index& rir, Index& rii,
1309 Index depth,
const Index vectorSize) {
1310 PacketBlock<Packet, 2> blockr, blocki;
1311 PacketBlock<PacketC, 4> cblock;
1313 for (; i + vectorSize <= depth; i += vectorSize) {
1314 EIGEN_IF_CONSTEXPR (StorageOrder ==
ColMajor) {
1315 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i + 0);
1316 cblock.packet[1] = lhs2.template loadPacket<PacketC>(0, i + 1);
1318 cblock.packet[2] = lhs2.template loadPacket<PacketC>(1, i + 0);
1319 cblock.packet[3] = lhs2.template loadPacket<PacketC>(1, i + 1);
1321 blockr.packet[0] = vec_mergeh(cblock.packet[0].v, cblock.packet[2].v);
1322 blockr.packet[1] = vec_mergeh(cblock.packet[1].v, cblock.packet[3].v);
1324 blocki.packet[0] = vec_mergel(cblock.packet[0].v, cblock.packet[2].v);
1325 blocki.packet[1] = vec_mergel(cblock.packet[1].v, cblock.packet[3].v);
1327 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i);
1328 cblock.packet[1] = lhs2.template loadPacket<PacketC>(1, i);
1330 cblock.packet[2] = lhs2.template loadPacket<PacketC>(0, i + 1);
1331 cblock.packet[3] = lhs2.template loadPacket<PacketC>(1, i + 1);
1333 blockr.packet[0] = vec_mergeh(cblock.packet[0].v, cblock.packet[1].v);
1334 blockr.packet[1] = vec_mergeh(cblock.packet[2].v, cblock.packet[3].v);
1336 blocki.packet[0] = vec_mergel(cblock.packet[0].v, cblock.packet[1].v);
1337 blocki.packet[1] = vec_mergel(cblock.packet[2].v, cblock.packet[3].v);
1340 EIGEN_IF_CONSTEXPR (Conjugate) {
1341 blocki.packet[0] = -blocki.packet[0];
1342 blocki.packet[1] = -blocki.packet[1];
1345 storeBlock<double, Packet, 2>(blockAt + rir, blockr);
1346 storeBlock<double, Packet, 2>(blockAt + rii, blocki);
1348 rir += 2 * vectorSize;
1349 rii += 2 * vectorSize;
1353 EIGEN_STRONG_INLINE
void operator()(std::complex<double>* blockA,
const DataMapper& lhs, Index depth, Index rows,
1354 Index stride, Index offset) {
1355 const Index vectorSize = quad_traits<double>::vectorsize;
1356 const Index vectorDelta = vectorSize * ((PanelMode) ? stride : depth);
1357 Index rir = ((PanelMode) ? (vectorSize * offset) : 0), rii;
1358 double* blockAt =
reinterpret_cast<double*
>(blockA);
1361 for (; j + vectorSize <= rows; j += vectorSize) {
1362 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
1365 rii = rir + vectorDelta;
1367 dhs_ccopy(blockAt, lhs2, i, rir, rii, depth, vectorSize);
1369 for (; i < depth; i++) {
1370 PacketBlock<Packet, 1> blockr, blocki;
1371 PacketBlock<PacketC, 2> cblock;
1373 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i);
1374 cblock.packet[1] = lhs2.template loadPacket<PacketC>(1, i);
1376 blockr.packet[0] = vec_mergeh(cblock.packet[0].v, cblock.packet[1].v);
1377 blocki.packet[0] = vec_mergel(cblock.packet[0].v, cblock.packet[1].v);
1379 EIGEN_IF_CONSTEXPR (Conjugate) {
1380 blocki.packet[0] = -blocki.packet[0];
1383 pstore<double>(blockAt + rir, blockr.packet[0]);
1384 pstore<double>(blockAt + rii, blocki.packet[0]);
1390 rir += ((PanelMode) ? (vectorSize * (2 * stride - depth)) : vectorDelta);
1394 EIGEN_IF_CONSTEXPR (PanelMode) rir += (offset * (rows - j - vectorSize));
1395 rii = rir + (((PanelMode) ? stride : depth) * (rows - j));
1397 for (Index i = 0; i < depth; i++) {
1399 for (; k < rows; k++) {
1400 blockAt[rir] = lhs(k, i).real();
1402 EIGEN_IF_CONSTEXPR (Conjugate)
1403 blockAt[rii] = -lhs(k, i).imag();
1405 blockAt[rii] = lhs(k, i).imag();
1416template <
typename DataMapper,
typename Packet,
typename PacketC,
int StorageOrder,
bool Conjugate,
bool PanelMode>
1417struct dhs_cpack<double, DataMapper, Packet, PacketC, StorageOrder, Conjugate, PanelMode, false> {
1418 EIGEN_ALWAYS_INLINE
void dhs_ccopy(
double* blockBt,
const DataMapper& rhs2, Index& i, Index& rir, Index& rii,
1419 Index depth,
const Index vectorSize) {
1420 for (; i < depth; i++) {
1421 PacketBlock<PacketC, 4> cblock;
1422 PacketBlock<Packet, 2> blockr, blocki;
1424 bload<DataMapper, PacketC, 2, ColMajor, false, 4>(cblock, rhs2, i, 0);
1426 blockr.packet[0] = vec_mergeh(cblock.packet[0].v, cblock.packet[1].v);
1427 blockr.packet[1] = vec_mergeh(cblock.packet[2].v, cblock.packet[3].v);
1429 blocki.packet[0] = vec_mergel(cblock.packet[0].v, cblock.packet[1].v);
1430 blocki.packet[1] = vec_mergel(cblock.packet[2].v, cblock.packet[3].v);
1432 EIGEN_IF_CONSTEXPR (Conjugate) {
1433 blocki.packet[0] = -blocki.packet[0];
1434 blocki.packet[1] = -blocki.packet[1];
1437 storeBlock<double, Packet, 2>(blockBt + rir, blockr);
1438 storeBlock<double, Packet, 2>(blockBt + rii, blocki);
1440 rir += 2 * vectorSize;
1441 rii += 2 * vectorSize;
1445 EIGEN_STRONG_INLINE
void operator()(std::complex<double>* blockB,
const DataMapper& rhs, Index depth, Index cols,
1446 Index stride, Index offset) {
1447 const Index vectorSize = quad_traits<double>::vectorsize;
1448 const Index vectorDelta = 2 * vectorSize * ((PanelMode) ? stride : depth);
1449 Index rir = ((PanelMode) ? (2 * vectorSize * offset) : 0), rii;
1450 double* blockBt =
reinterpret_cast<double*
>(blockB);
1453 for (; j + 2 * vectorSize <= cols; j += 2 * vectorSize) {
1454 const DataMapper rhs2 = rhs.getSubMapper(0, j);
1457 rii = rir + vectorDelta;
1459 dhs_ccopy(blockBt, rhs2, i, rir, rii, depth, vectorSize);
1461 rir += ((PanelMode) ? (2 * vectorSize * (2 * stride - depth)) : vectorDelta);
1464 EIGEN_IF_CONSTEXPR (PanelMode) rir -= (offset * (2 * vectorSize - 1));
1466 for (; j < cols; j++) {
1467 const DataMapper rhs2 = rhs.getSubMapper(0, j);
1468 rii = rir + ((PanelMode) ? stride : depth);
1470 for (Index i = 0; i < depth; i++) {
1471 blockBt[rir] = rhs2(i, 0).real();
1473 EIGEN_IF_CONSTEXPR (Conjugate)
1474 blockBt[rii] = -rhs2(i, 0).imag();
1476 blockBt[rii] = rhs2(i, 0).imag();
1482 rir += ((PanelMode) ? (2 * stride - depth) : depth);
1492template <
typename Packet,
bool NegativeAccumulate,
int N>
1493EIGEN_ALWAYS_INLINE
void pger_common(PacketBlock<Packet, N>* acc,
const Packet& lhsV,
const Packet* rhsV) {
1494 EIGEN_IF_CONSTEXPR (NegativeAccumulate) {
1495 for (
int M = 0; M < N; M++) {
1496 acc->packet[M] = vec_nmsub(lhsV, rhsV[M], acc->packet[M]);
1499 for (
int M = 0; M < N; M++) {
1500 acc->packet[M] = vec_madd(lhsV, rhsV[M], acc->packet[M]);
1505template <
int N,
typename Scalar,
typename Packet,
bool NegativeAccumulate, Index remaining_rows = 0>
1506EIGEN_ALWAYS_INLINE
void pger(PacketBlock<Packet, N>* acc,
const Scalar* lhs,
const Packet* rhsV) {
1508 EIGEN_IF_CONSTEXPR (remaining_rows > 0) {
1509 lhsV = ploadu_partial<Packet>(lhs, remaining_rows);
1511 lhsV = ploadLhs<Packet>(lhs);
1514 pger_common<Packet, NegativeAccumulate, N>(acc, lhsV, rhsV);
1519template <
int N,
typename Packet,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
1520EIGEN_ALWAYS_INLINE
void pgerc_common(PacketBlock<Packet, N>* accReal, PacketBlock<Packet, N>* accImag,
1521 const Packet& lhsV, Packet& lhsVi,
const Packet* rhsV,
const Packet* rhsVi) {
1522 pger_common<Packet, false, N>(accReal, lhsV, rhsV);
1523 EIGEN_IF_CONSTEXPR (LhsIsReal) {
1524 pger_common<Packet, ConjugateRhs, N>(accImag, lhsV, rhsVi);
1525 EIGEN_UNUSED_VARIABLE(lhsVi);
1527 EIGEN_IF_CONSTEXPR (!RhsIsReal) {
1528 pger_common<Packet, ConjugateLhs == ConjugateRhs, N>(accReal, lhsVi, rhsVi);
1529 pger_common<Packet, ConjugateRhs, N>(accImag, lhsV, rhsVi);
1531 EIGEN_UNUSED_VARIABLE(rhsVi);
1533 pger_common<Packet, ConjugateLhs, N>(accImag, lhsVi, rhsV);
1537template <
int N,
typename Scalar,
typename Packet,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal,
1538 Index remaining_rows = 0>
1539EIGEN_ALWAYS_INLINE
void pgerc(PacketBlock<Packet, N>* accReal, PacketBlock<Packet, N>* accImag,
const Scalar* lhs_ptr,
1540 const Scalar* lhs_ptr_imag,
const Packet* rhsV,
const Packet* rhsVi) {
1542 EIGEN_IF_CONSTEXPR (remaining_rows > 0) {
1543 lhsV = ploadu_partial<Packet>(lhs_ptr, remaining_rows);
1545 lhsV = ploadLhs<Packet>(lhs_ptr);
1548 EIGEN_IF_CONSTEXPR (!LhsIsReal) {
1549 EIGEN_IF_CONSTEXPR (remaining_rows > 0) {
1550 lhsVi = ploadu_partial<Packet>(lhs_ptr_imag, remaining_rows);
1552 lhsVi = ploadLhs<Packet>(lhs_ptr_imag);
1555 EIGEN_UNUSED_VARIABLE(lhs_ptr_imag);
1558 pgerc_common<N, Packet, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>(accReal, accImag, lhsV, lhsVi, rhsV, rhsVi);
1561template <
typename Packet>
1562EIGEN_ALWAYS_INLINE Packet ploadLhs(
const __UNPACK_TYPE__(Packet) * lhs) {
1563 return ploadu<Packet>(lhs);
1567template <
typename Packet,
int N>
1568EIGEN_ALWAYS_INLINE
void bsetzero(PacketBlock<Packet, N>& acc) {
1569 for (
int M = 0; M < N; M++) {
1570 acc.packet[M] = pset1<Packet>((__UNPACK_TYPE__(Packet))0);
1574template <
typename Packet,
int N>
1575EIGEN_ALWAYS_INLINE
void bscalec_common(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ,
1576 const Packet& pAlpha) {
1577 for (
int M = 0; M < N; M++) {
1578 acc.packet[M] = vec_mul(accZ.packet[M], pAlpha);
1582template <
typename Packet,
int N>
1583EIGEN_ALWAYS_INLINE
void band(PacketBlock<Packet, N>& acc,
const Packet& pMask) {
1584 for (
int M = 0; M < N; M++) {
1585 acc.packet[M] = pand<Packet>(acc.packet[M], pMask);
1590template <
typename Packet,
int N,
bool mask>
1591EIGEN_ALWAYS_INLINE
void bscalec(PacketBlock<Packet, N>& aReal, PacketBlock<Packet, N>& aImag,
const Packet& bReal,
1592 const Packet& bImag, PacketBlock<Packet, N>& cReal, PacketBlock<Packet, N>& cImag,
1593 const Packet& pMask) {
1594 EIGEN_IF_CONSTEXPR (mask && (
sizeof(__UNPACK_TYPE__(Packet)) ==
sizeof(
float))) {
1595 band<Packet, N>(aReal, pMask);
1596 band<Packet, N>(aImag, pMask);
1598 EIGEN_UNUSED_VARIABLE(pMask);
1601 bscalec_common<Packet, N>(cReal, aReal, bReal);
1603 bscalec_common<Packet, N>(cImag, aImag, bReal);
1605 pger_common<Packet, true, N>(&cReal, bImag, aImag.packet);
1607 pger_common<Packet, false, N>(&cImag, bImag, aReal.packet);
1613template <
typename DataMapper,
typename Packet, const Index accCols,
int StorageOrder,
bool Complex,
int N,
bool full>
1614EIGEN_ALWAYS_INLINE
void bload(PacketBlock<Packet, N*(Complex ? 2 : 1)>& acc,
const DataMapper& res, Index row,
1616 EIGEN_IF_CONSTEXPR (StorageOrder ==
RowMajor) {
1617 for (
int M = 0; M < N; M++) {
1618 acc.packet[M] = res.template loadPacket<Packet>(row + M, col);
1620 EIGEN_IF_CONSTEXPR (Complex) {
1621 for (
int M = 0; M < N; M++) {
1622 acc.packet[M + N] = res.template loadPacket<Packet>(row + M, col + accCols);
1626 for (
int M = 0; M < N; M++) {
1627 acc.packet[M] = res.template loadPacket<Packet>(row, col + M);
1629 EIGEN_IF_CONSTEXPR (Complex && full) {
1630 for (
int M = 0; M < N; M++) {
1631 acc.packet[M + N] = res.template loadPacket<Packet>(row + accCols, col + M);
1637template <
typename DataMapper,
typename Packet,
int N>
1638EIGEN_ALWAYS_INLINE
void bstore(PacketBlock<Packet, N>& acc,
const DataMapper& res, Index row) {
1639 for (
int M = 0; M < N; M++) {
1640 res.template storePacket<Packet>(row, M, acc.packet[M]);
1645template <
typename DataMapper,
typename Packet, const Index accCols,
bool Complex, Index N,
bool full>
1646EIGEN_ALWAYS_INLINE
void bload_partial(PacketBlock<Packet, N*(Complex ? 2 : 1)>& acc,
const DataMapper& res, Index row,
1648 EIGEN_IF_CONSTEXPR (Complex && full) {
1649 for (Index M = 0; M < N; M++) {
1650 acc.packet[M] = res.template loadPacket<Packet>(row, M);
1651 acc.packet[M + N] = res.template loadPacketPartial<Packet>(row + accCols, M, elements);
1654 for (Index M = 0; M < N; M++) {
1655 acc.packet[M] = res.template loadPacketPartial<Packet>(row, M, elements);
1660template <
typename DataMapper,
typename Packet, Index N>
1661EIGEN_ALWAYS_INLINE
void bstore_partial(PacketBlock<Packet, N>& acc,
const DataMapper& res, Index row, Index elements) {
1662 for (Index M = 0; M < N; M++) {
1663 res.template storePacketPartial<Packet>(row, M, acc.packet[M], elements);
1668#define USE_P10_AND_PVIPR2_0 (EIGEN_COMP_LLVM || (__GNUC__ >= 11))
1670#define USE_P10_AND_PVIPR2_0 0
1673#if !USE_P10_AND_PVIPR2_0
1674const static Packet4i mask4[4] = {{0, 0, 0, 0}, {-1, 0, 0, 0}, {-1, -1, 0, 0}, {-1, -1, -1, 0}};
1677template <
typename Packet>
1678EIGEN_ALWAYS_INLINE Packet bmask(
const Index remaining_rows) {
1679#if USE_P10_AND_PVIPR2_0
1681 return Packet(vec_reve(vec_genwm((1 << remaining_rows) - 1)));
1683 return Packet(vec_genwm((1 << remaining_rows) - 1));
1686 return Packet(mask4[remaining_rows]);
1691EIGEN_ALWAYS_INLINE Packet2d bmask<Packet2d>(
const Index remaining_rows) {
1692#if USE_P10_AND_PVIPR2_0
1693 Packet2d mask2 = Packet2d(vec_gendm(remaining_rows));
1695 return preverse(mask2);
1700 Packet2l ret = {-remaining_rows, 0};
1701 return Packet2d(ret);
1705template <
typename Packet,
int N>
1706EIGEN_ALWAYS_INLINE
void bscale(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ,
const Packet& pAlpha) {
1707 for (
int M = 0; M < N; M++) {
1708 acc.packet[M] = pmadd<Packet>(pAlpha, accZ.packet[M], acc.packet[M]);
1713template <
typename Packet,
int N,
bool mask>
1714EIGEN_ALWAYS_INLINE
void bscale(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ,
const Packet& pAlpha,
1715 const Packet& pMask) {
1716 EIGEN_IF_CONSTEXPR (mask) {
1717 band<Packet, N>(accZ, pMask);
1719 EIGEN_UNUSED_VARIABLE(pMask);
1722 bscale<Packet, N>(acc, accZ, pAlpha);
1725template <
typename Packet,
int N,
bool real>
1726EIGEN_ALWAYS_INLINE
void pbroadcastN(
const __UNPACK_TYPE__(Packet) * ap0,
const __UNPACK_TYPE__(Packet) * ap1,
1727 const __UNPACK_TYPE__(Packet) * ap2, Packet& a0, Packet& a1, Packet& a2,
1729 a0 = pset1<Packet>(ap0[0]);
1730 EIGEN_IF_CONSTEXPR (N == 4) {
1731 a1 = pset1<Packet>(ap0[1]);
1732 a2 = pset1<Packet>(ap0[2]);
1733 a3 = pset1<Packet>(ap0[3]);
1734 EIGEN_UNUSED_VARIABLE(ap1);
1735 EIGEN_UNUSED_VARIABLE(ap2);
1737 EIGEN_IF_CONSTEXPR (N > 1) {
1738 a1 = pset1<Packet>(ap1[0]);
1740 EIGEN_UNUSED_VARIABLE(a1);
1741 EIGEN_UNUSED_VARIABLE(ap1);
1743 EIGEN_IF_CONSTEXPR (N > 2) {
1744 a2 = pset1<Packet>(ap2[0]);
1746 EIGEN_UNUSED_VARIABLE(a2);
1747 EIGEN_UNUSED_VARIABLE(ap2);
1753EIGEN_ALWAYS_INLINE
void pbroadcastN<Packet4f, 4, true>(
const float* ap0,
const float*,
const float*, Packet4f& a0,
1754 Packet4f& a1, Packet4f& a2, Packet4f& a3) {
1755 pbroadcast4<Packet4f>(ap0, a0, a1, a2, a3);
1759EIGEN_ALWAYS_INLINE
void pbroadcastN<Packet4f, 4, false>(
const float* ap0,
const float* ap1,
const float* ap2,
1760 Packet4f& a0, Packet4f& a1, Packet4f& a2, Packet4f& a3) {
1761 pbroadcastN<Packet4f, 4, true>(ap0, ap1, ap2, a0, a1, a2, a3);
1765EIGEN_ALWAYS_INLINE
void pbroadcastN<Packet2d, 4, false>(
const double* ap0,
const double*,
const double*, Packet2d& a0,
1766 Packet2d& a1, Packet2d& a2, Packet2d& a3) {
1767 a1 = pload<Packet2d>(ap0);
1768 a3 = pload<Packet2d>(ap0 + 2);
1769 a0 = vec_splat(a1, 0);
1770 a1 = vec_splat(a1, 1);
1771 a2 = vec_splat(a3, 0);
1772 a3 = vec_splat(a3, 1);
1776template <
typename Packet,
typename Packetc,
int N,
bool full>
1777EIGEN_ALWAYS_INLINE
void bcouple_common(PacketBlock<Packet, N>& taccReal, PacketBlock<Packet, N>& taccImag,
1778 PacketBlock<Packetc, N>& acc1, PacketBlock<Packetc, N>& acc2) {
1779 for (
int M = 0; M < N; M++) {
1780 acc1.packet[M].v = vec_mergeh(taccReal.packet[M], taccImag.packet[M]);
1783 EIGEN_IF_CONSTEXPR (full) {
1784 for (
int M = 0; M < N; M++) {
1785 acc2.packet[M].v = vec_mergel(taccReal.packet[M], taccImag.packet[M]);
1790template <
typename Packet,
typename Packetc,
int N,
bool full>
1791EIGEN_ALWAYS_INLINE
void bcouple(PacketBlock<Packet, N>& taccReal, PacketBlock<Packet, N>& taccImag,
1792 PacketBlock<Packetc, N * 2>& tRes, PacketBlock<Packetc, N>& acc1,
1793 PacketBlock<Packetc, N>& acc2) {
1794 bcouple_common<Packet, Packetc, N, full>(taccReal, taccImag, acc1, acc2);
1796 for (
int M = 0; M < N; M++) {
1797 acc1.packet[M] = padd<Packetc>(tRes.packet[M], acc1.packet[M]);
1800 EIGEN_IF_CONSTEXPR (full) {
1801 for (
int M = 0; M < N; M++) {
1802 acc2.packet[M] = padd<Packetc>(tRes.packet[M + N], acc2.packet[M]);
1811#define MICRO_UNROLL(func) func(0) func(1) func(2) func(3) func(4) func(5) func(6) func(7)
1813#define MICRO_NORMAL_ROWS accRows == quad_traits<Scalar>::rows || accRows == 1
1815#define MICRO_NEW_ROWS ((MICRO_NORMAL_ROWS) ? accRows : 1)
1817#define MICRO_RHS(ptr, N) rhs_##ptr##N
1819#define MICRO_ZERO_PEEL(peel) \
1820 EIGEN_IF_CONSTEXPR ((PEEL_ROW > peel) && (peel != 0)) { \
1821 bsetzero<Packet, accRows>(accZero##peel); \
1823 EIGEN_UNUSED_VARIABLE(accZero##peel); \
1826#define MICRO_ADD(ptr, N) \
1827 EIGEN_IF_CONSTEXPR (MICRO_NORMAL_ROWS) { \
1828 MICRO_RHS(ptr, 0) += (accRows * N); \
1830 MICRO_RHS(ptr, 0) += N; \
1831 MICRO_RHS(ptr, 1) += N; \
1832 EIGEN_IF_CONSTEXPR (accRows == 3) { \
1833 MICRO_RHS(ptr, 2) += N; \
1837#define MICRO_ADD_ROWS(N) MICRO_ADD(ptr, N)
1839#define MICRO_BROADCAST1(peel, ptr, rhsV, real) \
1840 EIGEN_IF_CONSTEXPR (MICRO_NORMAL_ROWS) { \
1841 pbroadcastN<Packet, accRows, real>(MICRO_RHS(ptr, 0) + (accRows * peel), MICRO_RHS(ptr, 0), MICRO_RHS(ptr, 0), \
1842 rhsV##peel[0], rhsV##peel[1], rhsV##peel[2], rhsV##peel[3]); \
1844 pbroadcastN<Packet, accRows, real>(MICRO_RHS(ptr, 0) + peel, MICRO_RHS(ptr, 1) + peel, MICRO_RHS(ptr, 2) + peel, \
1845 rhsV##peel[0], rhsV##peel[1], rhsV##peel[2], rhsV##peel[3]); \
1848#define MICRO_BROADCAST(peel) MICRO_BROADCAST1(peel, ptr, rhsV, true)
1850#define MICRO_BROADCAST_EXTRA1(ptr, rhsV, real) \
1851 pbroadcastN<Packet, accRows, real>(MICRO_RHS(ptr, 0), MICRO_RHS(ptr, 1), MICRO_RHS(ptr, 2), rhsV[0], rhsV[1], \
1854#define MICRO_BROADCAST_EXTRA \
1856 MICRO_BROADCAST_EXTRA1(ptr, rhsV, true) \
1859#define MICRO_SRC2(ptr, N, M) \
1860 EIGEN_IF_CONSTEXPR (MICRO_NORMAL_ROWS) { \
1861 EIGEN_UNUSED_VARIABLE(strideB); \
1862 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr, 1)); \
1863 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr, 2)); \
1865 MICRO_RHS(ptr, 1) = rhs_base + N + M; \
1866 EIGEN_IF_CONSTEXPR (accRows == 3) { \
1867 MICRO_RHS(ptr, 2) = rhs_base + N * 2 + M; \
1869 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr, 2)); \
1873#define MICRO_SRC2_PTR MICRO_SRC2(ptr, strideB, 0)
1875#define MICRO_ZERO_PEEL_ROW MICRO_UNROLL(MICRO_ZERO_PEEL)
1877#define MICRO_WORK_PEEL(peel) \
1878 EIGEN_IF_CONSTEXPR (PEEL_ROW > peel) { \
1879 MICRO_BROADCAST(peel) \
1880 pger<accRows, Scalar, Packet, false>(&accZero##peel, lhs_ptr + (remaining_rows * peel), rhsV##peel); \
1882 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
1885#define MICRO_WORK_PEEL_ROW \
1886 Packet rhsV0[4], rhsV1[4], rhsV2[4], rhsV3[4], rhsV4[4], rhsV5[4], rhsV6[4], rhsV7[4]; \
1887 MICRO_UNROLL(MICRO_WORK_PEEL) \
1888 lhs_ptr += (remaining_rows * PEEL_ROW); \
1889 MICRO_ADD_ROWS(PEEL_ROW)
1891#define MICRO_ADD_PEEL(peel, sum) \
1892 EIGEN_IF_CONSTEXPR (PEEL_ROW > peel) { \
1893 for (Index i = 0; i < accRows; i++) { \
1894 accZero##sum.packet[i] += accZero##peel.packet[i]; \
1898#define MICRO_ADD_PEEL_ROW \
1899 MICRO_ADD_PEEL(4, 0) \
1900 MICRO_ADD_PEEL(5, 1) \
1901 MICRO_ADD_PEEL(6, 2) MICRO_ADD_PEEL(7, 3) MICRO_ADD_PEEL(2, 0) MICRO_ADD_PEEL(3, 1) MICRO_ADD_PEEL(1, 0)
1903#define MICRO_PREFETCHN1(ptr, N) \
1904 EIGEN_POWER_PREFETCH(MICRO_RHS(ptr, 0)); \
1905 EIGEN_IF_CONSTEXPR (N == 2 || N == 3) { \
1906 EIGEN_POWER_PREFETCH(MICRO_RHS(ptr, 1)); \
1907 EIGEN_IF_CONSTEXPR (N == 3) { \
1908 EIGEN_POWER_PREFETCH(MICRO_RHS(ptr, 2)); \
1912#define MICRO_PREFETCHN(N) MICRO_PREFETCHN1(ptr, N)
1914#define MICRO_COMPLEX_PREFETCHN(N) \
1915 MICRO_PREFETCHN1(ptr_real, N); \
1916 EIGEN_IF_CONSTEXPR (!RhsIsReal) { \
1917 MICRO_PREFETCHN1(ptr_imag, N); \
1920template <
typename Scalar,
typename Packet, const Index accRows, const Index remaining_rows, const Index load_rows = 0>
1921EIGEN_ALWAYS_INLINE
void MICRO_EXTRA_ROW(
const Scalar*& lhs_ptr,
const Scalar*& rhs_ptr0,
const Scalar*& rhs_ptr1,
1922 const Scalar*& rhs_ptr2, PacketBlock<Packet, accRows>& accZero) {
1923 MICRO_BROADCAST_EXTRA
1924 pger<accRows, Scalar, Packet, false, load_rows>(&accZero, lhs_ptr, rhsV);
1925 lhs_ptr += remaining_rows;
1928template <
typename Scalar,
typename Packet,
typename DataMapper,
const Index accRows,
const Index accCols,
1929 const Index remaining_rows>
1930EIGEN_ALWAYS_INLINE
void gemm_unrolled_row_iteration(
const DataMapper& res,
const Scalar* lhs_base,
1931 const Scalar* rhs_base, Index depth, Index strideA, Index offsetA,
1932 Index strideB, Index row, Index rows,
const Packet& pAlpha,
1933 const Packet& pMask) {
1934 const Scalar *rhs_ptr0 = rhs_base, *rhs_ptr1 =
nullptr, *rhs_ptr2 =
nullptr;
1935 const Scalar* lhs_ptr = lhs_base + row * strideA + remaining_rows * offsetA;
1936 PacketBlock<Packet, accRows> accZero0, accZero1, accZero2, accZero3, accZero4, accZero5, accZero6, accZero7, acc;
1939 bsetzero<Packet, accRows>(accZero0);
1941 const Index peel_depth = depth - (accCols - remaining_rows);
1943 if (peel_depth >= PEEL_ROW) {
1946 MICRO_PREFETCHN(accRows)
1947 EIGEN_POWER_PREFETCH(lhs_ptr);
1949 }
while ((k += PEEL_ROW) + PEEL_ROW <= peel_depth);
1952 for (; k < peel_depth; k++) {
1953 MICRO_EXTRA_ROW<Scalar, Packet, accRows, remaining_rows>(lhs_ptr, rhs_ptr0, rhs_ptr1, rhs_ptr2, accZero0);
1955 for (; k < depth; k++) {
1956 MICRO_EXTRA_ROW<Scalar, Packet, accRows, remaining_rows, remaining_rows>(lhs_ptr, rhs_ptr0, rhs_ptr1, rhs_ptr2,
1960 EIGEN_UNUSED_VARIABLE(rows);
1961 EIGEN_UNUSED_VARIABLE(pMask);
1962 bload_partial<DataMapper, Packet, 0, false, accRows>(acc, res, row, remaining_rows);
1963 bscale<Packet, accRows>(acc, accZero0, pAlpha);
1964 bstore_partial<DataMapper, Packet, accRows>(acc, res, row, remaining_rows);
1967#define MICRO_EXTRA(MICRO_EXTRA_UNROLL, value, is_col) \
1970 MICRO_EXTRA_UNROLL(1) \
1973 EIGEN_IF_CONSTEXPR (is_col || (sizeof(Scalar) == sizeof(float))) { \
1974 MICRO_EXTRA_UNROLL(2) \
1978 EIGEN_IF_CONSTEXPR (is_col || (sizeof(Scalar) == sizeof(float))) { \
1979 MICRO_EXTRA_UNROLL(3) \
1984#define MICRO_EXTRA_ROWS(N) \
1985 gemm_unrolled_row_iteration<Scalar, Packet, DataMapper, accRows, accCols, N>( \
1986 res, lhs_base, rhs_base, depth, strideA, offsetA, strideB, row, rows, pAlpha, pMask);
1988template <
typename Scalar,
typename Packet,
typename DataMapper, const Index accRows, const Index accCols>
1989EIGEN_ALWAYS_INLINE
void gemm_extra_row(
const DataMapper& res,
const Scalar* lhs_base,
const Scalar* rhs_base,
1990 Index depth, Index strideA, Index offsetA, Index strideB, Index row, Index rows,
1991 Index remaining_rows,
const Packet& pAlpha,
const Packet& pMask) {
1992 MICRO_EXTRA(MICRO_EXTRA_ROWS, remaining_rows,
false)
1995#define MICRO_UNROLL_WORK(func, func2, peel) \
1996 MICRO_UNROLL(func2); \
1997 func(0, peel) func(1, peel) func(2, peel) func(3, peel) func(4, peel) func(5, peel) func(6, peel) func(7, peel)
1999#define MICRO_WORK_ONE(iter, peel) \
2000 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2001 pger_common<Packet, false, accRows>(&accZero##iter, lhsV##iter, rhsV##peel); \
2004#define MICRO_TYPE_PEEL4(func, func2, peel) \
2005 EIGEN_IF_CONSTEXPR (PEEL > peel) { \
2006 Packet lhsV0, lhsV1, lhsV2, lhsV3, lhsV4, lhsV5, lhsV6, lhsV7; \
2007 MICRO_BROADCAST(peel) \
2008 MICRO_UNROLL_WORK(func, func2, peel) \
2010 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
2013#define MICRO_UNROLL_TYPE_PEEL(M, func, func1, func2) \
2014 Packet rhsV0[M], rhsV1[M], rhsV2[M], rhsV3[M], rhsV4[M], rhsV5[M], rhsV6[M], rhsV7[M]; \
2015 func(func1, func2, 0) func(func1, func2, 1) func(func1, func2, 2) func(func1, func2, 3) func(func1, func2, 4) \
2016 func(func1, func2, 5) func(func1, func2, 6) func(func1, func2, 7)
2018#define MICRO_UNROLL_TYPE_ONE(M, func, func1, func2) \
2020 func(func1, func2, 0)
2022#define MICRO_UNROLL_TYPE(MICRO_TYPE, size) \
2023 MICRO_TYPE(4, MICRO_TYPE_PEEL4, MICRO_WORK_ONE, MICRO_LOAD_ONE) \
2024 MICRO_ADD_ROWS(size)
2026#define MICRO_UNROLL_TYPE_PARTIAL(MICRO_TYPE, size) \
2027 MICRO_TYPE(4, MICRO_TYPE_PEEL4, MICRO_WORK_ONE, MICRO_LOAD_PARTIAL_ONE) \
2028 MICRO_ADD_ROWS(size)
2030#define MICRO_ONE_PEEL4 MICRO_UNROLL_TYPE(MICRO_UNROLL_TYPE_PEEL, PEEL)
2032#define MICRO_ONE4 MICRO_UNROLL_TYPE(MICRO_UNROLL_TYPE_ONE, 1)
2034#define MICRO_ONE4_PARTIAL MICRO_UNROLL_TYPE_PARTIAL(MICRO_UNROLL_TYPE_ONE, 1)
2036#define MICRO_DST_PTR_ONE(iter) \
2037 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2038 bsetzero<Packet, accRows>(accZero##iter); \
2040 EIGEN_UNUSED_VARIABLE(accZero##iter); \
2043#define MICRO_DST_PTR MICRO_UNROLL(MICRO_DST_PTR_ONE)
2045#define MICRO_SRC_PTR MICRO_UNROLL(MICRO_SRC_PTR_ONE)
2047#define MICRO_PREFETCH MICRO_UNROLL(MICRO_PREFETCH_ONE)
2049#define MICRO_STORE_ONE(iter) \
2050 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2051 EIGEN_IF_CONSTEXPR (MICRO_NORMAL_PARTIAL(iter)) { \
2052 bload<DataMapper, Packet, 0, ColMajor, false, accRows>(acc, res, row + iter * accCols, 0); \
2053 bscale<Packet, accRows>(acc, accZero##iter, pAlpha); \
2054 bstore<DataMapper, Packet, accRows>(acc, res, row + iter * accCols); \
2056 bload_partial<DataMapper, Packet, 0, false, accRows>(acc, res, row + iter * accCols, accCols2); \
2057 bscale<Packet, accRows>(acc, accZero##iter, pAlpha); \
2058 bstore_partial<DataMapper, Packet, accRows>(acc, res, row + iter * accCols, accCols2); \
2062#define MICRO_STORE MICRO_UNROLL(MICRO_STORE_ONE)
2064template <
int unroll_factor,
typename Scalar,
typename Packet,
typename DataMapper,
const Index accRows,
2065 const Index accCols,
bool full>
2066EIGEN_ALWAYS_INLINE
void gemm_unrolled_iteration(
const DataMapper& res,
const Scalar* lhs_base,
const Scalar* rhs_base,
2067 Index depth, Index strideA, Index offsetA, Index strideB, Index& row,
2068 const Packet& pAlpha, Index accCols2) {
2069 const Scalar *rhs_ptr0 = rhs_base, *rhs_ptr1 =
nullptr, *rhs_ptr2 =
nullptr;
2070 const Scalar *lhs_ptr0 =
nullptr, *lhs_ptr1 =
nullptr, *lhs_ptr2 =
nullptr, *lhs_ptr3 =
nullptr, *lhs_ptr4 =
nullptr,
2071 *lhs_ptr5 =
nullptr, *lhs_ptr6 =
nullptr, *lhs_ptr7 =
nullptr;
2072 PacketBlock<Packet, accRows> accZero0, accZero1, accZero2, accZero3, accZero4, accZero5, accZero6, accZero7;
2073 PacketBlock<Packet, accRows> acc;
2079 const Index peel_depth = full ? depth : (depth - (accCols - accCols2));
2081 for (; k + PEEL <= peel_depth; k += PEEL) {
2082 MICRO_PREFETCHN(accRows)
2086 for (; k < peel_depth; k++) {
2089 EIGEN_IF_CONSTEXPR (!full) {
2090 for (; k < depth; k++) {
2099#define MICRO_UNROLL_ITER2(N, M) \
2100 gemm_unrolled_iteration<N + ((M) ? 1 : 0), Scalar, Packet, DataMapper, accRows, accCols, !M>( \
2101 res3, lhs_base, rhs_base, depth, strideA, offsetA, strideB, row, pAlpha, M ? remaining_rows : accCols); \
2104template <
typename Scalar,
typename Packet,
typename DataMapper, const Index accRows, const Index accCols>
2105EIGEN_ALWAYS_INLINE
void gemm_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index depth,
2106 Index strideA, Index offsetA, Index strideB, Index offsetB, Index col, Index rows,
2107 Index remaining_rows,
const Packet& pAlpha,
const Packet& pMask) {
2108 const DataMapper res3 = res.getSubMapper(0, col);
2110 const Scalar* rhs_base = blockB + col * strideB + MICRO_NEW_ROWS * offsetB;
2111 const Scalar* lhs_base = blockA + accCols * offsetA;
2115 while (row + MAX_UNROLL * accCols <= rows) {
2116 MICRO_UNROLL_ITER2(MAX_UNROLL, 0);
2118 switch ((rows - row) / accCols) {
2121 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 7)
2126 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 6)
2131 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 5)
2136 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 4)
2141 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 3)
2146 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 2)
2151 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 1)
2159 if (remaining_rows > 0) {
2160 gemm_extra_row<Scalar, Packet, DataMapper, accRows, accCols>(res3, blockA, rhs_base, depth, strideA, offsetA,
2161 strideB, row, rows, remaining_rows, pAlpha, pMask);
2165#define MICRO_EXTRA_COLS(N) \
2166 gemm_cols<Scalar, Packet, DataMapper, N, accCols>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB, \
2167 col, rows, remaining_rows, pAlpha, pMask);
2169template <
typename Scalar,
typename Packet,
typename DataMapper, const Index accCols>
2170EIGEN_ALWAYS_INLINE
void gemm_extra_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index depth,
2171 Index strideA, Index offsetA, Index strideB, Index offsetB, Index col,
2172 Index rows, Index cols, Index remaining_rows,
const Packet& pAlpha,
2173 const Packet& pMask) {
2174 MICRO_EXTRA(MICRO_EXTRA_COLS, cols - col,
true)
2180template <
typename Scalar,
typename Packet,
typename RhsPacket,
typename DataMapper,
const Index accRows,
2181 const Index accCols>
2182EIGEN_STRONG_INLINE
void gemm(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index rows,
2183 Index depth, Index cols, Scalar alpha, Index strideA, Index strideB, Index offsetA,
2185 const Index remaining_rows = rows % accCols;
2187 if (strideA == -1) strideA = depth;
2188 if (strideB == -1) strideB = depth;
2190 const Packet pAlpha = pset1<Packet>(alpha);
2191 const Packet pMask = bmask<Packet>(remaining_rows);
2194 for (; col + accRows <= cols; col += accRows) {
2195 gemm_cols<Scalar, Packet, DataMapper, accRows, accCols>(res, blockA, blockB, depth, strideA, offsetA, strideB,
2196 offsetB, col, rows, remaining_rows, pAlpha, pMask);
2200 gemm_extra_cols<Scalar, Packet, DataMapper, accCols>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB,
2201 col, rows, cols, remaining_rows, pAlpha, pMask);
2205#define accColsC (accCols / 2)
2206#define advanceRows ((LhsIsReal) ? 1 : 2)
2207#define advanceCols ((RhsIsReal) ? 1 : 2)
2210#define PEEL_COMPLEX 3
2211#define PEEL_COMPLEX_ROW 3
2213#define MICRO_COMPLEX_UNROLL(func) func(0) func(1) func(2) func(3)
2215#define MICRO_COMPLEX_ZERO_PEEL(peel) \
2216 EIGEN_IF_CONSTEXPR ((PEEL_COMPLEX_ROW > peel) && (peel != 0)) { \
2217 bsetzero<Packet, accRows>(accReal##peel); \
2218 bsetzero<Packet, accRows>(accImag##peel); \
2220 EIGEN_UNUSED_VARIABLE(accReal##peel); \
2221 EIGEN_UNUSED_VARIABLE(accImag##peel); \
2224#define MICRO_COMPLEX_ADD_ROWS(N, used) \
2225 MICRO_ADD(ptr_real, N) \
2226 EIGEN_IF_CONSTEXPR (!RhsIsReal) { \
2227 MICRO_ADD(ptr_imag, N) \
2228 } else if (used) { \
2229 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr_imag, 0)); \
2230 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr_imag, 1)); \
2231 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr_imag, 2)); \
2234#define MICRO_COMPLEX_BROADCAST(peel) \
2235 MICRO_BROADCAST1(peel, ptr_real, rhsV, false) \
2236 EIGEN_IF_CONSTEXPR (!RhsIsReal) { \
2237 MICRO_BROADCAST1(peel, ptr_imag, rhsVi, false) \
2239 EIGEN_UNUSED_VARIABLE(rhsVi##peel); \
2242#define MICRO_COMPLEX_BROADCAST_EXTRA \
2243 Packet rhsV[4], rhsVi[4]; \
2244 MICRO_BROADCAST_EXTRA1(ptr_real, rhsV, false) \
2245 EIGEN_IF_CONSTEXPR (!RhsIsReal) { \
2246 MICRO_BROADCAST_EXTRA1(ptr_imag, rhsVi, false) \
2248 EIGEN_UNUSED_VARIABLE(rhsVi); \
2250 MICRO_COMPLEX_ADD_ROWS(1, true)
2252#define MICRO_COMPLEX_SRC2_PTR \
2253 MICRO_SRC2(ptr_real, strideB* advanceCols, 0) \
2254 EIGEN_IF_CONSTEXPR (!RhsIsReal) { \
2255 MICRO_RHS(ptr_imag, 0) = rhs_base + MICRO_NEW_ROWS * strideB; \
2256 MICRO_SRC2(ptr_imag, strideB* advanceCols, strideB) \
2258 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr_imag, 0)); \
2259 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr_imag, 1)); \
2260 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr_imag, 2)); \
2263#define MICRO_COMPLEX_ZERO_PEEL_ROW MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_ZERO_PEEL)
2265#define MICRO_COMPLEX_WORK_PEEL(peel) \
2266 EIGEN_IF_CONSTEXPR (PEEL_COMPLEX_ROW > peel) { \
2267 MICRO_COMPLEX_BROADCAST(peel) \
2268 pgerc<accRows, Scalar, Packet, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>( \
2269 &accReal##peel, &accImag##peel, lhs_ptr_real + (remaining_rows * peel), \
2270 lhs_ptr_imag + (remaining_rows * peel), rhsV##peel, rhsVi##peel); \
2272 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
2273 EIGEN_UNUSED_VARIABLE(rhsVi##peel); \
2276#define MICRO_COMPLEX_ADD_COLS(size) \
2277 lhs_ptr_real += (remaining_rows * size); \
2278 EIGEN_IF_CONSTEXPR (!LhsIsReal) \
2279 lhs_ptr_imag += (remaining_rows * size); \
2281 EIGEN_UNUSED_VARIABLE(lhs_ptr_imag);
2283#define MICRO_COMPLEX_WORK_PEEL_ROW \
2284 Packet rhsV0[4], rhsV1[4], rhsV2[4], rhsV3[4]; \
2285 Packet rhsVi0[4], rhsVi1[4], rhsVi2[4], rhsVi3[4]; \
2286 MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_WORK_PEEL) \
2287 MICRO_COMPLEX_ADD_COLS(PEEL_COMPLEX_ROW) \
2288 MICRO_COMPLEX_ADD_ROWS(PEEL_COMPLEX_ROW, false)
2290#define MICRO_COMPLEX_ADD_PEEL(peel, sum) \
2291 EIGEN_IF_CONSTEXPR (PEEL_COMPLEX_ROW > peel) { \
2292 for (Index i = 0; i < accRows; i++) { \
2293 accReal##sum.packet[i] += accReal##peel.packet[i]; \
2294 accImag##sum.packet[i] += accImag##peel.packet[i]; \
2298#define MICRO_COMPLEX_ADD_PEEL_ROW \
2299 MICRO_COMPLEX_ADD_PEEL(2, 0) MICRO_COMPLEX_ADD_PEEL(3, 1) MICRO_COMPLEX_ADD_PEEL(1, 0)
2301template <
typename Scalar,
typename Packet,
const Index accRows,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
2302 bool RhsIsReal,
const Index remaining_rows,
const Index load_rows = 0>
2303EIGEN_ALWAYS_INLINE
void MICRO_COMPLEX_EXTRA_ROW(
const Scalar*& lhs_ptr_real,
const Scalar*& lhs_ptr_imag,
2304 const Scalar*& rhs_ptr_real0,
const Scalar*& rhs_ptr_real1,
2305 const Scalar*& rhs_ptr_real2,
const Scalar*& rhs_ptr_imag0,
2306 const Scalar*& rhs_ptr_imag1,
const Scalar*& rhs_ptr_imag2,
2307 PacketBlock<Packet, accRows>& accReal,
2308 PacketBlock<Packet, accRows>& accImag) {
2309 MICRO_COMPLEX_BROADCAST_EXTRA
2310 pgerc<accRows, Scalar, Packet, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal, load_rows>(
2311 &accReal, &accImag, lhs_ptr_real, lhs_ptr_imag, rhsV, rhsVi);
2312 MICRO_COMPLEX_ADD_COLS(1)
2315template <
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
const Index accRows,
2316 const Index accCols,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal,
2317 const Index remaining_rows>
2318EIGEN_ALWAYS_INLINE
void gemm_unrolled_complex_row_iteration(
const DataMapper& res,
const Scalar* lhs_base,
2319 const Scalar* rhs_base, Index depth, Index strideA,
2320 Index offsetA, Index strideB, Index row, Index rows,
2321 const Packet& pAlphaReal,
const Packet& pAlphaImag,
2322 const Packet& pMask) {
2323 const Scalar *rhs_ptr_real0 = rhs_base, *rhs_ptr_real1 =
nullptr, *rhs_ptr_real2 =
nullptr;
2324 const Scalar *rhs_ptr_imag0 =
nullptr, *rhs_ptr_imag1 =
nullptr, *rhs_ptr_imag2 =
nullptr;
2325 const Scalar* lhs_ptr_real = lhs_base + advanceRows * row * strideA + remaining_rows * offsetA;
2326 const Scalar* lhs_ptr_imag =
nullptr;
2327 EIGEN_IF_CONSTEXPR (!LhsIsReal)
2328 lhs_ptr_imag = lhs_ptr_real + remaining_rows * strideA;
2330 EIGEN_UNUSED_VARIABLE(lhs_ptr_imag);
2331 PacketBlock<Packet, accRows> accReal0, accImag0, accReal1, accImag1, accReal2, accImag2, accReal3, accImag3;
2332 PacketBlock<Packet, accRows> taccReal, taccImag;
2333 PacketBlock<Packetc, accRows> acc0, acc1;
2334 PacketBlock<Packetc, accRows * 2> tRes;
2336 MICRO_COMPLEX_SRC2_PTR
2338 bsetzero<Packet, accRows>(accReal0);
2339 bsetzero<Packet, accRows>(accImag0);
2341 const Index peel_depth = depth - (accCols - remaining_rows);
2343 if (peel_depth >= PEEL_COMPLEX_ROW) {
2344 MICRO_COMPLEX_ZERO_PEEL_ROW
2346 MICRO_COMPLEX_PREFETCHN(accRows)
2347 EIGEN_POWER_PREFETCH(lhs_ptr_real);
2348 EIGEN_IF_CONSTEXPR (!LhsIsReal) {
2349 EIGEN_POWER_PREFETCH(lhs_ptr_imag);
2351 MICRO_COMPLEX_WORK_PEEL_ROW
2352 }
while ((k += PEEL_COMPLEX_ROW) + PEEL_COMPLEX_ROW <= peel_depth);
2353 MICRO_COMPLEX_ADD_PEEL_ROW
2355 for (; k < peel_depth; k++) {
2356 MICRO_COMPLEX_EXTRA_ROW<Scalar, Packet, accRows, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal, remaining_rows>(
2357 lhs_ptr_real, lhs_ptr_imag, rhs_ptr_real0, rhs_ptr_real1, rhs_ptr_real2, rhs_ptr_imag0, rhs_ptr_imag1,
2358 rhs_ptr_imag2, accReal0, accImag0);
2360 for (; k < depth; k++) {
2361 MICRO_COMPLEX_EXTRA_ROW<Scalar, Packet, accRows, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal, remaining_rows,
2362 remaining_rows>(lhs_ptr_real, lhs_ptr_imag, rhs_ptr_real0, rhs_ptr_real1, rhs_ptr_real2,
2363 rhs_ptr_imag0, rhs_ptr_imag1, rhs_ptr_imag2, accReal0, accImag0);
2366 EIGEN_UNUSED_VARIABLE(rows);
2367 constexpr bool full = (remaining_rows > accColsC);
2368 constexpr bool odd = (
sizeof(Scalar) ==
sizeof(float)) && (remaining_rows & 1);
2369 EIGEN_IF_CONSTEXPR (odd) {
2370 bload_partial<DataMapper, Packetc, accColsC, true, accRows, full>(tRes, res, row, 1);
2372 bload<DataMapper, Packetc, accColsC, ColMajor, true, accRows, full>(tRes, res, row, 0);
2374 bscalec<Packet, accRows, true>(accReal0, accImag0, pAlphaReal, pAlphaImag, taccReal, taccImag, pMask);
2375 bcouple<Packet, Packetc, accRows, full>(taccReal, taccImag, tRes, acc0, acc1);
2376 EIGEN_IF_CONSTEXPR (odd && !full) {
2377 bstore_partial<DataMapper, Packetc, accRows>(acc0, res, row + 0, 1);
2379 bstore<DataMapper, Packetc, accRows>(acc0, res, row + 0);
2381 EIGEN_IF_CONSTEXPR (full) {
2382 EIGEN_IF_CONSTEXPR (odd) {
2383 bstore_partial<DataMapper, Packetc, accRows>(acc1, res, row + accColsC, 1);
2385 bstore<DataMapper, Packetc, accRows>(acc1, res, row + accColsC);
2390#define MICRO_COMPLEX_EXTRA_ROWS(N) \
2391 gemm_unrolled_complex_row_iteration<Scalar, Packet, Packetc, DataMapper, accRows, accCols, ConjugateLhs, \
2392 ConjugateRhs, LhsIsReal, RhsIsReal, N>( \
2393 res, lhs_base, rhs_base, depth, strideA, offsetA, strideB, row, rows, pAlphaReal, pAlphaImag, pMask);
2395template <
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
const Index accRows,
2396 const Index accCols,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
2397EIGEN_ALWAYS_INLINE
void gemm_complex_extra_row(
const DataMapper& res,
const Scalar* lhs_base,
const Scalar* rhs_base,
2398 Index depth, Index strideA, Index offsetA, Index strideB, Index row,
2399 Index rows, Index remaining_rows,
const Packet& pAlphaReal,
2400 const Packet& pAlphaImag,
const Packet& pMask) {
2401 MICRO_EXTRA(MICRO_COMPLEX_EXTRA_ROWS, remaining_rows,
false)
2404#define MICRO_COMPLEX_UNROLL_WORK(func, func2, peel) \
2405 MICRO_COMPLEX_UNROLL(func2); \
2406 func(0, peel) func(1, peel) func(2, peel) func(3, peel)
2408#define MICRO_COMPLEX_WORK_ONE4(iter, peel) \
2409 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2410 pgerc_common<accRows, Packet, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>( \
2411 &accReal##iter, &accImag##iter, lhsV##iter, lhsVi##iter, rhsV##peel, rhsVi##peel); \
2414#define MICRO_COMPLEX_TYPE_PEEL4(func, func2, peel) \
2415 EIGEN_IF_CONSTEXPR (PEEL_COMPLEX > peel) { \
2416 Packet lhsV0, lhsV1, lhsV2, lhsV3; \
2417 Packet lhsVi0, lhsVi1, lhsVi2, lhsVi3; \
2418 MICRO_COMPLEX_BROADCAST(peel) \
2419 MICRO_COMPLEX_UNROLL_WORK(func, func2, peel) \
2421 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
2422 EIGEN_UNUSED_VARIABLE(rhsVi##peel); \
2425#define MICRO_COMPLEX_UNROLL_TYPE_PEEL(M, func, func1, func2) \
2426 Packet rhsV0[M], rhsV1[M], rhsV2[M], rhsV3[M]; \
2427 Packet rhsVi0[M], rhsVi1[M], rhsVi2[M], rhsVi3[M]; \
2428 func(func1, func2, 0) func(func1, func2, 1) func(func1, func2, 2) func(func1, func2, 3)
2430#define MICRO_COMPLEX_UNROLL_TYPE_ONE(M, func, func1, func2) \
2431 Packet rhsV0[M], rhsVi0[M]; \
2432 func(func1, func2, 0)
2434#define MICRO_COMPLEX_UNROLL_TYPE(MICRO_COMPLEX_TYPE, size) \
2435 MICRO_COMPLEX_TYPE(4, MICRO_COMPLEX_TYPE_PEEL4, MICRO_COMPLEX_WORK_ONE4, MICRO_COMPLEX_LOAD_ONE) \
2436 MICRO_COMPLEX_ADD_ROWS(size, false)
2438#define MICRO_COMPLEX_UNROLL_TYPE_PARTIAL(MICRO_COMPLEX_TYPE, size) \
2439 MICRO_COMPLEX_TYPE(4, MICRO_COMPLEX_TYPE_PEEL4, MICRO_COMPLEX_WORK_ONE4, MICRO_COMPLEX_LOAD_PARTIAL_ONE) \
2440 MICRO_COMPLEX_ADD_ROWS(size, false)
2442#define MICRO_COMPLEX_ONE_PEEL4 MICRO_COMPLEX_UNROLL_TYPE(MICRO_COMPLEX_UNROLL_TYPE_PEEL, PEEL_COMPLEX)
2444#define MICRO_COMPLEX_ONE4 MICRO_COMPLEX_UNROLL_TYPE(MICRO_COMPLEX_UNROLL_TYPE_ONE, 1)
2446#define MICRO_COMPLEX_ONE4_PARTIAL MICRO_COMPLEX_UNROLL_TYPE_PARTIAL(MICRO_COMPLEX_UNROLL_TYPE_ONE, 1)
2448#define MICRO_COMPLEX_DST_PTR_ONE(iter) \
2449 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2450 bsetzero<Packet, accRows>(accReal##iter); \
2451 bsetzero<Packet, accRows>(accImag##iter); \
2453 EIGEN_UNUSED_VARIABLE(accReal##iter); \
2454 EIGEN_UNUSED_VARIABLE(accImag##iter); \
2457#define MICRO_COMPLEX_DST_PTR MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_DST_PTR_ONE)
2459#define MICRO_COMPLEX_SRC_PTR MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_SRC_PTR_ONE)
2461#define MICRO_COMPLEX_PREFETCH MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_PREFETCH_ONE)
2463#define MICRO_COMPLEX_STORE_ONE(iter) \
2464 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2465 constexpr bool full = ((MICRO_NORMAL(iter)) || (accCols2 > accColsC)); \
2466 constexpr bool odd = !(MICRO_NORMAL(iter)) && (sizeof(Scalar) == sizeof(float)) && (accCols2 & 1); \
2467 EIGEN_IF_CONSTEXPR (odd) { \
2468 bload_partial<DataMapper, Packetc, accColsC, true, accRows, full>(tRes, res, row + iter * accCols, 1); \
2470 bload<DataMapper, Packetc, accColsC, ColMajor, true, accRows, full>(tRes, res, row + iter * accCols, 0); \
2472 bscalec<Packet, accRows, !(MICRO_NORMAL(iter))>(accReal##iter, accImag##iter, pAlphaReal, pAlphaImag, taccReal, \
2474 bcouple<Packet, Packetc, accRows, full>(taccReal, taccImag, tRes, acc0, acc1); \
2475 EIGEN_IF_CONSTEXPR (odd && !full) { \
2476 bstore_partial<DataMapper, Packetc, accRows>(acc0, res, row + iter * accCols + 0, 1); \
2478 bstore<DataMapper, Packetc, accRows>(acc0, res, row + iter * accCols + 0); \
2480 EIGEN_IF_CONSTEXPR (full) { \
2481 EIGEN_IF_CONSTEXPR (odd) { \
2482 bstore_partial<DataMapper, Packetc, accRows>(acc1, res, row + iter * accCols + accColsC, 1); \
2484 bstore<DataMapper, Packetc, accRows>(acc1, res, row + iter * accCols + accColsC); \
2489#define MICRO_COMPLEX_STORE MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_STORE_ONE)
2491template <
int unroll_factor,
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
2492 const Index accRows,
const Index accCols,
const Index accCols2,
bool ConjugateLhs,
bool ConjugateRhs,
2493 bool LhsIsReal,
bool RhsIsReal>
2494EIGEN_ALWAYS_INLINE
void gemm_complex_unrolled_iteration(
const DataMapper& res,
const Scalar* lhs_base,
2495 const Scalar* rhs_base, Index depth, Index strideA,
2496 Index offsetA, Index strideB, Index& row,
2497 const Packet& pAlphaReal,
const Packet& pAlphaImag,
2498 const Packet& pMask) {
2499 const Scalar *rhs_ptr_real0 = rhs_base, *rhs_ptr_real1 =
nullptr, *rhs_ptr_real2 =
nullptr;
2500 const Scalar *rhs_ptr_imag0 =
nullptr, *rhs_ptr_imag1 =
nullptr, *rhs_ptr_imag2 =
nullptr;
2501 const Index imag_delta = accCols * strideA;
2502 const Index imag_delta2 = accCols2 * strideA;
2503 const Scalar *lhs_ptr_real0 =
nullptr, *lhs_ptr_real1 =
nullptr;
2504 const Scalar *lhs_ptr_real2 =
nullptr, *lhs_ptr_real3 =
nullptr;
2505 PacketBlock<Packet, accRows> accReal0, accImag0, accReal1, accImag1;
2506 PacketBlock<Packet, accRows> accReal2, accImag2, accReal3, accImag3;
2507 PacketBlock<Packet, accRows> taccReal, taccImag;
2508 PacketBlock<Packetc, accRows> acc0, acc1;
2509 PacketBlock<Packetc, accRows * 2> tRes;
2511 MICRO_COMPLEX_SRC2_PTR
2512 MICRO_COMPLEX_SRC_PTR
2513 MICRO_COMPLEX_DST_PTR
2515 const Index peel_depth = depth - (accCols - accCols2);
2517 for (; k + PEEL_COMPLEX <= peel_depth; k += PEEL_COMPLEX) {
2518 MICRO_COMPLEX_PREFETCHN(accRows)
2519 MICRO_COMPLEX_PREFETCH
2520 MICRO_COMPLEX_ONE_PEEL4
2522 for (; k < peel_depth; k++) {
2525 EIGEN_IF_CONSTEXPR (accCols != accCols2) {
2526 for (; k < depth; k++) {
2527 MICRO_COMPLEX_ONE4_PARTIAL
2532 MICRO_COMPLEX_UPDATE
2535#define MICRO_COMPLEX_UNROLL_ITER2(N, M) \
2536 gemm_complex_unrolled_iteration<N + (M ? 1 : 0), Scalar, Packet, Packetc, DataMapper, accRows, accCols, \
2537 M ? M : accCols, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>( \
2538 res3, lhs_base, rhs_base, depth, strideA, offsetA, strideB, row, pAlphaReal, pAlphaImag, pMask); \
2541template <
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
const Index accRows,
2542 const Index accCols,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
2543EIGEN_ALWAYS_INLINE
void gemm_complex_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB,
2544 Index depth, Index strideA, Index offsetA, Index strideB, Index offsetB,
2545 Index col, Index rows, Index remaining_rows,
const Packet& pAlphaReal,
2546 const Packet& pAlphaImag,
const Packet& pMask) {
2547 const DataMapper res3 = res.getSubMapper(0, col);
2549 const Scalar* rhs_base = blockB + advanceCols * col * strideB + MICRO_NEW_ROWS * offsetB;
2550 const Scalar* lhs_base = blockA + accCols * offsetA;
2553#define MAX_COMPLEX_UNROLL 4
2554 while (row + MAX_COMPLEX_UNROLL * accCols <= rows) {
2555 MICRO_COMPLEX_UNROLL_ITER2(MAX_COMPLEX_UNROLL, 0);
2557 switch ((rows - row) / accCols) {
2558#if MAX_COMPLEX_UNROLL > 4
2560 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 4)
2563#if MAX_COMPLEX_UNROLL > 3
2565 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 3)
2568#if MAX_COMPLEX_UNROLL > 2
2570 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 2)
2573#if MAX_COMPLEX_UNROLL > 1
2575 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 1)
2581#undef MAX_COMPLEX_UNROLL
2583 if (remaining_rows > 0) {
2584 gemm_complex_extra_row<Scalar, Packet, Packetc, DataMapper, accRows, accCols, ConjugateLhs, ConjugateRhs, LhsIsReal,
2585 RhsIsReal>(res3, blockA, rhs_base, depth, strideA, offsetA, strideB, row, rows,
2586 remaining_rows, pAlphaReal, pAlphaImag, pMask);
2590#define MICRO_COMPLEX_EXTRA_COLS(N) \
2591 gemm_complex_cols<Scalar, Packet, Packetc, DataMapper, N, accCols, ConjugateLhs, ConjugateRhs, LhsIsReal, \
2592 RhsIsReal>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB, col, rows, \
2593 remaining_rows, pAlphaReal, pAlphaImag, pMask);
2595template <
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
const Index accCols,
2596 bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
2597EIGEN_ALWAYS_INLINE
void gemm_complex_extra_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB,
2598 Index depth, Index strideA, Index offsetA, Index strideB,
2599 Index offsetB, Index col, Index rows, Index cols, Index remaining_rows,
2600 const Packet& pAlphaReal,
const Packet& pAlphaImag,
2601 const Packet& pMask) {
2602 MICRO_EXTRA(MICRO_COMPLEX_EXTRA_COLS, cols - col,
true)
2605template <
typename LhsScalar,
typename RhsScalar,
typename Scalarc,
typename Scalar,
typename Packet,
typename Packetc,
2606 typename RhsPacket,
typename DataMapper,
const Index accRows,
const Index accCols,
bool ConjugateLhs,
2607 bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
2608EIGEN_STRONG_INLINE
void gemm_complex(
const DataMapper& res,
const LhsScalar* blockAc,
const RhsScalar* blockBc,
2609 Index rows, Index depth, Index cols, Scalarc alpha, Index strideA, Index strideB,
2610 Index offsetA, Index offsetB) {
2611 const Index remaining_rows = rows % accCols;
2613 if (strideA == -1) strideA = depth;
2614 if (strideB == -1) strideB = depth;
2616 const Packet pAlphaReal = pset1<Packet>(alpha.real());
2617 const Packet pAlphaImag = pset1<Packet>(alpha.imag());
2618 const Packet pMask = bmask<Packet>(remaining_rows);
2620 const Scalar* blockA = (Scalar*)blockAc;
2621 const Scalar* blockB = (Scalar*)blockBc;
2624 for (; col + accRows <= cols; col += accRows) {
2625 gemm_complex_cols<Scalar, Packet, Packetc, DataMapper, accRows, accCols, ConjugateLhs, ConjugateRhs, LhsIsReal,
2626 RhsIsReal>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB, col, rows,
2627 remaining_rows, pAlphaReal, pAlphaImag, pMask);
2631 gemm_complex_extra_cols<Scalar, Packet, Packetc, DataMapper, accCols, ConjugateLhs, ConjugateRhs, LhsIsReal,
2632 RhsIsReal>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB, col, rows, cols,
2633 remaining_rows, pAlphaReal, pAlphaImag, pMask);
2641EIGEN_ALWAYS_INLINE
bool supportsMMA() {
2642#if defined(EIGEN_ALTIVEC_MMA_ONLY)
2644#elif defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH) && defined(__BUILTIN_CPU_SUPPORTS__)
2645 return __builtin_cpu_supports(
"arch_3_1") && __builtin_cpu_supports(
"mma");
2651EIGEN_ALWAYS_INLINE Packet4f loadAndMultiplyF32(
const Packet4f& acc,
const Packet4f& pAlpha,
float* result) {
2652 Packet4f result_block = ploadu<Packet4f>(result);
2653 return pmadd(acc, pAlpha, result_block);
2656template <
bool lhsExtraRows>
2657EIGEN_ALWAYS_INLINE
void storeF32(
float*& result,
const Packet4f& result_block, Index rows, Index extra_rows) {
2658 EIGEN_IF_CONSTEXPR (lhsExtraRows) {
2659 pstoreu_partial(result, result_block, extra_rows);
2661 pstoreu(result, result_block);
2666template <
bool rhsExtraCols,
bool lhsExtraRows>
2667EIGEN_ALWAYS_INLINE
void storeResults(Packet4f (&acc)[4], Index rows,
const Packet4f& pAlpha,
float* result,
2668 Index extra_cols, Index extra_rows) {
2670 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
2672 Packet4f result_block = loadAndMultiplyF32(acc[x], pAlpha, result);
2673 storeF32<lhsExtraRows>(result, result_block, rows, extra_rows);
2674 }
while (++x < extra_cols);
2676 Packet4f result_block[4];
2677 float* result2 = result;
2679 result_block[x] = loadAndMultiplyF32(acc[x], pAlpha, result);
2684 storeF32<lhsExtraRows>(result2, result_block[x], rows, extra_rows);
2689EIGEN_ALWAYS_INLINE Packet4f oneConvertBF16Hi(
const Packet8us& data) {
2690 Packet8us z = pset1<Packet8us>(0);
2692 return reinterpret_cast<Packet4f
>(vec_mergeh(data, z));
2694 return reinterpret_cast<Packet4f
>(vec_mergeh(z, data));
2698EIGEN_ALWAYS_INLINE Packet4f oneConvertBF16Lo(
const Packet8us& data) {
2699 Packet8us z = pset1<Packet8us>(0);
2701 return reinterpret_cast<Packet4f
>(vec_mergel(data, z));
2703 return reinterpret_cast<Packet4f
>(vec_mergel(z, data));
2707template <Index N, Index M>
2708EIGEN_ALWAYS_INLINE
void storeConvertTwoBF16(
float* to, PacketBlock<Packet8bf, (N + 7) / 8>& block, Index extra = 0) {
2709 EIGEN_IF_CONSTEXPR (N < 4) {
2710 pstoreu_partial(to + 0, oneConvertBF16Hi(block.packet[0].m_val), extra);
2711 }
else EIGEN_IF_CONSTEXPR (N >= (M * 8 + 4)) {
2712 pstoreu(to + 0, oneConvertBF16Hi(block.packet[M].m_val));
2713 EIGEN_IF_CONSTEXPR (N >= 8) {
2714 pstoreu(to + 4, oneConvertBF16Lo(block.packet[M].m_val));
2720EIGEN_ALWAYS_INLINE
void storeConvertBlockBF16(
float* to, PacketBlock<Packet8bf, (N + 7) / 8>& block, Index extra) {
2721 storeConvertTwoBF16<N, 0>(to + 0, block, extra);
2722 EIGEN_IF_CONSTEXPR (N >= 16) {
2723 storeConvertTwoBF16<N, 1>(to + 8, block);
2725 EIGEN_IF_CONSTEXPR (N >= 32) {
2726 storeConvertTwoBF16<N, 2>(to + 16, block);
2727 storeConvertTwoBF16<N, 3>(to + 24, block);
2731template <
bool non_unit_str
ide, Index delta>
2732EIGEN_ALWAYS_INLINE Packet8bf loadBF16fromResult(bfloat16* src, Index resInc) {
2733 EIGEN_IF_CONSTEXPR (non_unit_stride) {
2734 return pgather<bfloat16, Packet8bf>(src + delta * resInc, resInc);
2736 return ploadu<Packet8bf>(src + delta);
2740static Packet16uc p16uc_MERGE16_32_1 = {0, 1, 16, 17, 2, 3, 18, 19, 0, 1, 16, 17, 2, 3, 18, 19};
2741static Packet16uc p16uc_MERGE16_32_2 = {4, 5, 20, 21, 6, 7, 22, 23, 4, 5, 20, 21, 6, 7, 22, 23};
2742static Packet16uc p16uc_MERGE16_32_3 = {8, 9, 24, 25, 10, 11, 26, 27, 8, 9, 24, 25, 10, 11, 26, 27};
2743static Packet16uc p16uc_MERGE16_32_4 = {12, 13, 28, 29, 14, 15, 30, 31, 12, 13, 28, 29, 14, 15, 30, 31};
2745static Packet16uc p16uc_MERGE16_32_5 = {0, 1, 16, 17, 16, 17, 16, 17, 0, 1, 16, 17, 16, 17, 16, 17};
2746static Packet16uc p16uc_MERGE16_32_6 = {2, 3, 18, 19, 18, 19, 18, 19, 2, 3, 18, 19, 18, 19, 18, 19};
2747static Packet16uc p16uc_MERGE16_32_7 = {4, 5, 20, 21, 20, 21, 20, 21, 4, 5, 20, 21, 20, 21, 20, 21};
2748static Packet16uc p16uc_MERGE16_32_8 = {6, 7, 22, 23, 22, 23, 22, 23, 6, 7, 22, 23, 22, 23, 22, 23};
2750EIGEN_ALWAYS_INLINE Packet4f oneConvertBF16Perm(
const Packet8us& data,
const Packet16uc& mask) {
2751 Packet8us z = pset1<Packet8us>(0);
2753 return reinterpret_cast<Packet4f
>(vec_perm(data, z, mask));
2755 return reinterpret_cast<Packet4f
>(vec_perm(z, data, mask));
2759template <
bool lhsExtraRows,
bool odd, Index size>
2760EIGEN_ALWAYS_INLINE
void convertArrayPointerBF16toF32DupOne(
float* result, Index rows,
const bfloat16* src,
2762 Packet4f dup[4 * 4];
2765 for (Index i = 0; i < size; i++) {
2766 data[i] = ploadu<Packet8bf>(src + rows * i);
2769 for (Index i = 0, j = 0; i < size; i++, j += 4) {
2770 dup[j + 0] = oneConvertBF16Perm(data[i].m_val, odd ? p16uc_MERGE16_32_5 : p16uc_MERGE16_32_1);
2771 dup[j + 1] = oneConvertBF16Perm(data[i].m_val, odd ? p16uc_MERGE16_32_6 : p16uc_MERGE16_32_2);
2772 dup[j + 2] = oneConvertBF16Perm(data[i].m_val, odd ? p16uc_MERGE16_32_7 : p16uc_MERGE16_32_3);
2773 dup[j + 3] = oneConvertBF16Perm(data[i].m_val, odd ? p16uc_MERGE16_32_8 : p16uc_MERGE16_32_4);
2776 for (Index j = 0; j < 4 * size; j += 4) {
2777 EIGEN_IF_CONSTEXPR (lhsExtraRows) {
2778 Packet4f z = pset1<Packet4f>(
float(0));
2781 pstoreu(result + (j + i) * 4, dup[j + i]);
2782 }
while (++i < extra_rows);
2784 pstoreu(result + (j + i) * 4, z);
2787 for (Index i = 0; i < 4; i++) {
2788 pstoreu(result + (j + i) * 4, dup[j + i]);
2794template <
bool lhsExtraRows>
2795EIGEN_ALWAYS_INLINE
void convertArrayPointerBF16toF32Dup(
float* result, Index cols, Index rows,
const bfloat16* src,
2796 Index delta, Index extra_rows) {
2799 for (; col + 4 * 2 <= cols; col += 4 * 2, result += 4 * 4 * 4, src += 4 * rows) {
2800 convertArrayPointerBF16toF32DupOne<lhsExtraRows, false, 4>(result, rows, src, extra_rows);
2802 for (; col + 2 <= cols; col += 2, result += 4 * 4, src += rows) {
2803 convertArrayPointerBF16toF32DupOne<lhsExtraRows, false, 1>(result, rows, src, extra_rows);
2806 convertArrayPointerBF16toF32DupOne<lhsExtraRows, true, 1>(result, rows, src - delta, extra_rows);
2810template <const Index size,
bool non_unit_str
ide>
2811EIGEN_ALWAYS_INLINE
void convertPointerBF16toF32(Index& i,
float* result, Index rows, bfloat16*& src, Index resInc) {
2812 while (i + size <= rows) {
2813 const Index count = size == 1 ? rows - i : size;
2814 PacketBlock<Packet8bf, (size + 7) / 8> r32;
2815 EIGEN_IF_CONSTEXPR (size < 8) {
2816 r32.packet[0] = pgather_partial<bfloat16, Packet8bf>(src, non_unit_stride ? resInc : 1, count);
2818 r32.packet[0] = loadBF16fromResult<non_unit_stride, 0>(src, resInc);
2820 EIGEN_IF_CONSTEXPR (size >= 16) {
2821 r32.packet[1] = loadBF16fromResult<non_unit_stride, 8>(src, resInc);
2823 EIGEN_IF_CONSTEXPR (size >= 32) {
2824 r32.packet[2] = loadBF16fromResult<non_unit_stride, 16>(src, resInc);
2825 r32.packet[3] = loadBF16fromResult<non_unit_stride, 24>(src, resInc);
2827 storeConvertBlockBF16<size>(result + i, r32, rows & 3);
2829 if (i < rows) src += count * resInc;
2830 EIGEN_IF_CONSTEXPR (size != 32) break;
2834template <
bool non_unit_stride>
2835EIGEN_ALWAYS_INLINE
void convertArrayPointerBF16toF32(
float* result, Index cols, Index rows, bfloat16* src,
2837 for (Index col = 0; col < cols; col++, result += rows) {
2839 bfloat16* src2 = src;
2840 convertPointerBF16toF32<32, non_unit_stride>(i, result, rows, src2, resInc);
2841 convertPointerBF16toF32<16, non_unit_stride>(i, result, rows, src2, resInc);
2842 convertPointerBF16toF32<8, non_unit_stride>(i, result, rows, src2, resInc);
2843 convertPointerBF16toF32<4, non_unit_stride>(i, result, rows, src2, resInc);
2844 convertPointerBF16toF32<1, non_unit_stride>(i, result, rows, src2, resInc);
2845 if (col + 1 < cols) src += rows * resInc;
2849template <Index num_acc, Index size = 4>
2850EIGEN_ALWAYS_INLINE
void zeroAccumulators(Packet4f (&acc)[num_acc][size]) {
2851 Packet4f z = pset1<Packet4f>(
float(0));
2853 for (Index k = 0; k < num_acc; k++) {
2854 for (Index j = 0; j < size; j++) {
2860template <Index num_acc>
2861EIGEN_ALWAYS_INLINE
void tranposeResults(Packet4f (&acc)[num_acc][4]) {
2862 for (Index i = 0; i < num_acc; i++) {
2863 Packet4ui t0, t1, t2, t3;
2864 t0 = vec_mergeh(
reinterpret_cast<Packet4ui
>(acc[i][0]),
reinterpret_cast<Packet4ui
>(acc[i][2]));
2865 t1 = vec_mergel(
reinterpret_cast<Packet4ui
>(acc[i][0]),
reinterpret_cast<Packet4ui
>(acc[i][2]));
2866 t2 = vec_mergeh(
reinterpret_cast<Packet4ui
>(acc[i][1]),
reinterpret_cast<Packet4ui
>(acc[i][3]));
2867 t3 = vec_mergel(
reinterpret_cast<Packet4ui
>(acc[i][1]),
reinterpret_cast<Packet4ui
>(acc[i][3]));
2868 acc[i][0] =
reinterpret_cast<Packet4f
>(vec_mergeh(t0, t2));
2869 acc[i][1] =
reinterpret_cast<Packet4f
>(vec_mergel(t0, t2));
2870 acc[i][2] =
reinterpret_cast<Packet4f
>(vec_mergeh(t1, t3));
2871 acc[i][3] =
reinterpret_cast<Packet4f
>(vec_mergel(t1, t3));
2875template <Index num_acc>
2876EIGEN_ALWAYS_INLINE
void addResults(Packet4f (&acc)[num_acc][4]) {
2877 for (Index i = 0, j = 0; j < num_acc; i++, j += 2) {
2878 for (Index x = 0, y = 0; x < 2; x++, y += 2) {
2879 for (Index w = 0, z = 0; w < 2; w++, z += 2) {
2880 acc[i][y + w] = acc[j + x][z + 0] + acc[j + x][z + 1];
2886template <Index num_acc,
bool rhsExtraCols,
bool lhsExtraRows, Index num_rhs>
2887EIGEN_ALWAYS_INLINE
void outputResultsVSX(Packet4f (&acc)[num_acc][4], Index rows,
const Packet4f& pAlpha,
2888 float* result,
const Index extra_cols, Index extra_rows) {
2889 tranposeResults<num_acc>(acc);
2890 addResults<num_acc>(acc);
2892 constexpr Index real_rhs = ((num_rhs / 2) - (rhsExtraCols ? 1 : 0));
2894 for (Index i = 0; i < real_rhs; i++, result += 4 * rows, k++) {
2895 storeResults<false, lhsExtraRows>(acc[k], rows, pAlpha, result, extra_cols, extra_rows);
2897 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
2898 storeResults<rhsExtraCols, lhsExtraRows>(acc[k], rows, pAlpha, result, extra_cols, extra_rows);
2903EIGEN_ALWAYS_INLINE
void loadTwoRhsFloat32(
const float* block, Index strideB, Index i, Packet4f& dhs0, Packet4f& dhs1) {
2904 dhs0 = ploadu<Packet4f>(block + strideB * i + 0);
2905 EIGEN_IF_CONSTEXPR (zero) {
2906 Packet4f dhs2 = pset1<Packet4f>(
float(0));
2907 dhs1 = vec_mergel(dhs0, dhs2);
2908 dhs0 = vec_mergeh(dhs0, dhs2);
2910 dhs1 = ploadu<Packet4f>(block + strideB * i + 4);
2914template <Index num_acc,
bool zero,
bool rhsExtraCols, Index num_rhs>
2915EIGEN_ALWAYS_INLINE
void KLoop(
const float* indexA,
const float* indexB, Packet4f (&acc)[num_acc][4], Index strideB,
2916 Index k, Index offsetB, Index extra_cols) {
2917 constexpr Index num_lhs = 4;
2918 Packet4f lhs[num_lhs], rhs[num_rhs];
2920 constexpr Index real_rhs = (num_rhs - (rhsExtraCols ? 2 : 0));
2921 for (Index i = 0; i < real_rhs; i += 2) {
2922 loadTwoRhsFloat32<zero>(indexB + k * 4, strideB, i, rhs[i + 0], rhs[i + 1]);
2924 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
2925 loadTwoRhsFloat32<zero>(indexB + k * extra_cols - offsetB, strideB, real_rhs, rhs[real_rhs + 0], rhs[real_rhs + 1]);
2928 indexA += 2 * k * 4;
2929 for (Index j = 0; j < num_lhs; j++) {
2930 lhs[j] = ploadu<Packet4f>(indexA + j * 4);
2933 for (Index j = 0; j < num_rhs; j++) {
2934 for (Index i = 0; i < num_lhs; i++) {
2935 acc[j][i] = pmadd(rhs[j], lhs[i], acc[j][i]);
2940template <const Index num_acc,
bool rhsExtraCols,
bool lhsExtraRows>
2941EIGEN_ALWAYS_INLINE
void colVSXLoopBodyIter(Index depth, Index rows,
const Packet4f& pAlpha,
const float* indexA,
2942 const float* indexB, Index strideB, Index offsetB,
float* result,
2943 const Index extra_cols,
const Index extra_rows) {
2944 constexpr Index num_rhs = num_acc;
2946 Packet4f acc[num_acc][4];
2948 zeroAccumulators<num_acc>(acc);
2951 for (k = 0; k + 2 <= depth; k += 2) {
2952 KLoop<num_acc, false, rhsExtraCols, num_rhs>(indexA, indexB, acc, strideB, k, offsetB, extra_cols);
2955 KLoop<num_acc, true, rhsExtraCols, num_rhs>(indexA, indexB, acc, strideB, k, offsetB, extra_cols);
2958 outputResultsVSX<num_acc, rhsExtraCols, lhsExtraRows, num_rhs>(acc, rows, pAlpha, result, extra_cols, extra_rows);
2962#define MAX_BFLOAT16_ACC_VSX 4
2964template <const Index num_acc,
bool rhsExtraCols,
bool lhsExtraRows>
2965void colVSXLoopBody(Index& col, Index depth, Index cols, Index rows,
const Packet4f& pAlpha,
const float* indexA,
2966 const float* indexB, Index strideB, Index offsetB,
float* result) {
2967 constexpr Index step = (num_acc * 4);
2968 const Index extra_cols = (rhsExtraCols) ? (cols & 3) : 0;
2969 const Index extra_rows = (lhsExtraRows) ? (rows & 3) : 0;
2970 constexpr bool multiIters = !rhsExtraCols && (num_acc == MAX_BFLOAT16_ACC_VSX);
2973 colVSXLoopBodyIter<num_acc * 2, rhsExtraCols, lhsExtraRows>(depth, rows, pAlpha, indexA, indexB, strideB, offsetB,
2974 result, extra_cols, extra_rows);
2976 indexB += strideB * (num_acc * 2);
2977 result += rows * step;
2978 }
while (multiIters && (step <= cols - (col += step)));
2981template <const Index num_acc,
bool rhsExtraCols,
bool lhsExtraRows>
2982EIGEN_ALWAYS_INLINE
void colVSXLoopBodyExtraN(Index col, Index depth, Index cols, Index rows,
const Packet4f& pAlpha,
2983 const float* indexA,
const float* blockB, Index strideB, Index offsetB,
2985 EIGEN_IF_CONSTEXPR (MAX_BFLOAT16_ACC_VSX > num_acc) {
2986 colVSXLoopBody<num_acc + (rhsExtraCols ? 1 : 0), rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA,
2987 blockB, strideB, offsetB, result);
2991template <
bool rhsExtraCols,
bool lhsExtraRows>
2992void colVSXLoopBodyExtra(Index col, Index depth, Index cols, Index rows,
const Packet4f& pAlpha,
const float* indexA,
2993 const float* blockB, Index strideB, Index offsetB,
float* result) {
2994 switch ((cols - col) >> 2) {
2996 colVSXLoopBodyExtraN<3, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
3000 colVSXLoopBodyExtraN<2, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
3004 colVSXLoopBodyExtraN<1, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
3008 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
3009 colVSXLoopBody<1, true, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB, offsetB, result);
3015template <Index size,
bool lhsExtraRows = false>
3016EIGEN_ALWAYS_INLINE
void colVSXLoops(Index depth, Index cols, Index rows,
const Packet4f& pAlpha,
3017 const bfloat16* indexA,
const float* indexA2,
const float* blockB2, Index strideA,
3018 Index strideB, Index offsetB,
float* result2) {
3019 Index delta_rows = 2 * (lhsExtraRows ? (rows & 3) : size);
3020 for (Index row = 0; row < size; row += 4) {
3021 convertArrayPointerBF16toF32Dup<lhsExtraRows>(
const_cast<float*
>(indexA2), strideA, delta_rows, indexA, row,
3024 const float* blockB = blockB2;
3025 float* result = result2 + row;
3028 if (cols >= (MAX_BFLOAT16_ACC_VSX * 4)) {
3029 colVSXLoopBody<MAX_BFLOAT16_ACC_VSX, false, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA2, blockB,
3030 strideB, 0, result);
3031 blockB += (strideB >> 1) * col;
3032 result += rows * col;
3035 colVSXLoopBodyExtra<true, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA2, blockB, strideB, offsetB,
3038 colVSXLoopBodyExtra<false, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA2, blockB, strideB, 0, result);
3043template <Index size>
3044EIGEN_ALWAYS_INLINE
void calcVSXColLoops(
const bfloat16*& indexA,
const float* indexA2, Index& row, Index depth,
3045 Index cols, Index rows,
const Packet4f& pAlpha,
const float* indexB,
3046 Index strideA, Index strideB, Index offsetA, Index offsetB, Index bigSuffix,
3048 if ((size == 16) || (rows & size)) {
3049 indexA += size * offsetA;
3050 colVSXLoops<size>(depth, cols, rows, pAlpha, indexA, indexA2, indexB, strideA, strideB, offsetB, result + row);
3052 indexA += bigSuffix * size / 16;
3056template <const Index size,
typename DataMapper>
3057EIGEN_ALWAYS_INLINE
void convertBF16toF32(Index& i,
float* result, Index rows,
const DataMapper& src) {
3058 constexpr Index extra = ((size < 4) ? 4 : size);
3059 while (i + size <= rows) {
3060 PacketBlock<Packet8bf, (size + 7) / 8> r32;
3061 r32.packet[0] = src.template loadPacket<Packet8bf>(i + 0);
3062 EIGEN_IF_CONSTEXPR (size >= 16) {
3063 r32.packet[1] = src.template loadPacket<Packet8bf>(i + 8);
3065 EIGEN_IF_CONSTEXPR (size >= 32) {
3066 r32.packet[2] = src.template loadPacket<Packet8bf>(i + 16);
3067 r32.packet[3] = src.template loadPacket<Packet8bf>(i + 24);
3069 storeConvertBlockBF16<size>(result + i, r32, rows & 3);
3071 EIGEN_IF_CONSTEXPR (size != 32) break;
3075template <typename DataMapper>
3076EIGEN_ALWAYS_INLINE
void convertArrayBF16toF32(
float* result, Index cols, Index rows, const DataMapper& src) {
3077 typedef typename DataMapper::LinearMapper LinearMapper;
3078 for (Index j = 0; j < cols; j++, result += rows) {
3079 const LinearMapper src2 = src.getLinearMapper(0, j);
3081 convertBF16toF32<32, LinearMapper>(i, result, rows, src2);
3082 convertBF16toF32<16, LinearMapper>(i, result, rows, src2);
3083 convertBF16toF32<8, LinearMapper>(i, result, rows, src2);
3084 convertBF16toF32<4, LinearMapper>(i, result, rows, src2);
3085 convertBF16toF32<1, LinearMapper>(i, result, rows, src2);
3089EIGEN_ALWAYS_INLINE Packet8bf convertF32toBF16VSX(
const float* res) {
3090 return F32ToBf16Both(ploadu<Packet4f>(res + 0), ploadu<Packet4f>(res + 4));
3093template <
typename DataMapper, const Index size>
3094EIGEN_ALWAYS_INLINE
void convertArrayF32toBF16ColVSX(
float* result, Index col, Index rows,
const DataMapper& res) {
3095 const DataMapper res2 = res.getSubMapper(0, col);
3097 float* result2 = result + col * rows;
3098 for (row = 0; row + 8 <= rows; row += 8, result2 += 8) {
3100 PacketBlock<Packet8bf, size> block;
3101 for (Index j = 0; j < size; j++) {
3102 block.packet[j] = convertF32toBF16VSX(result2 + j * rows);
3104 res2.template storePacketBlock<Packet8bf, size>(row, 0, block);
3108 for (Index j = 0; j < size; j++) {
3109 Packet8bf fp16 = convertF32toBF16VSX(result2 + j * rows);
3110 res2.template storePacketPartial<Packet8bf>(row, j, fp16, rows & 7);
3115template <
typename DataMapper>
3116EIGEN_ALWAYS_INLINE
void convertArrayF32toBF16VSX(
float* result, Index cols, Index rows,
const DataMapper& res) {
3118 for (col = 0; col + 4 <= cols; col += 4) {
3119 convertArrayF32toBF16ColVSX<DataMapper, 4>(result, col, rows, res);
3122 switch (cols - col) {
3124 convertArrayF32toBF16ColVSX<DataMapper, 1>(result, col, rows, res);
3127 convertArrayF32toBF16ColVSX<DataMapper, 2>(result, col, rows, res);
3130 convertArrayF32toBF16ColVSX<DataMapper, 3>(result, col, rows, res);
3135template <
typename DataMapper>
3136void gemmbfloat16(
const DataMapper& res,
const bfloat16* indexA,
const bfloat16* indexB, Index rows, Index depth,
3137 Index cols, bfloat16 alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3138 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
3139 const Packet4f pAlpha = pset1<Packet4f>(falpha);
3141 if (strideA == -1) strideA = depth;
3142 if (strideB == -1) strideB = depth;
3144 ei_declare_aligned_stack_constructed_variable(
float, result, cols* rows, 0);
3145 ei_declare_aligned_stack_constructed_variable(
float, indexB2, strideB* cols, 0);
3146 ei_declare_aligned_stack_constructed_variable(
float, indexA2, ((strideA + 1) & -2) * 4 * 2, 0);
3148 convertArrayBF16toF32<DataMapper>(result, cols, rows, res);
3149 convertArrayPointerBF16toF32(indexB2, cols, strideB,
const_cast<bfloat16*
>(indexB));
3151 Index bigSuffix = 2 * 8 * (strideA - offsetA);
3152 float* indexBF32 = indexB2 + 4 * offsetB;
3158 while (row + 16 <= rows) {
3159 calcVSXColLoops<16>(indexA, indexA2, row, depth, cols, rows, pAlpha, indexBF32, strideA, strideB, offsetA, offsetB,
3163 calcVSXColLoops<8>(indexA, indexA2, row, depth, cols, rows, pAlpha, indexBF32, strideA, strideB, offsetA, offsetB,
3166 calcVSXColLoops<4>(indexA, indexA2, row, depth, cols, rows, pAlpha, indexBF32, strideA, strideB, offsetA, offsetB,
3171 colVSXLoops<4, true>(depth, cols, rows, pAlpha, indexA, indexA2, indexBF32, strideA, strideB, offsetB,
3176 convertArrayF32toBF16VSX<DataMapper>(result, cols, rows, res);
3179#undef MAX_BFLOAT16_ACC_VSX
3181#include "MatrixVectorProduct.inc"
3186template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3187struct gemm_pack_lhs<double, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate, PanelMode> {
3188 void operator()(
double* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0, Index offset = 0);
3191template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3192void gemm_pack_lhs<double, Index, DataMapper, Pack1, Pack2, Packet, ColMajor, Conjugate, PanelMode>::operator()(
3193 double* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride, Index offset) {
3194 dhs_pack<double, DataMapper, Packet2d, ColMajor, PanelMode, true> pack;
3195 pack(blockA, lhs, depth, rows, stride, offset);
3198template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3199struct gemm_pack_lhs<double, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate, PanelMode> {
3200 void operator()(
double* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0, Index offset = 0);
3203template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3204void gemm_pack_lhs<double, Index, DataMapper, Pack1, Pack2, Packet, RowMajor, Conjugate, PanelMode>::operator()(
3205 double* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride, Index offset) {
3206 dhs_pack<double, DataMapper, Packet2d, RowMajor, PanelMode, true> pack;
3207 pack(blockA, lhs, depth, rows, stride, offset);
3210#if EIGEN_ALTIVEC_USE_CUSTOM_PACK
3211template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3212struct gemm_pack_rhs<double, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode> {
3213 void operator()(
double* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0, Index offset = 0);
3216template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3217void gemm_pack_rhs<double, Index, DataMapper, nr, ColMajor, Conjugate, PanelMode>::operator()(
3218 double* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3219 dhs_pack<double, DataMapper, Packet2d, ColMajor, PanelMode, false> pack;
3220 pack(blockB, rhs, depth, cols, stride, offset);
3223template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3224struct gemm_pack_rhs<double, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode> {
3225 void operator()(
double* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0, Index offset = 0);
3228template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3229void gemm_pack_rhs<double, Index, DataMapper, nr, RowMajor, Conjugate, PanelMode>::operator()(
3230 double* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3231 dhs_pack<double, DataMapper, Packet2d, RowMajor, PanelMode, false> pack;
3232 pack(blockB, rhs, depth, cols, stride, offset);
3235template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3236struct gemm_pack_rhs<bfloat16, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode> {
3237 void operator()(bfloat16* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0, Index offset = 0);
3240template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3241void gemm_pack_rhs<bfloat16, Index, DataMapper, nr, ColMajor, Conjugate, PanelMode>::operator()(
3242 bfloat16* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3243 dhs_pack<bfloat16, DataMapper, Packet8bf, ColMajor, PanelMode, false> pack;
3244 pack(blockB, rhs, depth, cols, stride, offset);
3247template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3248struct gemm_pack_rhs<bfloat16, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode> {
3249 void operator()(bfloat16* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0, Index offset = 0);
3252template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3253void gemm_pack_rhs<bfloat16, Index, DataMapper, nr, RowMajor, Conjugate, PanelMode>::operator()(
3254 bfloat16* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3255 dhs_pack<bfloat16, DataMapper, Packet8bf, RowMajor, PanelMode, false> pack;
3256 pack(blockB, rhs, depth, cols, stride, offset);
3260template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3261struct gemm_pack_lhs<bfloat16, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate, PanelMode> {
3262 void operator()(bfloat16* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0, Index offset = 0);
3265template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3266void gemm_pack_lhs<bfloat16, Index, DataMapper, Pack1, Pack2, Packet, ColMajor, Conjugate, PanelMode>::operator()(
3267 bfloat16* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride, Index offset) {
3268 dhs_pack<bfloat16, DataMapper, Packet8bf, ColMajor, PanelMode, true> pack;
3269 pack(blockA, lhs, depth, rows, stride, offset);
3272template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3273struct gemm_pack_lhs<bfloat16, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate, PanelMode> {
3274 void operator()(bfloat16* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0, Index offset = 0);
3277template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3278void gemm_pack_lhs<bfloat16, Index, DataMapper, Pack1, Pack2, Packet, RowMajor, Conjugate, PanelMode>::operator()(
3279 bfloat16* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride, Index offset) {
3280 dhs_pack<bfloat16, DataMapper, Packet8bf, RowMajor, PanelMode, true> pack;
3281 pack(blockA, lhs, depth, rows, stride, offset);
3284template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3285struct gemm_pack_lhs<float, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate, PanelMode> {
3286 void operator()(
float* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0, Index offset = 0);
3289template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3290void gemm_pack_lhs<float, Index, DataMapper, Pack1, Pack2, Packet, RowMajor, Conjugate, PanelMode>::operator()(
3291 float* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride, Index offset) {
3292 dhs_pack<float, DataMapper, Packet4f, RowMajor, PanelMode, true> pack;
3293 pack(blockA, lhs, depth, rows, stride, offset);
3296template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3297struct gemm_pack_lhs<float, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate, PanelMode> {
3298 void operator()(
float* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0, Index offset = 0);
3301template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3302void gemm_pack_lhs<float, Index, DataMapper, Pack1, Pack2, Packet, ColMajor, Conjugate, PanelMode>::operator()(
3303 float* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride, Index offset) {
3304 dhs_pack<float, DataMapper, Packet4f, ColMajor, PanelMode, true> pack;
3305 pack(blockA, lhs, depth, rows, stride, offset);
3308template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3309struct gemm_pack_lhs<std::complex<float>, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate, PanelMode> {
3310 void operator()(std::complex<float>* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
3314template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3315void gemm_pack_lhs<std::complex<float>, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate,
3316 PanelMode>::operator()(std::complex<float>* blockA,
const DataMapper& lhs, Index depth, Index rows,
3317 Index stride, Index offset) {
3318 dhs_cpack<float, DataMapper, Packet4f, Packet2cf, RowMajor, Conjugate, PanelMode, true> pack;
3319 pack(blockA, lhs, depth, rows, stride, offset);
3322template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3323struct gemm_pack_lhs<std::complex<float>, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate, PanelMode> {
3324 void operator()(std::complex<float>* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
3328template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3329void gemm_pack_lhs<std::complex<float>, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate,
3330 PanelMode>::operator()(std::complex<float>* blockA,
const DataMapper& lhs, Index depth, Index rows,
3331 Index stride, Index offset) {
3332 dhs_cpack<float, DataMapper, Packet4f, Packet2cf, ColMajor, Conjugate, PanelMode, true> pack;
3333 pack(blockA, lhs, depth, rows, stride, offset);
3336#if EIGEN_ALTIVEC_USE_CUSTOM_PACK
3337template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3338struct gemm_pack_rhs<float, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode> {
3339 void operator()(
float* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0, Index offset = 0);
3342template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3343void gemm_pack_rhs<float, Index, DataMapper, nr, ColMajor, Conjugate, PanelMode>::operator()(
3344 float* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3345 dhs_pack<float, DataMapper, Packet4f, ColMajor, PanelMode, false> pack;
3346 pack(blockB, rhs, depth, cols, stride, offset);
3349template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3350struct gemm_pack_rhs<float, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode> {
3351 void operator()(
float* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0, Index offset = 0);
3354template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3355void gemm_pack_rhs<float, Index, DataMapper, nr, RowMajor, Conjugate, PanelMode>::operator()(
3356 float* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3357 dhs_pack<float, DataMapper, Packet4f, RowMajor, PanelMode, false> pack;
3358 pack(blockB, rhs, depth, cols, stride, offset);
3362template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3363struct gemm_pack_rhs<std::complex<float>, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode> {
3364 void operator()(std::complex<float>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
3368template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3369void gemm_pack_rhs<std::complex<float>, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode>::operator()(
3370 std::complex<float>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3371 dhs_cpack<float, DataMapper, Packet4f, Packet2cf, ColMajor, Conjugate, PanelMode, false> pack;
3372 pack(blockB, rhs, depth, cols, stride, offset);
3375template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3376struct gemm_pack_rhs<std::complex<float>, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode> {
3377 void operator()(std::complex<float>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
3381template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3382void gemm_pack_rhs<std::complex<float>, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode>::operator()(
3383 std::complex<float>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3384 dhs_cpack<float, DataMapper, Packet4f, Packet2cf, RowMajor, Conjugate, PanelMode, false> pack;
3385 pack(blockB, rhs, depth, cols, stride, offset);
3388template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3389struct gemm_pack_lhs<std::complex<double>, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate, PanelMode> {
3390 void operator()(std::complex<double>* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
3394template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3395void gemm_pack_lhs<std::complex<double>, Index, DataMapper, Pack1, Pack2, Packet,
RowMajor, Conjugate,
3396 PanelMode>::operator()(std::complex<double>* blockA,
const DataMapper& lhs, Index depth, Index rows,
3397 Index stride, Index offset) {
3398 dhs_cpack<double, DataMapper, Packet2d, Packet1cd, RowMajor, Conjugate, PanelMode, true> pack;
3399 pack(blockA, lhs, depth, rows, stride, offset);
3402template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3403struct gemm_pack_lhs<std::complex<double>, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate, PanelMode> {
3404 void operator()(std::complex<double>* blockA,
const DataMapper& lhs, Index depth, Index rows, Index stride = 0,
3408template <
typename Index,
typename DataMapper,
int Pack1,
int Pack2,
typename Packet,
bool Conjugate,
bool PanelMode>
3409void gemm_pack_lhs<std::complex<double>, Index, DataMapper, Pack1, Pack2, Packet,
ColMajor, Conjugate,
3410 PanelMode>::operator()(std::complex<double>* blockA,
const DataMapper& lhs, Index depth, Index rows,
3411 Index stride, Index offset) {
3412 dhs_cpack<double, DataMapper, Packet2d, Packet1cd, ColMajor, Conjugate, PanelMode, true> pack;
3413 pack(blockA, lhs, depth, rows, stride, offset);
3416template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3417struct gemm_pack_rhs<std::complex<double>, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode> {
3418 void operator()(std::complex<double>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
3422template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3423void gemm_pack_rhs<std::complex<double>, Index, DataMapper, nr,
ColMajor, Conjugate, PanelMode>::operator()(
3424 std::complex<double>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3425 dhs_cpack<double, DataMapper, Packet2d, Packet1cd, ColMajor, Conjugate, PanelMode, false> pack;
3426 pack(blockB, rhs, depth, cols, stride, offset);
3429template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3430struct gemm_pack_rhs<std::complex<double>, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode> {
3431 void operator()(std::complex<double>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride = 0,
3435template <
typename Index,
typename DataMapper,
int nr,
bool Conjugate,
bool PanelMode>
3436void gemm_pack_rhs<std::complex<double>, Index, DataMapper, nr,
RowMajor, Conjugate, PanelMode>::operator()(
3437 std::complex<double>* blockB,
const DataMapper& rhs, Index depth, Index cols, Index stride, Index offset) {
3438 dhs_cpack<double, DataMapper, Packet2d, Packet1cd, RowMajor, Conjugate, PanelMode, false> pack;
3439 pack(blockB, rhs, depth, cols, stride, offset);
3443template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3444struct gebp_kernel<float, float, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3445 typedef typename quad_traits<float>::vectortype Packet;
3446 typedef typename quad_traits<float>::rhstype RhsPacket;
3448 void operator()(
const DataMapper& res,
const float* blockA,
const float* blockB, Index rows, Index depth, Index cols,
3449 float alpha, Index strideA = -1, Index strideB = -1, Index offsetA = 0, Index offsetB = 0);
3452template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3453void gebp_kernel<float, float, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3454 const DataMapper& res,
const float* blockA,
const float* blockB, Index rows, Index depth, Index cols,
float alpha,
3455 Index strideA, Index strideB, Index offsetA, Index offsetB) {
3456 const Eigen::Index accRows = quad_traits<float>::rows;
3457 const Eigen::Index accCols = quad_traits<float>::size;
3460 static void (*gemm_function)(
const DataMapper&,
const float*,
const float*, Eigen::Index, Eigen::Index, Eigen::Index,
3461 float, Eigen::Index, Eigen::Index, Eigen::Index, Eigen::Index) =
3462#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3463 (supportsMMA()) ? &Eigen::internal::gemmMMA<float, Packet, RhsPacket, DataMapper, accRows, accCols> :
3465 &Eigen::internal::gemm<float, Packet, RhsPacket, DataMapper, accRows, accCols>;
3466 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3469template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3470struct gebp_kernel<std::complex<float>, std::complex<float>, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3471 typedef Packet4f Packet;
3472 typedef Packet2cf Packetc;
3473 typedef Packet4f RhsPacket;
3475 void operator()(
const DataMapper& res,
const std::complex<float>* blockA,
const std::complex<float>* blockB,
3476 Index rows, Index depth, Index cols, std::complex<float> alpha, Index strideA = -1,
3477 Index strideB = -1, Index offsetA = 0, Index offsetB = 0);
3480template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3481void gebp_kernel<std::complex<float>, std::complex<float>, Index, DataMapper, mr, nr, ConjugateLhs,
3482 ConjugateRhs>::operator()(
const DataMapper& res,
const std::complex<float>* blockA,
3483 const std::complex<float>* blockB, Index rows, Index depth, Index cols,
3484 std::complex<float> alpha, Index strideA, Index strideB, Index offsetA,
3486 const Eigen::Index accRows = quad_traits<float>::rows;
3487 const Eigen::Index accCols = quad_traits<float>::size;
3488 static void (*gemm_function)(
const DataMapper&,
const std::complex<float>*,
const std::complex<float>*, Eigen::Index,
3489 Eigen::Index, Eigen::Index, std::complex<float>, Eigen::Index, Eigen::Index,
3490 Eigen::Index, Eigen::Index) =
3491#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3492 (supportsMMA()) ? &Eigen::internal::gemm_complexMMA<std::complex<float>, std::complex<float>, std::complex<float>,
3493 float, Packet, Packetc, RhsPacket, DataMapper, accRows,
3494 accCols, ConjugateLhs, ConjugateRhs,
false,
false>
3497 &Eigen::internal::gemm_complex<std::complex<float>, std::complex<float>, std::complex<float>,
3498 float, Packet, Packetc, RhsPacket, DataMapper, accRows, accCols,
3499 ConjugateLhs, ConjugateRhs,
false,
false>;
3500 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3503template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3504struct gebp_kernel<float, std::complex<float>, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3505 typedef Packet4f Packet;
3506 typedef Packet2cf Packetc;
3507 typedef Packet4f RhsPacket;
3509 void operator()(
const DataMapper& res,
const float* blockA,
const std::complex<float>* blockB, Index rows,
3510 Index depth, Index cols, std::complex<float> alpha, Index strideA = -1, Index strideB = -1,
3511 Index offsetA = 0, Index offsetB = 0);
3514template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3515void gebp_kernel<float, std::complex<float>, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3516 const DataMapper& res,
const float* blockA,
const std::complex<float>* blockB, Index rows, Index depth, Index cols,
3517 std::complex<float> alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3518 const Eigen::Index accRows = quad_traits<float>::rows;
3519 const Eigen::Index accCols = quad_traits<float>::size;
3520 static void (*gemm_function)(
const DataMapper&,
const float*,
const std::complex<float>*, Eigen::Index, Eigen::Index,
3521 Eigen::Index, std::complex<float>, Eigen::Index, Eigen::Index, Eigen::Index,
3523#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3524 (supportsMMA()) ? &Eigen::internal::gemm_complexMMA<float, std::complex<float>, std::complex<float>, float,
3525 Packet, Packetc, RhsPacket, DataMapper, accRows, accCols,
3526 ConjugateLhs, ConjugateRhs,
true,
false>
3529 &Eigen::internal::gemm_complex<float, std::complex<float>, std::complex<float>, float, Packet,
3530 Packetc, RhsPacket, DataMapper, accRows, accCols, ConjugateLhs,
3531 ConjugateRhs,
true,
false>;
3532 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3535template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3536struct gebp_kernel<std::complex<float>, float, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3537 typedef Packet4f Packet;
3538 typedef Packet2cf Packetc;
3539 typedef Packet4f RhsPacket;
3541 void operator()(
const DataMapper& res,
const std::complex<float>* blockA,
const float* blockB, Index rows,
3542 Index depth, Index cols, std::complex<float> alpha, Index strideA = -1, Index strideB = -1,
3543 Index offsetA = 0, Index offsetB = 0);
3546template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3547void gebp_kernel<std::complex<float>, float, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3548 const DataMapper& res,
const std::complex<float>* blockA,
const float* blockB, Index rows, Index depth, Index cols,
3549 std::complex<float> alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3550 const Eigen::Index accRows = quad_traits<float>::rows;
3551 const Eigen::Index accCols = quad_traits<float>::size;
3552 static void (*gemm_function)(
const DataMapper&,
const std::complex<float>*,
const float*, Eigen::Index, Eigen::Index,
3553 Eigen::Index, std::complex<float>, Eigen::Index, Eigen::Index, Eigen::Index,
3555#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3556 (supportsMMA()) ? &Eigen::internal::gemm_complexMMA<std::complex<float>, float, std::complex<float>, float,
3557 Packet, Packetc, RhsPacket, DataMapper, accRows, accCols,
3558 ConjugateLhs, ConjugateRhs,
false,
true>
3561 &Eigen::internal::gemm_complex<std::complex<float>, float, std::complex<float>, float, Packet,
3562 Packetc, RhsPacket, DataMapper, accRows, accCols, ConjugateLhs,
3563 ConjugateRhs,
false,
true>;
3564 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3567template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3568struct gebp_kernel<double, double, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3569 typedef typename quad_traits<double>::vectortype Packet;
3570 typedef typename quad_traits<double>::rhstype RhsPacket;
3572 void operator()(
const DataMapper& res,
const double* blockA,
const double* blockB, Index rows, Index depth,
3573 Index cols,
double alpha, Index strideA = -1, Index strideB = -1, Index offsetA = 0,
3577template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3578void gebp_kernel<double, double, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3579 const DataMapper& res,
const double* blockA,
const double* blockB, Index rows, Index depth, Index cols,
3580 double alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3581 const Eigen::Index accRows = quad_traits<double>::rows;
3582 const Eigen::Index accCols = quad_traits<double>::size;
3583 static void (*gemm_function)(
const DataMapper&,
const double*,
const double*, Eigen::Index, Eigen::Index,
3584 Eigen::Index, double, Eigen::Index, Eigen::Index, Eigen::Index, Eigen::Index) =
3585#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3586 (supportsMMA()) ? &Eigen::internal::gemmMMA<double, Packet, RhsPacket, DataMapper, accRows, accCols> :
3588 &Eigen::internal::gemm<double, Packet, RhsPacket, DataMapper, accRows, accCols>;
3589 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3592template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3593struct gebp_kernel<std::complex<double>, std::complex<double>, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3594 typedef quad_traits<double>::vectortype Packet;
3595 typedef Packet1cd Packetc;
3596 typedef quad_traits<double>::rhstype RhsPacket;
3598 void operator()(
const DataMapper& res,
const std::complex<double>* blockA,
const std::complex<double>* blockB,
3599 Index rows, Index depth, Index cols, std::complex<double> alpha, Index strideA = -1,
3600 Index strideB = -1, Index offsetA = 0, Index offsetB = 0);
3603template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3604void gebp_kernel<std::complex<double>, std::complex<double>, Index, DataMapper, mr, nr, ConjugateLhs,
3605 ConjugateRhs>::operator()(
const DataMapper& res,
const std::complex<double>* blockA,
3606 const std::complex<double>* blockB, Index rows, Index depth, Index cols,
3607 std::complex<double> alpha, Index strideA, Index strideB, Index offsetA,
3609 const Eigen::Index accRows = quad_traits<double>::rows;
3610 const Eigen::Index accCols = quad_traits<double>::size;
3611 static void (*gemm_function)(
const DataMapper&,
const std::complex<double>*,
const std::complex<double>*,
3612 Eigen::Index, Eigen::Index, Eigen::Index, std::complex<double>, Eigen::Index,
3613 Eigen::Index, Eigen::Index, Eigen::Index) =
3614#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3616 ? &Eigen::internal::gemm_complexMMA<std::complex<double>, std::complex<double>, std::complex<double>, double,
3617 Packet, Packetc, RhsPacket, DataMapper, accRows, accCols, ConjugateLhs,
3618 ConjugateRhs,
false,
false>
3621 &Eigen::internal::gemm_complex<std::complex<double>, std::complex<double>, std::complex<double>, double,
3622 Packet, Packetc, RhsPacket, DataMapper, accRows, accCols, ConjugateLhs,
3623 ConjugateRhs,
false,
false>;
3624 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3627template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3628struct gebp_kernel<std::complex<double>, double, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3629 typedef quad_traits<double>::vectortype Packet;
3630 typedef Packet1cd Packetc;
3631 typedef quad_traits<double>::rhstype RhsPacket;
3633 void operator()(
const DataMapper& res,
const std::complex<double>* blockA,
const double* blockB, Index rows,
3634 Index depth, Index cols, std::complex<double> alpha, Index strideA = -1, Index strideB = -1,
3635 Index offsetA = 0, Index offsetB = 0);
3638template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3639void gebp_kernel<std::complex<double>, double, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3640 const DataMapper& res,
const std::complex<double>* blockA,
const double* blockB, Index rows, Index depth,
3641 Index cols, std::complex<double> alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3642 const Eigen::Index accRows = quad_traits<double>::rows;
3643 const Eigen::Index accCols = quad_traits<double>::size;
3644 static void (*gemm_function)(
const DataMapper&,
const std::complex<double>*,
const double*, Eigen::Index,
3645 Eigen::Index, Eigen::Index, std::complex<double>, Eigen::Index, Eigen::Index,
3646 Eigen::Index, Eigen::Index) =
3647#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3648 (supportsMMA()) ? &Eigen::internal::gemm_complexMMA<std::complex<double>, double, std::complex<double>, double,
3649 Packet, Packetc, RhsPacket, DataMapper, accRows, accCols,
3650 ConjugateLhs, ConjugateRhs,
false,
true>
3653 &Eigen::internal::gemm_complex<std::complex<double>, double, std::complex<double>, double, Packet,
3654 Packetc, RhsPacket, DataMapper, accRows, accCols, ConjugateLhs,
3655 ConjugateRhs,
false,
true>;
3656 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3659template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3660struct gebp_kernel<double, std::complex<double>, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3661 typedef quad_traits<double>::vectortype Packet;
3662 typedef Packet1cd Packetc;
3663 typedef quad_traits<double>::rhstype RhsPacket;
3665 void operator()(
const DataMapper& res,
const double* blockA,
const std::complex<double>* blockB, Index rows,
3666 Index depth, Index cols, std::complex<double> alpha, Index strideA = -1, Index strideB = -1,
3667 Index offsetA = 0, Index offsetB = 0);
3670template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3671void gebp_kernel<double, std::complex<double>, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3672 const DataMapper& res,
const double* blockA,
const std::complex<double>* blockB, Index rows, Index depth,
3673 Index cols, std::complex<double> alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3674 const Eigen::Index accRows = quad_traits<double>::rows;
3675 const Eigen::Index accCols = quad_traits<double>::size;
3676 static void (*gemm_function)(
const DataMapper&,
const double*,
const std::complex<double>*, Eigen::Index,
3677 Eigen::Index, Eigen::Index, std::complex<double>, Eigen::Index, Eigen::Index,
3678 Eigen::Index, Eigen::Index) =
3679#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3680 (supportsMMA()) ? &Eigen::internal::gemm_complexMMA<double, std::complex<double>, std::complex<double>, double,
3681 Packet, Packetc, RhsPacket, DataMapper, accRows, accCols,
3682 ConjugateLhs, ConjugateRhs,
true,
false>
3685 &Eigen::internal::gemm_complex<double, std::complex<double>, std::complex<double>, double, Packet,
3686 Packetc, RhsPacket, DataMapper, accRows, accCols, ConjugateLhs,
3687 ConjugateRhs,
true,
false>;
3688 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3691template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3692struct gebp_kernel<bfloat16, bfloat16, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs> {
3693 typedef typename quad_traits<bfloat16>::vectortype Packet;
3694 typedef typename quad_traits<bfloat16>::rhstype RhsPacket;
3696 void operator()(
const DataMapper& res,
const bfloat16* blockA,
const bfloat16* blockB, Index rows, Index depth,
3697 Index cols, bfloat16 alpha, Index strideA = -1, Index strideB = -1, Index offsetA = 0,
3701template <
typename Index,
typename DataMapper,
int mr,
int nr,
bool ConjugateLhs,
bool ConjugateRhs>
3702void gebp_kernel<bfloat16, bfloat16, Index, DataMapper, mr, nr, ConjugateLhs, ConjugateRhs>::operator()(
3703 const DataMapper& res,
const bfloat16* blockA,
const bfloat16* blockB, Index rows, Index depth, Index cols,
3704 bfloat16 alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
3705 static void (*gemm_function)(
const DataMapper&,
const bfloat16*,
const bfloat16*, Eigen::Index, Eigen::Index,
3706 Eigen::Index, bfloat16, Eigen::Index, Eigen::Index, Eigen::Index, Eigen::Index) =
3707#ifdef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
3708 (supportsMMA()) ? &Eigen::internal::gemmMMAbfloat16<DataMapper> :
3710 &Eigen::internal::gemmbfloat16<DataMapper>;
3711 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321