Eigen  5.0.1
 
Loading...
Searching...
No Matches
MatrixProductMMA.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_MMA_ALTIVEC_H
13#define EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H
14
15// If using dynamic dispatch, set the CPU target.
16#if defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH)
17#pragma GCC push_options
18#pragma GCC target("cpu=power10,htm")
19#endif
20
21#ifdef __has_builtin
22#if !__has_builtin(__builtin_vsx_assemble_pair)
23#define __builtin_vsx_assemble_pair __builtin_mma_assemble_pair
24#endif
25#if !__has_builtin(__builtin_vsx_disassemble_pair)
26#define __builtin_vsx_disassemble_pair __builtin_mma_disassemble_pair
27#endif
28#endif
29
30// IWYU pragma: private
31#include "../../InternalHeaderCheck.h"
32
33#include "MatrixProductMMAbfloat16.h"
34
35namespace Eigen {
36
37namespace internal {
38
39#define accColsC (accCols / 2)
40
41EIGEN_ALWAYS_INLINE void bsetzeroMMA(__vector_quad* acc) { __builtin_mma_xxsetaccz(acc); }
42
43template <typename DataMapper, typename Packet, bool full>
44EIGEN_ALWAYS_INLINE void storeAccumulator(Index i, const DataMapper& data, const Packet& alpha, const Index elements,
45 __vector_quad* acc) {
46 PacketBlock<Packet, 4> result;
47 __builtin_mma_disassemble_acc(&result.packet, acc);
48
49 PacketBlock<Packet, 4> tRes;
50 if (full) {
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);
55 } else {
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);
59 }
60}
61
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);
71
72 PacketBlock<Packetc, 8> tRes;
73 EIGEN_IF_CONSTEXPR (odd) {
74 bload_partial<DataMapper, Packetc, accColsC, true, 4, full>(tRes, data, i, 1);
75 } else {
76 bload<DataMapper, Packetc, accColsC, ColMajor, true, 4, full>(tRes, data, i, 0);
77 }
78
79 PacketBlock<Packet, 4> taccReal, taccImag;
80 bscalec<Packet, 4, (accCols != accCols2)>(resultReal, resultImag, alphaReal, alphaImag, taccReal, taccImag, pMask);
81
82 PacketBlock<Packetc, 4> acc1, acc2;
83 bcouple<Packet, Packetc, 4, full>(taccReal, taccImag, tRes, acc1, acc2);
84
85 EIGEN_IF_CONSTEXPR (odd && !full) {
86 bstore_partial<DataMapper, Packetc, 4>(acc1, data, i, 1);
87 } else {
88 bstore<DataMapper, Packetc, 4>(acc1, data, i);
89 }
90 EIGEN_IF_CONSTEXPR (full) {
91 EIGEN_IF_CONSTEXPR (odd) {
92 bstore_partial<DataMapper, Packetc, 4>(acc2, data, i + accColsC, 1);
93 } else {
94 bstore<DataMapper, Packetc, 4>(acc2, data, i + accColsC);
95 }
96 }
97}
98
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);
103 } else {
104 __builtin_mma_xvf32gerpp(acc, (__vector unsigned char)a, (__vector unsigned char)b);
105 }
106}
107
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);
112 } else {
113 __builtin_mma_xvf64gerpp(acc, (__vector_pair)a, (__vector unsigned char)b);
114 }
115}
116
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);
121 if (LhsIsReal) {
122 pgerMMA<Packet, RhsPacket, ConjugateRhs>(accImag, rhsVi, lhsV);
123 EIGEN_UNUSED_VARIABLE(lhsVi);
124 } else {
125 if (!RhsIsReal) {
126 pgerMMA<Packet, RhsPacket, ConjugateLhs == ConjugateRhs>(accReal, rhsVi, lhsVi);
127 pgerMMA<Packet, RhsPacket, ConjugateRhs>(accImag, rhsVi, lhsV);
128 } else {
129 EIGEN_UNUSED_VARIABLE(rhsVi);
130 }
131 pgerMMA<Packet, RhsPacket, ConjugateLhs>(accImag, rhsV, lhsVi);
132 }
133}
134
135// This is necessary because ploadRhs for double returns a pair of vectors when MMA is enabled.
136template <typename Packet>
137EIGEN_ALWAYS_INLINE Packet ploadRhs(const __UNPACK_TYPE__(Packet) * rhs) {
138 return ploadu<Packet>(rhs);
139}
140
141template <typename Scalar, typename Packet>
142EIGEN_ALWAYS_INLINE void ploadRhsMMA(const Scalar* rhs, Packet& rhsV) {
143 rhsV = ploadRhs<Packet>(rhs);
144}
145
146template <>
147EIGEN_ALWAYS_INLINE void ploadRhsMMA(const double* rhs, __vector_pair& rhsV) {
148#if EIGEN_COMP_LLVM
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)));
152#else
153 rhsV = *reinterpret_cast<__vector_pair*>(const_cast<double*>(rhs));
154#endif
155}
156
157EIGEN_ALWAYS_INLINE void ploadLhsMMA(const double* lhs, __vector_pair& lhsV) { ploadRhsMMA(lhs, lhsV); }
158
159#define GEMM_MULTIPLE_COLS
160
161// Disable in GCC until unnecessary register moves are fixed
162// #if (EIGEN_COMP_LLVM || (__GNUC__ >= 11))
163#if EIGEN_COMP_LLVM
164#define VECTOR_PAIR_LOADS_LHS
165#endif
166
167// PEEL_MMA loop factor.
168#ifdef GEMM_MULTIPLE_COLS
169#define PEEL_MMA 8
170#else
171// Register spillage with GCC12+
172#if EIGEN_COMP_LLVM || (__GNUC__ < 12) || defined(VECTOR_PAIR_LOADS_LHS)
173#define PEEL_MMA 7
174#else
175#define PEEL_MMA 6
176#endif
177#endif
178
179#define MICRO_MMA_UNROLL(func) func(0) func(1) func(2) func(3) func(4) func(5) func(6) func(7)
180
181#define MICRO_MMA_WORK(func, type, peel) \
182 if (accItr == 1) { \
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) \
188 } else { \
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) \
191 }
192
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); \
196 }
197
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]); \
202 }
203
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; \
210 } else { \
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); \
215 } \
216 } else { \
217 EIGEN_UNUSED_VARIABLE(lhsV2##left); \
218 EIGEN_UNUSED_VARIABLE(plhsV##left); \
219 }
220
221#define MICRO_MMA_LOAD_TWO(left) MICRO_MMA_LOAD1_TWO(lhs_ptr, left)
222#endif
223
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) } \
227 }
228
229#define MICRO_MMA_LOAD_ONE_RHS1(peel, right) ploadRhsMMA(rhs_ptr##right + (accRows * peel), rhsV##right[peel]);
230
231#define MICRO_MMA_LOAD_ONE_RHS(peel) MICRO_MMA_UNROLL_ITER(MICRO_MMA_LOAD_ONE_RHS1, peel)
232
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) \
239 }
240
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)
251#else
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);
255
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) \
262 } else { \
263 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
264 MICRO_MMA_LOAD_ONE_RHS(peel1) \
265 MICRO_MMA_LOAD_ONE_RHS(peel2) \
266 } \
267 MICRO_MMA_UNROLL(funcl2) \
268 MICRO_MMA_WORK(funcw2, type, peel1) \
269 MICRO_MMA_WORK(funcw2, type, peel2) \
270 } else { \
271 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
272 MICRO_MMA_TYPE_PEEL(funcw1, funcl1, type, peel1) \
273 }
274
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)
282#endif
283
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)
287
288#define MICRO_MMA_UPDATE_RHS1(size, right) rhs_ptr##right += (accRows * size);
289
290#define MICRO_MMA_UPDATE_RHS(size) MICRO_MMA_UNROLL_ITER(MICRO_MMA_UPDATE_RHS1, size)
291
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)
295
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)
299
300#ifndef VECTOR_PAIR_LOADS_LHS
301#define MICRO_MMA_ONE_PEEL MICRO_MMA_UNROLL_TYPE(MICRO_MMA_UNROLL_TYPE_PEEL, PEEL_MMA)
302#else
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)
306
307#define MICRO_MMA_ONE_PEEL MICRO_MMA_UNROLL_TYPE2(MICRO_MMA_UNROLL_TYPE_PEEL2, PEEL_MMA)
308#endif
309
310#define MICRO_MMA_ONE MICRO_MMA_UNROLL_TYPE(MICRO_MMA_UNROLL_TYPE_ONE, 1)
311
312#define MICRO_MMA_ONE_PARTIAL MICRO_MMA_UNROLL_TYPE_PARTIAL(MICRO_MMA_UNROLL_TYPE_ONE, 1)
313
314#define MICRO_MMA_DST_PTR_ONE(iter) \
315 if (unroll_factor * accItr > iter) { \
316 bsetzeroMMA(&accZero##iter); \
317 } else { \
318 EIGEN_UNUSED_VARIABLE(accZero##iter); \
319 }
320
321#define MICRO_MMA_DST_PTR MICRO_MMA_UNROLL(MICRO_MMA_DST_PTR_ONE)
322
323#define MICRO_MMA_SRC_PTR MICRO_MMA_UNROLL(MICRO_SRC_PTR_ONE)
324
325#define MICRO_MMA_PREFETCH MICRO_MMA_UNROLL(MICRO_PREFETCH_ONE)
326
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); \
331 }
332
333#define MICRO_MMA_ITER_UNROLL(func) \
334 if (accItr == 1) { \
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) \
338 } else { \
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) \
340 }
341
342#define MICRO_MMA_STORE MICRO_MMA_ITER_UNROLL(MICRO_MMA_STORE_ONE)
343
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);
348
349#define MICRO_MMA_EXTRA_ROWS1(val, right) MICRO_MMA_EXTRA_ROWS(right);
350
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;
362
363 if (accItr > 1) {
364 rhs_ptr1 = rhs_base + (accRows * strideB);
365 } else {
366 EIGEN_UNUSED_VARIABLE(strideB);
367 EIGEN_UNUSED_VARIABLE(rhs_ptr1);
368 EIGEN_UNUSED_VARIABLE(res1);
369 }
370 if (accItr > 2) {
371 rhs_ptr2 = rhs_base + (2 * accRows * strideB);
372 rhs_ptr3 = rhs_base + (3 * accRows * strideB);
373 } else {
374 EIGEN_UNUSED_VARIABLE(rhs_ptr2);
375 EIGEN_UNUSED_VARIABLE(rhs_ptr3);
376 EIGEN_UNUSED_VARIABLE(res2);
377 EIGEN_UNUSED_VARIABLE(res3);
378 }
379
380 MICRO_MMA_SRC_PTR
381 MICRO_MMA_DST_PTR
382
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);
387 MICRO_MMA_PREFETCH
388 MICRO_MMA_ONE_PEEL
389 }
390 for (; k < peel_depth; k++) {
391 MICRO_MMA_ONE
392 }
393 EIGEN_IF_CONSTEXPR (!full) {
394 for (; k < depth; k++) {
395 MICRO_MMA_ONE_PARTIAL
396 }
397 }
398 MICRO_MMA_STORE
399
400 MICRO_UPDATE
401}
402
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); \
407 if (M) return;
408
409#define MICRO_MMA_ROWS(n) \
410 while (row + n * accCols <= rows) { \
411 MICRO_MMA_UNROLL_ITER2(n, 0); \
412 }
413
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;
423
424 const Scalar* rhs_base = blockB + col * strideB + accRows * offsetB;
425 const Scalar* lhs_base = blockA + accCols * offsetA;
426 Index row = 0;
427
428#define MAX_MMA_UNROLL 7
429
430#if MAX_MMA_UNROLL < 2
431 if (1) {
432#elif MAX_MMA_UNROLL < 4
433 if (accItr <= 2) {
434#else
435 if (accItr == 1) {
436#endif
437 MICRO_MMA_ROWS(MAX_MMA_UNROLL);
438 } else if (accItr == 2) {
439 MICRO_MMA_ROWS(4);
440 } else {
441 MICRO_MMA_ROWS(2);
442 }
443 switch ((rows - row) / accCols) {
444#if MAX_MMA_UNROLL > 7
445 case 7:
446 if (accItr == 1) {
447 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 7)
448 }
449 break;
450#endif
451#if MAX_MMA_UNROLL > 6
452 case 6:
453 if (accItr == 1) {
454 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 6)
455 }
456 break;
457#endif
458#if MAX_MMA_UNROLL > 5
459 case 5:
460 if (accItr == 1) {
461 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 5)
462 }
463 break;
464#endif
465#if MAX_MMA_UNROLL > 4
466 case 4:
467 if (accItr == 1) {
468 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 4)
469 }
470 break;
471#endif
472#if MAX_MMA_UNROLL > 3
473 case 3:
474 if (accItr <= 2) {
475 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 3)
476 }
477 break;
478#endif
479#if MAX_MMA_UNROLL > 2
480 case 2:
481 if (accItr <= 2) {
482 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 2)
483 }
484 break;
485#endif
486#if MAX_MMA_UNROLL > 1
487 case 1:
488 MICRO_UNROLL_ITER(MICRO_MMA_UNROLL_ITER2, 1)
489 break;
490#endif
491 default:
492 break;
493 }
494#undef MAX_MMA_UNROLL
495
496 if (remaining_rows > 0) {
497 MICRO_MMA_UNROLL_ITER(MICRO_MMA_EXTRA_ROWS1, 0)
498 }
499}
500
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); \
505 }
506
507template <typename Scalar, typename Packet, typename RhsPacket, typename DataMapper, const Index accRows,
508 const Index accCols>
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;
512
513 if (strideA == -1) strideA = depth;
514 if (strideB == -1) strideB = depth;
515
516 const Packet pAlpha = pset1<Packet>(alpha);
517 const Packet pMask = bmask<Packet>(remaining_rows);
518
519 typedef std::conditional_t<(sizeof(Scalar) == sizeof(float)), RhsPacket, __vector_pair> RhsPacket2;
520
521 Index col = 0;
522#ifdef GEMM_MULTIPLE_COLS
523 MICRO_MMA_COLS(4);
524 MICRO_MMA_COLS(2);
525#endif
526 MICRO_MMA_COLS(1);
527
528 if (col != 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);
531 }
532}
533
534#define advanceRows ((LhsIsReal) ? 1 : 2)
535#define advanceCols ((RhsIsReal) ? 1 : 2)
536
537// PEEL_COMPLEX_MMA loop factor.
538#ifdef GEMM_MULTIPLE_COLS
539#define PEEL_COMPLEX_MMA 4
540#else
541#define PEEL_COMPLEX_MMA 3
542#endif
543
544#define MICRO_COMPLEX_MMA_UNROLL(func) func(0) func(1) func(2) func(3)
545
546#define MICRO_COMPLEX_MMA_WORK(func, type, peel) \
547 if (accItr == 1) { \
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) \
551 } else { \
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) \
553 }
554
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]); \
559 }
560
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]); \
567 }
568
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); \
574 } else { \
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); \
578 } \
579 } else { \
580 EIGEN_UNUSED_VARIABLE(lhsVi2##left); \
581 EIGEN_UNUSED_VARIABLE(plhsVi##left); \
582 } \
583 MICRO_MMA_LOAD1_TWO(lhs_ptr_real, left)
584
585#define MICRO_COMPLEX_MMA_LOAD_TWO(left) MICRO_COMPLEX_MMA_LOAD1_TWO(lhs_ptr, left)
586#endif
587
588#define MICRO_COMPLEX_MMA_LOAD_RHS1(peel, right) \
589 ploadRhsMMA(rhs_ptr_real##right + (accRows * peel), rhsV##right[peel]); \
590 if (!RhsIsReal) { \
591 ploadRhsMMA(rhs_ptr_imag##right + (accRows * peel), rhsVi##right[peel]); \
592 }
593
594#define MICRO_COMPLEX_MMA_LOAD_ONE_RHS(peel) MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_LOAD_RHS1, peel)
595
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) \
603 }
604
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)
612#else
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); \
616 if (!RhsIsReal) { \
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); \
619 } else { \
620 EIGEN_UNUSED_VARIABLE(prhsVi##peel1); \
621 }
622
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) \
631 } else { \
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); \
636 } \
637 MICRO_COMPLEX_MMA_UNROLL(funcl2) \
638 MICRO_COMPLEX_MMA_WORK(funcw2, type, peel1) \
639 MICRO_COMPLEX_MMA_WORK(funcw2, type, peel2) \
640 } else { \
641 EIGEN_UNUSED_VARIABLE(prhsV##peel1); \
642 EIGEN_UNUSED_VARIABLE(prhsVi##peel1); \
643 MICRO_COMPLEX_MMA_TYPE_PEEL(funcw1, funcl1, type, peel1) \
644 }
645
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)
653#endif
654
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)
658
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);
662
663#define MICRO_COMPLEX_MMA_UPDATE_RHS(size) MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_UPDATE_RHS1, size)
664
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);
668
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);
672
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)
675#else
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);
680
681#define MICRO_COMPLEX_MMA_ONE_PEEL MICRO_COMPLEX_MMA_UNROLL_TYPE2(MICRO_COMPLEX_MMA_UNROLL_TYPE_PEEL2, PEEL_COMPLEX_MMA)
682#endif
683
684#define MICRO_COMPLEX_MMA_ONE MICRO_COMPLEX_MMA_UNROLL_TYPE(MICRO_COMPLEX_MMA_UNROLL_TYPE_ONE, 1)
685
686#define MICRO_COMPLEX_MMA_ONE_PARTIAL MICRO_COMPLEX_MMA_UNROLL_TYPE_PARTIAL(MICRO_COMPLEX_MMA_UNROLL_TYPE_ONE, 1)
687
688#define MICRO_COMPLEX_MMA_DST_PTR_ONE(iter) \
689 if (unroll_factor * accItr > iter) { \
690 bsetzeroMMA(&accReal##iter); \
691 bsetzeroMMA(&accImag##iter); \
692 } else { \
693 EIGEN_UNUSED_VARIABLE(accReal##iter); \
694 EIGEN_UNUSED_VARIABLE(accImag##iter); \
695 }
696
697#define MICRO_COMPLEX_MMA_DST_PTR MICRO_COMPLEX_MMA_UNROLL(MICRO_COMPLEX_MMA_DST_PTR_ONE)
698
699#define MICRO_COMPLEX_MMA_SRC_PTR MICRO_COMPLEX_MMA_UNROLL(MICRO_COMPLEX_SRC_PTR_ONE)
700
701#define MICRO_COMPLEX_MMA_PREFETCH MICRO_COMPLEX_MMA_UNROLL(MICRO_COMPLEX_PREFETCH_ONE)
702
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); \
707 }
708
709#define MICRO_COMPLEX_MMA_ITER_UNROLL(func) \
710 if (accItr == 1) { \
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) \
714 } else { \
715 func(0, 0, 0) func(1, 0, 1) func(2, 0, 2) func(3, 0, 3) \
716 }
717
718#define MICRO_COMPLEX_MMA_STORE MICRO_COMPLEX_MMA_ITER_UNROLL(MICRO_COMPLEX_MMA_STORE_ONE)
719
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, \
724 pAlphaImag, pMask);
725
726#define MICRO_COMPLEX_MMA_EXTRA_ROWS1(val, right) MICRO_COMPLEX_MMA_EXTRA_ROWS(right);
727
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;
741
742 if (!RhsIsReal) {
743 rhs_ptr_imag0 = rhs_base + accRows * strideB;
744 } else {
745 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag0);
746 }
747 if (accItr > 1) {
748 if (!RhsIsReal) {
749 rhs_ptr_real1 = rhs_base + (2 * accRows * strideB);
750 rhs_ptr_imag1 = rhs_base + (3 * accRows * strideB);
751 } else {
752 rhs_ptr_real1 = rhs_base + accRows * strideB;
753 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag1);
754 }
755 } else {
756 EIGEN_UNUSED_VARIABLE(rhs_ptr_real1);
757 EIGEN_UNUSED_VARIABLE(rhs_ptr_imag1);
758 EIGEN_UNUSED_VARIABLE(res1);
759 }
760 if (accItr > 2) {
761 if (!RhsIsReal) {
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);
766 } else {
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);
771 }
772 } else {
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);
779 }
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;
783
784 MICRO_COMPLEX_MMA_SRC_PTR
785 MICRO_COMPLEX_MMA_DST_PTR
786
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);
791 if (!RhsIsReal) {
792 EIGEN_POWER_PREFETCH(rhs_ptr_imag);
793 }
794 MICRO_COMPLEX_MMA_PREFETCH
795 MICRO_COMPLEX_MMA_ONE_PEEL
796 }
797 for (; k < peel_depth; k++) {
798 MICRO_COMPLEX_MMA_ONE
799 }
800 EIGEN_IF_CONSTEXPR (accCols != accCols2) {
801 for (; k < depth; k++) {
802 MICRO_COMPLEX_MMA_ONE_PARTIAL
803 }
804 }
805 MICRO_COMPLEX_MMA_STORE
806
807 MICRO_COMPLEX_UPDATE
808}
809
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); \
815 if (M) return;
816
817#define MICRO_COMPLEX_MMA_ROWS(n) \
818 while (row + n * accCols <= rows) { \
819 MICRO_COMPLEX_MMA_UNROLL_ITER2(n, 0); \
820 }
821
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;
833
834 const Scalar* rhs_base = blockB + advanceCols * col * strideB + accRows * offsetB;
835 const Scalar* lhs_base = blockA + accCols * offsetA;
836 Index row = 0;
837
838#define MAX_COMPLEX_MMA_UNROLL 4
839
840#if MAX_COMPLEX_MMA_UNROLL < 2
841 if (1) {
842#elif MAX_COMPLEX_MMA_UNROLL < 4
843 if (accItr <= 2) {
844#else
845 if (accItr == 1) {
846#endif
847 MICRO_COMPLEX_MMA_ROWS(MAX_COMPLEX_MMA_UNROLL);
848 } else if (accItr == 2) {
849 MICRO_COMPLEX_MMA_ROWS(2);
850 } else {
851 MICRO_COMPLEX_MMA_ROWS(1);
852 }
853 switch ((rows - row) / accCols) {
854#if MAX_COMPLEX_MMA_UNROLL > 3
855 case 3:
856 if (accItr == 1) {
857 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_MMA_UNROLL_ITER2, 3)
858 }
859 break;
860#endif
861#if MAX_COMPLEX_MMA_UNROLL > 2
862 case 2:
863 if (accItr == 1) {
864 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_MMA_UNROLL_ITER2, 2)
865 }
866 break;
867#endif
868#if MAX_COMPLEX_MMA_UNROLL > 1
869 case 1:
870 if (accItr <= 2) {
871 MICRO_COMPLEX_UNROLL_ITER(MICRO_COMPLEX_MMA_UNROLL_ITER2, 1)
872 }
873 break;
874#endif
875 default:
876 break;
877 }
878#undef MAX_COMPLEX_MMA_UNROLL
879
880 if (remaining_rows > 0) {
881 MICRO_MMA_UNROLL_ITER(MICRO_COMPLEX_MMA_EXTRA_ROWS1, 0)
882 }
883}
884
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); \
891 }
892
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;
899
900 if (strideA == -1) strideA = depth;
901 if (strideB == -1) strideB = depth;
902
903 const Packet pAlphaReal = pset1<Packet>(alpha.real());
904 const Packet pAlphaImag = pset1<Packet>(alpha.imag());
905 const Packet pMask = bmask<Packet>(remaining_rows);
906
907 const Scalar* blockA = (Scalar*)blockAc;
908 const Scalar* blockB = (Scalar*)blockBc;
909
910 typedef std::conditional_t<(sizeof(Scalar) == sizeof(float)), RhsPacket, __vector_pair> RhsPacket2;
911
912 Index col = 0;
913#ifdef GEMM_MULTIPLE_COLS
914 MICRO_COMPLEX_MMA_COLS(4);
915 MICRO_COMPLEX_MMA_COLS(2);
916#endif
917 MICRO_COMPLEX_MMA_COLS(1);
918
919 if (col != cols) {
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);
923 }
924}
925
926#undef accColsC
927#undef advanceRows
928#undef advanceCols
929
930} // end namespace internal
931
932} // end namespace Eigen
933
934#if defined(EIGEN_ALTIVEC_MMA_DYNAMIC_DISPATCH)
935#pragma GCC pop_options
936#endif
937
938#endif // EIGEN_MATRIX_PRODUCT_MMA_ALTIVEC_H