12#ifndef EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
13#define EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
16#if defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH)
17#pragma GCC push_options
18#pragma GCC target("cpu=power10,htm")
22#if !__has_builtin(__builtin_vsx_assemble_pair)
23#define __builtin_vsx_assemble_pair __builtin_mma_assemble_pair
25#if !__has_builtin(__builtin_vsx_disassemble_pair)
26#define __builtin_vsx_disassemble_pair __builtin_mma_disassemble_pair
31#include "../../InternalHeaderCheck.h"
33#include "MatrixProductMMAbfloat16.h"
39#define accColsC (accCols / 2)
41EIGEN_ALWAYS_INLINE
void bsetzeroMMA(__vector_quad* acc) { __builtin_mma_xxsetaccz(acc); }
43template <
typename DataMapper,
typename Packet,
bool full>
44EIGEN_ALWAYS_INLINE
void storeAccumulator(Index i,
const DataMapper& data,
const Packet& alpha,
const Index elements,
46 PacketBlock<Packet, 4> result;
47 __builtin_mma_disassemble_acc(&result.packet, acc);
49 PacketBlock<Packet, 4> tRes;
51 EIGEN_UNUSED_VARIABLE(elements);
52 bload<DataMapper, Packet, 0, ColMajor, false, 4>(tRes, data, i, 0);
53 bscale<Packet, 4>(tRes, result, alpha);
54 bstore<DataMapper, Packet, 4>(tRes, data, i);
56 bload_partial<DataMapper, Packet, 0, false, 4>(tRes, data, i, elements);
57 bscale<Packet, 4>(tRes, result, alpha);
58 bstore_partial<DataMapper, Packet, 4>(tRes, data, i, elements);
62template <
typename DataMapper,
typename Packet,
typename Packetc, const Index accCols, const Index accCols2>
63EIGEN_ALWAYS_INLINE
void storeComplexAccumulator(Index i,
const DataMapper& data,
const Packet& alphaReal,
64 const Packet& alphaImag,
const Packet& pMask, __vector_quad* accReal,
65 __vector_quad* accImag) {
66 constexpr bool full = (accCols2 > accColsC);
67 constexpr bool odd = (accCols != accCols2) && (
sizeof(__UNPACK_TYPE__(Packet)) ==
sizeof(
float)) && (accCols2 & 1);
68 PacketBlock<Packet, 4> resultReal, resultImag;
69 __builtin_mma_disassemble_acc(&resultReal.packet, accReal);
70 __builtin_mma_disassemble_acc(&resultImag.packet, accImag);
72 PacketBlock<Packetc, 8> tRes;
73 EIGEN_IF_CONSTEXPR (odd) {
74 bload_partial<DataMapper, Packetc, accColsC, true, 4, full>(tRes, data, i, 1);
76 bload<DataMapper, Packetc, accColsC, ColMajor, true, 4, full>(tRes, data, i, 0);
79 PacketBlock<Packet, 4> taccReal, taccImag;
80 bscalec<Packet, 4, (accCols != accCols2)>(resultReal, resultImag, alphaReal, alphaImag, taccReal, taccImag, pMask);
82 PacketBlock<Packetc, 4> acc1, acc2;
83 bcouple<Packet, Packetc, 4, full>(taccReal, taccImag, tRes, acc1, acc2);
85 EIGEN_IF_CONSTEXPR (odd && !full) {
86 bstore_partial<DataMapper, Packetc, 4>(acc1, data, i, 1);
88 bstore<DataMapper, Packetc, 4>(acc1, data, i);
90 EIGEN_IF_CONSTEXPR (full) {
91 EIGEN_IF_CONSTEXPR (odd) {
92 bstore_partial<DataMapper, Packetc, 4>(acc2, data, i + accColsC, 1);
94 bstore<DataMapper, Packetc, 4>(acc2, data, i + accColsC);
99template <
typename LhsPacket,
typename RhsPacket,
bool NegativeAccumulate>
100EIGEN_ALWAYS_INLINE
void pgerMMA(__vector_quad* acc,
const RhsPacket& a,
const LhsPacket& b) {
101 if (NegativeAccumulate) {
102 __builtin_mma_xvf32gernp(acc, (__vector
unsigned char)a, (__vector
unsigned char)b);
104 __builtin_mma_xvf32gerpp(acc, (__vector
unsigned char)a, (__vector
unsigned char)b);
108template <
typename LhsPacket,
typename RhsPacket,
bool NegativeAccumulate>
109EIGEN_ALWAYS_INLINE
void pgerMMA(__vector_quad* acc,
const __vector_pair& a,
const Packet2d& b) {
110 if (NegativeAccumulate) {
111 __builtin_mma_xvf64gernp(acc, (__vector_pair)a, (__vector
unsigned char)b);
113 __builtin_mma_xvf64gerpp(acc, (__vector_pair)a, (__vector
unsigned char)b);
117template <
typename Packet,
typename RhsPacket,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
118EIGEN_ALWAYS_INLINE
void pgercMMA(__vector_quad* accReal, __vector_quad* accImag,
const Packet& lhsV,
119 const Packet& lhsVi,
const RhsPacket& rhsV,
const RhsPacket& rhsVi) {
120 pgerMMA<Packet, RhsPacket, false>(accReal, rhsV, lhsV);
122 pgerMMA<Packet, RhsPacket, ConjugateRhs>(accImag, rhsVi, lhsV);
123 EIGEN_UNUSED_VARIABLE(lhsVi);
126 pgerMMA<Packet, RhsPacket, ConjugateLhs == ConjugateRhs>(accReal, rhsVi, lhsVi);
127 pgerMMA<Packet, RhsPacket, ConjugateRhs>(accImag, rhsVi, lhsV);
129 EIGEN_UNUSED_VARIABLE(rhsVi);
131 pgerMMA<Packet, RhsPacket, ConjugateLhs>(accImag, rhsV, lhsVi);
136template <
typename Packet>
137EIGEN_ALWAYS_INLINE Packet ploadRhs(
const __UNPACK_TYPE__(Packet) * rhs) {
138 return ploadu<Packet>(rhs);
141template <
typename Scalar,
typename Packet>
142EIGEN_ALWAYS_INLINE
void ploadRhsMMA(
const Scalar* rhs, Packet& rhsV) {
143 rhsV = ploadRhs<Packet>(rhs);
147EIGEN_ALWAYS_INLINE
void ploadRhsMMA(
const double* rhs, __vector_pair& rhsV) {
149 __builtin_vsx_assemble_pair(
150 &rhsV,
reinterpret_cast<__vector
unsigned char>(ploadRhs<Packet2d>(rhs + (
sizeof(Packet2d) /
sizeof(
double)))),
151 reinterpret_cast<__vector
unsigned char>(ploadRhs<Packet2d>(rhs)));
153 rhsV = *
reinterpret_cast<__vector_pair*
>(
const_cast<double*
>(rhs));
157EIGEN_ALWAYS_INLINE
void ploadLhsMMA(
const double* lhs, __vector_pair& lhsV) { ploadRhsMMA(lhs, lhsV); }
159#define GEMM_MULTIPLE_COLS
164#define VECTOR_PAIR_LOADS_LHS
168#ifdef GEMM_MULTIPLE_COLS
172#if EIGEN_COMP_LLVM || (__GNUC__ < 12) || defined(VECTOR_PAIR_LOADS_LHS)
179#define MICRO_MMA_UNROLL(func) func(0) func(1) func(2) func(3) func(4) func(5) func(6) func(7)
181#define MICRO_MMA_WORK(func, type, peel) \
183 func(0, type, peel, 0, 0) func(1, type, peel, 1, 0) func(2, type, peel, 2, 0) func(3, type, peel, 3, 0) \
184 func(4, type, peel, 4, 0) func(5, type, peel, 5, 0) func(6, type, peel, 6, 0) func(7, type, peel, 7, 0) \
185 } else if (accItr == 2) { \
186 func(0, type, peel, 0, 0) func(1, type, peel, 0, 1) func(2, type, peel, 1, 0) func(3, type, peel, 1, 1) \
187 func(4, type, peel, 2, 0) func(5, type, peel, 2, 1) func(6, type, peel, 3, 0) func(7, type, peel, 3, 1) \
189 func(0, type, peel, 0, 0) func(1, type, peel, 0, 1) func(2, type, peel, 0, 2) func(3, type, peel, 0, 3) \
190 func(4, type, peel, 1, 0) func(5, type, peel, 1, 1) func(6, type, peel, 1, 2) func(7, type, peel, 1, 3) \
193#define MICRO_MMA_WORK_ONE(iter, type, peel, left, right) \
194 if (unroll_factor > left) { \
195 pgerMMA<Packet, type, false>(&accZero##iter, rhsV##right[peel], lhsV##left); \
198#ifdef VECTOR_PAIR_LOADS_LHS
199#define MICRO_MMA_WORK_TWO(iter, type, peel, left, right) \
200 if (unroll_factor > left) { \
201 pgerMMA<Packet, type, false>(&accZero##iter, rhsV##right[peel], lhsV2##left.packet[peel & 1]); \
204#define MICRO_MMA_LOAD1_TWO(lhs_ptr, left) \
205 if (unroll_factor > left) { \
206 if (MICRO_NORMAL(left)) { \
207 ploadLhsMMA(reinterpret_cast<const double*>(lhs_ptr##left), plhsV##left); \
208 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(&lhsV2##left.packet), &plhsV##left); \
209 lhs_ptr##left += accCols * 2; \
211 lhsV2##left.packet[0] = ploadLhs<Packet>(lhs_ptr##left); \
212 lhsV2##left.packet[1] = ploadLhs<Packet>(lhs_ptr##left + accCols2); \
213 lhs_ptr##left += accCols2 * 2; \
214 EIGEN_UNUSED_VARIABLE(plhsV##left); \
217 EIGEN_UNUSED_VARIABLE(lhsV2##left); \
218 EIGEN_UNUSED_VARIABLE(plhsV##left); \
221#define MICRO_MMA_LOAD_TWO(left) MICRO_MMA_LOAD1_TWO(lhs_ptr, left)
224#define MICRO_MMA_UNROLL_ITER(func, val) \
225 func(val, 0) if (accItr > 1) { \
226 func(val, 1) if (accItr > 2) { func(val, 2) func(val, 3) } \
229#define MICRO_MMA_LOAD_ONE_RHS1(peel, right) ploadRhsMMA(rhs_ptr##right + (accRows * peel), rhsV##right[peel]);
231#define MICRO_MMA_LOAD_ONE_RHS(peel) MICRO_MMA_UNROLL_ITER(MICRO_MMA_LOAD_ONE_RHS1, peel)
233#define MICRO_MMA_TYPE_PEEL(funcw, funcl, type, peel) \
234 if (PEEL_MMA > peel) { \
235 Packet lhsV0, lhsV1, lhsV2, lhsV3, lhsV4, lhsV5, lhsV6, lhsV7; \
236 MICRO_MMA_LOAD_ONE_RHS(peel) \
237 MICRO_MMA_UNROLL(funcl) \
238 MICRO_MMA_WORK(funcw, type, peel) \
241#ifndef VECTOR_PAIR_LOADS_LHS
242#define MICRO_MMA_UNROLL_TYPE_PEEL(funcw, funcl, type) \
243 type rhsV0[8], rhsV1[(accItr > 1) ? 8 : 1], rhsV2[(accItr > 2) ? 8 : 1], rhsV3[(accItr > 2) ? 8 : 1]; \
244 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 0) \
245 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 1) \
246 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 2) \
247 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 3) \
248 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 4) \
249 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 5) \
250 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 6) MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 7)
252#define MICRO_MMA_LOAD_TWO_RHS(peel1, right) \
253 ploadRhsMMA(reinterpret_cast<const double*>(rhs_ptr##right + (accRows * peel1)), prhsV##peel1); \
254 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(&rhsV##right[peel1]), &prhsV##peel1);
256#define MICRO_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, peel1, peel2) \
257 if (PEEL_MMA > peel2) { \
258 PacketBlock<Packet, 2> lhsV20, lhsV21, lhsV22, lhsV23, lhsV24, lhsV25, lhsV26, lhsV27; \
259 __vector_pair plhsV0, plhsV1, plhsV2, plhsV3, plhsV4, plhsV5, plhsV6, plhsV7; \
260 if (sizeof(type) == 16) { \
261 MICRO_MMA_UNROLL_ITER(MICRO_MMA_LOAD_TWO_RHS, peel1) \
263 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
264 MICRO_MMA_LOAD_ONE_RHS(peel1) \
265 MICRO_MMA_LOAD_ONE_RHS(peel2) \
267 MICRO_MMA_UNROLL(funcl2) \
268 MICRO_MMA_WORK(funcw2, type, peel1) \
269 MICRO_MMA_WORK(funcw2, type, peel2) \
271 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
272 MICRO_MMA_TYPE_PEEL(funcw1, funcl1, type, peel1) \
275#define MICRO_MMA_UNROLL_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type) \
276 type rhsV0[8], rhsV1[(accItr > 1) ? 8 : 1], rhsV2[(accItr > 2) ? 8 : 1], rhsV3[(accItr > 2) ? 8 : 1]; \
277 __vector_pair prhsV0, prhsV2, prhsV4, prhsV6; \
278 MICRO_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, 0, 1) \
279 MICRO_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, 2, 3) \
280 MICRO_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, 4, 5) \
281 MICRO_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, 6, 7)
284#define MICRO_MMA_UNROLL_TYPE_ONE(funcw, funcl, type) \
285 type rhsV0[1], rhsV1[1], rhsV2[1], rhsV3[1]; \
286 MICRO_MMA_TYPE_PEEL(funcw, funcl, type, 0)
288#define MICRO_MMA_UPDATE_RHS1(size, right) rhs_ptr##right += (accRows * size);
290#define MICRO_MMA_UPDATE_RHS(size) MICRO_MMA_UNROLL_ITER(MICRO_MMA_UPDATE_RHS1, size)
292#define MICRO_MMA_UNROLL_TYPE(MICRO_MMA_TYPE, size) \
293 MICRO_MMA_TYPE(MICRO_MMA_WORK_ONE, MICRO_LOAD_ONE, RhsPacket) \
294 MICRO_MMA_UPDATE_RHS(size)
296#define MICRO_MMA_UNROLL_TYPE_PARTIAL(MICRO_MMA_TYPE, size) \
297 MICRO_MMA_TYPE(MICRO_MMA_WORK_ONE, MICRO_LOAD_PARTIAL_ONE, RhsPacket) \
298 MICRO_MMA_UPDATE_RHS(size)
300#ifndef VECTOR_PAIR_LOADS_LHS
301#define MICRO_MMA_ONE_PEEL MICRO_MMA_UNROLL_TYPE(MICRO_MMA_UNROLL_TYPE_PEEL, PEEL_MMA)
303#define MICRO_MMA_UNROLL_TYPE2(MICRO_MMA_TYPE, size) \
304 MICRO_MMA_TYPE(MICRO_MMA_WORK_ONE, MICRO_LOAD_ONE, MICRO_MMA_WORK_TWO, MICRO_MMA_LOAD_TWO, RhsPacket) \
305 MICRO_MMA_UPDATE_RHS(size)
307#define MICRO_MMA_ONE_PEEL MICRO_MMA_UNROLL_TYPE2(MICRO_MMA_UNROLL_TYPE_PEEL2, PEEL_MMA)
310#define MICRO_MMA_ONE MICRO_MMA_UNROLL_TYPE(MICRO_MMA_UNROLL_TYPE_ONE, 1)
312#define MICRO_MMA_ONE_PARTIAL MICRO_MMA_UNROLL_TYPE_PARTIAL(MICRO_MMA_UNROLL_TYPE_ONE, 1)
314#define MICRO_MMA_DST_PTR_ONE(iter) \
315 if (unroll_factor * accItr > iter) { \
316 bsetzeroMMA(&accZero##iter); \
318 EIGEN_UNUSED_VARIABLE(accZero##iter); \
321#define MICRO_MMA_DST_PTR MICRO_MMA_UNROLL(MICRO_MMA_DST_PTR_ONE)
323#define MICRO_MMA_SRC_PTR MICRO_MMA_UNROLL(MICRO_SRC_PTR_ONE)
325#define MICRO_MMA_PREFETCH MICRO_MMA_UNROLL(MICRO_PREFETCH_ONE)
327#define MICRO_MMA_STORE_ONE(iter, left, right) \
328 if (unroll_factor > left) { \
329 storeAccumulator<DataMapper, Packet, MICRO_NORMAL_PARTIAL(left)>(row + left * accCols, res##right, pAlpha, \
330 accCols2, &accZero##iter); \
333#define MICRO_MMA_ITER_UNROLL(func) \
335 func(0, 0, 0) func(1, 1, 0) func(2, 2, 0) func(3, 3, 0) func(4, 4, 0) func(5, 5, 0) func(6, 6, 0) func(7, 7, 0) \
336 } else if (accItr == 2) { \
337 func(0, 0, 0) func(1, 0, 1) func(2, 1, 0) func(3, 1, 1) func(4, 2, 0) func(5, 2, 1) func(6, 3, 0) func(7, 3, 1) \
339 func(0, 0, 0) func(1, 0, 1) func(2, 0, 2) func(3, 0, 3) func(4, 1, 0) func(5, 1, 1) func(6, 1, 2) func(7, 1, 3) \
342#define MICRO_MMA_STORE MICRO_MMA_ITER_UNROLL(MICRO_MMA_STORE_ONE)
344#define MICRO_MMA_EXTRA_ROWS(right) \
345 gemm_extra_row<Scalar, Packet, DataMapper, accRows, accCols>( \
346 res3##right, blockA, rhs_base + right * accRows * strideB, depth, strideA, offsetA, strideB, row, rows, \
347 remaining_rows, pAlpha, pMask);
349#define MICRO_MMA_EXTRA_ROWS1(val, right) MICRO_MMA_EXTRA_ROWS(right);
351template <
int unroll_factor,
typename Scalar,
typename Packet,
typename RhsPacket,
typename DataMapper,
352 const Index accRows,
const Index accCols,
bool full,
const Index accItr>
353EIGEN_ALWAYS_INLINE
void gemm_unrolled_MMA_iteration(
const DataMapper& res0,
const DataMapper& res1,
354 const DataMapper& res2,
const DataMapper& res3,
355 const Scalar* lhs_base,
const Scalar* rhs_base, Index depth,
356 Index strideA, Index strideB, Index offsetA, Index& row,
357 const Packet& pAlpha, Index accCols2) {
358 const Scalar *rhs_ptr0 = rhs_base, *rhs_ptr1 =
nullptr, *rhs_ptr2 =
nullptr, *rhs_ptr3 =
nullptr;
359 const Scalar *lhs_ptr0 =
nullptr, *lhs_ptr1 =
nullptr, *lhs_ptr2 =
nullptr, *lhs_ptr3 =
nullptr, *lhs_ptr4 =
nullptr,
360 *lhs_ptr5 =
nullptr, *lhs_ptr6 =
nullptr, *lhs_ptr7 =
nullptr;
361 __vector_quad accZero0, accZero1, accZero2, accZero3, accZero4, accZero5, accZero6, accZero7;
364 rhs_ptr1 = rhs_base + (accRows * strideB);
366 EIGEN_UNUSED_VARIABLE(strideB);
367 EIGEN_UNUSED_VARIABLE(rhs_ptr1);
368 EIGEN_UNUSED_VARIABLE(res1);
371 rhs_ptr2 = rhs_base + (2 * accRows * strideB);
372 rhs_ptr3 = rhs_base + (3 * accRows * strideB);
374 EIGEN_UNUSED_VARIABLE(rhs_ptr2);
375 EIGEN_UNUSED_VARIABLE(rhs_ptr3);
376 EIGEN_UNUSED_VARIABLE(res2);
377 EIGEN_UNUSED_VARIABLE(res3);
383 const Index peel_depth = full ? depth : (depth - (accCols - accCols2));
384 Index k = 0, depth2 = peel_depth - PEEL_MMA;
385 for (; k <= depth2; k += PEEL_MMA) {
386 EIGEN_POWER_PREFETCH(rhs_ptr);
390 for (; k < peel_depth; k++) {
393 EIGEN_IF_CONSTEXPR (!full) {
394 for (; k < depth; k++) {
395 MICRO_MMA_ONE_PARTIAL
403#define MICRO_MMA_UNROLL_ITER2(N, M) \
404 gemm_unrolled_MMA_iteration<N + (M ? 1 : 0), Scalar, Packet, RhsPacket, DataMapper, accRows, accCols, !M, accItr>( \
405 res30, res31, res32, res33, lhs_base, rhs_base, depth, strideA, strideB, offsetA, row, pAlpha, \
406 M ? remaining_rows : accCols); \
409#define MICRO_MMA_ROWS(n) \
410 while (row + n * accCols <= rows) { \
411 MICRO_MMA_UNROLL_ITER2(n, 0); \
414template <
typename Scalar,
typename Packet,
typename RhsPacket,
typename DataMapper,
const Index accRows,
415 const Index accCols,
const Index accItr>
416EIGEN_ALWAYS_INLINE
void gemmMMA_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index depth,
417 Index strideA, Index offsetA, Index strideB, Index offsetB, Index col, Index rows,
418 Index remaining_rows,
const Packet& pAlpha,
const Packet& pMask) {
419 const DataMapper res30 = res.getSubMapper(0, col);
420 const DataMapper res31 = (accItr > 1) ? res30.getSubMapper(0, accRows * 1) : res30;
421 const DataMapper res32 = (accItr > 2) ? res30.getSubMapper(0, accRows * 2) : res30;
422 const DataMapper res33 = (accItr > 2) ? res30.getSubMapper(0, accRows * 3) : res30;
424 const Scalar* rhs_base = blockB + col * strideB + accRows * offsetB;
425 const Scalar* lhs_base = blockA + accCols * offsetA;
428#define MAX_MMA_UNROLL 7
430#if MAX_MMA_UNROLL < 2
432#elif MAX_MMA_UNROLL < 4
437 MICRO_MMA_ROWS(MAX_MMA_UNROLL);
438 }
else if (accItr == 2) {
443 switch ((rows - row) / accCols) {
444#if MAX_MMA_UNROLL > 7
447 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 7)
451#if MAX_MMA_UNROLL > 6
454 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 6)
458#if MAX_MMA_UNROLL > 5
461 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 5)
465#if MAX_MMA_UNROLL > 4
468 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 4)
472#if MAX_MMA_UNROLL > 3
475 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 3)
479#if MAX_MMA_UNROLL > 2
482 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 2)
486#if MAX_MMA_UNROLL > 1
488 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 1)
496 if (remaining_rows > 0) {
497 MICRO_MMA_UNROLL_ITER(MICRO_MMA_EXTRA_ROWS1, 0)
501#define MICRO_MMA_COLS(n) \
502 for (; col + n * accRows <= cols; col += n * accRows) { \
503 gemmMMA_cols<Scalar, Packet, RhsPacket2, DataMapper, accRows, accCols, n>( \
504 res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB, col, rows, remaining_rows, pAlpha, pMask); \
507template <
typename Scalar,
typename Packet,
typename RhsPacket,
typename DataMapper,
const Index accRows,
509void gemmMMA(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index rows, Index depth, Index cols,
510 Scalar alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
511 const Index remaining_rows = rows % accCols;
513 if (strideA == -1) strideA = depth;
514 if (strideB == -1) strideB = depth;
516 const Packet pAlpha = pset1<Packet>(alpha);
517 const Packet pMask = bmask<Packet>(remaining_rows);
519 typedef std::conditional_t<(
sizeof(Scalar) ==
sizeof(float)), RhsPacket, __vector_pair> RhsPacket2;
522#ifdef GEMM_MULTIPLE_COLS
529 gemm_extra_cols<Scalar, Packet, DataMapper, accCols>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB,
530 col, rows, cols, remaining_rows, pAlpha, pMask);
534#define advanceRows ((LhsIsReal) ? 1 : 2)
535#define advanceCols ((RhsIsReal) ? 1 : 2)
538#ifdef GEMM_MULTIPLE_COLS
539#define PEEL_COMPLEX_MMA 4
541#define PEEL_COMPLEX_MMA 3
544#define MICRO_COMPLEX_MMA_UNROLL(func) func(0) func(1) func(2) func(3)
546#define MICRO_COMPLEX_MMA_WORK(func, type, peel) \
548 func(0, type, peel, 0, 0) func(1, type, peel, 1, 0) func(2, type, peel, 2, 0) func(3, type, peel, 3, 0) \
549 } else if (accItr == 2) { \
550 func(0, type, peel, 0, 0) func(1, type, peel, 0, 1) func(2, type, peel, 1, 0) func(3, type, peel, 1, 1) \
552 func(0, type, peel, 0, 0) func(1, type, peel, 0, 1) func(2, type, peel, 0, 2) func(3, type, peel, 0, 3) \
555#define MICRO_COMPLEX_MMA_WORK_ONE(iter, type, peel, left, right) \
556 if (unroll_factor > left) { \
557 pgercMMA<Packet, type, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>( \
558 &accReal##iter, &accImag##iter, lhsV##left, lhsVi##left, rhsV##right[peel], rhsVi##right[peel]); \
561#ifdef VECTOR_PAIR_LOADS_LHS
562#define MICRO_COMPLEX_MMA_WORK_TWO(iter, type, peel, left, right) \
563 if (unroll_factor > left) { \
564 pgercMMA<Packet, type, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal>( \
565 &accReal##iter, &accImag##iter, lhsV2##left.packet[peel & 1], lhsVi2##left.packet[peel & 1], \
566 rhsV##right[peel], rhsVi##right[peel]); \
569#define MICRO_COMPLEX_MMA_LOAD1_TWO(lhs_ptr, left) \
570 if (!LhsIsReal && (unroll_factor > left)) { \
571 if (MICRO_NORMAL(left)) { \
572 ploadLhsMMA(reinterpret_cast<const double*>(lhs_ptr_real##left + imag_delta), plhsVi##left); \
573 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(&lhsVi2##left.packet), &plhsVi##left); \
575 lhsVi2##left.packet[0] = ploadLhs<Packet>(lhs_ptr_real##left + imag_delta2); \
576 lhsVi2##left.packet[1] = ploadLhs<Packet>(lhs_ptr_real##left + imag_delta2 + accCols2); \
577 EIGEN_UNUSED_VARIABLE(plhsVi##left); \
580 EIGEN_UNUSED_VARIABLE(lhsVi2##left); \
581 EIGEN_UNUSED_VARIABLE(plhsVi##left); \
583 MICRO_MMA_LOAD1_TWO(lhs_ptr_real, left)
585#define MICRO_COMPLEX_MMA_LOAD_TWO(left) MICRO_COMPLEX_MMA_LOAD1_TWO(lhs_ptr, left)
588#define MICRO_COMPLEX_MMA_LOAD_RHS1(peel, right) \
589 ploadRhsMMA(rhs_ptr_real##right + (accRows * peel), rhsV##right[peel]); \
591 ploadRhsMMA(rhs_ptr_imag##right + (accRows * peel), rhsVi##right[peel]); \
594#define MICRO_COMPLEX_MMA_LOAD_ONE_RHS(peel) MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_LOAD_RHS1, peel)
596#define MICRO_COMPLEX_MMA_TYPE_PEEL(funcw, funcl, type, peel) \
597 if (PEEL_COMPLEX_MMA > peel) { \
598 Packet lhsV0, lhsV1, lhsV2, lhsV3; \
599 Packet lhsVi0, lhsVi1, lhsVi2, lhsVi3; \
600 MICRO_COMPLEX_MMA_LOAD_ONE_RHS(peel) \
601 MICRO_COMPLEX_MMA_UNROLL(funcl) \
602 MICRO_COMPLEX_MMA_WORK(funcw, type, peel) \
605#ifndef VECTOR_PAIR_LOADS_LHS
606#define MICRO_COMPLEX_MMA_UNROLL_TYPE_PEEL(funcw, funcl, type) \
607 type rhsV0[4], rhsVi0[4], rhsV1[(accItr > 1) ? 4 : 1], rhsVi1[(accItr > 1) ? 4 : 1], rhsV2[(accItr > 2) ? 4 : 1], \
608 rhsVi2[(accItr > 2) ? 4 : 1], rhsV3[(accItr > 2) ? 4 : 1], rhsVi3[(accItr > 2) ? 4 : 1]; \
609 MICRO_COMPLEX_MMA_TYPE_PEEL(funcw, funcl, type, 0) \
610 MICRO_COMPLEX_MMA_TYPE_PEEL(funcw, funcl, type, 1) \
611 MICRO_COMPLEX_MMA_TYPE_PEEL(funcw, funcl, type, 2) MICRO_COMPLEX_MMA_TYPE_PEEL(funcw, funcl, type, 3)
613#define MICRO_COMPLEX_MMA_LOAD_TWO_RHS(peel1, right) \
614 ploadRhsMMA(reinterpret_cast<const double*>(rhs_ptr_real##right + (accRows * peel1)), prhsV##peel1); \
615 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(&rhsV##right[peel1]), &prhsV##peel1); \
617 ploadRhsMMA(reinterpret_cast<const double*>(rhs_ptr_imag##right + (accRows * peel1)), prhsVi##peel1); \
618 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(&rhsVi##right[peel1]), &prhsVi##peel1); \
620 EIGEN_UNUSED_VARIABLE(prhsVi##peel1); \
623#define MICRO_COMPLEX_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, peel1, peel2) \
624 if (PEEL_COMPLEX_MMA > peel2) { \
625 PacketBlock<Packet, 2> lhsV20, lhsV21, lhsV22, lhsV23; \
626 PacketBlock<Packet, 2> lhsVi20, lhsVi21, lhsVi22, lhsVi23; \
627 __vector_pair plhsV0, plhsV1, plhsV2, plhsV3; \
628 __vector_pair plhsVi0, plhsVi1, plhsVi2, plhsVi3; \
629 if (sizeof(type) == 16) { \
630 MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_LOAD_TWO_RHS, peel1) \
632 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
633 EIGEN_UNUSED_VARIABLE(prhsVi##peel1); \
634 MICRO_COMPLEX_MMA_LOAD_ONE_RHS(peel1); \
635 MICRO_COMPLEX_MMA_LOAD_ONE_RHS(peel2); \
637 MICRO_COMPLEX_MMA_UNROLL(funcl2) \
638 MICRO_COMPLEX_MMA_WORK(funcw2, type, peel1) \
639 MICRO_COMPLEX_MMA_WORK(funcw2, type, peel2) \
641 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
642 EIGEN_UNUSED_VARIABLE(prhsVi##peel1); \
643 MICRO_COMPLEX_MMA_TYPE_PEEL(funcw1, funcl1, type, peel1) \
646#define MICRO_COMPLEX_MMA_UNROLL_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type) \
647 type rhsV0[4], rhsVi0[4], rhsV1[(accItr > 1) ? 4 : 1], rhsVi1[(accItr > 1) ? 4 : 1], rhsV2[(accItr > 2) ? 4 : 1], \
648 rhsVi2[(accItr > 2) ? 4 : 1], rhsV3[(accItr > 2) ? 4 : 1], rhsVi3[(accItr > 2) ? 4 : 1]; \
649 __vector_pair prhsV0, prhsV2; \
650 __vector_pair prhsVi0, prhsVi2; \
651 MICRO_COMPLEX_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, 0, 1) \
652 MICRO_COMPLEX_MMA_TYPE_PEEL2(funcw1, funcl1, funcw2, funcl2, type, 2, 3)
655#define MICRO_COMPLEX_MMA_UNROLL_TYPE_ONE(funcw, funcl, type) \
656 type rhsV0[1], rhsVi0[1], rhsV1[1], rhsVi1[1], rhsV2[1], rhsVi2[1], rhsV3[1], rhsVi3[1]; \
657 MICRO_COMPLEX_MMA_TYPE_PEEL(funcw, funcl, type, 0)
659#define MICRO_COMPLEX_MMA_UPDATE_RHS1(size, right) \
660 rhs_ptr_real##right += (accRows * size); \
661 if (!RhsIsReal) rhs_ptr_imag##right += (accRows * size);
663#define MICRO_COMPLEX_MMA_UPDATE_RHS(size) MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_UPDATE_RHS1, size)
665#define MICRO_COMPLEX_MMA_UNROLL_TYPE(MICRO_COMPLEX_MMA_TYPE, size) \
666 MICRO_COMPLEX_MMA_TYPE(MICRO_COMPLEX_MMA_WORK_ONE, MICRO_COMPLEX_LOAD_ONE, RhsPacket) \
667 MICRO_COMPLEX_MMA_UPDATE_RHS(size);
669#define MICRO_COMPLEX_MMA_UNROLL_TYPE_PARTIAL(MICRO_COMPLEX_MMA_TYPE, size) \
670 MICRO_COMPLEX_MMA_TYPE(MICRO_COMPLEX_MMA_WORK_ONE, MICRO_COMPLEX_LOAD_PARTIAL_ONE, RhsPacket) \
671 MICRO_COMPLEX_MMA_UPDATE_RHS(size);
673#ifndef VECTOR_PAIR_LOADS_LHS
674#define MICRO_COMPLEX_MMA_ONE_PEEL MICRO_COMPLEX_MMA_UNROLL_TYPE(MICRO_COMPLEX_MMA_UNROLL_TYPE_PEEL, PEEL_COMPLEX_MMA)
676#define MICRO_COMPLEX_MMA_UNROLL_TYPE2(MICRO_COMPLEX_MMA_TYPE, size) \
677 MICRO_COMPLEX_MMA_TYPE(MICRO_COMPLEX_MMA_WORK_ONE, MICRO_COMPLEX_LOAD_ONE, MICRO_COMPLEX_MMA_WORK_TWO, \
678 MICRO_COMPLEX_MMA_LOAD_TWO, RhsPacket) \
679 MICRO_COMPLEX_MMA_UPDATE_RHS(size);
681#define MICRO_COMPLEX_MMA_ONE_PEEL MICRO_COMPLEX_MMA_UNROLL_TYPE2(MICRO_COMPLEX_MMA_UNROLL_TYPE_PEEL2, PEEL_COMPLEX_MMA)
684#define MICRO_COMPLEX_MMA_ONE MICRO_COMPLEX_MMA_UNROLL_TYPE(MICRO_COMPLEX_MMA_UNROLL_TYPE_ONE, 1)
686#define MICRO_COMPLEX_MMA_ONE_PARTIAL MICRO_COMPLEX_MMA_UNROLL_TYPE_PARTIAL(MICRO_COMPLEX_MMA_UNROLL_TYPE_ONE, 1)
688#define MICRO_COMPLEX_MMA_DST_PTR_ONE(iter) \
689 if (unroll_factor * accItr > iter) { \
690 bsetzeroMMA(&accReal##iter); \
691 bsetzeroMMA(&accImag##iter); \
693 EIGEN_UNUSED_VARIABLE(accReal##iter); \
694 EIGEN_UNUSED_VARIABLE(accImag##iter); \
697#define MICRO_COMPLEX_MMA_DST_PTR MICRO_COMPLEX_MMA_UNROLL(MICRO_COMPLEX_MMA_DST_PTR_ONE)
699#define MICRO_COMPLEX_MMA_SRC_PTR MICRO_COMPLEX_MMA_UNROLL(MICRO_COMPLEX_SRC_PTR_ONE)
701#define MICRO_COMPLEX_MMA_PREFETCH MICRO_COMPLEX_MMA_UNROLL(MICRO_COMPLEX_PREFETCH_ONE)
703#define MICRO_COMPLEX_MMA_STORE_ONE(iter, left, right) \
704 if (unroll_factor > left) { \
705 storeComplexAccumulator<DataMapper, Packet, Packetc, accCols, (unroll_factor != (left + 1)) ? accCols : accCols2>( \
706 row + left * accCols, res##right, pAlphaReal, pAlphaImag, pMask, &accReal##iter, &accImag##iter); \
709#define MICRO_COMPLEX_MMA_ITER_UNROLL(func) \
711 func(0, 0, 0) func(1, 1, 0) func(2, 2, 0) func(3, 3, 0) \
712 } else if (accItr == 2) { \
713 func(0, 0, 0) func(1, 0, 1) func(2, 1, 0) func(3, 1, 1) \
715 func(0, 0, 0) func(1, 0, 1) func(2, 0, 2) func(3, 0, 3) \
718#define MICRO_COMPLEX_MMA_STORE MICRO_COMPLEX_MMA_ITER_UNROLL(MICRO_COMPLEX_MMA_STORE_ONE)
720#define MICRO_COMPLEX_MMA_EXTRA_ROWS(right) \
721 gemm_complex_extra_row<Scalar, Packet, Packetc, DataMapper, accRows, accCols, ConjugateLhs, ConjugateRhs, LhsIsReal, \
722 RhsIsReal>(res3##right, blockA, rhs_base + right * accRows * (RhsIsReal ? 1 : 2) * strideB, \
723 depth, strideA, offsetA, strideB, row, rows, remaining_rows, pAlphaReal, \
726#define MICRO_COMPLEX_MMA_EXTRA_ROWS1(val, right) MICRO_COMPLEX_MMA_EXTRA_ROWS(right);
728template <
int unroll_factor,
typename Scalar,
typename Packet,
typename Packetc,
typename RhsPacket,
729 typename DataMapper,
const Index accRows,
const Index accCols,
const Index accCols2,
bool ConjugateLhs,
730 bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal,
const Index accItr>
731EIGEN_ALWAYS_INLINE
void gemm_complex_unrolled_MMA_iteration(
const DataMapper& res0,
const DataMapper& res1,
732 const DataMapper& res2,
const DataMapper& res3,
733 const Scalar* lhs_base,
const Scalar* rhs_base,
734 Index depth, Index strideA, Index offsetA, Index strideB,
735 Index& row,
const Packet& pAlphaReal,
736 const Packet& pAlphaImag,
const Packet& pMask) {
737 const Scalar *rhs_ptr_real0 = rhs_base, *rhs_ptr_real1 =
nullptr, *rhs_ptr_real2 =
nullptr, *rhs_ptr_real3 =
nullptr;
738 const Scalar *rhs_ptr_imag0 =
nullptr, *rhs_ptr_imag1 =
nullptr, *rhs_ptr_imag2 =
nullptr, *rhs_ptr_imag3 =
nullptr;
739 const Index imag_delta = accCols * strideA;
740 const Index imag_delta2 = accCols2 * strideA;
743 rhs_ptr_imag0 = rhs_base + accRows * strideB;
745 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag0);
749 rhs_ptr_real1 = rhs_base + (2 * accRows * strideB);
750 rhs_ptr_imag1 = rhs_base + (3 * accRows * strideB);
752 rhs_ptr_real1 = rhs_base + accRows * strideB;
753 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag1);
756 EIGEN_UNUSED_VARIABLE(rhs_ptr_real1);
757 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag1);
758 EIGEN_UNUSED_VARIABLE(res1);
762 rhs_ptr_real2 = rhs_base + (4 * accRows * strideB);
763 rhs_ptr_imag2 = rhs_base + (5 * accRows * strideB);
764 rhs_ptr_real3 = rhs_base + (6 * accRows * strideB);
765 rhs_ptr_imag3 = rhs_base + (7 * accRows * strideB);
767 rhs_ptr_real2 = rhs_base + (2 * accRows * strideB);
768 rhs_ptr_real3 = rhs_base + (3 * accRows * strideB);
769 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag2);
770 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag3);
773 EIGEN_UNUSED_VARIABLE(rhs_ptr_real2);
774 EIGEN_UNUSED_VARIABLE(rhs_ptr_real3);
775 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag2);
776 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag3);
777 EIGEN_UNUSED_VARIABLE(res2);
778 EIGEN_UNUSED_VARIABLE(res3);
780 const Scalar *lhs_ptr_real0 =
nullptr, *lhs_ptr_real1 =
nullptr;
781 const Scalar *lhs_ptr_real2 =
nullptr, *lhs_ptr_real3 =
nullptr;
782 __vector_quad accReal0, accImag0, accReal1, accImag1, accReal2, accImag2, accReal3, accImag3;
784 MICRO_COMPLEX_MMA_SRC_PTR
785 MICRO_COMPLEX_MMA_DST_PTR
787 const Index peel_depth = depth - (accCols - accCols2);
788 Index k = 0, depth2 = peel_depth - PEEL_COMPLEX_MMA;
789 for (; k <= depth2; k += PEEL_COMPLEX_MMA) {
790 EIGEN_POWER_PREFETCH(rhs_ptr_real);
792 EIGEN_POWER_PREFETCH(rhs_ptr_imag);
794 MICRO_COMPLEX_MMA_PREFETCH
795 MICRO_COMPLEX_MMA_ONE_PEEL
797 for (; k < peel_depth; k++) {
798 MICRO_COMPLEX_MMA_ONE
800 EIGEN_IF_CONSTEXPR (accCols != accCols2) {
801 for (; k < depth; k++) {
802 MICRO_COMPLEX_MMA_ONE_PARTIAL
805 MICRO_COMPLEX_MMA_STORE
810#define MICRO_COMPLEX_MMA_UNROLL_ITER2(N, M) \
811 gemm_complex_unrolled_MMA_iteration<N + (M ? 1 : 0), Scalar, Packet, Packetc, RhsPacket, DataMapper, accRows, \
812 accCols, M ? M : accCols, ConjugateLhs, ConjugateRhs, LhsIsReal, RhsIsReal, \
813 accItr>(res30, res31, res32, res33, lhs_base, rhs_base, depth, strideA, offsetA, \
814 strideB, row, pAlphaReal, pAlphaImag, pMask); \
817#define MICRO_COMPLEX_MMA_ROWS(n) \
818 while (row + n * accCols <= rows) { \
819 MICRO_COMPLEX_MMA_UNROLL_ITER2(n, 0); \
822template <
typename Scalar,
typename Packet,
typename Packetc,
typename RhsPacket,
typename DataMapper,
823 const Index accRows,
const Index accCols,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
824 bool RhsIsReal,
const Index accItr>
825EIGEN_ALWAYS_INLINE
void gemmMMA_complex_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB,
826 Index depth, Index strideA, Index offsetA, Index strideB, Index offsetB,
827 Index col, Index rows, Index remaining_rows,
const Packet& pAlphaReal,
828 const Packet& pAlphaImag,
const Packet& pMask) {
829 const DataMapper res30 = res.getSubMapper(0, col);
830 const DataMapper res31 = (accItr > 1) ? res30.getSubMapper(0, accRows * 1) : res30;
831 const DataMapper res32 = (accItr > 2) ? res30.getSubMapper(0, accRows * 2) : res30;
832 const DataMapper res33 = (accItr > 2) ? res30.getSubMapper(0, accRows * 3) : res30;
834 const Scalar* rhs_base = blockB + advanceCols * col * strideB + accRows * offsetB;
835 const Scalar* lhs_base = blockA + accCols * offsetA;
838#define MAX_COMPLEX_MMA_UNROLL 4
840#if MAX_COMPLEX_MMA_UNROLL < 2
842#elif MAX_COMPLEX_MMA_UNROLL < 4
847 MICRO_COMPLEX_MMA_ROWS(MAX_COMPLEX_MMA_UNROLL);
848 }
else if (accItr == 2) {
849 MICRO_COMPLEX_MMA_ROWS(2);
851 MICRO_COMPLEX_MMA_ROWS(1);
853 switch ((rows - row) / accCols) {
854#if MAX_COMPLEX_MMA_UNROLL > 3
857 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_MMA_UNROLL_ITER2, 3)
861#if MAX_COMPLEX_MMA_UNROLL > 2
864 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_MMA_UNROLL_ITER2, 2)
868#if MAX_COMPLEX_MMA_UNROLL > 1
871 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_MMA_UNROLL_ITER2, 1)
878#undef MAX_COMPLEX_MMA_UNROLL
880 if (remaining_rows > 0) {
881 MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_EXTRA_ROWS1, 0)
885#define MICRO_COMPLEX_MMA_COLS(n) \
886 for (; col + n * accRows <= cols; col += n * accRows) { \
887 gemmMMA_complex_cols<Scalar, Packet, Packetc, RhsPacket2, DataMapper, accRows, accCols, ConjugateLhs, \
888 ConjugateRhs, LhsIsReal, RhsIsReal, n>(res, blockA, blockB, depth, strideA, offsetA, strideB, \
889 offsetB, col, rows, remaining_rows, pAlphaReal, \
890 pAlphaImag, pMask); \
893template <
typename LhsScalar,
typename RhsScalar,
typename Scalarc,
typename Scalar,
typename Packet,
typename Packetc,
894 typename RhsPacket,
typename DataMapper,
const Index accRows,
const Index accCols,
bool ConjugateLhs,
895 bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
896void gemm_complexMMA(
const DataMapper& res,
const LhsScalar* blockAc,
const RhsScalar* blockBc, Index rows, Index depth,
897 Index cols, Scalarc alpha, Index strideA, Index strideB, Index offsetA, Index offsetB) {
898 const Index remaining_rows = rows % accCols;
900 if (strideA == -1) strideA = depth;
901 if (strideB == -1) strideB = depth;
903 const Packet pAlphaReal = pset1<Packet>(alpha.real());
904 const Packet pAlphaImag = pset1<Packet>(alpha.imag());
905 const Packet pMask = bmask<Packet>(remaining_rows);
907 const Scalar* blockA = (Scalar*)blockAc;
908 const Scalar* blockB = (Scalar*)blockBc;
910 typedef std::conditional_t<(
sizeof(Scalar) ==
sizeof(float)), RhsPacket, __vector_pair> RhsPacket2;
913#ifdef GEMM_MULTIPLE_COLS
914 MICRO_COMPLEX_MMA_COLS(4);
915 MICRO_COMPLEX_MMA_COLS(2);
917 MICRO_COMPLEX_MMA_COLS(1);
920 gemm_complex_extra_cols<Scalar, Packet, Packetc, DataMapper, accCols, ConjugateLhs, ConjugateRhs, LhsIsReal,
921 RhsIsReal>(res, blockA, blockB, depth, strideA, offsetA, strideB, offsetB, col, rows, cols,
922 remaining_rows, pAlphaReal, pAlphaImag, pMask);
934#if defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH)
935#pragma GCC pop_options