Eigen  5.0.1
 
Loading...
Searching...
No Matches
MatrixProduct.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2020 Everton Constantino (everton.constantino@ibm.com)
5// Copyright (C) 2021 Chip Kerchner (chip.kerchner@ibm.com)
6//
7// This Source Code Form is subject to the terms of the Mozilla
8// Public License v. 2.0. If a copy of the MPL was not distributed
9// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
10// SPDX-License-Identifier: MPL-2.0
11
12#ifndef EIGEN_MATRIX_PRODUCT_ALTIVEC_H
13#define EIGEN_MATRIX_PRODUCT_ALTIVEC_H
14
15#ifndef EIGEN_ALTIVEC_USE_CUSTOM_PACK
16#define EIGEN_ALTIVEC_USE_CUSTOM_PACK 1
17#endif
18
19#if !defined(EIGEN_ALTIVEC_DISABLE_MMA)
20#define EIGEN_ALTIVEC_DISABLE_MMA 0
21#endif
22
23// Check for MMA builtin support.
24#if !EIGEN_ALTIVEC_DISABLE_MMA && defined(__has_builtin)
25#if __has_builtin(__builtin_mma_assemble_acc)
26#define EIGEN_ALTIVEC_MMA_SUPPORT
27#endif
28#endif
29
30// Check if and how we should actually use MMA if supported.
31#if defined(EIGEN_ALTIVEC_MMA_SUPPORT)
32
33#if !defined(EIGEN_ALTIVEC_ENABLE_MMA_DYNAMIC_DISPATCH)
34#define EIGEN_ALTIVEC_ENABLE_MMA_DYNAMIC_DISPATCH 0
35#endif
36
37// Check if we want to enable dynamic dispatch. Not supported by LLVM.
38#if EIGEN_ALTIVEC_ENABLE_MMA_DYNAMIC_DISPATCH && !EIGEN_COMP_LLVM
39#define EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH 1
40// Otherwise, use MMA by default if available.
41#elif defined(__MMA__)
42#define EIGEN_ALTIVEC_MMA_ONLY 1
43#endif
44
45#endif // EIGEN_ALTIVEC_MMA_SUPPORT
46
47#include "MatrixProductCommon.h"
48
49#if defined(EIGEN_ALTIVEC_MMA_ONLY) || defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH)
50#include "MatrixProductMMA.h"
51#endif
52
53// IWYU pragma: private
54#include "../../InternalHeaderCheck.h"
55
56namespace Eigen {
57
58namespace internal {
59
60/**************************
61 * Constants and typedefs *
62 **************************/
63template <typename Scalar>
64struct quad_traits {
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 };
69};
70
71template <>
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 };
77};
78
79template <>
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 };
85};
86
87// MatrixProduct decomposes real/imaginary vectors into a real vector and an imaginary vector, this turned out
88// to be faster than Eigen's usual approach of having real/imaginary pairs on a single vector. These constants then
89// are responsible to extract from convert between Eigen's and MatrixProduct approach.
90
91const static Packet16uc p16uc_GETREAL32 = {0, 1, 2, 3, 8, 9, 10, 11, 16, 17, 18, 19, 24, 25, 26, 27};
92
93const static Packet16uc p16uc_GETIMAG32 = {4, 5, 6, 7, 12, 13, 14, 15, 20, 21, 22, 23, 28, 29, 30, 31};
94
95const static Packet16uc p16uc_GETREAL32b = {0, 1, 2, 3, 16, 17, 18, 19, 8, 9, 10, 11, 24, 25, 26, 27};
96
97const static Packet16uc p16uc_GETIMAG32b = {4, 5, 6, 7, 20, 21, 22, 23, 12, 13, 14, 15, 28, 29, 30, 31};
98
99/*********************************************
100 * Single precision real and complex packing *
101 * *******************************************/
102
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;
121 if (i < j) {
122 v.real(dt(j, i).real());
123 v.imag(-dt(j, i).imag());
124 } else if (i > j) {
125 v.real(dt(i, j).real());
126 v.imag(dt(i, j).imag());
127 } else {
128 v.real(dt(i, j).real());
129 v.imag((Scalar)0.0);
130 }
131 return v;
132}
133
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);
142
143 Index rir = 0, rii, j = 0;
144 for (; j + vectorSize <= cols; j += vectorSize) {
145 rii = rir + vectorDelta;
146
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);
150
151 blockBf[rir + k] = v.real();
152 blockBf[rii + k] = v.imag();
153 }
154 rir += vectorSize;
155 rii += vectorSize;
156 }
157
158 rir += vectorDelta;
159 }
160
161 for (; j < cols; j++) {
162 rii = rir + rows;
163
164 for (Index i = k2; i < depth; i++) {
165 std::complex<Scalar> v = getAdjointVal<Scalar, StorageOrder>(i, j, rhs);
166
167 blockBf[rir] = v.real();
168 blockBf[rii] = v.imag();
169
170 rir += 1;
171 rii += 1;
172 }
173
174 rir += rows;
175 }
176}
177
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);
186
187 Index rir = 0, rii, j = 0;
188 for (; j + vectorSize <= rows; j += vectorSize) {
189 rii = rir + vectorDelta;
190
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);
194
195 blockAf[rir + k] = v.real();
196 blockAf[rii + k] = v.imag();
197 }
198 rir += vectorSize;
199 rii += vectorSize;
200 }
201
202 rir += vectorDelta;
203 }
204
205 if (j < rows) {
206 rii = rir + ((rows - j) * depth);
207
208 for (Index i = 0; i < depth; i++) {
209 Index k = j;
210 for (; k < rows; k++) {
211 std::complex<Scalar> v = getAdjointVal<Scalar, StorageOrder>(k, i, lhs);
212
213 blockAf[rir] = v.real();
214 blockAf[rii] = v.imag();
215
216 rir += 1;
217 rii += 1;
218 }
219 }
220 }
221}
222
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;
229
230 Index ri = 0, j = 0;
231 for (; j + N * vectorSize <= cols; j += N * vectorSize) {
232 Index i = k2;
233 for (; i < depth; i++) {
234 for (Index k = 0; k < N * vectorSize; k++) {
235 if (i <= j + k)
236 blockB[ri + k] = rhs(j + k, i);
237 else
238 blockB[ri + k] = rhs(i, j + k);
239 }
240 ri += N * vectorSize;
241 }
242 }
243
244 for (; j < cols; j++) {
245 for (Index i = k2; i < depth; i++) {
246 if (j <= i)
247 blockB[ri] = rhs(i, j);
248 else
249 blockB[ri] = rhs(j, i);
250 ri += 1;
251 }
252 }
253}
254
255template <typename Scalar, int StorageOrder>
256EIGEN_STRONG_INLINE void symm_pack_lhs_helper(Scalar* blockA, const Scalar* _lhs, Index lhsStride, Index cols,
257 Index rows) {
258 const Index depth = cols;
259 const_blas_data_mapper<Scalar, Index, StorageOrder> lhs(_lhs, lhsStride);
260 const Index vectorSize = quad_traits<Scalar>::vectorsize;
261
262 Index ri = 0, j = 0;
263 for (; j + vectorSize <= rows; j += vectorSize) {
264 Index i = 0;
265
266 for (; i < depth; i++) {
267 for (Index k = 0; k < vectorSize; k++) {
268 if (i <= j + k)
269 blockA[ri + k] = lhs(j + k, i);
270 else
271 blockA[ri + k] = lhs(i, j + k);
272 }
273 ri += vectorSize;
274 }
275 }
276
277 if (j < rows) {
278 for (Index i = 0; i < depth; i++) {
279 Index k = j;
280 for (; k < rows; k++) {
281 if (i <= k)
282 blockA[ri] = lhs(k, i);
283 else
284 blockA[ri] = lhs(i, k);
285 ri += 1;
286 }
287 }
288 }
289}
290
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,
294 Index k2) {
295 symm_pack_complex_rhs_helper<float, StorageOrder, 1>(blockB, _rhs, rhsStride, rows, cols, k2);
296 }
297};
298
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,
302 Index rows) {
303 symm_pack_complex_lhs_helper<float, StorageOrder>(blockA, _lhs, lhsStride, cols, rows);
304 }
305};
306
307// *********** symm_pack std::complex<float64> ***********
308
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);
314 }
315};
316
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,
320 Index rows) {
321 symm_pack_complex_lhs_helper<double, StorageOrder>(blockA, _lhs, lhsStride, cols, rows);
322 }
323};
324
325// *********** symm_pack float32 ***********
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);
330 }
331};
332
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);
337 }
338};
339
340// *********** symm_pack float64 ***********
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);
345 }
346};
347
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);
352 }
353};
354
365
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]);
373 }
374 EIGEN_IF_CONSTEXPR (N > 3) {
375 pstore<Scalar>(to + (3 * size), block.packet[3]);
376 }
377}
378
379// General template for lhs & rhs complex packing.
380template <typename Scalar, typename DataMapper, typename Packet, typename PacketC, int StorageOrder, bool Conjugate,
381 bool PanelMode, bool UseLhs>
382struct dhs_cpack {
383 template <bool transpose>
384 EIGEN_ALWAYS_INLINE void dhs_cblock(PacketBlock<PacketC, 8>& cblock, PacketBlock<Packet, 4>& block,
385 const Packet16uc& permute) {
386 if (transpose) {
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);
391
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])));
402#else
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));
407#endif
408
409 block.packet[0] = t0;
410 block.packet[1] = t1;
411 block.packet[2] = t2;
412 block.packet[3] = t3;
413 } else {
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);
418 }
419 }
420
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;
425
426 for (; i + vectorSize <= depth; i += vectorSize) {
427 EIGEN_IF_CONSTEXPR (UseLhs) {
428 bload<DataMapper, PacketC, 2, StorageOrder, true, 4>(cblock, lhs2, 0, i);
429 } else {
430 bload<DataMapper, PacketC, 2, StorageOrder, true, 4>(cblock, lhs2, i, 0);
431 }
432
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);
436 } else {
437 dhs_cblock<false>(cblock, blockr, p16uc_GETREAL32);
438 dhs_cblock<false>(cblock, blocki, p16uc_GETIMAG32);
439 }
440
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];
446 }
447
448 storeBlock<Scalar, Packet, 4>(blockAt + rir, blockr);
449 storeBlock<Scalar, Packet, 4>(blockAt + rii, blocki);
450
451 rir += 4 * vectorSize;
452 rii += 4 * vectorSize;
453 }
454 }
455
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);
462 Index j = 0;
463
464 for (; j + vectorSize <= rows; j += vectorSize) {
465 const DataMapper lhs2 = UseLhs ? lhs.getSubMapper(j, 0) : lhs.getSubMapper(0, j);
466 Index i = 0;
467
468 rii = rir + vectorDelta;
469
470 dhs_ccopy(blockAt, lhs2, i, rir, rii, depth, vectorSize);
471
472 for (; i < depth; i++) {
473 PacketBlock<Packet, 1> blockr, blocki;
474 PacketBlock<PacketC, 2> cblock;
475
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);
480 } else {
481 cblock.packet[0] = lhs2.template loadPacket<PacketC>(i, 0);
482 cblock.packet[1] = lhs2.template loadPacket<PacketC>(i, 2);
483 }
484 } else {
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));
488 } else {
489 cblock.packet[0] = pload2(lhs2(i, 0), lhs2(i, 1));
490 cblock.packet[1] = pload2(lhs2(i, 2), lhs2(i, 3));
491 }
492 }
493
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);
496
497 EIGEN_IF_CONSTEXPR (Conjugate) {
498 blocki.packet[0] = -blocki.packet[0];
499 }
500
501 pstore<Scalar>(blockAt + rir, blockr.packet[0]);
502 pstore<Scalar>(blockAt + rii, blocki.packet[0]);
503
504 rir += vectorSize;
505 rii += vectorSize;
506 }
507
508 rir += ((PanelMode) ? (vectorSize * (2 * stride - depth)) : vectorDelta);
509 }
510
511 EIGEN_IF_CONSTEXPR (!UseLhs) {
512 EIGEN_IF_CONSTEXPR (PanelMode) rir -= (offset * (vectorSize - 1));
513
514 for (; j < rows; j++) {
515 const DataMapper lhs2 = lhs.getSubMapper(0, j);
516 rii = rir + ((PanelMode) ? stride : depth);
517
518 for (Index i = 0; i < depth; i++) {
519 blockAt[rir] = lhs2(i, 0).real();
520
521 EIGEN_IF_CONSTEXPR (Conjugate)
522 blockAt[rii] = -lhs2(i, 0).imag();
523 else
524 blockAt[rii] = lhs2(i, 0).imag();
525
526 rir += 1;
527 rii += 1;
528 }
529
530 rir += ((PanelMode) ? (2 * stride - depth) : depth);
531 }
532 } else {
533 if (j < rows) {
534 EIGEN_IF_CONSTEXPR (PanelMode) rir += (offset * (rows - j - vectorSize));
535 rii = rir + (((PanelMode) ? stride : depth) * (rows - j));
536
537 for (Index i = 0; i < depth; i++) {
538 Index k = j;
539 for (; k < rows; k++) {
540 blockAt[rir] = lhs(k, i).real();
541
542 EIGEN_IF_CONSTEXPR (Conjugate)
543 blockAt[rii] = -lhs(k, i).imag();
544 else
545 blockAt[rii] = lhs(k, i).imag();
546
547 rir += 1;
548 rii += 1;
549 }
550 }
551 }
552 }
553 }
554};
555
556// General template for lhs & rhs packing.
557template <typename Scalar, typename DataMapper, typename Packet, int StorageOrder, bool PanelMode, bool UseLhs>
558struct dhs_pack {
559 template <Index n>
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];
563
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);
568 } else {
569 bload<DataMapper, Packet, 4, StorageOrder, false, 4>(block[k], lhs2, i + k * vectorSize, 0);
570 }
571 }
572
573 EIGEN_IF_CONSTEXPR (((StorageOrder == RowMajor) && UseLhs) || ((StorageOrder == ColMajor) && !UseLhs)) {
574 for (Index k = 0; k < n; k++) {
575 ptranspose(block[k]);
576 }
577 }
578
579 for (Index k = 0; k < n; k++) {
580 storeBlock<Scalar, Packet, 4>(blockA + ri + k * 4 * vectorSize, block[k]);
581 }
582
583 ri += n * 4 * vectorSize;
584 }
585 }
586
587 EIGEN_STRONG_INLINE void operator()(Scalar* blockA, const DataMapper& lhs, Index depth, Index rows, Index stride,
588 Index offset) {
589 const Index vectorSize = quad_traits<Scalar>::vectorsize;
590 Index ri = 0, j = 0;
591
592 for (; j + vectorSize <= rows; j += vectorSize) {
593 const DataMapper lhs2 = UseLhs ? lhs.getSubMapper(j, 0) : lhs.getSubMapper(0, j);
594 Index i = 0;
595
596 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * offset;
597
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);
601
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);
609 } else {
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);
614 }
615 } else {
616 Packet lhsV;
617 EIGEN_IF_CONSTEXPR (UseLhs) {
618 lhsV = lhs2.template loadPacket<Packet>(0, i);
619 } else {
620 lhsV = lhs2.template loadPacket<Packet>(i, 0);
621 }
622 pstore<Scalar>(blockA + ri, lhsV);
623 }
624
625 ri += vectorSize;
626 }
627
628 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * (stride - offset - depth);
629 }
630
631 EIGEN_IF_CONSTEXPR (!UseLhs) {
632 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset;
633
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);
638 ri += 1;
639 }
640
641 EIGEN_IF_CONSTEXPR (PanelMode) ri += stride - depth;
642 }
643 } else {
644 if (j < rows) {
645 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (rows - j);
646
647 for (Index i = 0; i < depth; i++) {
648 Index k = j;
649 for (; k < rows; k++) {
650 blockA[ri] = lhs(k, i);
651 ri += 1;
652 }
653 }
654 }
655 }
656 }
657};
658
659// General template for lhs packing, float64 specialization.
660template <typename DataMapper, int StorageOrder, bool PanelMode>
661struct dhs_pack<double, DataMapper, Packet2d, StorageOrder, PanelMode, true> {
662 template <Index n>
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];
666
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);
672 } else {
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);
675 }
676 }
677
678 EIGEN_IF_CONSTEXPR (StorageOrder == RowMajor) {
679 for (Index k = 0; k < n; k++) {
680 ptranspose(block[k]);
681 }
682 }
683
684 for (Index k = 0; k < n; k++) {
685 storeBlock<double, Packet2d, 2>(blockA + ri + k * 2 * vectorSize, block[k]);
686 }
687
688 ri += n * 2 * vectorSize;
689 }
690 }
691
692 EIGEN_STRONG_INLINE void operator()(double* blockA, const DataMapper& lhs, Index depth, Index rows, Index stride,
693 Index offset) {
694 const Index vectorSize = quad_traits<double>::vectorsize;
695 Index ri = 0, j = 0;
696
697 for (; j + vectorSize <= rows; j += vectorSize) {
698 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
699 Index i = 0;
700
701 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * offset;
702
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);
706
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);
711 } else {
712 Packet2d lhsV = lhs2.template loadPacket<Packet2d>(0, i);
713 pstore<double>(blockA + ri, lhsV);
714 }
715
716 ri += vectorSize;
717 }
718
719 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * (stride - offset - depth);
720 }
721
722 if (j < rows) {
723 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (rows - j);
724
725 for (Index i = 0; i < depth; i++) {
726 Index k = j;
727 for (; k < rows; k++) {
728 blockA[ri] = lhs(k, i);
729 ri += 1;
730 }
731 }
732 }
733 }
734};
735
736// General template for rhs packing, float64 specialization.
737template <typename DataMapper, int StorageOrder, bool PanelMode>
738struct dhs_pack<double, DataMapper, Packet2d, StorageOrder, PanelMode, false> {
739 template <Index n>
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];
744
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);
752 } else {
753 block3[k].packet[0] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 0, 0); //[a1 a2]
754 block3[k].packet[1] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 0, 2); //[a3 a4]
755 block3[k].packet[2] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 1, 0); //[b1 b2]
756 block3[k].packet[3] = rhs2.template loadPacket<Packet2d>(i + k * vectorSize + 1, 2); //[b3 b4]
757 }
758 }
759
760 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
761 for (Index k = 0; k < n; k++) {
762 ptranspose(block1[k]);
763 ptranspose(block2[k]);
764 }
765 }
766
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]);
773 } else {
774 storeBlock<double, Packet2d, 4>(blockB + ri + k * 4 * vectorSize, block3[k]);
775 }
776 }
777
778 ri += n * 4 * vectorSize;
779 }
780 }
781
782 EIGEN_STRONG_INLINE void operator()(double* blockB, const DataMapper& rhs, Index depth, Index cols, Index stride,
783 Index offset) {
784 const Index vectorSize = quad_traits<double>::vectorsize;
785 Index ri = 0, j = 0;
786
787 for (; j + 2 * vectorSize <= cols; j += 2 * vectorSize) {
788 const DataMapper rhs2 = rhs.getSubMapper(0, j);
789 Index i = 0;
790
791 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (2 * vectorSize);
792
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);
796
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);
801
802 ri += vectorSize;
803
804 blockB[ri + 0] = rhs2(i, 2);
805 blockB[ri + 1] = rhs2(i, 3);
806 } else {
807 Packet2d rhsV = rhs2.template loadPacket<Packet2d>(i, 0);
808 pstore<double>(blockB + ri, rhsV);
809
810 ri += vectorSize;
811
812 rhsV = rhs2.template loadPacket<Packet2d>(i, 2);
813 pstore<double>(blockB + ri, rhsV);
814 }
815 ri += vectorSize;
816 }
817
818 EIGEN_IF_CONSTEXPR (PanelMode) ri += (2 * vectorSize) * (stride - offset - depth);
819 }
820
821 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset;
822
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);
827 ri += 1;
828 }
829
830 EIGEN_IF_CONSTEXPR (PanelMode) ri += stride - depth;
831 }
832 }
833};
834
835// General template for lhs packing, bfloat16 specialization.
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,
839 Index offset) {
840 const Index vectorSize = quad_traits<bfloat16>::vectorsize;
841 Index ri = 0, j = 0;
842
843 for (; j + 2 * vectorSize <= rows; j += 2 * vectorSize) {
844 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
845 Index i = 0;
846
847 EIGEN_IF_CONSTEXPR (PanelMode) ri += 2 * vectorSize * offset;
848
849 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
850 for (; i + 2 <= depth; i += 2) {
851 PacketBlock<Packet8bf, 4> block;
852
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);
857
858 Packet8bf t0, t1;
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;
865
866 storeBlock<bfloat16, Packet8bf, 4>(blockA + ri, block);
867
868 ri += 2 * 2 * vectorSize;
869 }
870 if (depth & 1) {
871 PacketBlock<Packet8bf, 2> block;
872
873 block.packet[0] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
874 block.packet[1] = lhs2.template loadPacket<Packet8bf>(1 * vectorSize, i + 0);
875
876 storeBlock<bfloat16, Packet8bf, 2>(blockA + ri, block);
877
878 ri += 2 * vectorSize;
879 }
880 } else {
881 for (; i + vectorSize <= depth; i += vectorSize) {
882 PacketBlock<Packet8bf, 8> block1, block2;
883
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);
886
887 Packet4ui v1[8], v2[8];
888
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));
921
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])));
955#else
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));
972#endif
973
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]);
979 }
980
981 ri += 2 * vectorSize * vectorSize;
982 }
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);
987 }
988
989 ri += 2 * 2 * vectorSize;
990 }
991 if (depth & 1) {
992 for (Index M = 0; M < 2 * vectorSize; M++) {
993 blockA[ri + M] = lhs2(M, i);
994 }
995 ri += 2 * vectorSize;
996 }
997 }
998
999 EIGEN_IF_CONSTEXPR (PanelMode) ri += 2 * vectorSize * (stride - offset - depth);
1000 }
1001 for (; j + vectorSize <= rows; j += vectorSize) {
1002 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
1003 Index i = 0;
1004
1005 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * offset;
1006
1007 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1008 for (; i + 2 <= depth; i += 2) {
1009 PacketBlock<Packet8bf, 2> block;
1010
1011 block.packet[0] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
1012 block.packet[1] = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 1);
1013
1014 Packet8bf t0;
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;
1018
1019 storeBlock<bfloat16, Packet8bf, 2>(blockA + ri, block);
1020
1021 ri += 2 * vectorSize;
1022 }
1023 if (depth & 1) {
1024 Packet8bf lhsV = lhs2.template loadPacket<Packet8bf>(0 * vectorSize, i + 0);
1025 pstore<bfloat16>(blockA + ri, lhsV);
1026
1027 ri += vectorSize;
1028 }
1029 } else {
1030 for (; i + vectorSize <= depth; i += vectorSize) {
1031 PacketBlock<Packet8bf, 8> block1;
1032
1033 bload<DataMapper, Packet8bf, 8, StorageOrder, false, 8>(block1, lhs2, 0 * vectorSize, i);
1034
1035 Packet4ui v1[8];
1036
1037 // This is transposing and interleaving data
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));
1054
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])));
1072#else
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));
1081#endif
1082
1083 for (Index M = 0; M < 8; M++) {
1084 pstore<bfloat16>(blockA + ri + (vectorSize * M), block1.packet[M]);
1085 }
1086
1087 ri += vectorSize * vectorSize;
1088 }
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);
1093 }
1094
1095 ri += 2 * vectorSize;
1096 }
1097 if (depth & 1) {
1098 for (Index M = 0; M < vectorSize; M++) {
1099 blockA[ri + M] = lhs2(M, i);
1100 }
1101
1102 ri += vectorSize;
1103 }
1104 }
1105
1106 EIGEN_IF_CONSTEXPR (PanelMode) ri += vectorSize * (stride - offset - depth);
1107 }
1108 if (j + 4 <= rows) {
1109 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
1110 Index i = 0;
1111
1112 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * offset;
1113
1114 for (; i + 2 <= depth; i += 2) {
1115 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1116 PacketBlock<Packet8bf, 2> block;
1117
1118 block.packet[0] = lhs2.template loadPacketPartial<Packet8bf>(0, i + 0, 4);
1119 block.packet[1] = lhs2.template loadPacketPartial<Packet8bf>(0, i + 1, 4);
1120
1121 block.packet[0] = vec_mergeh(block.packet[0].m_val, block.packet[1].m_val);
1122
1123 pstore<bfloat16>(blockA + ri, block.packet[0]);
1124 } else {
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);
1133 }
1134
1135 ri += 2 * 4;
1136 }
1137 if (depth & 1) {
1138 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1139 Packet8bf lhsV = lhs2.template loadPacketPartial<Packet8bf>(0, i + 0, 4);
1140
1141 pstore_partial<bfloat16>(blockA + ri, lhsV, 4);
1142 } else {
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);
1147 }
1148
1149 ri += 4;
1150 }
1151
1152 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * (stride - offset - depth);
1153 j += 4;
1154 }
1155
1156 if (j < rows) {
1157 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (rows - j);
1158
1159 Index i = 0;
1160 for (; i + 2 <= depth; i += 2) {
1161 Index k = j;
1162 for (; k < rows; k++) {
1163 blockA[ri + 0] = lhs(k, i + 0);
1164 blockA[ri + 1] = lhs(k, i + 1);
1165 ri += 2;
1166 }
1167 }
1168 if (depth & 1) {
1169 for (; j < rows; j++) {
1170 blockA[ri] = lhs(j, i);
1171 ri += 1;
1172 }
1173 }
1174 }
1175 }
1176};
1177
1178// General template for rhs packing, bfloat16 specialization.
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,
1182 Index offset) {
1183 const Index vectorSize = quad_traits<bfloat16>::vectorsize;
1184 Index ri = 0, j = 0;
1185
1186 for (; j + 4 <= cols; j += 4) {
1187 const DataMapper rhs2 = rhs.getSubMapper(0, j);
1188 Index i = 0;
1189
1190 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * offset;
1191
1192 for (; i + vectorSize <= depth; i += vectorSize) {
1193 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1194 PacketBlock<Packet8bf, 4> block;
1195
1196 bload<DataMapper, Packet8bf, 4, StorageOrder, false, 4>(block, rhs2, i, 0);
1197
1198 Packet4ui t0, t1, t2, t3;
1199
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));
1208
1209#ifdef EIGEN_VECTORIZE_VSX
1210 block.packet[0] =
1211 reinterpret_cast<Packet8us>(vec_mergeh(reinterpret_cast<Packet2ul>(t0), reinterpret_cast<Packet2ul>(t2)));
1212 block.packet[1] =
1213 reinterpret_cast<Packet8us>(vec_mergel(reinterpret_cast<Packet2ul>(t0), reinterpret_cast<Packet2ul>(t2)));
1214 block.packet[2] =
1215 reinterpret_cast<Packet8us>(vec_mergeh(reinterpret_cast<Packet2ul>(t1), reinterpret_cast<Packet2ul>(t3)));
1216 block.packet[3] =
1217 reinterpret_cast<Packet8us>(vec_mergel(reinterpret_cast<Packet2ul>(t1), reinterpret_cast<Packet2ul>(t3)));
1218#else
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));
1223#endif
1224
1225 storeBlock<bfloat16, Packet8bf, 4>(blockB + ri, block);
1226 } else {
1227 PacketBlock<Packet8bf, 8> block;
1228
1229 for (int M = 0; M < 8; M++) {
1230 block.packet[M] = rhs2.template loadPacketPartial<Packet8bf>(i + M, 0, 4);
1231 }
1232
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);
1237
1238 const Index size = 16 / sizeof(bfloat16);
1239
1240 for (int M = 0; M < 4; M++) {
1241 pstore<bfloat16>(blockB + ri + (M * size), block.packet[M]);
1242 }
1243 }
1244
1245 ri += 4 * vectorSize;
1246 }
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);
1257 } else {
1258 PacketBlock<Packet8bf, 2> block;
1259
1260 for (int M = 0; M < 2; M++) {
1261 block.packet[M] = rhs2.template loadPacketPartial<Packet8bf>(i + M, 0, 4);
1262 }
1263
1264 block.packet[0] = vec_mergeh(block.packet[0].m_val, block.packet[1].m_val);
1265
1266 pstore<bfloat16>(blockB + ri, block.packet[0]);
1267 }
1268
1269 ri += 4 * 2;
1270 }
1271 if (depth & 1) {
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);
1276
1277 ri += 4;
1278 }
1279
1280 EIGEN_IF_CONSTEXPR (PanelMode) ri += 4 * (stride - offset - depth);
1281 }
1282
1283 if (j < cols) {
1284 EIGEN_IF_CONSTEXPR (PanelMode) ri += offset * (cols - j);
1285
1286 Index i = 0;
1287 for (; i + 2 <= depth; i += 2) {
1288 Index k = j;
1289 for (; k < cols; k++) {
1290 blockB[ri + 0] = rhs(i + 0, k);
1291 blockB[ri + 1] = rhs(i + 1, k);
1292 ri += 2;
1293 }
1294 }
1295 if (depth & 1) {
1296 for (; j < cols; j++) {
1297 blockB[ri] = rhs(i, j);
1298 ri += 1;
1299 }
1300 }
1301 }
1302 }
1303};
1304
1305// General template for lhs complex packing, float64 specialization.
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;
1312
1313 for (; i + vectorSize <= depth; i += vectorSize) {
1314 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1315 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i + 0); //[a1 a1i]
1316 cblock.packet[1] = lhs2.template loadPacket<PacketC>(0, i + 1); //[b1 b1i]
1317
1318 cblock.packet[2] = lhs2.template loadPacket<PacketC>(1, i + 0); //[a2 a2i]
1319 cblock.packet[3] = lhs2.template loadPacket<PacketC>(1, i + 1); //[b2 b2i]
1320
1321 blockr.packet[0] = vec_mergeh(cblock.packet[0].v, cblock.packet[2].v); //[a1 a2]
1322 blockr.packet[1] = vec_mergeh(cblock.packet[1].v, cblock.packet[3].v); //[b1 b2]
1323
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);
1326 } else {
1327 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i); //[a1 a1i]
1328 cblock.packet[1] = lhs2.template loadPacket<PacketC>(1, i); //[a2 a2i]
1329
1330 cblock.packet[2] = lhs2.template loadPacket<PacketC>(0, i + 1); //[b1 b1i]
1331 cblock.packet[3] = lhs2.template loadPacket<PacketC>(1, i + 1); //[b2 b2i]
1332
1333 blockr.packet[0] = vec_mergeh(cblock.packet[0].v, cblock.packet[1].v); //[a1 a2]
1334 blockr.packet[1] = vec_mergeh(cblock.packet[2].v, cblock.packet[3].v); //[b1 b2]
1335
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);
1338 }
1339
1340 EIGEN_IF_CONSTEXPR (Conjugate) {
1341 blocki.packet[0] = -blocki.packet[0];
1342 blocki.packet[1] = -blocki.packet[1];
1343 }
1344
1345 storeBlock<double, Packet, 2>(blockAt + rir, blockr);
1346 storeBlock<double, Packet, 2>(blockAt + rii, blocki);
1347
1348 rir += 2 * vectorSize;
1349 rii += 2 * vectorSize;
1350 }
1351 }
1352
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);
1359 Index j = 0;
1360
1361 for (; j + vectorSize <= rows; j += vectorSize) {
1362 const DataMapper lhs2 = lhs.getSubMapper(j, 0);
1363 Index i = 0;
1364
1365 rii = rir + vectorDelta;
1366
1367 dhs_ccopy(blockAt, lhs2, i, rir, rii, depth, vectorSize);
1368
1369 for (; i < depth; i++) {
1370 PacketBlock<Packet, 1> blockr, blocki;
1371 PacketBlock<PacketC, 2> cblock;
1372
1373 cblock.packet[0] = lhs2.template loadPacket<PacketC>(0, i);
1374 cblock.packet[1] = lhs2.template loadPacket<PacketC>(1, i);
1375
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);
1378
1379 EIGEN_IF_CONSTEXPR (Conjugate) {
1380 blocki.packet[0] = -blocki.packet[0];
1381 }
1382
1383 pstore<double>(blockAt + rir, blockr.packet[0]);
1384 pstore<double>(blockAt + rii, blocki.packet[0]);
1385
1386 rir += vectorSize;
1387 rii += vectorSize;
1388 }
1389
1390 rir += ((PanelMode) ? (vectorSize * (2 * stride - depth)) : vectorDelta);
1391 }
1392
1393 if (j < rows) {
1394 EIGEN_IF_CONSTEXPR (PanelMode) rir += (offset * (rows - j - vectorSize));
1395 rii = rir + (((PanelMode) ? stride : depth) * (rows - j));
1396
1397 for (Index i = 0; i < depth; i++) {
1398 Index k = j;
1399 for (; k < rows; k++) {
1400 blockAt[rir] = lhs(k, i).real();
1401
1402 EIGEN_IF_CONSTEXPR (Conjugate)
1403 blockAt[rii] = -lhs(k, i).imag();
1404 else
1405 blockAt[rii] = lhs(k, i).imag();
1406
1407 rir += 1;
1408 rii += 1;
1409 }
1410 }
1411 }
1412 }
1413};
1414
1415// General template for rhs complex packing, float64 specialization.
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;
1423
1424 bload<DataMapper, PacketC, 2, ColMajor, false, 4>(cblock, rhs2, i, 0);
1425
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);
1428
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);
1431
1432 EIGEN_IF_CONSTEXPR (Conjugate) {
1433 blocki.packet[0] = -blocki.packet[0];
1434 blocki.packet[1] = -blocki.packet[1];
1435 }
1436
1437 storeBlock<double, Packet, 2>(blockBt + rir, blockr);
1438 storeBlock<double, Packet, 2>(blockBt + rii, blocki);
1439
1440 rir += 2 * vectorSize;
1441 rii += 2 * vectorSize;
1442 }
1443 }
1444
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);
1451 Index j = 0;
1452
1453 for (; j + 2 * vectorSize <= cols; j += 2 * vectorSize) {
1454 const DataMapper rhs2 = rhs.getSubMapper(0, j);
1455 Index i = 0;
1456
1457 rii = rir + vectorDelta;
1458
1459 dhs_ccopy(blockBt, rhs2, i, rir, rii, depth, vectorSize);
1460
1461 rir += ((PanelMode) ? (2 * vectorSize * (2 * stride - depth)) : vectorDelta);
1462 }
1463
1464 EIGEN_IF_CONSTEXPR (PanelMode) rir -= (offset * (2 * vectorSize - 1));
1465
1466 for (; j < cols; j++) {
1467 const DataMapper rhs2 = rhs.getSubMapper(0, j);
1468 rii = rir + ((PanelMode) ? stride : depth);
1469
1470 for (Index i = 0; i < depth; i++) {
1471 blockBt[rir] = rhs2(i, 0).real();
1472
1473 EIGEN_IF_CONSTEXPR (Conjugate)
1474 blockBt[rii] = -rhs2(i, 0).imag();
1475 else
1476 blockBt[rii] = rhs2(i, 0).imag();
1477
1478 rir += 1;
1479 rii += 1;
1480 }
1481
1482 rir += ((PanelMode) ? (2 * stride - depth) : depth);
1483 }
1484 }
1485};
1486
1487/**************
1488 * GEMM utils *
1489 **************/
1490
1491// 512-bits rank1-update of acc. It can either positive or negative accumulate (useful for complex gemm).
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]);
1497 }
1498 } else {
1499 for (int M = 0; M < N; M++) {
1500 acc->packet[M] = vec_madd(lhsV, rhsV[M], acc->packet[M]);
1501 }
1502 }
1503}
1504
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) {
1507 Packet lhsV;
1508 EIGEN_IF_CONSTEXPR (remaining_rows > 0) {
1509 lhsV = ploadu_partial<Packet>(lhs, remaining_rows);
1510 } else {
1511 lhsV = ploadLhs<Packet>(lhs);
1512 }
1513
1514 pger_common<Packet, NegativeAccumulate, N>(acc, lhsV, rhsV);
1515}
1516
1517// 512-bits rank1-update of complex acc. It takes decoupled accumulators as entries. It also takes care of mixed types
1518// real * complex and complex * real.
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);
1526 } else {
1527 EIGEN_IF_CONSTEXPR (!RhsIsReal) {
1528 pger_common<Packet, ConjugateLhs == ConjugateRhs, N>(accReal, lhsVi, rhsVi);
1529 pger_common<Packet, ConjugateRhs, N>(accImag, lhsV, rhsVi);
1530 } else {
1531 EIGEN_UNUSED_VARIABLE(rhsVi);
1532 }
1533 pger_common<Packet, ConjugateLhs, N>(accImag, lhsVi, rhsV);
1534 }
1535}
1536
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) {
1541 Packet lhsV;
1542 EIGEN_IF_CONSTEXPR (remaining_rows > 0) {
1543 lhsV = ploadu_partial<Packet>(lhs_ptr, remaining_rows);
1544 } else {
1545 lhsV = ploadLhs<Packet>(lhs_ptr);
1546 }
1547 Packet lhsVi;
1548 EIGEN_IF_CONSTEXPR (!LhsIsReal) {
1549 EIGEN_IF_CONSTEXPR (remaining_rows > 0) {
1550 lhsVi = ploadu_partial<Packet>(lhs_ptr_imag, remaining_rows);
1551 } else {
1552 lhsVi = ploadLhs<Packet>(lhs_ptr_imag);
1553 }
1554 } else {
1555 EIGEN_UNUSED_VARIABLE(lhs_ptr_imag);
1556 }
1557
1558 pgerc_common<N, Packet, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>(accReal, accImag, lhsV, lhsVi, rhsV, rhsVi);
1559}
1560
1561template <typename Packet>
1562EIGEN_ALWAYS_INLINE Packet ploadLhs(const __UNPACK_TYPE__(Packet) * lhs) {
1563 return ploadu<Packet>(lhs);
1564}
1565
1566// Zero the accumulator on PacketBlock.
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);
1571 }
1572}
1573
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);
1579 }
1580}
1581
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);
1586 }
1587}
1588
1589// Complex version of PacketBlock scaling.
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);
1597 } else {
1598 EIGEN_UNUSED_VARIABLE(pMask);
1599 }
1600
1601 bscalec_common<Packet, N>(cReal, aReal, bReal);
1602
1603 bscalec_common<Packet, N>(cImag, aImag, bReal);
1604
1605 pger_common<Packet, true, N>(&cReal, bImag, aImag.packet);
1606
1607 pger_common<Packet, false, N>(&cImag, bImag, aReal.packet);
1608}
1609
1610// Load a PacketBlock, the N parameters make tuning gemm easier so we can add more accumulators as needed.
1611//
1612// full = operate (load) on the entire PacketBlock or only half
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,
1615 Index col) {
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);
1619 }
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);
1623 }
1624 }
1625 } else {
1626 for (int M = 0; M < N; M++) {
1627 acc.packet[M] = res.template loadPacket<Packet>(row, col + M);
1628 }
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);
1632 }
1633 }
1634 }
1635}
1636
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]);
1641 }
1642}
1643
1644// When Complex && full, the first N packets are full loads and `elements` counts only the second half at row + accCols.
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,
1647 Index elements) {
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);
1652 }
1653 } else {
1654 for (Index M = 0; M < N; M++) {
1655 acc.packet[M] = res.template loadPacketPartial<Packet>(row, M, elements);
1656 }
1657 }
1658}
1659
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);
1664 }
1665}
1666
1667#ifdef _ARCH_PWR10
1668#define USE_P10_AND_PVIPR2_0 (EIGEN_COMP_LLVM || (__GNUC__ >= 11))
1669#else
1670#define USE_P10_AND_PVIPR2_0 0
1671#endif
1672
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}};
1675#endif
1676
1677template <typename Packet>
1678EIGEN_ALWAYS_INLINE Packet bmask(const Index remaining_rows) {
1679#if USE_P10_AND_PVIPR2_0
1680#ifdef _BIG_ENDIAN
1681 return Packet(vec_reve(vec_genwm((1 << remaining_rows) - 1)));
1682#else
1683 return Packet(vec_genwm((1 << remaining_rows) - 1));
1684#endif
1685#else
1686 return Packet(mask4[remaining_rows]);
1687#endif
1688}
1689
1690template <>
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));
1694#ifdef _BIG_ENDIAN
1695 return preverse(mask2);
1696#else
1697 return mask2;
1698#endif
1699#else
1700 Packet2l ret = {-remaining_rows, 0};
1701 return Packet2d(ret);
1702#endif
1703}
1704
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]);
1709 }
1710}
1711
1712// Scale the PacketBlock vectors by alpha.
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);
1718 } else {
1719 EIGEN_UNUSED_VARIABLE(pMask);
1720 }
1721
1722 bscale<Packet, N>(acc, accZ, pAlpha);
1723}
1724
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,
1728 Packet& a3) {
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);
1736 } else {
1737 EIGEN_IF_CONSTEXPR (N > 1) {
1738 a1 = pset1<Packet>(ap1[0]);
1739 } else {
1740 EIGEN_UNUSED_VARIABLE(a1);
1741 EIGEN_UNUSED_VARIABLE(ap1);
1742 }
1743 EIGEN_IF_CONSTEXPR (N > 2) {
1744 a2 = pset1<Packet>(ap2[0]);
1745 } else {
1746 EIGEN_UNUSED_VARIABLE(a2);
1747 EIGEN_UNUSED_VARIABLE(ap2);
1748 }
1749 }
1750}
1751
1752template <>
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);
1756}
1757
1758template <>
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);
1762}
1763
1764template <>
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);
1773}
1774
1775// Grab two decouples real/imaginary PacketBlocks and return two coupled (real/imaginary pairs) PacketBlocks.
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]);
1781 }
1782
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]);
1786 }
1787 }
1788}
1789
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);
1795
1796 for (int M = 0; M < N; M++) {
1797 acc1.packet[M] = padd<Packetc>(tRes.packet[M], acc1.packet[M]);
1798 }
1799
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]);
1803 }
1804 }
1805}
1806
1807// PEEL loop factor.
1808#define PEEL 7
1809#define PEEL_ROW 7
1810
1811#define MICRO_UNROLL(func) func(0) func(1) func(2) func(3) func(4) func(5) func(6) func(7)
1812
1813#define MICRO_NORMAL_ROWS accRows == quad_traits<Scalar>::rows || accRows == 1
1814
1815#define MICRO_NEW_ROWS ((MICRO_NORMAL_ROWS) ? accRows : 1)
1816
1817#define MICRO_RHS(ptr, N) rhs_##ptr##N
1818
1819#define MICRO_ZERO_PEEL(peel) \
1820 EIGEN_IF_CONSTEXPR ((PEEL_ROW > peel) && (peel != 0)) { \
1821 bsetzero<Packet, accRows>(accZero##peel); \
1822 } else { \
1823 EIGEN_UNUSED_VARIABLE(accZero##peel); \
1824 }
1825
1826#define MICRO_ADD(ptr, N) \
1827 EIGEN_IF_CONSTEXPR (MICRO_NORMAL_ROWS) { \
1828 MICRO_RHS(ptr, 0) += (accRows * N); \
1829 } else { \
1830 MICRO_RHS(ptr, 0) += N; \
1831 MICRO_RHS(ptr, 1) += N; \
1832 EIGEN_IF_CONSTEXPR (accRows == 3) { \
1833 MICRO_RHS(ptr, 2) += N; \
1834 } \
1835 }
1836
1837#define MICRO_ADD_ROWS(N) MICRO_ADD(ptr, N)
1838
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]); \
1843 } else { \
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]); \
1846 }
1847
1848#define MICRO_BROADCAST(peel) MICRO_BROADCAST1(peel, ptr, rhsV, true)
1849
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], \
1852 rhsV[2], rhsV[3]);
1853
1854#define MICRO_BROADCAST_EXTRA \
1855 Packet rhsV[4]; \
1856 MICRO_BROADCAST_EXTRA1(ptr, rhsV, true) \
1857 MICRO_ADD_ROWS(1)
1858
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)); \
1864 } else { \
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; \
1868 } else { \
1869 EIGEN_UNUSED_VARIABLE(MICRO_RHS(ptr, 2)); \
1870 } \
1871 }
1872
1873#define MICRO_SRC2_PTR MICRO_SRC2(ptr, strideB, 0)
1874
1875#define MICRO_ZERO_PEEL_ROW MICRO_UNROLL(MICRO_ZERO_PEEL)
1876
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); \
1881 } else { \
1882 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
1883 }
1884
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)
1890
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]; \
1895 } \
1896 }
1897
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)
1902
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)); \
1909 } \
1910 }
1911
1912#define MICRO_PREFETCHN(N) MICRO_PREFETCHN1(ptr, N)
1913
1914#define MICRO_COMPLEX_PREFETCHN(N) \
1915 MICRO_PREFETCHN1(ptr_real, N); \
1916 EIGEN_IF_CONSTEXPR (!RhsIsReal) { \
1917 MICRO_PREFETCHN1(ptr_imag, N); \
1918 }
1919
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;
1926}
1927
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;
1937
1938 MICRO_SRC2_PTR
1939 bsetzero<Packet, accRows>(accZero0);
1940
1941 const Index peel_depth = depth - (accCols - remaining_rows);
1942 Index k = 0;
1943 if (peel_depth >= PEEL_ROW) {
1944 MICRO_ZERO_PEEL_ROW
1945 do {
1946 MICRO_PREFETCHN(accRows)
1947 EIGEN_POWER_PREFETCH(lhs_ptr);
1948 MICRO_WORK_PEEL_ROW
1949 } while ((k += PEEL_ROW) + PEEL_ROW <= peel_depth);
1950 MICRO_ADD_PEEL_ROW
1951 }
1952 for (; k < peel_depth; k++) {
1953 MICRO_EXTRA_ROW<Scalar, Packet, accRows, remaining_rows>(lhs_ptr, rhs_ptr0, rhs_ptr1, rhs_ptr2, accZero0);
1954 }
1955 for (; k < depth; k++) {
1956 MICRO_EXTRA_ROW<Scalar, Packet, accRows, remaining_rows, remaining_rows>(lhs_ptr, rhs_ptr0, rhs_ptr1, rhs_ptr2,
1957 accZero0);
1958 }
1959
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);
1965}
1966
1967#define MICRO_EXTRA(MICRO_EXTRA_UNROLL, value, is_col) \
1968 switch (value) { \
1969 default: \
1970 MICRO_EXTRA_UNROLL(1) \
1971 break; \
1972 case 2: \
1973 EIGEN_IF_CONSTEXPR (is_col || (sizeof(Scalar) == sizeof(float))) { \
1974 MICRO_EXTRA_UNROLL(2) \
1975 } \
1976 break; \
1977 case 3: \
1978 EIGEN_IF_CONSTEXPR (is_col || (sizeof(Scalar) == sizeof(float))) { \
1979 MICRO_EXTRA_UNROLL(3) \
1980 } \
1981 break; \
1982 }
1983
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);
1987
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)
1993}
1994
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)
1998
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); \
2002 }
2003
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) \
2009 } else { \
2010 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
2011 }
2012
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)
2017
2018#define MICRO_UNROLL_TYPE_ONE(M, func, func1, func2) \
2019 Packet rhsV0[M]; \
2020 func(func1, func2, 0)
2021
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)
2025
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)
2029
2030#define MICRO_ONE_PEEL4 MICRO_UNROLL_TYPE(MICRO_UNROLL_TYPE_PEEL, PEEL)
2031
2032#define MICRO_ONE4 MICRO_UNROLL_TYPE(MICRO_UNROLL_TYPE_ONE, 1)
2033
2034#define MICRO_ONE4_PARTIAL MICRO_UNROLL_TYPE_PARTIAL(MICRO_UNROLL_TYPE_ONE, 1)
2035
2036#define MICRO_DST_PTR_ONE(iter) \
2037 EIGEN_IF_CONSTEXPR (unroll_factor > iter) { \
2038 bsetzero<Packet, accRows>(accZero##iter); \
2039 } else { \
2040 EIGEN_UNUSED_VARIABLE(accZero##iter); \
2041 }
2042
2043#define MICRO_DST_PTR MICRO_UNROLL(MICRO_DST_PTR_ONE)
2044
2045#define MICRO_SRC_PTR MICRO_UNROLL(MICRO_SRC_PTR_ONE)
2046
2047#define MICRO_PREFETCH MICRO_UNROLL(MICRO_PREFETCH_ONE)
2048
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); \
2055 } else { \
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); \
2059 } \
2060 }
2061
2062#define MICRO_STORE MICRO_UNROLL(MICRO_STORE_ONE)
2063
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;
2074
2075 MICRO_SRC2_PTR
2076 MICRO_SRC_PTR
2077 MICRO_DST_PTR
2078
2079 const Index peel_depth = full ? depth : (depth - (accCols - accCols2));
2080 Index k = 0;
2081 for (; k + PEEL <= peel_depth; k += PEEL) {
2082 MICRO_PREFETCHN(accRows)
2083 MICRO_PREFETCH
2084 MICRO_ONE_PEEL4
2085 }
2086 for (; k < peel_depth; k++) {
2087 MICRO_ONE4
2088 }
2089 EIGEN_IF_CONSTEXPR (!full) {
2090 for (; k < depth; k++) {
2091 MICRO_ONE4_PARTIAL
2092 }
2093 }
2094 MICRO_STORE
2095
2096 MICRO_UPDATE
2097}
2098
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); \
2102 if (M) return;
2103
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);
2109
2110 const Scalar* rhs_base = blockB + col * strideB + MICRO_NEW_ROWS * offsetB;
2111 const Scalar* lhs_base = blockA + accCols * offsetA;
2112 Index row = 0;
2113
2114#define MAX_UNROLL 7
2115 while (row + MAX_UNROLL * accCols <= rows) {
2116 MICRO_UNROLL_ITER2(MAX_UNROLL, 0);
2117 }
2118 switch ((rows - row) / accCols) {
2119#if MAX_UNROLL > 7
2120 case 7:
2121 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 7)
2122 break;
2123#endif
2124#if MAX_UNROLL > 6
2125 case 6:
2126 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 6)
2127 break;
2128#endif
2129#if MAX_UNROLL > 5
2130 case 5:
2131 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 5)
2132 break;
2133#endif
2134#if MAX_UNROLL > 4
2135 case 4:
2136 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 4)
2137 break;
2138#endif
2139#if MAX_UNROLL > 3
2140 case 3:
2141 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 3)
2142 break;
2143#endif
2144#if MAX_UNROLL > 2
2145 case 2:
2146 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 2)
2147 break;
2148#endif
2149#if MAX_UNROLL > 1
2150 case 1:
2151 MICRO_UNROLL_ITER(MICRO_UNROLL_ITER2, 1)
2152 break;
2153#endif
2154 default:
2155 break;
2156 }
2157#undef MAX_UNROLL
2158
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);
2162 }
2163}
2164
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);
2168
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)
2175}
2176
2177/****************
2178 * GEMM kernels *
2179 * **************/
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,
2184 Index offsetB) {
2185 const Index remaining_rows = rows % accCols;
2186
2187 if (strideA == -1) strideA = depth;
2188 if (strideB == -1) strideB = depth;
2189
2190 const Packet pAlpha = pset1<Packet>(alpha);
2191 const Packet pMask = bmask<Packet>(remaining_rows);
2192
2193 Index col = 0;
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);
2197 }
2198
2199 if (col != cols) {
2200 gemm_extra_cols<Scalar, Packet, DataMapper, accCols>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB,
2201 col, rows, cols, remaining_rows, pAlpha, pMask);
2202 }
2203}
2204
2205#define accColsC (accCols / 2)
2206#define advanceRows ((LhsIsReal) ? 1 : 2)
2207#define advanceCols ((RhsIsReal) ? 1 : 2)
2208
2209// PEEL_COMPLEX loop factor.
2210#define PEEL_COMPLEX 3
2211#define PEEL_COMPLEX_ROW 3
2212
2213#define MICRO_COMPLEX_UNROLL(func) func(0) func(1) func(2) func(3)
2214
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); \
2219 } else { \
2220 EIGEN_UNUSED_VARIABLE(accReal##peel); \
2221 EIGEN_UNUSED_VARIABLE(accImag##peel); \
2222 }
2223
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)); \
2232 }
2233
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) \
2238 } else { \
2239 EIGEN_UNUSED_VARIABLE(rhsVi##peel); \
2240 }
2241
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) \
2247 } else { \
2248 EIGEN_UNUSED_VARIABLE(rhsVi); \
2249 } \
2250 MICRO_COMPLEX_ADD_ROWS(1, true)
2251
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) \
2257 } else { \
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)); \
2261 }
2262
2263#define MICRO_COMPLEX_ZERO_PEEL_ROW MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_ZERO_PEEL)
2264
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); \
2271 } else { \
2272 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
2273 EIGEN_UNUSED_VARIABLE(rhsVi##peel); \
2274 }
2275
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); \
2280 else \
2281 EIGEN_UNUSED_VARIABLE(lhs_ptr_imag);
2282
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)
2289
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]; \
2295 } \
2296 }
2297
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)
2300
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)
2313}
2314
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;
2329 else
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;
2335
2336 MICRO_COMPLEX_SRC2_PTR
2337
2338 bsetzero<Packet, accRows>(accReal0);
2339 bsetzero<Packet, accRows>(accImag0);
2340
2341 const Index peel_depth = depth - (accCols - remaining_rows);
2342 Index k = 0;
2343 if (peel_depth >= PEEL_COMPLEX_ROW) {
2344 MICRO_COMPLEX_ZERO_PEEL_ROW
2345 do {
2346 MICRO_COMPLEX_PREFETCHN(accRows)
2347 EIGEN_POWER_PREFETCH(lhs_ptr_real);
2348 EIGEN_IF_CONSTEXPR (!LhsIsReal) {
2349 EIGEN_POWER_PREFETCH(lhs_ptr_imag);
2350 }
2351 MICRO_COMPLEX_WORK_PEEL_ROW
2352 } while ((k += PEEL_COMPLEX_ROW) + PEEL_COMPLEX_ROW <= peel_depth);
2353 MICRO_COMPLEX_ADD_PEEL_ROW
2354 }
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);
2359 }
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);
2364 }
2365
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);
2371 } else {
2372 bload<DataMapper, Packetc, accColsC, ColMajor, true, accRows, full>(tRes, res, row, 0);
2373 }
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);
2378 } else {
2379 bstore<DataMapper, Packetc, accRows>(acc0, res, row + 0);
2380 }
2381 EIGEN_IF_CONSTEXPR (full) {
2382 EIGEN_IF_CONSTEXPR (odd) {
2383 bstore_partial<DataMapper, Packetc, accRows>(acc1, res, row + accColsC, 1);
2384 } else {
2385 bstore<DataMapper, Packetc, accRows>(acc1, res, row + accColsC);
2386 }
2387 }
2388}
2389
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);
2394
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)
2402}
2403
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)
2407
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); \
2412 }
2413
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) \
2420 } else { \
2421 EIGEN_UNUSED_VARIABLE(rhsV##peel); \
2422 EIGEN_UNUSED_VARIABLE(rhsVi##peel); \
2423 }
2424
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)
2429
2430#define MICRO_COMPLEX_UNROLL_TYPE_ONE(M, func, func1, func2) \
2431 Packet rhsV0[M], rhsVi0[M]; \
2432 func(func1, func2, 0)
2433
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)
2437
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)
2441
2442#define MICRO_COMPLEX_ONE_PEEL4 MICRO_COMPLEX_UNROLL_TYPE(MICRO_COMPLEX_UNROLL_TYPE_PEEL, PEEL_COMPLEX)
2443
2444#define MICRO_COMPLEX_ONE4 MICRO_COMPLEX_UNROLL_TYPE(MICRO_COMPLEX_UNROLL_TYPE_ONE, 1)
2445
2446#define MICRO_COMPLEX_ONE4_PARTIAL MICRO_COMPLEX_UNROLL_TYPE_PARTIAL(MICRO_COMPLEX_UNROLL_TYPE_ONE, 1)
2447
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); \
2452 } else { \
2453 EIGEN_UNUSED_VARIABLE(accReal##iter); \
2454 EIGEN_UNUSED_VARIABLE(accImag##iter); \
2455 }
2456
2457#define MICRO_COMPLEX_DST_PTR MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_DST_PTR_ONE)
2458
2459#define MICRO_COMPLEX_SRC_PTR MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_SRC_PTR_ONE)
2460
2461#define MICRO_COMPLEX_PREFETCH MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_PREFETCH_ONE)
2462
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); \
2469 } else { \
2470 bload<DataMapper, Packetc, accColsC, ColMajor, true, accRows, full>(tRes, res, row + iter * accCols, 0); \
2471 } \
2472 bscalec<Packet, accRows, !(MICRO_NORMAL(iter))>(accReal##iter, accImag##iter, pAlphaReal, pAlphaImag, taccReal, \
2473 taccImag, pMask); \
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); \
2477 } else { \
2478 bstore<DataMapper, Packetc, accRows>(acc0, res, row + iter * accCols + 0); \
2479 } \
2480 EIGEN_IF_CONSTEXPR (full) { \
2481 EIGEN_IF_CONSTEXPR (odd) { \
2482 bstore_partial<DataMapper, Packetc, accRows>(acc1, res, row + iter * accCols + accColsC, 1); \
2483 } else { \
2484 bstore<DataMapper, Packetc, accRows>(acc1, res, row + iter * accCols + accColsC); \
2485 } \
2486 } \
2487 }
2488
2489#define MICRO_COMPLEX_STORE MICRO_COMPLEX_UNROLL(MICRO_COMPLEX_STORE_ONE)
2490
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;
2510
2511 MICRO_COMPLEX_SRC2_PTR
2512 MICRO_COMPLEX_SRC_PTR
2513 MICRO_COMPLEX_DST_PTR
2514
2515 const Index peel_depth = depth - (accCols - accCols2);
2516 Index k = 0;
2517 for (; k + PEEL_COMPLEX <= peel_depth; k += PEEL_COMPLEX) {
2518 MICRO_COMPLEX_PREFETCHN(accRows)
2519 MICRO_COMPLEX_PREFETCH
2520 MICRO_COMPLEX_ONE_PEEL4
2521 }
2522 for (; k < peel_depth; k++) {
2523 MICRO_COMPLEX_ONE4
2524 }
2525 EIGEN_IF_CONSTEXPR (accCols != accCols2) {
2526 for (; k < depth; k++) {
2527 MICRO_COMPLEX_ONE4_PARTIAL
2528 }
2529 }
2530 MICRO_COMPLEX_STORE
2531
2532 MICRO_COMPLEX_UPDATE
2533}
2534
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); \
2539 if (M) return;
2540
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);
2548
2549 const Scalar* rhs_base = blockB + advanceCols * col * strideB + MICRO_NEW_ROWS * offsetB;
2550 const Scalar* lhs_base = blockA + accCols * offsetA;
2551 Index row = 0;
2552
2553#define MAX_COMPLEX_UNROLL 4
2554 while (row + MAX_COMPLEX_UNROLL * accCols <= rows) {
2555 MICRO_COMPLEX_UNROLL_ITER2(MAX_COMPLEX_UNROLL, 0);
2556 }
2557 switch ((rows - row) / accCols) {
2558#if MAX_COMPLEX_UNROLL > 4
2559 case 4:
2560 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 4)
2561 break;
2562#endif
2563#if MAX_COMPLEX_UNROLL > 3
2564 case 3:
2565 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 3)
2566 break;
2567#endif
2568#if MAX_COMPLEX_UNROLL > 2
2569 case 2:
2570 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 2)
2571 break;
2572#endif
2573#if MAX_COMPLEX_UNROLL > 1
2574 case 1:
2575 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_UNROLL_ITER2, 1)
2576 break;
2577#endif
2578 default:
2579 break;
2580 }
2581#undef MAX_COMPLEX_UNROLL
2582
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);
2587 }
2588}
2589
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);
2594
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)
2603}
2604
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;
2612
2613 if (strideA == -1) strideA = depth;
2614 if (strideB == -1) strideB = depth;
2615
2616 const Packet pAlphaReal = pset1<Packet>(alpha.real());
2617 const Packet pAlphaImag = pset1<Packet>(alpha.imag());
2618 const Packet pMask = bmask<Packet>(remaining_rows);
2619
2620 const Scalar* blockA = (Scalar*)blockAc;
2621 const Scalar* blockB = (Scalar*)blockBc;
2622
2623 Index col = 0;
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);
2628 }
2629
2630 if (col != cols) {
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);
2634 }
2635}
2636
2637#undef accColsC
2638#undef advanceCols
2639#undef advanceRows
2640
2641EIGEN_ALWAYS_INLINE bool supportsMMA() {
2642#if defined(EIGEN_ALTIVEC_MMA_ONLY)
2643 return true;
2644#elif defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH) && defined(__BUILTIN_CPU_SUPPORTS__)
2645 return __builtin_cpu_supports("arch_3_1") && __builtin_cpu_supports("mma");
2646#else
2647 return false; // No dynamic dispatch for LLVM or older GCC
2648#endif
2649}
2650
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);
2654}
2655
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);
2660 } else {
2661 pstoreu(result, result_block);
2662 }
2663 result += rows;
2664}
2665
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) {
2669 Index x = 0;
2670 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
2671 do {
2672 Packet4f result_block = loadAndMultiplyF32(acc[x], pAlpha, result);
2673 storeF32<lhsExtraRows>(result, result_block, rows, extra_rows);
2674 } while (++x < extra_cols);
2675 } else {
2676 Packet4f result_block[4];
2677 float* result2 = result;
2678 do {
2679 result_block[x] = loadAndMultiplyF32(acc[x], pAlpha, result);
2680 result += rows;
2681 } while (++x < 4);
2682 x = 0;
2683 do {
2684 storeF32<lhsExtraRows>(result2, result_block[x], rows, extra_rows);
2685 } while (++x < 4);
2686 }
2687}
2688
2689EIGEN_ALWAYS_INLINE Packet4f oneConvertBF16Hi(const Packet8us& data) {
2690 Packet8us z = pset1<Packet8us>(0);
2691#ifdef _BIG_ENDIAN
2692 return reinterpret_cast<Packet4f>(vec_mergeh(data, z));
2693#else
2694 return reinterpret_cast<Packet4f>(vec_mergeh(z, data));
2695#endif
2696}
2697
2698EIGEN_ALWAYS_INLINE Packet4f oneConvertBF16Lo(const Packet8us& data) {
2699 Packet8us z = pset1<Packet8us>(0);
2700#ifdef _BIG_ENDIAN
2701 return reinterpret_cast<Packet4f>(vec_mergel(data, z));
2702#else
2703 return reinterpret_cast<Packet4f>(vec_mergel(z, data));
2704#endif
2705}
2706
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));
2715 }
2716 }
2717}
2718
2719template <Index N>
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);
2724 }
2725 EIGEN_IF_CONSTEXPR (N >= 32) {
2726 storeConvertTwoBF16<N, 2>(to + 16, block);
2727 storeConvertTwoBF16<N, 3>(to + 24, block);
2728 }
2729}
2730
2731template <bool non_unit_stride, 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);
2735 } else {
2736 return ploadu<Packet8bf>(src + delta);
2737 }
2738}
2739
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};
2744
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};
2749
2750EIGEN_ALWAYS_INLINE Packet4f oneConvertBF16Perm(const Packet8us& data, const Packet16uc& mask) {
2751 Packet8us z = pset1<Packet8us>(0);
2752#ifdef _BIG_ENDIAN
2753 return reinterpret_cast<Packet4f>(vec_perm(data, z, mask));
2754#else
2755 return reinterpret_cast<Packet4f>(vec_perm(z, data, mask));
2756#endif
2757}
2758
2759template <bool lhsExtraRows, bool odd, Index size>
2760EIGEN_ALWAYS_INLINE void convertArrayPointerBF16toF32DupOne(float* result, Index rows, const bfloat16* src,
2761 Index extra_rows) {
2762 Packet4f dup[4 * 4];
2763 Packet8bf data[4];
2764
2765 for (Index i = 0; i < size; i++) {
2766 data[i] = ploadu<Packet8bf>(src + rows * i);
2767 }
2768
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);
2774 }
2775
2776 for (Index j = 0; j < 4 * size; j += 4) {
2777 EIGEN_IF_CONSTEXPR (lhsExtraRows) {
2778 Packet4f z = pset1<Packet4f>(float(0));
2779 Index i = 0;
2780 do {
2781 pstoreu(result + (j + i) * 4, dup[j + i]);
2782 } while (++i < extra_rows);
2783 do {
2784 pstoreu(result + (j + i) * 4, z);
2785 } while (++i < 4);
2786 } else {
2787 for (Index i = 0; i < 4; i++) {
2788 pstoreu(result + (j + i) * 4, dup[j + i]);
2789 }
2790 }
2791 }
2792}
2793
2794template <bool lhsExtraRows>
2795EIGEN_ALWAYS_INLINE void convertArrayPointerBF16toF32Dup(float* result, Index cols, Index rows, const bfloat16* src,
2796 Index delta, Index extra_rows) {
2797 Index col = 0;
2798 src += delta * 2;
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);
2801 }
2802 for (; col + 2 <= cols; col += 2, result += 4 * 4, src += rows) {
2803 convertArrayPointerBF16toF32DupOne<lhsExtraRows, false, 1>(result, rows, src, extra_rows);
2804 }
2805 if (cols & 1) {
2806 convertArrayPointerBF16toF32DupOne<lhsExtraRows, true, 1>(result, rows, src - delta, extra_rows);
2807 }
2808}
2809
2810template <const Index size, bool non_unit_stride>
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);
2817 } else {
2818 r32.packet[0] = loadBF16fromResult<non_unit_stride, 0>(src, resInc);
2819 }
2820 EIGEN_IF_CONSTEXPR (size >= 16) {
2821 r32.packet[1] = loadBF16fromResult<non_unit_stride, 8>(src, resInc);
2822 }
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);
2826 }
2827 storeConvertBlockBF16<size>(result + i, r32, rows & 3);
2828 i += count;
2829 if (i < rows) src += count * resInc;
2830 EIGEN_IF_CONSTEXPR (size != 32) break;
2831 }
2832}
2833
2834template <bool non_unit_stride>
2835EIGEN_ALWAYS_INLINE void convertArrayPointerBF16toF32(float* result, Index cols, Index rows, bfloat16* src,
2836 Index resInc) {
2837 for (Index col = 0; col < cols; col++, result += rows) {
2838 Index i = 0;
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;
2846 }
2847}
2848
2849template <Index num_acc, Index size = 4>
2850EIGEN_ALWAYS_INLINE void zeroAccumulators(Packet4f (&acc)[num_acc][size]) {
2851 Packet4f z = pset1<Packet4f>(float(0));
2852
2853 for (Index k = 0; k < num_acc; k++) {
2854 for (Index j = 0; j < size; j++) {
2855 acc[k][j] = z;
2856 }
2857 }
2858}
2859
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));
2872 }
2873}
2874
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];
2881 }
2882 }
2883 }
2884}
2885
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);
2891
2892 constexpr Index real_rhs = ((num_rhs / 2) - (rhsExtraCols ? 1 : 0));
2893 Index k = 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);
2896 }
2897 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
2898 storeResults<rhsExtraCols, lhsExtraRows>(acc[k], rows, pAlpha, result, extra_cols, extra_rows);
2899 }
2900}
2901
2902template <bool zero>
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);
2909 } else {
2910 dhs1 = ploadu<Packet4f>(block + strideB * i + 4);
2911 }
2912}
2913
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];
2919
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]);
2923 }
2924 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
2925 loadTwoRhsFloat32<zero>(indexB + k * extra_cols - offsetB, strideB, real_rhs, rhs[real_rhs + 0], rhs[real_rhs + 1]);
2926 }
2927
2928 indexA += 2 * k * 4;
2929 for (Index j = 0; j < num_lhs; j++) {
2930 lhs[j] = ploadu<Packet4f>(indexA + j * 4);
2931 }
2932
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]);
2936 }
2937 }
2938}
2939
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;
2945
2946 Packet4f acc[num_acc][4];
2947
2948 zeroAccumulators<num_acc>(acc);
2949
2950 Index k;
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);
2953 }
2954 if (depth & 1) {
2955 KLoop<num_acc, true, rhsExtraCols, num_rhs>(indexA, indexB, acc, strideB, k, offsetB, extra_cols);
2956 }
2957
2958 outputResultsVSX<num_acc, rhsExtraCols, lhsExtraRows, num_rhs>(acc, rows, pAlpha, result, extra_cols, extra_rows);
2959}
2960
2961// No more than 4 (uses 2X the accumulators or 8X the number of VSX registers)
2962#define MAX_BFLOAT16_ACC_VSX 4
2963
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); // each accumulator has 4 elements
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);
2971
2972 do {
2973 colVSXLoopBodyIter<num_acc * 2, rhsExtraCols, lhsExtraRows>(depth, rows, pAlpha, indexA, indexB, strideB, offsetB,
2974 result, extra_cols, extra_rows);
2975
2976 indexB += strideB * (num_acc * 2);
2977 result += rows * step;
2978 } while (multiIters && (step <= cols - (col += step)));
2979}
2980
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,
2984 float* result) {
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);
2988 }
2989}
2990
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) {
2995 case 3:
2996 colVSXLoopBodyExtraN<3, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
2997 offsetB, result);
2998 break;
2999 case 2:
3000 colVSXLoopBodyExtraN<2, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
3001 offsetB, result);
3002 break;
3003 case 1:
3004 colVSXLoopBodyExtraN<1, rhsExtraCols, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB,
3005 offsetB, result);
3006 break;
3007 default:
3008 EIGEN_IF_CONSTEXPR (rhsExtraCols) {
3009 colVSXLoopBody<1, true, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA, blockB, strideB, offsetB, result);
3010 }
3011 break;
3012 }
3013}
3014
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,
3022 rows & 3);
3023
3024 const float* blockB = blockB2;
3025 float* result = result2 + row;
3026
3027 Index col = 0;
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;
3033 }
3034 if (cols & 3) {
3035 colVSXLoopBodyExtra<true, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA2, blockB, strideB, offsetB,
3036 result);
3037 } else {
3038 colVSXLoopBodyExtra<false, lhsExtraRows>(col, depth, cols, rows, pAlpha, indexA2, blockB, strideB, 0, result);
3039 }
3040 }
3041}
3042
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,
3047 float* result) {
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);
3051 row += size;
3052 indexA += bigSuffix * size / 16;
3053 }
3054}
3055
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);
3064 }
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);
3068 }
3069 storeConvertBlockBF16<size>(result + i, r32, rows & 3);
3070 i += extra;
3071 EIGEN_IF_CONSTEXPR (size != 32) break;
3072 }
3073}
3074
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);
3080 Index i = 0;
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);
3086 }
3087}
3088
3089EIGEN_ALWAYS_INLINE Packet8bf convertF32toBF16VSX(const float* res) {
3090 return F32ToBf16Both(ploadu<Packet4f>(res + 0), ploadu<Packet4f>(res + 4));
3091}
3092
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);
3096 Index row;
3097 float* result2 = result + col * rows;
3098 for (row = 0; row + 8 <= rows; row += 8, result2 += 8) {
3099 // get and save block
3100 PacketBlock<Packet8bf, size> block;
3101 for (Index j = 0; j < size; j++) {
3102 block.packet[j] = convertF32toBF16VSX(result2 + j * rows);
3103 }
3104 res2.template storePacketBlock<Packet8bf, size>(row, 0, block);
3105 }
3106 // extra rows
3107 if (row < rows) {
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);
3111 }
3112 }
3113}
3114
3115template <typename DataMapper>
3116EIGEN_ALWAYS_INLINE void convertArrayF32toBF16VSX(float* result, Index cols, Index rows, const DataMapper& res) {
3117 Index col;
3118 for (col = 0; col + 4 <= cols; col += 4) {
3119 convertArrayF32toBF16ColVSX<DataMapper, 4>(result, col, rows, res);
3120 }
3121 // extra cols
3122 switch (cols - col) {
3123 case 1:
3124 convertArrayF32toBF16ColVSX<DataMapper, 1>(result, col, rows, res);
3125 break;
3126 case 2:
3127 convertArrayF32toBF16ColVSX<DataMapper, 2>(result, col, rows, res);
3128 break;
3129 case 3:
3130 convertArrayF32toBF16ColVSX<DataMapper, 3>(result, col, rows, res);
3131 break;
3132 }
3133}
3134
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);
3140
3141 if (strideA == -1) strideA = depth;
3142 if (strideB == -1) strideB = depth;
3143
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);
3147
3148 convertArrayBF16toF32<DataMapper>(result, cols, rows, res);
3149 convertArrayPointerBF16toF32(indexB2, cols, strideB, const_cast<bfloat16*>(indexB));
3150
3151 Index bigSuffix = 2 * 8 * (strideA - offsetA);
3152 float* indexBF32 = indexB2 + 4 * offsetB;
3153 offsetB *= 3;
3154 strideB *= 2;
3155
3156 Index row = 0;
3157 // LHS (8x16) block
3158 while (row + 16 <= rows) {
3159 calcVSXColLoops<16>(indexA, indexA2, row, depth, cols, rows, pAlpha, indexBF32, strideA, strideB, offsetA, offsetB,
3160 bigSuffix, result);
3161 }
3162 // LHS (8x8) block
3163 calcVSXColLoops<8>(indexA, indexA2, row, depth, cols, rows, pAlpha, indexBF32, strideA, strideB, offsetA, offsetB,
3164 bigSuffix, result);
3165 // LHS (8x4) block
3166 calcVSXColLoops<4>(indexA, indexA2, row, depth, cols, rows, pAlpha, indexBF32, strideA, strideB, offsetA, offsetB,
3167 bigSuffix, result);
3168 // extra rows
3169 if (rows & 3) {
3170 // This index is the beginning of remaining block.
3171 colVSXLoops<4, true>(depth, cols, rows, pAlpha, indexA, indexA2, indexBF32, strideA, strideB, offsetB,
3172 result + row);
3173 }
3174
3175 // Convert back to bfloat16
3176 convertArrayF32toBF16VSX<DataMapper>(result, cols, rows, res);
3177}
3178
3179#undef MAX_BFLOAT16_ACC_VSX
3180
3181#include "MatrixVectorProduct.inc"
3182
3183/************************************
3184 * ppc64le template specializations *
3185 * **********************************/
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);
3189};
3190
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);
3196}
3197
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);
3201};
3202
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);
3208}
3209
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);
3214};
3215
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);
3221}
3222
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);
3226};
3227
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);
3233}
3234
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);
3238};
3239
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);
3245}
3246
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);
3250};
3251
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);
3257}
3258#endif
3259
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);
3263};
3264
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);
3270}
3271
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);
3275};
3276
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);
3282}
3283
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);
3287};
3288
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);
3294}
3295
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);
3299};
3300
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);
3306}
3307
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,
3311 Index offset = 0);
3312};
3313
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);
3320}
3321
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,
3325 Index offset = 0);
3326};
3327
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);
3334}
3335
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);
3340};
3341
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);
3347}
3348
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);
3352};
3353
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);
3359}
3360#endif
3361
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,
3365 Index offset = 0);
3366};
3367
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);
3373}
3374
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,
3378 Index offset = 0);
3379};
3380
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);
3386}
3387
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,
3391 Index offset = 0);
3392};
3393
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);
3400}
3401
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,
3405 Index offset = 0);
3406};
3407
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);
3414}
3415
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,
3419 Index offset = 0);
3420};
3421
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);
3427}
3428
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,
3432 Index offset = 0);
3433};
3434
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);
3440}
3441
3442// ********* gebp specializations *********
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;
3447
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);
3450};
3451
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;
3458 // The kernels below take Eigen::Index, which need not be the Index this kernel is instantiated with: a Tensor
3459 // contraction instantiates gebp_kernel with the tensor's StorageIndex.
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> :
3464#endif
3465 &Eigen::internal::gemm<float, Packet, RhsPacket, DataMapper, accRows, accCols>;
3466 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3467}
3468
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;
3474
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);
3478};
3479
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,
3485 Index offsetB) {
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>
3495 :
3496#endif
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);
3501}
3502
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;
3508
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);
3512};
3513
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,
3522 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>
3527 :
3528#endif
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);
3533}
3534
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;
3540
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);
3544};
3545
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,
3554 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>
3559 :
3560#endif
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);
3565}
3566
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;
3571
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,
3574 Index offsetB = 0);
3575};
3576
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> :
3587#endif
3588 &Eigen::internal::gemm<double, Packet, RhsPacket, DataMapper, accRows, accCols>;
3589 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3590}
3591
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;
3597
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);
3601};
3602
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,
3608 Index offsetB) {
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
3615 (supportsMMA())
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>
3619 :
3620#endif
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);
3625}
3626
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;
3632
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);
3636};
3637
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>
3651 :
3652#endif
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);
3657}
3658
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;
3664
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);
3668};
3669
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>
3683 :
3684#endif
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);
3689}
3690
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;
3695
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,
3698 Index offsetB = 0);
3699};
3700
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> :
3709#endif
3710 &Eigen::internal::gemmbfloat16<DataMapper>;
3711 gemm_function(res, blockA, blockB, rows, depth, cols, alpha, strideA, strideB, offsetA, offsetB);
3712}
3713} // end namespace internal
3714
3715} // end namespace Eigen
3716
3717#endif // EIGEN_MATRIX_PRODUCT_ALTIVEC_H
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321