Eigen  5.0.1
 
Loading...
Searching...
No Matches
MatrixVectorProduct.inc
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2021 Chip Kerchner (chip.kerchner@ibm.com)
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_MATRIX_VECTOR_PRODUCT_ALTIVEC_H
12#define EIGEN_MATRIX_VECTOR_PRODUCT_ALTIVEC_H
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17#if defined(__MMA__) && !EIGEN_ALTIVEC_DISABLE_MMA
18#if EIGEN_COMP_LLVM || (__GNUC__ > 10 || __GNUC_MINOR__ >= 3)
19#define USE_GEMV_MMA
20#endif
21
22#if !EIGEN_COMP_LLVM && (__GNUC__ < 11)
23// Only allow one vector_pair in buggy gcc - gcc 10.x has a bug
24#define GCC_ONE_VECTORPAIR_BUG
25#endif
26#endif
27
28// #define USE_SLOWER_GEMV_MMA // MMA is currently not as fast as VSX in complex double GEMV (revisit when gcc is
29// improved)
30
31// #define EIGEN_POWER_USE_GEMV_PREFETCH
32#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
33#define EIGEN_POWER_GEMV_PREFETCH(p) prefetch(p)
34#else
35#define EIGEN_POWER_GEMV_PREFETCH(p)
36#endif
37
38#ifdef __has_builtin
39#if !__has_builtin(__builtin_vsx_assemble_pair)
40#define __builtin_vsx_assemble_pair __builtin_mma_assemble_pair
41#endif
42#if !__has_builtin(__builtin_vsx_disassemble_pair)
43#define __builtin_vsx_disassemble_pair __builtin_mma_disassemble_pair
44#endif
45#endif
46
47#if EIGEN_COMP_LLVM
48#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
49 __builtin_vsx_assemble_pair(&dst, (__vector unsigned char)src2, (__vector unsigned char)src1)
50#else
51#if (__GNUC__ <= 10)
52#if (__GNUC_MINOR__ > 3)
53#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
54 __builtin_vsx_assemble_pair(&dst, (__vector unsigned char)src2, (__vector unsigned char)src1)
55#else
56#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
57 __builtin_vsx_assemble_pair(&dst, (__vector unsigned char)src1, (__vector unsigned char)src2)
58#endif
59#else
60#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
61 __builtin_vsx_build_pair(&dst, (__vector unsigned char)src1, (__vector unsigned char)src2)
62#endif
63#endif
64
65#define GEMV_IS_COMPLEX_COMPLEX ((sizeof(LhsPacket) == 16) && (sizeof(RhsPacket) == 16))
66#define GEMV_IS_FLOAT (ResPacketSize == (16 / sizeof(float)))
67#define GEMV_IS_SCALAR (sizeof(ResPacket) != 16)
68#define GEMV_IS_COMPLEX_FLOAT (ResPacketSize == (16 / sizeof(std::complex<float>)))
69
71template <typename ResPacket, typename ResScalar>
72EIGEN_ALWAYS_INLINE void storeMaddData(ResScalar* res, const ResPacket& palpha, const ResPacket& data) {
73 pstoreu(res, pmadd(data, palpha, ploadu<ResPacket>(res)));
74}
75
76template <typename ResScalar>
77EIGEN_ALWAYS_INLINE void storeMaddData(ResScalar* res, ResScalar& alpha, ResScalar& data) {
78 *res += (alpha * data);
79}
80
81#define GEMV_UNROLL(func, N) func(0, N) func(1, N) func(2, N) func(3, N) func(4, N) func(5, N) func(6, N) func(7, N)
82
83#define GEMV_UNROLL_HALF(func, N) func(0, 0, 1, N) func(1, 2, 3, N) func(2, 4, 5, N) func(3, 6, 7, N)
84
85#define GEMV_GETN(N) (((N) * ResPacketSize) >> 2)
86
87#define GEMV_LOADPACKET_COL(iter) lhs.template load<LhsPacket, LhsAlignment>(i + ((iter) * LhsPacketSize), j)
88
89#ifdef USE_GEMV_MMA
90#define GEMV_UNROLL3(func, N, which) \
91 func(0, N, which) func(1, N, which) func(2, N, which) func(3, N, which) func(4, N, which) func(5, N, which) \
92 func(6, N, which) func(7, N, which)
93
94#define GEMV_UNUSED_VAR(iter, N, which) \
95 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) <= iter) { \
96 EIGEN_UNUSED_VARIABLE(which##iter); \
97 }
98
99#define GEMV_UNUSED_EXTRA_VAR(iter, N, which) \
100 EIGEN_IF_CONSTEXPR (N <= iter) { \
101 EIGEN_UNUSED_VARIABLE(which##iter); \
102 }
103
104#define GEMV_UNUSED_EXTRA(N, which) GEMV_UNROLL3(GEMV_UNUSED_EXTRA_VAR, N, which)
105
106#define GEMV_UNUSED(N, which) GEMV_UNROLL3(GEMV_UNUSED_VAR, N, which)
107
108#define GEMV_INIT_MMA(iter, N) \
109 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
110 __builtin_mma_xxsetaccz(&e##iter); \
111 }
112
113#if EIGEN_COMP_LLVM
114#define GEMV_LOADPAIR_COL_MMA(iter1, iter2) \
115 GEMV_BUILDPAIR_MMA(b##iter1, GEMV_LOADPACKET_COL(iter2), GEMV_LOADPACKET_COL((iter2) + 1));
116#else
117#define GEMV_LOADPAIR_COL_MMA(iter1, iter2) \
118 const LhsScalar& src##iter1 = lhs(i + ((iter1 * 32) / sizeof(LhsScalar)), j); \
119 b##iter1 = *reinterpret_cast<__vector_pair*>(const_cast<LhsScalar*>(&src##iter1));
120#endif
121
122#define GEMV_LOAD1A_COL_MMA(iter, N) \
123 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
124 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
125 g##iter = GEMV_LOADPACKET_COL(iter); \
126 EIGEN_UNUSED_VARIABLE(b##iter); \
127 } else { \
128 GEMV_LOADPAIR_COL_MMA(iter, iter << 1) \
129 EIGEN_UNUSED_VARIABLE(g##iter); \
130 } \
131 } else { \
132 EIGEN_UNUSED_VARIABLE(b##iter); \
133 EIGEN_UNUSED_VARIABLE(g##iter); \
134 }
135
136#define GEMV_WORK1A_COL_MMA(iter, N) \
137 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
138 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
139 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter, a0, g##iter); \
140 } else { \
141 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter, b##iter, a0); \
142 } \
143 }
144
145#define GEMV_LOAD1B_COL_MMA(iter1, iter2, iter3, N) \
146 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter1) { \
147 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
148 GEMV_LOADPAIR_COL_MMA(iter2, iter2) \
149 EIGEN_UNUSED_VARIABLE(b##iter3); \
150 } else { \
151 GEMV_LOADPAIR_COL_MMA(iter2, iter2 << 1) \
152 GEMV_LOADPAIR_COL_MMA(iter3, iter3 << 1) \
153 } \
154 } else { \
155 EIGEN_UNUSED_VARIABLE(b##iter2); \
156 EIGEN_UNUSED_VARIABLE(b##iter3); \
157 } \
158 EIGEN_UNUSED_VARIABLE(g##iter2); \
159 EIGEN_UNUSED_VARIABLE(g##iter3);
160
161#define GEMV_WORK1B_COL_MMA(iter1, iter2, iter3, N) \
162 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter1) { \
163 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
164 LhsPacket h[2]; \
165 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(h), &b##iter2); \
166 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter2, a0, h[0]); \
167 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter3, a0, h[1]); \
168 } else { \
169 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter2, b##iter2, a0); \
170 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter3, b##iter3, a0); \
171 } \
172 }
173
174#if EIGEN_COMP_LLVM
175#define GEMV_LOAD_COL_MMA(N) \
176 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > 1) { \
177 GEMV_UNROLL_HALF(GEMV_LOAD1B_COL_MMA, (N >> 1)) \
178 } else { \
179 GEMV_UNROLL(GEMV_LOAD1A_COL_MMA, N) \
180 }
181
182#define GEMV_WORK_COL_MMA(N) \
183 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > 1) { \
184 GEMV_UNROLL_HALF(GEMV_WORK1B_COL_MMA, (N >> 1)) \
185 } else { \
186 GEMV_UNROLL(GEMV_WORK1A_COL_MMA, N) \
187 }
188#else
189#define GEMV_LOAD_COL_MMA(N) GEMV_UNROLL(GEMV_LOAD1A_COL_MMA, N)
190
191#define GEMV_WORK_COL_MMA(N) GEMV_UNROLL(GEMV_WORK1A_COL_MMA, N)
192#endif
193
194#define GEMV_DISASSEMBLE_MMA(iter, N) \
195 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
196 __builtin_mma_disassemble_acc(&result##iter.packet, &e##iter); \
197 EIGEN_IF_CONSTEXPR (!GEMV_IS_FLOAT) { \
198 result##iter.packet[0][1] = result##iter.packet[1][0]; \
199 result##iter.packet[2][1] = result##iter.packet[3][0]; \
200 } \
201 }
202
203#define GEMV_LOADPAIR2_COL_MMA(iter1, iter2) \
204 b##iter1 = *reinterpret_cast<__vector_pair*>(res + i + ((iter2) * ResPacketSize));
205
206#define GEMV_LOAD2_COL_MMA(iter1, iter2, iter3, N) \
207 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter1) { \
208 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
209 GEMV_LOADPAIR2_COL_MMA(iter2, iter2); \
210 EIGEN_UNUSED_VARIABLE(b##iter3); \
211 } else { \
212 GEMV_LOADPAIR2_COL_MMA(iter2, iter2 << 1); \
213 GEMV_LOADPAIR2_COL_MMA(iter3, iter3 << 1); \
214 } \
215 } else { \
216 EIGEN_UNUSED_VARIABLE(b##iter2); \
217 EIGEN_UNUSED_VARIABLE(b##iter3); \
218 }
219
220#if EIGEN_COMP_LLVM
221#define GEMV_WORKPAIR2_COL_MMA(iter2, iter3, iter4) \
222 ResPacket f##iter2[2]; \
223 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(f##iter2), &b##iter2); \
224 f##iter2[0] = pmadd(result##iter2.packet[0], palpha, f##iter2[0]); \
225 f##iter2[1] = pmadd(result##iter3.packet[(iter2 == iter3) ? 2 : 0], palpha, f##iter2[1]); \
226 GEMV_BUILDPAIR_MMA(b##iter2, f##iter2[0], f##iter2[1]);
227#else
228#define GEMV_WORKPAIR2_COL_MMA(iter2, iter3, iter4) \
229 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
230 __asm__("xvmaddasp %0,%x1,%x3\n\txvmaddasp %L0,%x2,%x3" \
231 : "+&d"(b##iter2) \
232 : "wa"(result##iter3.packet[0]), "wa"(result##iter2.packet[0]), "wa"(palpha)); \
233 } else { \
234 __asm__("xvmaddadp %0,%x1,%x3\n\txvmaddadp %L0,%x2,%x3" \
235 : "+&d"(b##iter2) \
236 : "wa"(result##iter2.packet[2]), "wa"(result##iter2.packet[0]), "wa"(palpha)); \
237 }
238#endif
239
240#define GEMV_WORK2_COL_MMA(iter1, iter2, iter3, N) \
241 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter1) { \
242 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
243 GEMV_WORKPAIR2_COL_MMA(iter2, iter3, iter2); \
244 } else { \
245 GEMV_WORKPAIR2_COL_MMA(iter2, iter2, iter2 << 1); \
246 GEMV_WORKPAIR2_COL_MMA(iter3, iter3, iter3 << 1); \
247 } \
248 }
249
250#define GEMV_STOREPAIR2_COL_MMA(iter1, iter2) \
251 *reinterpret_cast<__vector_pair*>(res + i + ((iter2) * ResPacketSize)) = b##iter1;
252
253#define GEMV_STORE_COL_MMA(iter, N) \
254 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
255 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
256 storeMaddData<ResPacket, ResScalar>(res + i + (iter * ResPacketSize), palpha, result##iter.packet[0]); \
257 } else { \
258 GEMV_LOADPAIR2_COL_MMA(iter, iter << 1) \
259 GEMV_WORKPAIR2_COL_MMA(iter, iter, iter << 1) \
260 GEMV_STOREPAIR2_COL_MMA(iter, iter << 1) \
261 } \
262 }
263
264#define GEMV_STORE2_COL_MMA(iter1, iter2, iter3, N) \
265 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter1) { \
266 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
267 GEMV_STOREPAIR2_COL_MMA(iter2, iter2); \
268 } else { \
269 GEMV_STOREPAIR2_COL_MMA(iter2, iter2 << 1) \
270 GEMV_STOREPAIR2_COL_MMA(iter3, iter3 << 1) \
271 } \
272 }
273
274#define GEMV_PROCESS_COL_ONE_MMA(N) \
275 GEMV_UNROLL(GEMV_INIT_MMA, N) \
276 Index j = j2; \
277 __vector_pair b0, b1, b2, b3, b4, b5, b6, b7; \
278 do { \
279 LhsPacket g0, g1, g2, g3, g4, g5, g6, g7; \
280 RhsPacket a0 = pset1<RhsPacket>(rhs2(j, 0)); \
281 GEMV_UNROLL(GEMV_PREFETCH, N) \
282 GEMV_LOAD_COL_MMA(N) \
283 GEMV_WORK_COL_MMA(N) \
284 } while (++j < jend); \
285 GEMV_UNROLL(GEMV_DISASSEMBLE_MMA, N) \
286 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) <= 1) { \
287 GEMV_UNROLL(GEMV_STORE_COL_MMA, N) \
288 } else { \
289 GEMV_UNROLL_HALF(GEMV_LOAD2_COL_MMA, (N >> 1)) \
290 GEMV_UNROLL_HALF(GEMV_WORK2_COL_MMA, (N >> 1)) GEMV_UNROLL_HALF(GEMV_STORE2_COL_MMA, (N >> 1)) \
291 } \
292 i += (ResPacketSize * N);
293#endif
294
295#define GEMV_INIT(iter, N) \
296 EIGEN_IF_CONSTEXPR (N > iter) { \
297 c##iter = pset1<ResPacket>(ResScalar(0)); \
298 } else { \
299 EIGEN_UNUSED_VARIABLE(c##iter); \
300 }
301
302#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
303#define GEMV_PREFETCH(iter, N) \
304 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > ((iter >> 1) + ((N >> 1) * (iter & 1)))) { \
305 lhs.prefetch(i + (iter * LhsPacketSize) + prefetch_dist, j); \
306 }
307#else
308#define GEMV_PREFETCH(iter, N)
309#endif
310
311#define GEMV_WORK_COL(iter, N) \
312 EIGEN_IF_CONSTEXPR (N > iter) { \
313 c##iter = pcj.pmadd(GEMV_LOADPACKET_COL(iter), a0, c##iter); \
314 }
315
316#define GEMV_STORE_COL(iter, N) \
317 EIGEN_IF_CONSTEXPR (N > iter) { \
318 pstoreu(res + i + (iter * ResPacketSize), \
319 pmadd(c##iter, palpha, ploadu<ResPacket>(res + i + (iter * ResPacketSize)))); \
320 }
321
323#define GEMV_PROCESS_COL_ONE(N) \
324 GEMV_UNROLL(GEMV_INIT, N) \
325 Index j = j2; \
326 do { \
327 RhsPacket a0 = pset1<RhsPacket>(rhs2(j, 0)); \
328 GEMV_UNROLL(GEMV_PREFETCH, N) \
329 GEMV_UNROLL(GEMV_WORK_COL, N) \
330 } while (++j < jend); \
331 GEMV_UNROLL(GEMV_STORE_COL, N) \
332 i += (ResPacketSize * N);
333
334#ifdef USE_GEMV_MMA
335#define GEMV_PROCESS_COL(N) GEMV_PROCESS_COL_ONE_MMA(N)
336#else
337#define GEMV_PROCESS_COL(N) GEMV_PROCESS_COL_ONE(N)
338#endif
339
341#ifdef USE_GEMV_MMA
342template <typename LhsPacket, typename RhsPacket, bool accumulate>
343EIGEN_ALWAYS_INLINE void pger_vecMMA_acc(__vector_quad* acc, const RhsPacket& a, const LhsPacket& b) {
344 if (accumulate) {
345 __builtin_mma_xvf32gerpp(acc, (__vector unsigned char)a, (__vector unsigned char)b);
346 } else {
347 __builtin_mma_xvf32ger(acc, (__vector unsigned char)a, (__vector unsigned char)b);
348 }
349}
350
352template <typename LhsPacket, typename RhsPacket, bool accumulate>
353EIGEN_ALWAYS_INLINE void pger_vecMMA_acc(__vector_quad* acc, __vector_pair& a, const LhsPacket& b) {
354 if (accumulate) {
355 __builtin_mma_xvf64gerpp(acc, a, (__vector unsigned char)b);
356 } else {
357 __builtin_mma_xvf64ger(acc, a, (__vector unsigned char)b);
358 }
359}
360#endif
361
362template <typename LhsScalar, typename LhsMapper, typename RhsScalar, typename RhsMapper, typename ResScalar>
363EIGEN_STRONG_INLINE void gemv_col(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs, ResScalar* res,
364 Index resIncr, ResScalar alpha) {
365 typedef gemv_traits<LhsScalar, RhsScalar> Traits;
366
367 typedef typename Traits::LhsPacket LhsPacket;
368 typedef typename Traits::RhsPacket RhsPacket;
369 typedef typename Traits::ResPacket ResPacket;
370
371 EIGEN_UNUSED_VARIABLE(resIncr);
372 eigen_internal_assert(resIncr == 1);
373
374 // The following copy tells the compiler that lhs's attributes are not modified outside this function
375 // This helps GCC to generate proper code.
376 LhsMapper lhs(alhs);
377 RhsMapper rhs2(rhs);
378
379 conj_helper<LhsScalar, RhsScalar, false, false> cj;
380 conj_helper<LhsPacket, RhsPacket, false, false> pcj;
381
382 const Index lhsStride = lhs.stride();
383 // LhsAlignment stays Unaligned; enabling aligned reads would require
384 // propagating the Mapper's Alignment through the run() template.
385 enum {
386 LhsAlignment = Unaligned,
387 ResPacketSize = Traits::ResPacketSize,
388 LhsPacketSize = Traits::LhsPacketSize,
389 RhsPacketSize = Traits::RhsPacketSize,
390 };
391
392#ifndef GCC_ONE_VECTORPAIR_BUG
393 const Index n8 = rows - 8 * ResPacketSize + 1;
394 const Index n4 = rows - 4 * ResPacketSize + 1;
395 const Index n2 = rows - 2 * ResPacketSize + 1;
396#endif
397 const Index n1 = rows - 1 * ResPacketSize + 1;
398#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
399 const Index prefetch_dist = 64 * LhsPacketSize;
400#endif
401
402 // TODO: improve the following heuristic:
403 const Index block_cols = cols < 128 ? cols : (lhsStride * sizeof(LhsScalar) < 16000 ? 16 : 8);
404 ResPacket palpha = pset1<ResPacket>(alpha);
405
406 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
407 Index jend = numext::mini(j2 + block_cols, cols);
408 Index i = 0;
409 ResPacket c0, c1, c2, c3, c4, c5, c6, c7;
410#ifdef USE_GEMV_MMA
411 __vector_quad e0, e1, e2, e3, e4, e5, e6, e7;
412 PacketBlock<ResPacket, 4> result0, result1, result2, result3, result4, result5, result6, result7;
413 GEMV_UNUSED(8, e)
414 GEMV_UNUSED(8, result)
415 GEMV_UNUSED_EXTRA(1, c)
416#endif
417#ifndef GCC_ONE_VECTORPAIR_BUG
418 while (i < n8) {
419 GEMV_PROCESS_COL(8)
420 }
421 if (i < n4) {
422 GEMV_PROCESS_COL(4)
423 }
424 if (i < n2) {
425 GEMV_PROCESS_COL(2)
426 }
427 if (i < n1)
428#else
429 while (i < n1)
430#endif
431 {
432 GEMV_PROCESS_COL_ONE(1)
433 }
434 for (; i < rows; ++i) {
435 ResScalar d0(0);
436 Index j = j2;
437 do {
438 d0 += cj.pmul(lhs(i, j), rhs2(j, 0));
439 } while (++j < jend);
440 res[i] += alpha * d0;
441 }
442 }
443}
444
445template <bool extraRows>
446EIGEN_ALWAYS_INLINE void outputVecCol(const Packet4f& acc, float* result, const Packet4f& pAlpha, Index extra_rows) {
447 Packet4f d0 = ploadu<Packet4f>(result);
448 d0 = pmadd(acc, pAlpha, d0);
449 EIGEN_IF_CONSTEXPR (extraRows) {
450 pstoreu_partial(result, d0, extra_rows);
451 } else {
452 pstoreu(result, d0);
453 }
454}
455
456template <Index num_acc, bool extraRows, Index size>
457EIGEN_ALWAYS_INLINE void outputVecColResults(Packet4f (&acc)[num_acc][size], float* result, const Packet4f& pAlpha,
458 Index extra_rows) {
459 constexpr Index real_acc = (num_acc - (extraRows ? 1 : 0));
460 for (Index k = 0; k < real_acc; k++) {
461 outputVecCol<false>(acc[k][0], result + k * 4, pAlpha, extra_rows);
462 }
463 EIGEN_IF_CONSTEXPR (extraRows) {
464 outputVecCol<true>(acc[real_acc][0], result + real_acc * 4, pAlpha, extra_rows);
465 }
466}
467
468static Packet16uc p16uc_MERGE16_32_V1 = {0, 1, 16, 17, 0, 1, 16, 17, 0, 1, 16, 17, 0, 1, 16, 17};
469static Packet16uc p16uc_MERGE16_32_V2 = {2, 3, 18, 19, 2, 3, 18, 19, 2, 3, 18, 19, 2, 3, 18, 19};
470
471template <Index num_acc, typename LhsMapper, bool zero>
472EIGEN_ALWAYS_INLINE void loadVecLoopVSX(Index k, LhsMapper& lhs, Packet4f (&a0)[num_acc][2]) {
473 Packet8bf c0 = lhs.template loadPacket<Packet8bf>(k * 4, 0);
474 Packet8bf b1;
475 EIGEN_IF_CONSTEXPR (!zero) {
476 b1 = lhs.template loadPacket<Packet8bf>(k * 4, 1);
477
478 a0[k + 0][1] = oneConvertBF16Hi(b1.m_val);
479 }
480 a0[k + 0][0] = oneConvertBF16Hi(c0.m_val);
481
482 if (num_acc > (k + 1)) {
483 a0[k + 1][0] = oneConvertBF16Lo(c0.m_val);
484 EIGEN_IF_CONSTEXPR (!zero) {
485 a0[k + 1][1] = oneConvertBF16Lo(b1.m_val);
486 }
487 }
488}
489
490template <Index num_acc, bool zero>
491EIGEN_ALWAYS_INLINE void multVecVSX(Packet4f (&acc)[num_acc][2], Packet4f (&a0)[num_acc][2], Packet4f (&b0)[2]) {
492 for (Index k = 0; k < num_acc; k++) {
493 for (Index i = 0; i < (zero ? 1 : 2); i++) {
494 acc[k][i] = pmadd(b0[i], a0[k][i], acc[k][i]);
495 }
496 }
497}
498
499template <typename RhsMapper, bool linear>
500struct loadColData_impl {
501 // linear == false
502 static EIGEN_ALWAYS_INLINE Packet8bf run(RhsMapper& rhs, Index j) {
503 const Index n = unpacket_traits<Packet8bf>::size;
504 EIGEN_ALIGN16 bfloat16 to[n];
505 LOAD_STORE_UNROLL_16
506 for (Index i = 0; i < n; i++) {
507 to[i] = rhs(j + i, 0);
508 }
509 return pload<Packet8bf>(to);
510 }
511};
512
513template <typename RhsMapper>
514struct loadColData_impl<RhsMapper, true> {
515 // linear == true
516 static EIGEN_ALWAYS_INLINE Packet8bf run(RhsMapper& rhs, Index j) {
517 return rhs.template loadPacket<Packet8bf>(j + 0, 0);
518 }
519};
520
521template <typename RhsMapper, bool linear>
522EIGEN_ALWAYS_INLINE Packet8bf loadColData(RhsMapper& rhs, Index j) {
523 return loadColData_impl<RhsMapper, linear>::run(rhs, j);
524}
525
526template <Index num_acc, typename LhsMapper, typename RhsMapper, bool zero, bool linear>
527EIGEN_ALWAYS_INLINE void vecColLoopVSX(Index j, LhsMapper& lhs, RhsMapper& rhs, Packet4f (&acc)[num_acc][2]) {
528 Packet4f a0[num_acc][2], b0[2];
529 Packet8bf b2 = loadColData<RhsMapper, linear>(rhs, j);
530
531 b0[0] = oneConvertBF16Perm(b2.m_val, p16uc_MERGE16_32_V1);
532 EIGEN_IF_CONSTEXPR (!zero) {
533 b0[1] = oneConvertBF16Perm(b2.m_val, p16uc_MERGE16_32_V2);
534 }
535
536 using LhsSubMapper = typename LhsMapper::SubMapper;
537
538 LhsSubMapper lhs2 = lhs.getSubMapper(0, j);
539 for (Index k = 0; k < num_acc; k += 2) {
540 loadVecLoopVSX<num_acc, LhsSubMapper, zero>(k, lhs2, a0);
541 }
542
543 multVecVSX<num_acc, zero>(acc, a0, b0);
544}
545
546template <Index num_acc>
547EIGEN_ALWAYS_INLINE void addResultsVSX(Packet4f (&acc)[num_acc][2]) {
548 for (Index i = 0; i < num_acc; i++) {
549 acc[i][0] = acc[i][0] + acc[i][1];
550 }
551}
552
553// Uses 2X the accumulators or 4X the number of VSX registers
554#define MAX_BFLOAT16_VEC_ACC_VSX 8
555
556template <const Index num_acc, typename LhsMapper, typename RhsMapper, bool extraRows, bool linear>
557void colVSXVecColLoopBody(Index& row, Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
558 float* result) {
559 constexpr Index step = (num_acc * 4);
560 const Index extra_rows = (extraRows) ? (rows & 3) : 0;
561 constexpr bool multiIters = !extraRows && (num_acc == MAX_BFLOAT16_VEC_ACC_VSX);
562
563 do {
564 Packet4f acc[num_acc][2];
565
566 zeroAccumulators<num_acc, 2>(acc);
567
568 using LhsSubMapper = typename LhsMapper::SubMapper;
569
570 LhsSubMapper lhs2 = lhs.getSubMapper(row, 0);
571 for (Index j = 0; j + 2 <= cend; j += 2) {
572 vecColLoopVSX<num_acc, LhsSubMapper, RhsMapper, false, linear>(j, lhs2, rhs, acc);
573 }
574 if (cend & 1) {
575 vecColLoopVSX<num_acc, LhsSubMapper, RhsMapper, true, linear>(cend - 1, lhs2, rhs, acc);
576 }
577
578 addResultsVSX<num_acc>(acc);
579
580 outputVecColResults<num_acc, extraRows, 2>(acc, result, pAlpha, extra_rows);
581
582 result += step;
583 } while (multiIters && (step <= rows - (row += step)));
584}
585
586template <const Index num_acc, typename LhsMapper, typename RhsMapper, bool extraRows, bool linear>
587EIGEN_ALWAYS_INLINE void colVSXVecColLoopBodyExtraN(Index& row, Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
588 const Packet4f& pAlpha, float* result) {
589 EIGEN_IF_CONSTEXPR (MAX_BFLOAT16_VEC_ACC_VSX > num_acc) {
590 colVSXVecColLoopBody<num_acc + (extraRows ? 1 : 0), LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs,
591 rhs, pAlpha, result);
592 }
593}
594
595template <typename LhsMapper, typename RhsMapper, bool extraRows, bool linear>
596EIGEN_ALWAYS_INLINE void colVSXVecColLoopBodyExtra(Index& row, Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
597 const Packet4f& pAlpha, float* result) {
598 switch ((rows - row) >> 2) {
599 case 7:
600 colVSXVecColLoopBodyExtraN<7, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
601 break;
602 case 6:
603 colVSXVecColLoopBodyExtraN<6, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
604 break;
605 case 5:
606 colVSXVecColLoopBodyExtraN<5, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
607 break;
608 case 4:
609 colVSXVecColLoopBodyExtraN<4, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
610 break;
611 case 3:
612 colVSXVecColLoopBodyExtraN<3, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
613 break;
614 case 2:
615 colVSXVecColLoopBodyExtraN<2, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
616 break;
617 case 1:
618 colVSXVecColLoopBodyExtraN<1, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
619 break;
620 default:
621 EIGEN_IF_CONSTEXPR (extraRows) {
622 colVSXVecColLoopBody<1, LhsMapper, RhsMapper, true, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
623 }
624 break;
625 }
626}
627
628template <typename LhsMapper, typename RhsMapper, bool linear>
629EIGEN_ALWAYS_INLINE void calcVSXVecColLoops(Index cend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
630 const Packet4f& pAlpha, float* result) {
631 Index row = 0;
632 if (rows >= (MAX_BFLOAT16_VEC_ACC_VSX * 4)) {
633 colVSXVecColLoopBody<MAX_BFLOAT16_VEC_ACC_VSX, LhsMapper, RhsMapper, false, linear>(row, cend, rows, lhs, rhs,
634 pAlpha, result);
635 result += row;
636 }
637 if (rows & 3) {
638 colVSXVecColLoopBodyExtra<LhsMapper, RhsMapper, true, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
639 } else {
640 colVSXVecColLoopBodyExtra<LhsMapper, RhsMapper, false, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
641 }
642}
643
644template <const Index size, bool inc, Index delta>
645EIGEN_ALWAYS_INLINE void storeBF16fromResult(bfloat16* dst, const Packet8bf& data, Index resInc, Index extra) {
646 if (inc) {
647 EIGEN_IF_CONSTEXPR (size < 8) {
648 pscatter_partial(dst + delta * resInc, data, resInc, extra);
649 } else {
650 pscatter(dst + delta * resInc, data, resInc);
651 }
652 } else {
653 EIGEN_IF_CONSTEXPR (size < 8) {
654 pstoreu_partial(dst + delta, data, extra);
655 } else {
656 pstoreu(dst + delta, data);
657 }
658 }
659}
660
661template <const Index size, bool inc = false>
662EIGEN_ALWAYS_INLINE void convertPointerF32toBF16VSX(Index& i, float* result, Index rows, bfloat16*& dst,
663 Index resInc = 1) {
664 constexpr Index extra = ((size < 8) ? 8 : size);
665 while (i + size <= rows) {
666 PacketBlock<Packet8bf, (size + 7) / 8> r32;
667 r32.packet[0] = convertF32toBF16VSX(result + i + 0);
668 EIGEN_IF_CONSTEXPR (size >= 16) {
669 r32.packet[1] = convertF32toBF16VSX(result + i + 8);
670 }
671 EIGEN_IF_CONSTEXPR (size >= 32) {
672 r32.packet[2] = convertF32toBF16VSX(result + i + 16);
673 r32.packet[3] = convertF32toBF16VSX(result + i + 24);
674 }
675 storeBF16fromResult<size, inc, 0>(dst, r32.packet[0], resInc, rows & 7);
676 EIGEN_IF_CONSTEXPR (size >= 16) {
677 storeBF16fromResult<size, inc, 8>(dst, r32.packet[1], resInc);
678 }
679 EIGEN_IF_CONSTEXPR (size >= 32) {
680 storeBF16fromResult<size, inc, 16>(dst, r32.packet[2], resInc);
681 storeBF16fromResult<size, inc, 24>(dst, r32.packet[3], resInc);
682 }
683 i += extra;
684 dst += extra * resInc;
685 EIGEN_IF_CONSTEXPR (size != 32) break;
686 }
687}
688
689template <bool inc = false>
690EIGEN_ALWAYS_INLINE void convertArrayPointerF32toBF16VSX(float* result, Index rows, bfloat16* dst, Index resInc = 1) {
691 Index i = 0;
692 convertPointerF32toBF16VSX<32, inc>(i, result, rows, dst, resInc);
693 convertPointerF32toBF16VSX<16, inc>(i, result, rows, dst, resInc);
694 convertPointerF32toBF16VSX<8, inc>(i, result, rows, dst, resInc);
695 convertPointerF32toBF16VSX<1, inc>(i, result, rows, dst, resInc);
696}
697
698template <typename RhsMapper, typename LhsMapper, typename = void>
699struct UseStride : std::false_type {
700 static EIGEN_ALWAYS_INLINE void run(Index j2, Index jend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
701 const Packet4f& pAlpha, float* result) {
702 using RhsSubMapper = typename RhsMapper::SubMapper;
703
704 RhsSubMapper rhs2 = rhs.getSubMapper(j2, 0);
705 calcVSXVecColLoops<LhsMapper, RhsSubMapper, false>(jend - j2, rows, lhs, rhs2, pAlpha, result);
706 }
707};
708
709template <typename RhsMapper, typename LhsMapper>
710struct UseStride<RhsMapper, LhsMapper,
711 std::enable_if_t<std::is_member_function_pointer<decltype(&RhsMapper::stride)>::value>>
712 : std::true_type {
713 static EIGEN_ALWAYS_INLINE void run(Index j2, Index jend, Index rows, LhsMapper& lhs, RhsMapper& rhs,
714 const Packet4f& pAlpha, float* result) {
715 using RhsSubMapper = typename RhsMapper::SubMapper;
716
717 RhsSubMapper rhs2 = rhs.getSubMapper(j2, 0);
718 if (rhs.stride() == 1) {
719 calcVSXVecColLoops<LhsMapper, RhsSubMapper, true>(jend - j2, rows, lhs, rhs2, pAlpha, result);
720 } else {
721 calcVSXVecColLoops<LhsMapper, RhsSubMapper, false>(jend - j2, rows, lhs, rhs2, pAlpha, result);
722 }
723 }
724};
725
726template <typename LhsMapper, typename RhsMapper>
727void gemv_bfloat16_col(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs, bfloat16* res,
728 Index resIncr, bfloat16 alpha) {
729 EIGEN_UNUSED_VARIABLE(resIncr);
730 eigen_internal_assert(resIncr == 1);
731
732 // The following copy tells the compiler that lhs's attributes are not modified outside this function
733 // This helps GCC to generate proper code.
734 LhsMapper lhs(alhs);
735 RhsMapper rhs2(rhs);
736
737 const Index lhsStride = lhs.stride();
738
739 // TODO: improve the following heuristic:
740 const Index block_cols = cols < 128 ? cols : (lhsStride * sizeof(bfloat16) < 16000 ? 16 : 8);
741 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
742 Packet4f pAlpha = pset1<Packet4f>(falpha);
743
744 ei_declare_aligned_stack_constructed_variable(float, result, rows, 0);
745
746 convertArrayPointerBF16toF32(result, 1, rows, res);
747
748 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
749 Index jend = numext::mini(j2 + block_cols, cols);
750
751 using LhsSubMapper = typename LhsMapper::SubMapper;
752
753 LhsSubMapper lhs2 = lhs.getSubMapper(0, j2);
754 UseStride<RhsMapper, LhsSubMapper>::run(j2, jend, rows, lhs2, rhs2, pAlpha, result);
755 }
756
757 convertArrayPointerF32toBF16VSX(result, rows, res);
758}
759
760template <Index num_acc, Index size>
761EIGEN_ALWAYS_INLINE void outputVecResults(Packet4f (&acc)[num_acc][size], float* result, const Packet4f& pAlpha) {
762 constexpr Index extra = num_acc & 3;
763
764 for (Index k = 0; k < num_acc; k += 4) {
765 Packet4f d0 = ploadu<Packet4f>(result + k);
766 d0 = pmadd(acc[k + 0][0], pAlpha, d0);
767
768 if (num_acc > (k + 3)) {
769 pstoreu(result + k, d0);
770 } else {
771 if (extra == 3) {
772 pstoreu_partial(result + k, d0, extra);
773 } else {
774 memcpy((void*)(result + k), (void*)(&d0), sizeof(float) * extra);
775 }
776 }
777 }
778}
779
780template <Index num_acc>
781EIGEN_ALWAYS_INLINE void preduxVecResults2VSX(Packet4f (&acc)[num_acc][2], Index k) {
782 if (num_acc > (k + 1)) {
783 acc[k][1] = vec_mergel(acc[k + 0][0], acc[k + 1][0]);
784 acc[k][0] = vec_mergeh(acc[k + 0][0], acc[k + 1][0]);
785 acc[k][0] = acc[k][0] + acc[k][1];
786 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 8);
787 } else {
788 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 8);
789#ifdef _BIG_ENDIAN
790 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 12);
791#else
792 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 4);
793#endif
794 }
795}
796
797template <Index num_acc>
798EIGEN_ALWAYS_INLINE void preduxVecResultsVSX(Packet4f (&acc)[num_acc][2]) {
799 for (Index k = 0; k < num_acc; k += 4) {
800 preduxVecResults2VSX<num_acc>(acc, k + 0);
801 if (num_acc > (k + 2)) {
802 preduxVecResults2VSX<num_acc>(acc, k + 2);
803#ifdef EIGEN_VECTORIZE_VSX
804 acc[k + 0][0] = reinterpret_cast<Packet4f>(
805 vec_mergeh(reinterpret_cast<Packet2ul>(acc[k + 0][0]), reinterpret_cast<Packet2ul>(acc[k + 2][0])));
806#else
807 acc[k + 0][0] = reinterpret_cast<Packet4f>(vec_perm(acc[k + 0][0], acc[k + 2][0], p16uc_TRANSPOSE64_HI));
808#endif
809 }
810 }
811}
812
813#ifndef _ARCH_PWR9
814EIGEN_ALWAYS_INLINE Packet8us loadPacketPartialZero(const Packet8us& data, Index extra_cols) {
815 Packet16uc shift = pset1<Packet16uc>(8 * 2 * (8 - extra_cols));
816#ifdef _BIG_ENDIAN
817 return reinterpret_cast<Packet8us>(vec_slo(vec_sro(reinterpret_cast<Packet16uc>(data), shift), shift));
818#else
819 return reinterpret_cast<Packet8us>(vec_sro(vec_slo(reinterpret_cast<Packet16uc>(data), shift), shift));
820#endif
821}
822#endif
823
824template <Index num_acc, typename LhsMapper, typename RhsMapper, bool extra>
825EIGEN_ALWAYS_INLINE void multVSXVecLoop(Packet4f (&acc)[num_acc][2], const LhsMapper& lhs, RhsMapper& rhs, Index j,
826 Index extra_cols) {
827 Packet4f a0[num_acc][2], b0[2];
828 Packet8bf a1, b1;
829
830 if (extra) {
831 b1 = rhs.template loadPacketPartial<Packet8bf>(j, extra_cols);
832#ifndef _ARCH_PWR9
833 b1 = loadPacketPartialZero(b1.m_val, extra_cols);
834#endif
835 } else {
836 b1 = rhs.template loadPacket<Packet8bf>(j);
837 }
838 b0[0] = oneConvertBF16Hi(b1.m_val);
839 b0[1] = oneConvertBF16Lo(b1.m_val);
840
841 const LhsMapper lhs2 = lhs.getSubMapper(0, j);
842 for (Index k = 0; k < num_acc; k++) {
843 if (extra) {
844 a1 = lhs2.template loadPacketPartial<Packet8bf>(k, 0, extra_cols);
845#ifndef _ARCH_PWR9
846 a1 = loadPacketPartialZero(a1.m_val, extra_cols);
847#endif
848 } else {
849 a1 = lhs2.template loadPacket<Packet8bf>(k, 0);
850 }
851 a0[k][0] = oneConvertBF16Hi(a1.m_val);
852 a0[k][1] = oneConvertBF16Lo(a1.m_val);
853 }
854
855 multVecVSX<num_acc, false>(acc, a0, b0);
856}
857
858template <Index num_acc, typename LhsMapper, typename RhsMapper>
859EIGEN_ALWAYS_INLINE void vecVSXLoop(Index cols, const LhsMapper& lhs, RhsMapper& rhs, Packet4f (&acc)[num_acc][2],
860 Index extra_cols) {
861 Index j = 0;
862 for (; j + 8 <= cols; j += 8) {
863 multVSXVecLoop<num_acc, LhsMapper, RhsMapper, false>(acc, lhs, rhs, j, extra_cols);
864 }
865
866 if (extra_cols) {
867 multVSXVecLoop<num_acc, LhsMapper, RhsMapper, true>(acc, lhs, rhs, j, extra_cols);
868 }
869}
870
871template <const Index num_acc, typename LhsMapper, typename RhsMapper>
872void colVSXVecLoopBody(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
873 float* result) {
874 constexpr bool multiIters = (num_acc == MAX_BFLOAT16_VEC_ACC_VSX);
875 const Index extra_cols = (cols & 7);
876
877 do {
878 Packet4f acc[num_acc][2];
879
880 zeroAccumulators<num_acc, 2>(acc);
881
882 const LhsMapper lhs2 = lhs.getSubMapper(row, 0);
883 vecVSXLoop<num_acc, LhsMapper, RhsMapper>(cols, lhs2, rhs, acc, extra_cols);
884
885 addResultsVSX<num_acc>(acc);
886
887 preduxVecResultsVSX<num_acc>(acc);
888
889 outputVecResults<num_acc, 2>(acc, result, pAlpha);
890
891 result += num_acc;
892 } while (multiIters && (num_acc <= rows - (row += num_acc)));
893}
894
895template <const Index num_acc, typename LhsMapper, typename RhsMapper>
896EIGEN_ALWAYS_INLINE void colVSXVecLoopBodyExtraN(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs,
897 const Packet4f& pAlpha, float* result) {
898 EIGEN_IF_CONSTEXPR (MAX_BFLOAT16_VEC_ACC_VSX > num_acc) {
899 colVSXVecLoopBody<num_acc, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
900 }
901}
902
903template <typename LhsMapper, typename RhsMapper>
904EIGEN_ALWAYS_INLINE void colVSXVecLoopBodyExtra(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs,
905 const Packet4f& pAlpha, float* result) {
906 switch (rows - row) {
907 case 7:
908 colVSXVecLoopBodyExtraN<7, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
909 break;
910 case 6:
911 colVSXVecLoopBodyExtraN<6, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
912 break;
913 case 5:
914 colVSXVecLoopBodyExtraN<5, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
915 break;
916 case 4:
917 colVSXVecLoopBodyExtraN<4, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
918 break;
919 case 3:
920 colVSXVecLoopBodyExtraN<3, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
921 break;
922 case 2:
923 colVSXVecLoopBodyExtraN<2, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
924 break;
925 case 1:
926 colVSXVecLoopBodyExtraN<1, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
927 break;
928 }
929}
930
931template <typename LhsMapper, typename RhsMapper>
932EIGEN_ALWAYS_INLINE void calcVSXVecLoops(Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs, const Packet4f& pAlpha,
933 float* result) {
934 Index row = 0;
935 if (rows >= MAX_BFLOAT16_VEC_ACC_VSX) {
936 colVSXVecLoopBody<MAX_BFLOAT16_VEC_ACC_VSX, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
937 result += row;
938 }
939 colVSXVecLoopBodyExtra<LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
940}
941
942template <typename LhsMapper, typename RhsMapper>
943EIGEN_STRONG_INLINE void gemv_bfloat16_row(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs,
944 bfloat16* res, Index resIncr, bfloat16 alpha) {
945 typedef typename RhsMapper::LinearMapper LinearMapper;
946
947 // The following copy tells the compiler that lhs's attributes are not modified outside this function
948 // This helps GCC to generate proper code.
949 LhsMapper lhs(alhs);
950 LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
951
952 eigen_internal_assert(rhs.stride() == 1);
953
954 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
955 const Packet4f pAlpha = pset1<Packet4f>(falpha);
956
957 ei_declare_aligned_stack_constructed_variable(float, result, rows, 0);
958 if (resIncr == 1) {
959 convertArrayPointerBF16toF32(result, 1, rows, res);
960 } else {
961 convertArrayPointerBF16toF32<true>(result, 1, rows, res, resIncr);
962 }
963 calcVSXVecLoops<LhsMapper, LinearMapper>(cols, rows, lhs, rhs2, pAlpha, result);
964 if (resIncr == 1) {
965 convertArrayPointerF32toBF16VSX(result, rows, res);
966 } else {
967 convertArrayPointerF32toBF16VSX<true>(result, rows, res, resIncr);
968 }
969}
970
971#undef MAX_BFLOAT16_VEC_ACC_VSX
972
973const Packet16uc p16uc_COMPLEX32_XORFLIP = {0x44, 0x55, 0x66, 0x77, 0x00, 0x11, 0x22, 0x33,
974 0xcc, 0xdd, 0xee, 0xff, 0x88, 0x99, 0xaa, 0xbb};
975const Packet16uc p16uc_COMPLEX64_XORFLIP = {0x88, 0x99, 0xaa, 0xbb, 0xcc, 0xdd, 0xee, 0xff,
976 0x00, 0x11, 0x22, 0x33, 0x44, 0x55, 0x66, 0x77};
977
978#ifdef _BIG_ENDIAN
979const Packet16uc p16uc_COMPLEX32_CONJ_XOR = {0x00, 0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00,
980 0x00, 0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00};
981const Packet16uc p16uc_COMPLEX64_CONJ_XOR = {0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
982 0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00};
983const Packet16uc p16uc_COMPLEX32_CONJ_XOR2 = {0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
984 0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00};
985const Packet16uc p16uc_COMPLEX64_CONJ_XOR2 = {0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
986 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00};
987const Packet16uc p16uc_COMPLEX32_NEGATE = {0x80, 0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00,
988 0x80, 0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00};
989const Packet16uc p16uc_COMPLEX64_NEGATE = {0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
990 0x80, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00};
991#else
992const Packet16uc p16uc_COMPLEX32_CONJ_XOR = {0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80,
993 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80};
994const Packet16uc p16uc_COMPLEX64_CONJ_XOR = {0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00,
995 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80};
996const Packet16uc p16uc_COMPLEX32_CONJ_XOR2 = {0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00, 0x00,
997 0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00, 0x00};
998const Packet16uc p16uc_COMPLEX64_CONJ_XOR2 = {0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80,
999 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00};
1000const Packet16uc p16uc_COMPLEX32_NEGATE = {0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00, 0x80,
1001 0x00, 0x00, 0x00, 0x80, 0x00, 0x00, 0x00, 0x80};
1002const Packet16uc p16uc_COMPLEX64_NEGATE = {0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80,
1003 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x80};
1004#endif
1005
1006#ifdef _BIG_ENDIAN
1007#define COMPLEX_DELTA 0
1008#else
1009#define COMPLEX_DELTA 2
1010#endif
1011
1013EIGEN_ALWAYS_INLINE Packet2cf pconj2(const Packet2cf& a) {
1014 return Packet2cf(pxor(a.v, reinterpret_cast<Packet4f>(p16uc_COMPLEX32_CONJ_XOR)));
1015}
1016
1017EIGEN_ALWAYS_INLINE Packet1cd pconj2(const Packet1cd& a) {
1018 return Packet1cd(pxor(a.v, reinterpret_cast<Packet2d>(p16uc_COMPLEX64_CONJ_XOR)));
1019}
1020
1022EIGEN_ALWAYS_INLINE Packet2cf pconjinv(const Packet2cf& a) {
1023#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1024 return Packet2cf(Packet4f(vec_neg(Packet2d(a.v))));
1025#else
1026 return Packet2cf(pxor(a.v, reinterpret_cast<Packet4f>(p16uc_COMPLEX32_CONJ_XOR2)));
1027#endif
1028}
1029
1030EIGEN_ALWAYS_INLINE Packet1cd pconjinv(const Packet1cd& a) {
1031 return Packet1cd(pxor(a.v, reinterpret_cast<Packet2d>(p16uc_COMPLEX64_CONJ_XOR2)));
1032}
1033
1034#if defined(_ARCH_PWR8) && (!EIGEN_COMP_LLVM || __clang_major__ >= 12)
1035#define PERMXOR_GOOD // Clang had a bug with vec_permxor and endianness prior to version 12
1036#endif
1037
1039EIGEN_ALWAYS_INLINE Packet2cf pcplxflipconj(const Packet2cf& a) {
1040#ifdef PERMXOR_GOOD
1041 return Packet2cf(Packet4f(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX32_CONJ_XOR, p16uc_COMPLEX32_XORFLIP)));
1042#else
1043 return pcplxflip(pconj2(a));
1044#endif
1045}
1046
1047EIGEN_ALWAYS_INLINE Packet1cd pcplxflipconj(const Packet1cd& a) {
1048#ifdef PERMXOR_GOOD
1049 return Packet1cd(Packet2d(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX64_CONJ_XOR, p16uc_COMPLEX64_XORFLIP)));
1050#else
1051 return pcplxflip(pconj2(a));
1052#endif
1053}
1054
1056EIGEN_ALWAYS_INLINE Packet2cf pcplxconjflip(const Packet2cf& a) {
1057#ifdef PERMXOR_GOOD
1058 return Packet2cf(Packet4f(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX32_CONJ_XOR2, p16uc_COMPLEX32_XORFLIP)));
1059#else
1060 return pconj2(pcplxflip(a));
1061#endif
1062}
1063
1064EIGEN_ALWAYS_INLINE Packet1cd pcplxconjflip(const Packet1cd& a) {
1065#ifdef PERMXOR_GOOD
1066 return Packet1cd(Packet2d(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX64_CONJ_XOR2, p16uc_COMPLEX64_XORFLIP)));
1067#else
1068 return pconj2(pcplxflip(a));
1069#endif
1070}
1071
1073EIGEN_ALWAYS_INLINE Packet2cf pnegate2(const Packet2cf& a) {
1074#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1075 return Packet2cf(vec_neg(a.v));
1076#else
1077 return Packet2cf(pxor(a.v, reinterpret_cast<Packet4f>(p16uc_COMPLEX32_NEGATE)));
1078#endif
1079}
1080
1081EIGEN_ALWAYS_INLINE Packet1cd pnegate2(const Packet1cd& a) {
1082#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1083 return Packet1cd(vec_neg(a.v));
1084#else
1085 return Packet1cd(pxor(a.v, reinterpret_cast<Packet2d>(p16uc_COMPLEX64_NEGATE)));
1086#endif
1087}
1088
1090EIGEN_ALWAYS_INLINE Packet2cf pcplxflipnegate(const Packet2cf& a) {
1091#ifdef PERMXOR_GOOD
1092 return Packet2cf(Packet4f(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX32_NEGATE, p16uc_COMPLEX32_XORFLIP)));
1093#else
1094 return pcplxflip(pnegate2(a));
1095#endif
1096}
1097
1098EIGEN_ALWAYS_INLINE Packet1cd pcplxflipnegate(const Packet1cd& a) {
1099#ifdef PERMXOR_GOOD
1100 return Packet1cd(Packet2d(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX64_NEGATE, p16uc_COMPLEX64_XORFLIP)));
1101#else
1102 return pcplxflip(pnegate2(a));
1103#endif
1104}
1105
1107EIGEN_ALWAYS_INLINE Packet2cf pcplxflip2(const Packet2cf& a) {
1108 return Packet2cf(Packet4f(vec_perm(Packet16uc(a.v), Packet16uc(a.v), p16uc_COMPLEX32_XORFLIP)));
1109}
1110
1111EIGEN_ALWAYS_INLINE Packet1cd pcplxflip2(const Packet1cd& a) {
1112#ifdef EIGEN_VECTORIZE_VSX
1113 return Packet1cd(__builtin_vsx_xxpermdi(a.v, a.v, 2));
1114#else
1115 return Packet1cd(Packet2d(vec_perm(Packet16uc(a.v), Packet16uc(a.v), p16uc_COMPLEX64_XORFLIP)));
1116#endif
1117}
1118
1120EIGEN_ALWAYS_INLINE Packet4f pload_complex_half(std::complex<float>* src) {
1121 Packet4f t;
1122#ifdef EIGEN_VECTORIZE_VSX
1123 // Load float64/two float32 (doubleword alignment)
1124 __asm__("lxsdx %x0,%y1" : "=wa"(t) : "Z"(*src));
1125#else
1126 *reinterpret_cast<std::complex<float>*>(reinterpret_cast<float*>(&t) + COMPLEX_DELTA) = *src;
1127#endif
1128 return t;
1129}
1130
1132template <typename RhsScalar>
1133EIGEN_ALWAYS_INLINE void pload_realimag(RhsScalar* src, Packet4f& r, Packet4f& i) {
1134#ifdef _ARCH_PWR9
1135 __asm__("lxvwsx %x0,%y1" : "=wa"(r) : "Z"(*(reinterpret_cast<float*>(src) + 0)));
1136 __asm__("lxvwsx %x0,%y1" : "=wa"(i) : "Z"(*(reinterpret_cast<float*>(src) + 1)));
1137#else
1138 Packet4f t = pload_complex_half(src);
1139 r = vec_splat(t, COMPLEX_DELTA + 0);
1140 i = vec_splat(t, COMPLEX_DELTA + 1);
1141#endif
1142}
1143
1144template <typename RhsScalar>
1145EIGEN_ALWAYS_INLINE void pload_realimag(RhsScalar* src, Packet2d& r, Packet2d& i) {
1146#ifdef EIGEN_VECTORIZE_VSX
1147 __asm__("lxvdsx %x0,%y1" : "=wa"(r) : "Z"(*(reinterpret_cast<double*>(src) + 0)));
1148 __asm__("lxvdsx %x0,%y1" : "=wa"(i) : "Z"(*(reinterpret_cast<double*>(src) + 1)));
1149#else
1150 Packet2d t = ploadu<Packet2d>(reinterpret_cast<double*>(src));
1151 r = vec_splat(t, 0);
1152 i = vec_splat(t, 1);
1153#endif
1154}
1155
1156#ifndef EIGEN_VECTORIZE_POWER8_VECTOR
1157const Packet16uc p16uc_MERGEE = {0x00, 0x01, 0x02, 0x03, 0x10, 0x11, 0x12, 0x13,
1158 0x08, 0x09, 0x0A, 0x0B, 0x18, 0x19, 0x1A, 0x1B};
1159
1160const Packet16uc p16uc_MERGEO = {0x04, 0x05, 0x06, 0x07, 0x14, 0x15, 0x16, 0x17,
1161 0x0C, 0x0D, 0x0E, 0x0F, 0x1C, 0x1D, 0x1E, 0x1F};
1162#endif
1163
1165template <typename RhsScalar>
1166EIGEN_ALWAYS_INLINE void pload_realimag_row(RhsScalar* src, Packet4f& r, Packet4f& i) {
1167 Packet4f t = ploadu<Packet4f>(reinterpret_cast<float*>(src));
1168#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1169 r = vec_mergee(t, t);
1170 i = vec_mergeo(t, t);
1171#else
1172 r = vec_perm(t, t, p16uc_MERGEE);
1173 i = vec_perm(t, t, p16uc_MERGEO);
1174#endif
1175}
1176
1177template <typename RhsScalar>
1178EIGEN_ALWAYS_INLINE void pload_realimag_row(RhsScalar* src, Packet2d& r, Packet2d& i) {
1179 return pload_realimag(src, r, i);
1180}
1181
1183EIGEN_ALWAYS_INLINE Packet4f pload_realimag_combine(std::complex<float>* src) {
1184#ifdef EIGEN_VECTORIZE_VSX
1185 Packet4f ret;
1186 __asm__("lxvdsx %x0,%y1" : "=wa"(ret) : "Z"(*(reinterpret_cast<double*>(src) + 0)));
1187 return ret;
1188#else
1189 return Packet4f(ploaddup<Packet2d>(reinterpret_cast<double*>(src)));
1190#endif
1191}
1192
1193EIGEN_ALWAYS_INLINE Packet2d pload_realimag_combine(std::complex<double>* src) { return ploadu<Packet1cd>(src).v; }
1194
1196EIGEN_ALWAYS_INLINE Packet4f pload_realimag_combine_row(std::complex<float>* src) { return ploadu<Packet2cf>(src).v; }
1197
1198EIGEN_ALWAYS_INLINE Packet2d pload_realimag_combine_row(std::complex<double>* src) { return ploadu<Packet1cd>(src).v; }
1199
1201template <typename ResPacket>
1202EIGEN_ALWAYS_INLINE Packet4f pload_complex(std::complex<float>* src) {
1203 EIGEN_IF_CONSTEXPR (GEMV_IS_SCALAR) {
1204 return pload_complex_half(src);
1205 } else {
1206 return ploadu<Packet4f>(reinterpret_cast<float*>(src));
1207 }
1208}
1209
1210template <typename ResPacket>
1211EIGEN_ALWAYS_INLINE Packet2d pload_complex(std::complex<double>* src) {
1212 return ploadu<Packet2d>(reinterpret_cast<double*>(src));
1213}
1214
1216template <typename ResPacket>
1217EIGEN_ALWAYS_INLINE Packet4f pload_complex(Packet2cf* src) {
1218 return src->v;
1219}
1220
1221template <typename ResPacket>
1222EIGEN_ALWAYS_INLINE Packet2d pload_complex(Packet1cd* src) {
1223 return src->v;
1224}
1225
1227EIGEN_ALWAYS_INLINE Packet4f pload_complex_full(std::complex<float>* src) {
1228 return Packet4f(ploaddup<Packet2d>(reinterpret_cast<double*>(src)));
1229}
1230
1231EIGEN_ALWAYS_INLINE Packet2d pload_complex_full(std::complex<double>* src) { return ploadu<Packet1cd>(src).v; }
1232
1234EIGEN_ALWAYS_INLINE Packet4f pload_complex_full_row(std::complex<float>* src) { return ploadu<Packet2cf>(src).v; }
1235
1236EIGEN_ALWAYS_INLINE Packet2d pload_complex_full_row(std::complex<double>* src) { return pload_complex_full(src); }
1237
1239EIGEN_ALWAYS_INLINE Packet4f pload_real(float* src) { return pset1<Packet4f>(*src); }
1240
1241EIGEN_ALWAYS_INLINE Packet2d pload_real(double* src) { return pset1<Packet2d>(*src); }
1242
1243EIGEN_ALWAYS_INLINE Packet4f pload_real(const Packet4f& src) { return src; }
1244
1245EIGEN_ALWAYS_INLINE Packet2d pload_real(const Packet2d& src) { return src; }
1246
1248EIGEN_ALWAYS_INLINE Packet4f pload_real_full(float* src) {
1249 Packet4f ret = ploadu<Packet4f>(src);
1250 return vec_mergeh(ret, ret);
1251}
1252
1253EIGEN_ALWAYS_INLINE Packet2d pload_real_full(double* src) { return pload_real(src); }
1254
1255EIGEN_ALWAYS_INLINE Packet4f pload_real_full(std::complex<float>* src) {
1256 return pload_complex_full(src); // Just for compilation
1257}
1258
1259EIGEN_ALWAYS_INLINE Packet2d pload_real_full(std::complex<double>* src) {
1260 return pload_complex_full(src); // Just for compilation
1261}
1262
1264template <typename ResPacket>
1265EIGEN_ALWAYS_INLINE Packet4f pload_real_row(float* src) {
1266 EIGEN_IF_CONSTEXPR (GEMV_IS_SCALAR) {
1267 return pload_real_full(src);
1268 } else {
1269 return ploadu<Packet4f>(src);
1270 }
1271}
1272
1273template <typename ResPacket>
1274EIGEN_ALWAYS_INLINE Packet2d pload_real_row(double* src) {
1275 return pload_real(src);
1276}
1277
1278// Never executed: predux_complex only forms these calls when ResPacket is a
1279// scalar, where GEMV_IS_COMPLEX_COMPLEX is false. They exist because under C++14
1280// EIGEN_IF_CONSTEXPR is a plain if, so that arm still has to compile, and its
1281// scalar operand arrives as a const reference.
1282EIGEN_ALWAYS_INLINE Packet2cf padd(const Packet2cf& a, const std::complex<float>& b) {
1283 EIGEN_UNUSED_VARIABLE(b);
1284 return a;
1285}
1286
1287EIGEN_ALWAYS_INLINE Packet1cd padd(const Packet1cd& a, const std::complex<double>& b) {
1288 EIGEN_UNUSED_VARIABLE(b);
1289 return a;
1290}
1291
1293template <typename Scalar, typename ResScalar>
1294EIGEN_ALWAYS_INLINE Scalar pset1_realimag(ResScalar& alpha, int which, int conj) {
1295 return (which) ? ((conj) ? -alpha.real() : alpha.real()) : ((conj) ? -alpha.imag() : alpha.imag());
1296}
1297
1299template <typename Scalar, typename ResScalar, typename ResPacket, int which>
1300EIGEN_ALWAYS_INLINE Packet2cf pset1_complex(std::complex<float>& alpha) {
1301 Packet2cf ret;
1302 ret.v[COMPLEX_DELTA + 0] = pset1_realimag<Scalar, ResScalar>(alpha, (which & 0x01), (which & 0x04));
1303 ret.v[COMPLEX_DELTA + 1] = pset1_realimag<Scalar, ResScalar>(alpha, (which & 0x02), (which & 0x08));
1304 ret.v[2 - COMPLEX_DELTA] = ret.v[COMPLEX_DELTA + 0];
1305 ret.v[3 - COMPLEX_DELTA] = ret.v[COMPLEX_DELTA + 1];
1306 return ret;
1307}
1308
1309template <typename Scalar, typename ResScalar, typename ResPacket, int which>
1310EIGEN_ALWAYS_INLINE Packet1cd pset1_complex(std::complex<double>& alpha) {
1311 Packet1cd ret;
1312 ret.v[0] = pset1_realimag<Scalar, ResScalar>(alpha, (which & 0x01), (which & 0x04));
1313 ret.v[1] = pset1_realimag<Scalar, ResScalar>(alpha, (which & 0x02), (which & 0x08));
1314 return ret;
1315}
1316
1318template <typename Packet>
1319EIGEN_ALWAYS_INLINE Packet pset_zero() {
1320 return pset1<Packet>(__UNPACK_TYPE__(Packet)(0));
1321}
1322
1323template <>
1324EIGEN_ALWAYS_INLINE Packet2cf pset_zero<Packet2cf>() {
1325 return Packet2cf(pset1<Packet4f>(float(0)));
1326}
1327
1328template <>
1329EIGEN_ALWAYS_INLINE Packet1cd pset_zero<Packet1cd>() {
1330 return Packet1cd(pset1<Packet2d>(double(0)));
1331}
1332
1334template <typename Packet, typename LhsPacket, typename RhsPacket>
1335EIGEN_ALWAYS_INLINE Packet pset_init(const Packet& c1) {
1336 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) {
1337 EIGEN_UNUSED_VARIABLE(c1);
1338 return pset_zero<Packet>();
1339 } else {
1340 return c1; // Intentionally left uninitialized
1341 }
1342}
1343
1344template <typename PResPacket, typename ResPacket, typename ResScalar, typename Scalar>
1345struct alpha_store {
1346 alpha_store(ResScalar& alpha) {
1347 separate.r = pset1_complex<Scalar, ResScalar, ResPacket, 0x3>(alpha);
1348 separate.i = pset1_complex<Scalar, ResScalar, ResPacket, 0x0>(alpha);
1349 }
1350 struct ri {
1351 PResPacket r;
1352 PResPacket i;
1353 } separate;
1354};
1355
1357template <typename ScalarPacket, typename AlphaData>
1358EIGEN_ALWAYS_INLINE ScalarPacket pmadd_complex(const ScalarPacket& c0, const ScalarPacket& c2, const ScalarPacket& c4,
1359 AlphaData& b0) {
1360 return pmadd(c2, b0.separate.i.v, pmadd(c0, b0.separate.r.v, c4));
1361}
1362
1364template <typename Scalar, typename ScalarPacket, typename PResPacket, typename ResPacket, typename ResScalar,
1365 typename AlphaData>
1366EIGEN_ALWAYS_INLINE void pstoreu_pmadd_complex(const PResPacket& c0, AlphaData& b0, ResScalar* res) {
1367 PResPacket c2 = pcplxflipconj(c0);
1368 EIGEN_IF_CONSTEXPR (GEMV_IS_SCALAR) {
1369 ScalarPacket c4 = ploadu<ScalarPacket>(reinterpret_cast<Scalar*>(res));
1370 ScalarPacket c3 = pmadd_complex<ScalarPacket, AlphaData>(c0.v, c2.v, c4, b0);
1371 pstoreu(reinterpret_cast<Scalar*>(res), c3);
1372 } else {
1373 ScalarPacket c4 = pload_complex<ResPacket>(res);
1374 PResPacket c3 = PResPacket(pmadd_complex<ScalarPacket, AlphaData>(c0.v, c2.v, c4, b0));
1375 pstoreu(res, c3);
1376 }
1377}
1378
1379template <typename ScalarPacket, typename PResPacket, typename ResPacket, typename ResScalar, typename AlphaData,
1380 Index ResPacketSize, Index iter2>
1381EIGEN_ALWAYS_INLINE void pstoreu_pmadd_complex(const PResPacket& c0, const PResPacket& c1, AlphaData& b0,
1382 ResScalar* res) {
1383 PResPacket c2 = pcplxflipconj(c0);
1384 PResPacket c3 = pcplxflipconj(c1);
1385#if !defined(_ARCH_PWR10)
1386 ScalarPacket c4 = pload_complex<ResPacket>(res + (iter2 * ResPacketSize));
1387 ScalarPacket c5 = pload_complex<ResPacket>(res + ((iter2 + 1) * ResPacketSize));
1388 PResPacket c6 = PResPacket(pmadd_complex<ScalarPacket, AlphaData>(c0.v, c2.v, c4, b0));
1389 PResPacket c7 = PResPacket(pmadd_complex<ScalarPacket, AlphaData>(c1.v, c3.v, c5, b0));
1390 pstoreu(res + (iter2 * ResPacketSize), c6);
1391 pstoreu(res + ((iter2 + 1) * ResPacketSize), c7);
1392#else
1393 __vector_pair a = *reinterpret_cast<__vector_pair*>(res + (iter2 * ResPacketSize));
1394#if EIGEN_COMP_LLVM
1395 PResPacket c6[2];
1396 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(c6), &a);
1397 c6[0] = PResPacket(pmadd_complex<ScalarPacket, AlphaData>(c0.v, c2.v, c6[0].v, b0));
1398 c6[1] = PResPacket(pmadd_complex<ScalarPacket, AlphaData>(c1.v, c3.v, c6[1].v, b0));
1399 GEMV_BUILDPAIR_MMA(a, c6[0].v, c6[1].v);
1400#else
1401 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) {
1402 __asm__("xvmaddasp %L0,%x1,%x2\n\txvmaddasp %0,%x1,%x3" : "+&d"(a) : "wa"(b0.separate.r.v), "wa"(c0.v), "wa"(c1.v));
1403 __asm__("xvmaddasp %L0,%x1,%x2\n\txvmaddasp %0,%x1,%x3" : "+&d"(a) : "wa"(b0.separate.i.v), "wa"(c2.v), "wa"(c3.v));
1404 } else {
1405 __asm__("xvmaddadp %L0,%x1,%x2\n\txvmaddadp %0,%x1,%x3" : "+&d"(a) : "wa"(b0.separate.r.v), "wa"(c0.v), "wa"(c1.v));
1406 __asm__("xvmaddadp %L0,%x1,%x2\n\txvmaddadp %0,%x1,%x3" : "+&d"(a) : "wa"(b0.separate.i.v), "wa"(c2.v), "wa"(c3.v));
1407 }
1408#endif
1409 *reinterpret_cast<__vector_pair*>(res + (iter2 * ResPacketSize)) = a;
1410#endif
1411}
1412
1414template <typename Scalar, typename LhsScalar, typename LhsMapper, typename LhsPacket>
1415EIGEN_ALWAYS_INLINE LhsPacket loadLhsPacket(LhsMapper& lhs, Index i, Index j) {
1416 EIGEN_IF_CONSTEXPR (sizeof(Scalar) == sizeof(LhsScalar)) {
1417 const LhsScalar& src = lhs(i + 0, j);
1418 return LhsPacket(pload_real_full(const_cast<LhsScalar*>(&src)));
1419 }
1420 return lhs.template load<LhsPacket, Unaligned>(i + 0, j);
1421}
1422
1424template <typename ComplexPacket, typename RealPacket, bool ConjugateLhs, bool ConjugateRhs, bool Negate>
1425EIGEN_ALWAYS_INLINE RealPacket pmadd_complex_complex(const RealPacket& a, const RealPacket& b, const RealPacket& c) {
1426 EIGEN_IF_CONSTEXPR (ConjugateLhs && ConjugateRhs) {
1427 return vec_madd(a, pconj2(ComplexPacket(b)).v, c);
1428 } else EIGEN_IF_CONSTEXPR (Negate && !ConjugateLhs && ConjugateRhs) {
1429 return vec_nmsub(a, b, c);
1430 } else {
1431 return vec_madd(a, b, c);
1432 }
1433}
1434
1436template <typename ComplexPacket, typename RealPacket, bool Conjugate>
1437EIGEN_ALWAYS_INLINE RealPacket pmadd_complex_real(const RealPacket& a, const RealPacket& b, const RealPacket& c) {
1438 EIGEN_IF_CONSTEXPR (Conjugate) {
1439 return vec_madd(a, pconj2(ComplexPacket(b)).v, c);
1440 } else {
1441 return vec_madd(a, b, c);
1442 }
1443}
1444
1445template <typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket, bool ConjugateLhs,
1446 bool ConjugateRhs, int StorageOrder>
1447EIGEN_ALWAYS_INLINE void gemv_mult_generic(const LhsPacket& a0, RhsScalar* b, PResPacket& c0) {
1448 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs> pcj;
1449 RhsPacket b0;
1450 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1451 b0 = pset1<RhsPacket>(*b);
1452 } else {
1453 b0 = ploadu<RhsPacket>(b);
1454 }
1455 c0 = pcj.pmadd(a0, b0, c0);
1456}
1457
1459template <typename ScalarPacket, typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket,
1460 typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1461EIGEN_ALWAYS_INLINE void gemv_mult_complex_complex(LhsPacket& a0, RhsScalar* b, PResPacket& c0, ResPacket& c1) {
1462 ScalarPacket br, bi;
1463 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1464 pload_realimag<RhsScalar>(b, br, bi);
1465 } else {
1466 pload_realimag_row<RhsScalar>(b, br, bi);
1467 }
1468 EIGEN_IF_CONSTEXPR (ConjugateLhs && !ConjugateRhs) a0 = pconj2(a0);
1469 LhsPacket a1 = pcplxflipconj(a0);
1470 ScalarPacket cr = pmadd_complex_complex<LhsPacket, ScalarPacket, ConjugateLhs, ConjugateRhs, false>(a0.v, br, c0.v);
1471 ScalarPacket ci = pmadd_complex_complex<LhsPacket, ScalarPacket, ConjugateLhs, ConjugateRhs, true>(a1.v, bi, c1.v);
1472 c1 = ResPacket(ci);
1473 c0 = PResPacket(cr);
1474}
1475
1477template <typename ScalarPacket, typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket,
1478 typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1479EIGEN_ALWAYS_INLINE void gemv_mult_real_complex(const LhsPacket& a0, RhsScalar* b, PResPacket& c0) {
1480 ScalarPacket b0;
1481 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1482 b0 = pload_complex_full(b);
1483 } else {
1484 b0 = pload_complex_full_row(b);
1485 }
1486 ScalarPacket cri = pmadd_complex_real<PResPacket, ScalarPacket, ConjugateRhs>(a0, b0, c0.v);
1487 c0 = PResPacket(cri);
1488}
1489
1491template <typename ScalarPacket, typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket,
1492 typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1493EIGEN_ALWAYS_INLINE void gemv_mult_complex_real(LhsPacket& a0, RhsScalar* b, PResPacket& c0) {
1494 ScalarPacket a1 = pload_complex<ResPacket>(&a0);
1495 ScalarPacket b0;
1496 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1497 b0 = pload_real(b);
1498 } else {
1499 b0 = pload_real_row<ResPacket>(b);
1500 }
1501 ScalarPacket cri = pmadd_complex_real<PResPacket, ScalarPacket, ConjugateLhs>(a1, b0, c0.v);
1502 c0 = PResPacket(cri);
1503}
1504
1505#define GEMV_MULT_COMPLEX_COMPLEX(LhsType, RhsType, ResType) \
1506 template <typename ScalarPacket, typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket, \
1507 typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder> \
1508 EIGEN_ALWAYS_INLINE void gemv_mult_complex(LhsType& a0, RhsType* b, ResType& c0, ResType& c1) { \
1509 gemv_mult_complex_complex<ScalarPacket, LhsPacket, RhsScalar, RhsPacket, PResPacket, ResPacket, ConjugateLhs, \
1510 ConjugateRhs, StorageOrder>(a0, b, c0, c1); \
1511 }
1512
1513GEMV_MULT_COMPLEX_COMPLEX(Packet2cf, std::complex<float>, Packet2cf)
1514GEMV_MULT_COMPLEX_COMPLEX(Packet1cd, std::complex<double>, Packet1cd)
1515
1516#define GEMV_MULT_REAL_COMPLEX(LhsType, RhsType, ResType) \
1517 template <typename ScalarPacket, typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket, \
1518 typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder> \
1519 EIGEN_ALWAYS_INLINE void gemv_mult_complex(LhsType& a0, RhsType* b, ResType& c0, RhsType&) { \
1520 gemv_mult_real_complex<ScalarPacket, LhsPacket, RhsScalar, RhsPacket, PResPacket, ResPacket, ConjugateLhs, \
1521 ConjugateRhs, StorageOrder>(a0, b, c0); \
1522 }
1523
1524GEMV_MULT_REAL_COMPLEX(float, std::complex<float>, Packet2cf)
1525GEMV_MULT_REAL_COMPLEX(double, std::complex<double>, Packet1cd)
1526GEMV_MULT_REAL_COMPLEX(Packet4f, std::complex<float>, Packet2cf)
1527GEMV_MULT_REAL_COMPLEX(Packet2d, std::complex<double>, Packet1cd)
1528
1529#define GEMV_MULT_COMPLEX_REAL(LhsType, RhsType, ResType1, ResType2) \
1530 template <typename ScalarPacket, typename LhsPacket, typename RhsScalar, typename RhsPacket, typename PResPacket, \
1531 typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder> \
1532 EIGEN_ALWAYS_INLINE void gemv_mult_complex(LhsType& a0, RhsType* b, ResType1& c0, ResType2&) { \
1533 gemv_mult_complex_real<ScalarPacket, LhsPacket, RhsScalar, RhsPacket, PResPacket, ResPacket, ConjugateLhs, \
1534 ConjugateRhs, StorageOrder>(a0, b, c0); \
1535 }
1536
1537GEMV_MULT_COMPLEX_REAL(Packet2cf, float, Packet2cf, std::complex<float>)
1538GEMV_MULT_COMPLEX_REAL(Packet1cd, double, Packet1cd, std::complex<double>)
1539GEMV_MULT_COMPLEX_REAL(std::complex<float>, float, Packet2cf, std::complex<float>)
1540GEMV_MULT_COMPLEX_REAL(std::complex<double>, double, Packet1cd, std::complex<double>)
1541
1542#ifdef USE_GEMV_MMA
1544template <typename T>
1545EIGEN_ALWAYS_INLINE T convertReal(T a) {
1546 return a;
1547}
1548
1549EIGEN_ALWAYS_INLINE Packet4f convertReal(const Packet2cf& a) { return a.v; }
1550
1551EIGEN_ALWAYS_INLINE Packet2d convertReal(const Packet1cd& a) { return a.v; }
1552
1554template <typename T>
1555EIGEN_ALWAYS_INLINE T convertComplex(T a) {
1556 return a;
1557}
1558
1559EIGEN_ALWAYS_INLINE Packet2cf convertComplex(const Packet4f& a) { return Packet2cf(a); }
1560
1561EIGEN_ALWAYS_INLINE Packet1cd convertComplex(const Packet2d& a) { return Packet1cd(a); }
1562
1564template <typename ScalarPacket, typename LhsPacket, typename SLhsPacket, typename ResPacket>
1565EIGEN_ALWAYS_INLINE void pload_complex_MMA(SLhsPacket& a) {
1566 a = SLhsPacket(pload_complex<ResPacket>(&a));
1567}
1568
1569template <typename ScalarPacket, typename LhsPacket, typename SLhsPacket, typename ResPacket>
1570EIGEN_ALWAYS_INLINE void pload_complex_MMA(__vector_pair&) {
1571 // Pass thru
1572}
1573
1575template <typename LhsPacket, typename RhsPacket, bool NegativeAccumulate>
1576EIGEN_ALWAYS_INLINE void pger_vecMMA(__vector_quad* acc, const RhsPacket& a, const LhsPacket& b) {
1577 EIGEN_IF_CONSTEXPR (NegativeAccumulate) {
1578 __builtin_mma_xvf32gernp(acc, (__vector unsigned char)a, (__vector unsigned char)b);
1579 } else {
1580 __builtin_mma_xvf32gerpp(acc, (__vector unsigned char)a, (__vector unsigned char)b);
1581 }
1582}
1583
1585template <typename LhsPacket, typename RhsPacket, bool NegativeAccumulate>
1586EIGEN_ALWAYS_INLINE void pger_vecMMA(__vector_quad* acc, __vector_pair& a, const Packet2d& b) {
1587 EIGEN_IF_CONSTEXPR (NegativeAccumulate) {
1588 __builtin_mma_xvf64gernp(acc, (__vector_pair)a, (__vector unsigned char)b);
1589 } else {
1590 __builtin_mma_xvf64gerpp(acc, (__vector_pair)a, (__vector unsigned char)b);
1591 }
1592}
1593
1594template <typename LhsPacket, typename RhsPacket, bool NegativeAccumulate>
1595EIGEN_ALWAYS_INLINE void pger_vecMMA(__vector_quad*, __vector_pair&, const Packet4f&) {
1596 // Just for compilation
1597}
1598
1600template <typename RealPacket, typename LhsPacket, bool ConjugateLhs, bool ConjugateRhs, bool Negate>
1601EIGEN_ALWAYS_INLINE void pmadd_complex_complex_MMA(const LhsPacket& a, const RealPacket& b, __vector_quad* c) {
1602 EIGEN_IF_CONSTEXPR (ConjugateLhs && ConjugateRhs) {
1603 RealPacket b2 = pconj2(convertComplex(b)).v;
1604 return pger_vecMMA<RealPacket, RealPacket, false>(c, b2, a.v);
1605 } else EIGEN_IF_CONSTEXPR (Negate && !ConjugateLhs && ConjugateRhs) {
1606 return pger_vecMMA<RealPacket, RealPacket, true>(c, b, a.v);
1607 } else {
1608 return pger_vecMMA<RealPacket, RealPacket, false>(c, b, a.v);
1609 }
1610}
1611
1612template <typename RealPacket, typename LhsPacket, bool ConjugateLhs, bool ConjugateRhs, bool Negate>
1613EIGEN_ALWAYS_INLINE void pmadd_complex_complex_MMA(__vector_pair& a, const RealPacket& b, __vector_quad* c) {
1614 EIGEN_IF_CONSTEXPR (ConjugateLhs && ConjugateRhs) {
1615 RealPacket b2 = pconj2(convertComplex(b)).v;
1616 return pger_vecMMA<RealPacket, __vector_pair, false>(c, a, b2);
1617 } else EIGEN_IF_CONSTEXPR (Negate && !ConjugateLhs && ConjugateRhs) {
1618 return pger_vecMMA<RealPacket, __vector_pair, true>(c, a, b);
1619 } else {
1620 return pger_vecMMA<RealPacket, __vector_pair, false>(c, a, b);
1621 }
1622}
1623
1625template <typename RealPacket, typename LhsPacket, bool Conjugate, int StorageOrder>
1626EIGEN_ALWAYS_INLINE void pmadd_complex_real_MMA(const LhsPacket& a, const RealPacket& b, __vector_quad* c) {
1627 RealPacket a2 = convertReal(a);
1628 EIGEN_IF_CONSTEXPR (Conjugate) {
1629 RealPacket b2 = pconj2(convertComplex(b)).v;
1630 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1631 return pger_vecMMA<RealPacket, RealPacket, false>(c, b2, a2);
1632 } else {
1633 return pger_vecMMA<RealPacket, RealPacket, false>(c, a2, b2);
1634 }
1635 } else {
1636 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1637 return pger_vecMMA<RealPacket, RealPacket, false>(c, b, a2);
1638 } else {
1639 return pger_vecMMA<RealPacket, RealPacket, false>(c, a2, b);
1640 }
1641 }
1642}
1643
1645template <typename RealPacket, typename LhsPacket, bool Conjugate, int StorageOrder>
1646EIGEN_ALWAYS_INLINE void pmadd_complex_real_MMA(__vector_pair& a, const RealPacket& b, __vector_quad* c) {
1647 EIGEN_IF_CONSTEXPR (Conjugate) {
1648 RealPacket b2 = pconj2(convertComplex(b)).v;
1649 return pger_vecMMA<RealPacket, __vector_pair, false>(c, a, b2);
1650 } else {
1651 return pger_vecMMA<RealPacket, __vector_pair, false>(c, a, b);
1652 }
1653}
1654
1656template <typename ScalarPacket, typename LhsPacket, typename SLhsPacket, typename RhsScalar, typename ResPacket,
1657 bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1658EIGEN_ALWAYS_INLINE void gemv_mult_complex_complex_MMA(SLhsPacket& a0, RhsScalar* b, __vector_quad* c0) {
1659 ScalarPacket b0;
1660 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1661 b0 = pload_realimag_combine(b);
1662 } else {
1663 b0 = pload_realimag_combine_row(b);
1664 }
1665 pmadd_complex_complex_MMA<ScalarPacket, LhsPacket, ConjugateLhs, ConjugateRhs, false>(a0, b0, c0);
1666}
1667
1669template <typename ScalarPacket, typename LhsPacket, typename SLhsPacket, typename RhsScalar, typename ResPacket,
1670 bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1671EIGEN_ALWAYS_INLINE void gemv_mult_complex_real_MMA(SLhsPacket& a0, RhsScalar* b, __vector_quad* c0) {
1672 pload_complex_MMA<ScalarPacket, LhsPacket, SLhsPacket, ResPacket>(a0);
1673 ScalarPacket b0;
1674 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1675 b0 = pload_real(b);
1676 } else {
1677 b0 = pload_real_row<ResPacket>(b);
1678 }
1679 pmadd_complex_real_MMA<ScalarPacket, LhsPacket, ConjugateLhs, ColMajor>(a0, b0, c0);
1680}
1681
1683template <typename ScalarPacket, typename LhsPacket, typename SLhsPacket, typename RhsScalar, typename ResPacket,
1684 bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1685EIGEN_ALWAYS_INLINE void gemv_mult_real_complex_MMA(SLhsPacket& a0, RhsScalar* b, __vector_quad* c0) {
1686 ScalarPacket b0;
1687 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1688 b0 = pload_complex_full(b);
1689 } else {
1690 b0 = pload_complex_full_row(b);
1691 }
1692 pmadd_complex_real_MMA<ScalarPacket, LhsPacket, ConjugateRhs,
1693 (sizeof(RhsScalar) == sizeof(std::complex<float>)) ? StorageOrder : ColMajor>(a0, b0, c0);
1694}
1695
1696#define GEMV_MULT_COMPLEX_COMPLEX_MMA(LhsType, RhsType) \
1697 template <typename ScalarPacket, typename LhsScalar, typename LhsPacket, typename SLhsPacket, typename RhsScalar, \
1698 typename RhsPacket, typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder> \
1699 EIGEN_ALWAYS_INLINE void gemv_mult_complex_MMA(LhsType& a0, RhsType* b, __vector_quad* c0) { \
1700 gemv_mult_complex_complex_MMA<ScalarPacket, LhsPacket, SLhsPacket, RhsScalar, ResPacket, ConjugateLhs, \
1701 ConjugateRhs, StorageOrder>(a0, b, c0); \
1702 }
1703
1704GEMV_MULT_COMPLEX_COMPLEX_MMA(Packet2cf, std::complex<float>)
1705GEMV_MULT_COMPLEX_COMPLEX_MMA(__vector_pair, std::complex<float>)
1706GEMV_MULT_COMPLEX_COMPLEX_MMA(Packet1cd, std::complex<double>)
1707
1708
1709template <typename ScalarPacket, typename LhsScalar, typename LhsPacket, typename SLhsPacket, typename RhsScalar,
1710 typename RhsPacket, typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder>
1711EIGEN_ALWAYS_INLINE void gemv_mult_complex_MMA(__vector_pair& a0, std::complex<double>* b, __vector_quad* c0) {
1712 EIGEN_IF_CONSTEXPR (sizeof(LhsScalar) == 16) {
1713 gemv_mult_complex_complex_MMA<ScalarPacket, LhsPacket, SLhsPacket, RhsScalar, ResPacket, ConjugateLhs, ConjugateRhs,
1714 StorageOrder>(a0, b, c0);
1715 } else {
1716 gemv_mult_real_complex_MMA<ScalarPacket, LhsPacket, SLhsPacket, RhsScalar, ResPacket, ConjugateLhs, ConjugateRhs,
1717 StorageOrder>(a0, b, c0);
1718 }
1719}
1720
1721#define GEMV_MULT_REAL_COMPLEX_MMA(LhsType, RhsType) \
1722 template <typename ScalarPacket, typename LhsScalar, typename LhsPacket, typename SLhsPacket, typename RhsScalar, \
1723 typename RhsPacket, typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder> \
1724 EIGEN_ALWAYS_INLINE void gemv_mult_complex_MMA(LhsType& a0, RhsType* b, __vector_quad* c0) { \
1725 gemv_mult_real_complex_MMA<ScalarPacket, LhsPacket, SLhsPacket, RhsScalar, ResPacket, ConjugateLhs, ConjugateRhs, \
1726 StorageOrder>(a0, b, c0); \
1727 }
1728
1729GEMV_MULT_REAL_COMPLEX_MMA(Packet4f, std::complex<float>)
1730GEMV_MULT_REAL_COMPLEX_MMA(Packet2d, std::complex<double>)
1731
1732#define GEMV_MULT_COMPLEX_REAL_MMA(LhsType, RhsType) \
1733 template <typename ScalarPacket, typename LhsScalar, typename LhsPacket, typename SLhsPacket, typename RhsScalar, \
1734 typename RhsPacket, typename ResPacket, bool ConjugateLhs, bool ConjugateRhs, int StorageOrder> \
1735 EIGEN_ALWAYS_INLINE void gemv_mult_complex_MMA(LhsType& a0, RhsType* b, __vector_quad* c0) { \
1736 gemv_mult_complex_real_MMA<ScalarPacket, LhsPacket, SLhsPacket, RhsScalar, ResPacket, ConjugateLhs, ConjugateRhs, \
1737 StorageOrder>(a0, b, c0); \
1738 }
1739
1740GEMV_MULT_COMPLEX_REAL_MMA(Packet2cf, float)
1741GEMV_MULT_COMPLEX_REAL_MMA(Packet1cd, double)
1742GEMV_MULT_COMPLEX_REAL_MMA(__vector_pair, float)
1743GEMV_MULT_COMPLEX_REAL_MMA(__vector_pair, double)
1744
1745
1746template <typename Scalar, typename ScalarPacket, typename LhsPacket, typename RhsPacket, bool ConjugateLhs,
1747 bool ConjugateRhs>
1748EIGEN_ALWAYS_INLINE void disassembleResults2(__vector_quad* c0, PacketBlock<ScalarPacket, 4>& result0) {
1749 __builtin_mma_disassemble_acc(&result0.packet, c0);
1750 EIGEN_IF_CONSTEXPR (sizeof(LhsPacket) == 16) {
1751 EIGEN_IF_CONSTEXPR (sizeof(RhsPacket) == 16) {
1752 ScalarPacket tmp0, tmp2;
1753 tmp2 = vec_mergeh(result0.packet[2], result0.packet[3]);
1754 tmp0 = vec_mergeh(result0.packet[0], result0.packet[1]);
1755 result0.packet[3] = vec_mergel(result0.packet[3], result0.packet[2]);
1756 result0.packet[1] = vec_mergel(result0.packet[1], result0.packet[0]);
1757 result0.packet[2] = tmp2;
1758 result0.packet[0] = tmp0;
1759
1760 EIGEN_IF_CONSTEXPR (ConjugateLhs) {
1761 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
1762 result0.packet[2] = pconj2(convertComplex(result0.packet[2])).v;
1763 } else EIGEN_IF_CONSTEXPR (ConjugateRhs) {
1764 result0.packet[1] = pconj2(convertComplex(result0.packet[1])).v;
1765 result0.packet[3] = pconj2(convertComplex(result0.packet[3])).v;
1766 } else {
1767 result0.packet[1] = pconjinv(convertComplex(result0.packet[1])).v;
1768 result0.packet[3] = pconjinv(convertComplex(result0.packet[3])).v;
1769 }
1770 result0.packet[0] = vec_add(result0.packet[0], result0.packet[1]);
1771 result0.packet[2] = vec_add(result0.packet[2], result0.packet[3]);
1772 } else {
1773 result0.packet[0][1] = result0.packet[1][1];
1774 result0.packet[2][1] = result0.packet[3][1];
1775 }
1776 }
1777}
1778
1779template <typename Scalar, typename ScalarPacket, typename LhsPacket, typename RhsPacket, bool ConjugateLhs,
1780 bool ConjugateRhs>
1781EIGEN_ALWAYS_INLINE void disassembleResults4(__vector_quad* c0, PacketBlock<ScalarPacket, 4>& result0) {
1782 __builtin_mma_disassemble_acc(&result0.packet, c0);
1783 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) {
1784 EIGEN_IF_CONSTEXPR (ConjugateLhs) {
1785 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
1786 result0.packet[1] = pcplxflip2(convertComplex(result0.packet[1])).v;
1787 } else {
1788 EIGEN_IF_CONSTEXPR (ConjugateRhs) {
1789 result0.packet[1] = pcplxconjflip(convertComplex(result0.packet[1])).v;
1790 } else {
1791 result0.packet[1] = pcplxflipconj(convertComplex(result0.packet[1])).v;
1792 }
1793 }
1794 result0.packet[0] = vec_add(result0.packet[0], result0.packet[1]);
1795 } else EIGEN_IF_CONSTEXPR (sizeof(LhsPacket) == sizeof(std::complex<float>)) {
1796 EIGEN_IF_CONSTEXPR (ConjugateLhs) {
1797 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
1798 }
1799 } else {
1800 result0.packet[0] = vec_mergee(result0.packet[0], result0.packet[1]);
1801 }
1802}
1803
1804template <typename Scalar, typename ScalarPacket, int ResPacketSize, typename LhsPacket, typename RhsPacket,
1805 bool ConjugateLhs, bool ConjugateRhs>
1806EIGEN_ALWAYS_INLINE void disassembleResults(__vector_quad* c0, PacketBlock<ScalarPacket, 4>& result0) {
1807 EIGEN_IF_CONSTEXPR (!GEMV_IS_COMPLEX_FLOAT) {
1808 disassembleResults2<Scalar, ScalarPacket, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>(c0, result0);
1809 } else {
1810 disassembleResults4<Scalar, ScalarPacket, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>(c0, result0);
1811 }
1812}
1813#endif
1814
1815#define GEMV_GETN_COMPLEX(N) (((N) * ResPacketSize) >> 1)
1816
1817#define GEMV_LOADPACKET_COL_COMPLEX(iter) \
1818 loadLhsPacket<Scalar, LhsScalar, LhsMapper, PLhsPacket>(lhs, i + ((iter) * ResPacketSize), j)
1819
1820#define GEMV_LOADPACKET_COL_COMPLEX_DATA(iter) convertReal(GEMV_LOADPACKET_COL_COMPLEX(iter))
1821
1822#ifdef USE_GEMV_MMA
1823#define GEMV_INIT_COL_COMPLEX_MMA(iter, N) \
1824 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
1825 __builtin_mma_xxsetaccz(&e0##iter); \
1826 }
1827
1828#if EIGEN_COMP_LLVM
1829#define GEMV_LOADPAIR_COL_COMPLEX_MMA(iter1, iter2) \
1830 GEMV_BUILDPAIR_MMA(a##iter1, GEMV_LOADPACKET_COL_COMPLEX_DATA(iter2), \
1831 GEMV_LOADPACKET_COL_COMPLEX_DATA((iter2) + 1)); \
1832 EIGEN_UNUSED_VARIABLE(f##iter1);
1833#else
1834#define GEMV_LOADPAIR_COL_COMPLEX_MMA(iter1, iter2) \
1835 EIGEN_IF_CONSTEXPR (sizeof(LhsPacket) == 16) { \
1836 const LhsScalar& src = lhs(i + ((32 * iter1) / sizeof(LhsScalar)), j); \
1837 a##iter1 = *reinterpret_cast<__vector_pair*>(const_cast<LhsScalar*>(&src)); \
1838 EIGEN_UNUSED_VARIABLE(f##iter1); \
1839 } else { \
1840 f##iter1 = lhs.template load<PLhsPacket, Unaligned>(i + ((iter2) * ResPacketSize), j); \
1841 GEMV_BUILDPAIR_MMA(a##iter1, vec_splat(convertReal(f##iter1), 0), vec_splat(convertReal(f##iter1), 1)); \
1842 }
1843#endif
1844
1845#define GEMV_LOAD1_COL_COMPLEX_MMA(iter, N) \
1846 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
1847 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
1848 f##iter = GEMV_LOADPACKET_COL_COMPLEX(iter); \
1849 EIGEN_UNUSED_VARIABLE(a##iter); \
1850 } else { \
1851 GEMV_LOADPAIR_COL_COMPLEX_MMA(iter, iter << 1) \
1852 } \
1853 } else { \
1854 EIGEN_UNUSED_VARIABLE(a##iter); \
1855 EIGEN_UNUSED_VARIABLE(f##iter); \
1856 }
1857
1858#define GEMV_WORK1_COL_COMPLEX_MMA(iter, N) \
1859 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
1860 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
1861 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, PLhsPacket, RhsScalar, RhsPacket, ResPacket, \
1862 ConjugateLhs, ConjugateRhs, ColMajor>(f##iter, b, &e0##iter); \
1863 } else { \
1864 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, __vector_pair, RhsScalar, RhsPacket, ResPacket, \
1865 ConjugateLhs, ConjugateRhs, ColMajor>(a##iter, b, &e0##iter); \
1866 } \
1867 }
1868
1869#define GEMV_LOADPAIR2_COL_COMPLEX_MMA(iter1, iter2) \
1870 GEMV_BUILDPAIR_MMA(a##iter1, GEMV_LOADPACKET_COL_COMPLEX_DATA(iter2), GEMV_LOADPACKET_COL_COMPLEX_DATA((iter2) + 1));
1871
1872#define GEMV_LOAD2_COL_COMPLEX_MMA(iter1, iter2, iter3, N) \
1873 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter1) { \
1874 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
1875 GEMV_LOADPAIR2_COL_COMPLEX_MMA(iter2, iter2); \
1876 EIGEN_UNUSED_VARIABLE(a##iter3); \
1877 } else { \
1878 GEMV_LOADPAIR2_COL_COMPLEX_MMA(iter2, iter2 << 1); \
1879 GEMV_LOADPAIR2_COL_COMPLEX_MMA(iter3, iter3 << 1); \
1880 } \
1881 } else { \
1882 EIGEN_UNUSED_VARIABLE(a##iter2); \
1883 EIGEN_UNUSED_VARIABLE(a##iter3); \
1884 } \
1885 EIGEN_UNUSED_VARIABLE(f##iter2); \
1886 EIGEN_UNUSED_VARIABLE(f##iter3);
1887
1888#define GEMV_WORK2_COL_COMPLEX_MMA(iter1, iter2, iter3, N) \
1889 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter1) { \
1890 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
1891 PLhsPacket g[2]; \
1892 __builtin_vsx_disassemble_pair(reinterpret_cast<void*>(g), &a##iter2); \
1893 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, PLhsPacket, RhsScalar, RhsPacket, ResPacket, \
1894 ConjugateLhs, ConjugateRhs, ColMajor>(g[0], b, &e0##iter2); \
1895 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, PLhsPacket, RhsScalar, RhsPacket, ResPacket, \
1896 ConjugateLhs, ConjugateRhs, ColMajor>(g[1], b, &e0##iter3); \
1897 } else { \
1898 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, __vector_pair, RhsScalar, RhsPacket, ResPacket, \
1899 ConjugateLhs, ConjugateRhs, ColMajor>(a##iter2, b, &e0##iter2); \
1900 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, __vector_pair, RhsScalar, RhsPacket, ResPacket, \
1901 ConjugateLhs, ConjugateRhs, ColMajor>(a##iter3, b, &e0##iter3); \
1902 } \
1903 }
1904
1905#if EIGEN_COMP_LLVM
1906#define GEMV_LOAD_COL_COMPLEX_MMA(N) \
1907 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > 1) { \
1908 GEMV_UNROLL_HALF(GEMV_LOAD2_COL_COMPLEX_MMA, (N >> 1)) \
1909 } else { \
1910 GEMV_UNROLL(GEMV_LOAD1_COL_COMPLEX_MMA, N) \
1911 }
1912
1913#define GEMV_WORK_COL_COMPLEX_MMA(N) \
1914 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > 1) { \
1915 GEMV_UNROLL_HALF(GEMV_WORK2_COL_COMPLEX_MMA, (N >> 1)) \
1916 } else { \
1917 GEMV_UNROLL(GEMV_WORK1_COL_COMPLEX_MMA, N) \
1918 }
1919#else
1920#define GEMV_LOAD_COL_COMPLEX_MMA(N) GEMV_UNROLL(GEMV_LOAD1_COL_COMPLEX_MMA, N)
1921
1922#define GEMV_WORK_COL_COMPLEX_MMA(N) GEMV_UNROLL(GEMV_WORK1_COL_COMPLEX_MMA, N)
1923#endif
1924
1925#define GEMV_DISASSEMBLE_COMPLEX_MMA(iter) \
1926 disassembleResults<Scalar, ScalarPacket, ResPacketSize, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>( \
1927 &e0##iter, result0##iter);
1928
1929#define GEMV_STORE_COL_COMPLEX_MMA(iter, N) \
1930 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
1931 GEMV_DISASSEMBLE_COMPLEX_MMA(iter); \
1932 c0##iter = PResPacket(result0##iter.packet[0]); \
1933 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
1934 pstoreu_pmadd_complex<Scalar, ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData>( \
1935 c0##iter, alpha_data, res + i + (iter * ResPacketSize)); \
1936 } else { \
1937 pstoreu_pmadd_complex<Scalar, ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData>( \
1938 c0##iter, alpha_data, res + i + ((iter << 1) * ResPacketSize)); \
1939 c0##iter = PResPacket(result0##iter.packet[2]); \
1940 pstoreu_pmadd_complex<Scalar, ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData>( \
1941 c0##iter, alpha_data, res + i + (((iter << 1) + 1) * ResPacketSize)); \
1942 } \
1943 }
1944
1945#define GEMV_STORE2_COL_COMPLEX_MMA(iter1, iter2, iter3, N) \
1946 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter1) { \
1947 GEMV_DISASSEMBLE_COMPLEX_MMA(iter2); \
1948 GEMV_DISASSEMBLE_COMPLEX_MMA(iter3); \
1949 c0##iter2 = PResPacket(result0##iter2.packet[0]); \
1950 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
1951 c0##iter3 = PResPacket(result0##iter3.packet[0]); \
1952 pstoreu_pmadd_complex<ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData, ResPacketSize, iter2>( \
1953 c0##iter2, c0##iter3, alpha_data, res + i); \
1954 } else { \
1955 c0##iter3 = PResPacket(result0##iter2.packet[2]); \
1956 pstoreu_pmadd_complex<ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData, ResPacketSize, iter2 << 1>( \
1957 c0##iter2, c0##iter3, alpha_data, res + i); \
1958 c0##iter2 = PResPacket(result0##iter3.packet[0]); \
1959 c0##iter3 = PResPacket(result0##iter3.packet[2]); \
1960 pstoreu_pmadd_complex<ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData, ResPacketSize, iter3 << 1>( \
1961 c0##iter2, c0##iter3, alpha_data, res + i); \
1962 } \
1963 }
1964
1965#define GEMV_PROCESS_COL_COMPLEX_ONE_MMA(N) \
1966 GEMV_UNROLL(GEMV_INIT_COL_COMPLEX_MMA, N) \
1967 Index j = j2; \
1968 do { \
1969 const RhsScalar& b1 = rhs2(j, 0); \
1970 RhsScalar* b = const_cast<RhsScalar*>(&b1); \
1971 GEMV_UNROLL(GEMV_PREFETCH, N) \
1972 GEMV_LOAD_COL_COMPLEX_MMA(N) \
1973 GEMV_WORK_COL_COMPLEX_MMA(N) \
1974 } while (++j < jend); \
1975 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) <= 2) { \
1976 GEMV_UNROLL(GEMV_STORE_COL_COMPLEX_MMA, N) \
1977 } else { \
1978 GEMV_UNROLL_HALF(GEMV_STORE2_COL_COMPLEX_MMA, (N >> 1)) \
1979 } \
1980 i += (ResPacketSize * N);
1981#endif
1982
1983#define GEMV_INIT_COMPLEX(iter, N) \
1984 EIGEN_IF_CONSTEXPR (N > iter) { \
1985 c0##iter = pset_zero<PResPacket>(); \
1986 c1##iter = pset_init<ResPacket, LhsPacket, RhsPacket>(c1##iter); \
1987 } else { \
1988 EIGEN_UNUSED_VARIABLE(c0##iter); \
1989 EIGEN_UNUSED_VARIABLE(c1##iter); \
1990 }
1991
1992#define GEMV_WORK_COL_COMPLEX(iter, N) \
1993 EIGEN_IF_CONSTEXPR (N > iter) { \
1994 f##iter = GEMV_LOADPACKET_COL_COMPLEX(iter); \
1995 gemv_mult_complex<ScalarPacket, PLhsPacket, RhsScalar, RhsPacket, PResPacket, ResPacket, ConjugateLhs, \
1996 ConjugateRhs, ColMajor>(f##iter, b, c0##iter, c1##iter); \
1997 } else { \
1998 EIGEN_UNUSED_VARIABLE(f##iter); \
1999 }
2000
2001#define GEMV_STORE_COL_COMPLEX(iter, N) \
2002 EIGEN_IF_CONSTEXPR (N > iter) { \
2003 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) { \
2004 c0##iter = padd(c0##iter, c1##iter); \
2005 } \
2006 pstoreu_pmadd_complex<Scalar, ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData>( \
2007 c0##iter, alpha_data, res + i + (iter * ResPacketSize)); \
2008 }
2009
2011#define GEMV_PROCESS_COL_COMPLEX_ONE(N) \
2012 GEMV_UNROLL(GEMV_INIT_COMPLEX, N) \
2013 Index j = j2; \
2014 do { \
2015 const RhsScalar& b1 = rhs2(j, 0); \
2016 RhsScalar* b = const_cast<RhsScalar*>(&b1); \
2017 GEMV_UNROLL(GEMV_PREFETCH, N) \
2018 GEMV_UNROLL(GEMV_WORK_COL_COMPLEX, N) \
2019 } while (++j < jend); \
2020 GEMV_UNROLL(GEMV_STORE_COL_COMPLEX, N) \
2021 i += (ResPacketSize * N);
2022
2023#if defined(USE_GEMV_MMA) && (EIGEN_COMP_LLVM || defined(USE_SLOWER_GEMV_MMA))
2024#define USE_GEMV_COL_COMPLEX_MMA
2025#endif
2026
2027#ifdef USE_GEMV_COL_COMPLEX_MMA
2028#define GEMV_PROCESS_COL_COMPLEX(N) GEMV_PROCESS_COL_COMPLEX_ONE_MMA(N)
2029#else
2030#if defined(USE_GEMV_MMA) && (__GNUC__ > 10)
2031#define GEMV_PROCESS_COL_COMPLEX(N) \
2032 EIGEN_IF_CONSTEXPR (sizeof(Scalar) != sizeof(LhsPacket)) { \
2033 GEMV_PROCESS_COL_COMPLEX_ONE_MMA(N) \
2034 } else { \
2035 GEMV_PROCESS_COL_COMPLEX_ONE(N) \
2036 }
2037#else
2038#define GEMV_PROCESS_COL_COMPLEX(N) GEMV_PROCESS_COL_COMPLEX_ONE(N)
2039#endif
2040#endif
2041
2042template <typename Scalar, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, bool LhsIsReal,
2043 typename RhsScalar, typename RhsMapper, bool ConjugateRhs, bool RhsIsReal, typename ResScalar>
2044EIGEN_STRONG_INLINE void gemv_complex_col(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs,
2045 ResScalar* res, Index resIncr, ResScalar alpha) {
2046 typedef gemv_traits<LhsScalar, RhsScalar> Traits;
2047
2048 typedef typename Traits::LhsPacket LhsPacket;
2049 typedef typename Traits::RhsPacket RhsPacket;
2050 typedef typename Traits::ResPacket ResPacket;
2051
2052 typedef typename packet_traits<Scalar>::type ScalarPacket;
2053 typedef typename packet_traits<LhsScalar>::type PLhsPacket;
2054 typedef typename packet_traits<ResScalar>::type PResPacket;
2055 typedef gemv_traits<ResPacket, ResPacket> PTraits;
2056
2057 EIGEN_UNUSED_VARIABLE(resIncr);
2058 eigen_internal_assert(resIncr == 1);
2059
2060 // The following copy tells the compiler that lhs's attributes are not modified outside this function
2061 // This helps GCC to generate proper code.
2062 LhsMapper lhs(alhs);
2063 RhsMapper rhs2(rhs);
2064
2065 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
2066
2067 const Index lhsStride = lhs.stride();
2068 // LhsAlignment stays Unaligned; enabling aligned reads would require
2069 // propagating the Mapper's Alignment through the run() template.
2070 enum {
2071 LhsAlignment = Unaligned,
2072 ResPacketSize = PTraits::ResPacketSize,
2073 LhsPacketSize = PTraits::LhsPacketSize,
2074 RhsPacketSize = PTraits::RhsPacketSize,
2075 };
2076#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
2077 const Index prefetch_dist = 64 * LhsPacketSize;
2078#endif
2079
2080#ifndef GCC_ONE_VECTORPAIR_BUG
2081 const Index n8 = rows - 8 * ResPacketSize + 1;
2082 const Index n4 = rows - 4 * ResPacketSize + 1;
2083 const Index n2 = rows - 2 * ResPacketSize + 1;
2084#endif
2085 const Index n1 = rows - 1 * ResPacketSize + 1;
2086
2087 // TODO: improve the following heuristic:
2088 const Index block_cols = cols < 128 ? cols : (lhsStride * sizeof(LhsScalar) < 16000 ? 16 : 8);
2089
2090 typedef alpha_store<PResPacket, ResPacket, ResScalar, Scalar> AlphaData;
2091 AlphaData alpha_data(alpha);
2092
2093 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
2094 Index jend = numext::mini(j2 + block_cols, cols);
2095 Index i = 0;
2096 PResPacket c00, c01, c02, c03, c04, c05, c06, c07;
2097 ResPacket c10, c11, c12, c13, c14, c15, c16, c17;
2098 PLhsPacket f0, f1, f2, f3, f4, f5, f6, f7;
2099#ifdef USE_GEMV_MMA
2100 __vector_quad e00, e01, e02, e03, e04, e05, e06, e07;
2101 __vector_pair a0, a1, a2, a3, a4, a5, a6, a7;
2102 PacketBlock<ScalarPacket, 4> result00, result01, result02, result03, result04, result05, result06, result07;
2103 GEMV_UNUSED(8, e0)
2104 GEMV_UNUSED(8, result0)
2105 GEMV_UNUSED(8, a)
2106 GEMV_UNUSED(8, f)
2107#if !defined(GCC_ONE_VECTORPAIR_BUG) && defined(USE_GEMV_COL_COMPLEX_MMA)
2108 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX || !GEMV_IS_COMPLEX_FLOAT)
2109#endif
2110#endif
2111#ifndef GCC_ONE_VECTORPAIR_BUG
2112 {
2113 while (i < n8) {
2114 GEMV_PROCESS_COL_COMPLEX(8)
2115 }
2116 }
2117 while (i < n4) {
2118 GEMV_PROCESS_COL_COMPLEX(4)
2119 }
2120 if (i < n2) {
2121 GEMV_PROCESS_COL_COMPLEX(2)
2122 }
2123 if (i < n1)
2124#else
2125 while (i < n1)
2126#endif
2127 {
2128 GEMV_PROCESS_COL_COMPLEX_ONE(1)
2129 }
2130 for (; i < rows; ++i) {
2131 ResScalar d0(0);
2132 Index j = j2;
2133 do {
2134 d0 += cj.pmul(lhs(i, j), rhs2(j, 0));
2135 } while (++j < jend);
2136 res[i] += alpha * d0;
2137 }
2138 }
2139}
2140
2141template <typename Scalar, int N>
2142struct ScalarBlock {
2143 Scalar scalar[N];
2144};
2145
2146#ifdef USE_GEMV_MMA
2147static Packet16uc p16uc_ELEMENT_3 = {0x0c, 0x0d, 0x0e, 0x0f, 0x1c, 0x1d, 0x1e, 0x1f,
2148 0x0c, 0x0d, 0x0e, 0x0f, 0x1c, 0x1d, 0x1e, 0x1f};
2149
2151template <typename ResScalar, typename ResPacket>
2152EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_real(__vector_quad* acc0, __vector_quad* acc1) {
2153 PacketBlock<ResPacket, 4> result0, result1;
2154 __builtin_mma_disassemble_acc(&result0.packet, acc0);
2155 __builtin_mma_disassemble_acc(&result1.packet, acc1);
2156 result0.packet[0] = vec_mergeh(result0.packet[0], result1.packet[0]);
2157 result0.packet[1] = vec_mergeo(result0.packet[1], result1.packet[1]);
2158 result0.packet[2] = vec_mergel(result0.packet[2], result1.packet[2]);
2159 result0.packet[3] = vec_perm(result0.packet[3], result1.packet[3], p16uc_ELEMENT_3);
2160 result0.packet[0] =
2161 vec_add(vec_add(result0.packet[0], result0.packet[2]), vec_add(result0.packet[1], result0.packet[3]));
2162 return *reinterpret_cast<ScalarBlock<ResScalar, 2>*>(&result0.packet[0]);
2163}
2164
2165template <>
2166EIGEN_ALWAYS_INLINE ScalarBlock<double, 2> predux_real<double, Packet2d>(__vector_quad* acc0, __vector_quad* acc1) {
2167 PacketBlock<Packet2d, 4> result0, result1;
2168 __builtin_mma_disassemble_acc(&result0.packet, acc0);
2169 __builtin_mma_disassemble_acc(&result1.packet, acc1);
2170 result0.packet[0] =
2171 vec_add(vec_mergeh(result0.packet[0], result1.packet[0]), vec_mergel(result0.packet[1], result1.packet[1]));
2172 return *reinterpret_cast<ScalarBlock<double, 2>*>(&result0.packet[0]);
2173}
2174
2176template <typename LhsPacket, typename RhsPacket, bool ConjugateLhs, bool ConjugateRhs>
2177EIGEN_ALWAYS_INLINE ScalarBlock<std::complex<float>, 2> addComplexResults(PacketBlock<Packet4f, 4>& result0,
2178 const PacketBlock<Packet4f, 4>& result1) {
2179 ScalarBlock<std::complex<float>, 2> cc0;
2180 result0.packet[0] = reinterpret_cast<Packet4f>(
2181 vec_mergeh(reinterpret_cast<Packet2d>(result0.packet[0]), reinterpret_cast<Packet2d>(result1.packet[0])));
2182 result0.packet[2] = reinterpret_cast<Packet4f>(
2183 vec_mergel(reinterpret_cast<Packet2d>(result0.packet[2]), reinterpret_cast<Packet2d>(result1.packet[2])));
2184 result0.packet[0] = vec_add(result0.packet[0], result0.packet[2]);
2185 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) {
2186 result0.packet[1] = reinterpret_cast<Packet4f>(
2187 vec_mergeh(reinterpret_cast<Packet2d>(result0.packet[1]), reinterpret_cast<Packet2d>(result1.packet[1])));
2188 result0.packet[3] = reinterpret_cast<Packet4f>(
2189 vec_mergel(reinterpret_cast<Packet2d>(result0.packet[3]), reinterpret_cast<Packet2d>(result1.packet[3])));
2190 result0.packet[1] = vec_add(result0.packet[1], result0.packet[3]);
2191 EIGEN_IF_CONSTEXPR (ConjugateLhs) {
2192 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
2193 result0.packet[1] = pcplxflip2(convertComplex(result0.packet[1])).v;
2194 } else EIGEN_IF_CONSTEXPR (ConjugateRhs) {
2195 result0.packet[1] = pcplxconjflip(convertComplex(result0.packet[1])).v;
2196 } else {
2197 result0.packet[1] = pcplxflipconj(convertComplex(result0.packet[1])).v;
2198 }
2199 result0.packet[0] = vec_add(result0.packet[0], result0.packet[1]);
2200 } else {
2201 EIGEN_IF_CONSTEXPR (ConjugateLhs && (sizeof(LhsPacket) == sizeof(std::complex<float>))) {
2202 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
2203 }
2204 }
2205 cc0.scalar[0].real(result0.packet[0][0]);
2206 cc0.scalar[0].imag(result0.packet[0][1]);
2207 cc0.scalar[1].real(result0.packet[0][2]);
2208 cc0.scalar[1].imag(result0.packet[0][3]);
2209 return cc0;
2210}
2211
2212template <typename LhsPacket, typename RhsPacket, bool ConjugateLhs, bool ConjugateRhs>
2213EIGEN_ALWAYS_INLINE ScalarBlock<std::complex<double>, 2> addComplexResults(PacketBlock<Packet2d, 4>&,
2214 const PacketBlock<Packet2d, 4>&) {
2215 ScalarBlock<std::complex<double>, 2> cc0;
2216 EIGEN_UNUSED_VARIABLE(cc0);
2217 return cc0; // Just for compilation
2218}
2219
2221template <typename ResScalar, typename ResPacket, typename LhsPacket, typename RhsPacket, bool ConjugateLhs,
2222 bool ConjugateRhs>
2223EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_complex(__vector_quad* acc0, __vector_quad* acc1) {
2224 PacketBlock<ResPacket, 4> result0, result1;
2225 __builtin_mma_disassemble_acc(&result0.packet, acc0);
2226 __builtin_mma_disassemble_acc(&result1.packet, acc1);
2227 return addComplexResults<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>(result0, result1);
2228}
2229
2230template <typename ResScalar, typename ResPacket>
2231EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_real(__vector_quad* acc0) {
2232 PacketBlock<ResPacket, 4> result0;
2233 __builtin_mma_disassemble_acc(&result0.packet, acc0);
2234 result0.packet[0] =
2235 vec_add(vec_mergeh(result0.packet[0], result0.packet[2]), vec_mergel(result0.packet[1], result0.packet[3]));
2236 return *reinterpret_cast<ScalarBlock<ResScalar, 2>*>(&result0.packet[0]);
2237}
2238
2239template <typename ResScalar, typename ResPacket, typename LhsPacket, typename RhsPacket, bool ConjugateLhs,
2240 bool ConjugateRhs>
2241EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_complex(__vector_quad* acc0) {
2242 ScalarBlock<ResScalar, 2> cc0;
2243 PacketBlock<ResPacket, 4> result0;
2244 __builtin_mma_disassemble_acc(&result0.packet, acc0);
2245 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) {
2246 EIGEN_IF_CONSTEXPR (ConjugateLhs) {
2247 result0.packet[1] = pconjinv(convertComplex(result0.packet[1])).v;
2248 result0.packet[3] = pconjinv(convertComplex(result0.packet[3])).v;
2249 } else EIGEN_IF_CONSTEXPR (ConjugateRhs) {
2250 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
2251 result0.packet[2] = pconj2(convertComplex(result0.packet[2])).v;
2252 } else {
2253 result0.packet[1] = pconj2(convertComplex(result0.packet[1])).v;
2254 result0.packet[3] = pconj2(convertComplex(result0.packet[3])).v;
2255 }
2256 result0.packet[0] = vec_add(result0.packet[0], __builtin_vsx_xxpermdi(result0.packet[1], result0.packet[1], 2));
2257 result0.packet[2] = vec_add(result0.packet[2], __builtin_vsx_xxpermdi(result0.packet[3], result0.packet[3], 2));
2258 } else {
2259 result0.packet[0] = __builtin_vsx_xxpermdi(result0.packet[0], result0.packet[1], 1);
2260 result0.packet[2] = __builtin_vsx_xxpermdi(result0.packet[2], result0.packet[3], 1);
2261 }
2262 cc0.scalar[0].real(result0.packet[0][0]);
2263 cc0.scalar[0].imag(result0.packet[0][1]);
2264 cc0.scalar[1].real(result0.packet[2][0]);
2265 cc0.scalar[1].imag(result0.packet[2][1]);
2266 return cc0;
2267}
2268#endif
2269
2270template <typename ResScalar, typename ResPacket>
2271EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_real(const ResPacket& a, const ResPacket& b) {
2272 ScalarBlock<ResScalar, 2> cc0;
2273 cc0.scalar[0] = predux(a);
2274 cc0.scalar[1] = predux(b);
2275 return cc0;
2276}
2277
2278template <typename ResScalar, typename ResPacket>
2279EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_complex(const ResPacket& a, const ResPacket& b) {
2280 return predux_real<ResScalar, ResPacket>(a, b);
2281}
2282
2283#define GEMV_UNROLL_ROW(func, N) func(0, N) func(1, N) func(2, N) func(3, N) func(4, N) func(5, N) func(6, N) func(7, N)
2284
2285#define GEMV_UNROLL_ROW_HALF(func, N) func(0, 0, 1, N) func(1, 2, 3, N) func(2, 4, 5, N) func(3, 6, 7, N)
2286
2287#define GEMV_LOADPACKET_ROW(iter) lhs.template load<LhsPacket, Unaligned>(i + (iter), j)
2288
2289#ifdef USE_GEMV_MMA
2290#define GEMV_UNROLL3_ROW(func, N, which) \
2291 func(0, N, which) func(1, N, which) func(2, N, which) func(3, N, which) func(4, N, which) func(5, N, which) \
2292 func(6, N, which) func(7, N, which)
2293
2294#define GEMV_UNUSED_ROW(N, which) GEMV_UNROLL3_ROW(GEMV_UNUSED_VAR, N, which)
2295
2296#define GEMV_INIT_ROW(iter, N) \
2297 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
2298 __builtin_mma_xxsetaccz(&c##iter); \
2299 }
2300
2301#define GEMV_LOADPAIR_ROW(iter1, iter2) \
2302 GEMV_BUILDPAIR_MMA(b##iter1, GEMV_LOADPACKET_ROW(iter2), GEMV_LOADPACKET_ROW((iter2) + 1));
2303
2304#define GEMV_WORK_ROW(iter, N) \
2305 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
2306 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
2307 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&c##iter, a0, GEMV_LOADPACKET_ROW(iter)); \
2308 } else { \
2309 __vector_pair b##iter; \
2310 GEMV_LOADPAIR_ROW(iter, iter << 1) \
2311 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&c##iter, b##iter, a0); \
2312 } \
2313 }
2314
2315#define GEMV_PREDUX2(iter1, iter2, iter3, N) \
2316 EIGEN_IF_CONSTEXPR (N > iter1) { \
2317 EIGEN_IF_CONSTEXPR (GEMV_IS_FLOAT) { \
2318 cc##iter1 = predux_real<ResScalar, ResPacket>(&c##iter2, &c##iter3); \
2319 } else { \
2320 cc##iter1 = predux_real<ResScalar, ResPacket>(&c##iter1); \
2321 } \
2322 } else { \
2323 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2324 }
2325#else
2326#define GEMV_INIT_ROW(iter, N) \
2327 EIGEN_IF_CONSTEXPR (N > iter) { \
2328 c##iter = pset1<ResPacket>(ResScalar(0)); \
2329 } else { \
2330 EIGEN_UNUSED_VARIABLE(c##iter); \
2331 }
2332
2333#define GEMV_WORK_ROW(iter, N) \
2334 EIGEN_IF_CONSTEXPR (N > iter) { \
2335 c##iter = pcj.pmadd(GEMV_LOADPACKET_ROW(iter), a0, c##iter); \
2336 }
2337
2338#define GEMV_PREDUX2(iter1, iter2, iter3, N) \
2339 EIGEN_IF_CONSTEXPR (N > iter1) { \
2340 cc##iter1 = predux_real<ResScalar, ResPacket>(c##iter2, c##iter3); \
2341 } else { \
2342 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2343 }
2344#endif
2345
2346#define GEMV_MULT(iter1, iter2, iter3, N) \
2347 EIGEN_IF_CONSTEXPR (N > iter1) { \
2348 cc##iter1.scalar[0] += cj.pmul(lhs(i + iter2, j), a0); \
2349 cc##iter1.scalar[1] += cj.pmul(lhs(i + iter3, j), a0); \
2350 }
2351
2352#define GEMV_STORE_ROW(iter1, iter2, iter3, N) \
2353 EIGEN_IF_CONSTEXPR (N > iter1) { \
2354 storeMaddData<ResScalar>(res + ((i + iter2) * resIncr), alpha, cc##iter1.scalar[0]); \
2355 storeMaddData<ResScalar>(res + ((i + iter3) * resIncr), alpha, cc##iter1.scalar[1]); \
2356 }
2357
2359#define GEMV_PROCESS_ROW(N) \
2360 for (; i < n##N; i += N) { \
2361 GEMV_UNROLL_ROW(GEMV_INIT_ROW, N) \
2362 Index j = 0; \
2363 for (; j + LhsPacketSize <= cols; j += LhsPacketSize) { \
2364 RhsPacket a0 = rhs2.template load<RhsPacket, Unaligned>(j); \
2365 GEMV_UNROLL_ROW(GEMV_WORK_ROW, N) \
2366 } \
2367 GEMV_UNROLL_ROW_HALF(GEMV_PREDUX2, (N >> 1)) \
2368 for (; j < cols; ++j) { \
2369 RhsScalar a0 = rhs2(j); \
2370 GEMV_UNROLL_ROW_HALF(GEMV_MULT, (N >> 1)) \
2371 } \
2372 GEMV_UNROLL_ROW_HALF(GEMV_STORE_ROW, (N >> 1)) \
2373 }
2374
2375template <typename LhsScalar, typename LhsMapper, typename RhsScalar, typename RhsMapper, typename ResScalar>
2376EIGEN_STRONG_INLINE void gemv_row(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs, ResScalar* res,
2377 Index resIncr, ResScalar alpha) {
2378 typedef gemv_traits<LhsScalar, RhsScalar> Traits;
2379
2380 typedef typename Traits::LhsPacket LhsPacket;
2381 typedef typename Traits::RhsPacket RhsPacket;
2382 typedef typename Traits::ResPacket ResPacket;
2383
2384 // The following copy tells the compiler that lhs's attributes are not modified outside this function
2385 // This helps GCC to generate proper code.
2386 LhsMapper lhs(alhs);
2387 typename RhsMapper::LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
2388
2389 eigen_internal_assert(rhs.stride() == 1);
2390 conj_helper<LhsScalar, RhsScalar, false, false> cj;
2391 conj_helper<LhsPacket, RhsPacket, false, false> pcj;
2392
2393 // TODO: fine tune the following heuristic. The rationale is that if the matrix is very large,
2394 // processing 8 rows at once might be counter productive wrt cache.
2395#ifndef GCC_ONE_VECTORPAIR_BUG
2396 const Index n8 = lhs.stride() * sizeof(LhsScalar) > 32000 ? (rows - 7) : (rows - 7);
2397 const Index n4 = rows - 3;
2398 const Index n2 = rows - 1;
2399#endif
2400
2401 // LhsAlignment stays Unaligned; enabling aligned reads would require
2402 // propagating the Mapper's Alignment through the run() template.
2403 enum {
2404 LhsAlignment = Unaligned,
2405 ResPacketSize = Traits::ResPacketSize,
2406 LhsPacketSize = Traits::LhsPacketSize,
2407 RhsPacketSize = Traits::RhsPacketSize,
2408 };
2409
2410 Index i = 0;
2411#ifdef USE_GEMV_MMA
2412 __vector_quad c0, c1, c2, c3, c4, c5, c6, c7;
2413 GEMV_UNUSED_ROW(8, c)
2414#else
2415 ResPacket c0, c1, c2, c3, c4, c5, c6, c7;
2416#endif
2417#ifndef GCC_ONE_VECTORPAIR_BUG
2418 ScalarBlock<ResScalar, 2> cc0, cc1, cc2, cc3;
2419 GEMV_PROCESS_ROW(8)
2420 GEMV_PROCESS_ROW(4)
2421 GEMV_PROCESS_ROW(2)
2422#endif
2423 for (; i < rows; ++i) {
2424 ResPacket d0 = pset1<ResPacket>(ResScalar(0));
2425 Index j = 0;
2426 for (; j + LhsPacketSize <= cols; j += LhsPacketSize) {
2427 RhsPacket b0 = rhs2.template load<RhsPacket, Unaligned>(j);
2428
2429 d0 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 0, j), b0, d0);
2430 }
2431 ResScalar dd0 = predux(d0);
2432 for (; j < cols; ++j) {
2433 dd0 += cj.pmul(lhs(i, j), rhs2(j));
2434 }
2435 res[i * resIncr] += alpha * dd0;
2436 }
2437}
2438
2439#define EIGEN_POWER_GEMV_REAL_SPECIALIZE_COL(Scalar) \
2440 template <typename Index, typename LhsMapper, bool ConjugateLhs, typename RhsMapper, bool ConjugateRhs, int Version> \
2441 struct general_matrix_vector_product<Index, Scalar, LhsMapper, ColMajor, ConjugateLhs, Scalar, RhsMapper, \
2442 ConjugateRhs, Version> { \
2443 typedef typename ScalarBinaryOpTraits<Scalar, Scalar>::ReturnType ResScalar; \
2444 \
2445 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE static void run(Index rows, Index cols, const LhsMapper& lhs, \
2446 const RhsMapper& rhs, ResScalar* res, Index resIncr, \
2447 ResScalar alpha) { \
2448 if (numext::is_exactly_zero(alpha)) return; \
2449 gemv_col<Scalar, LhsMapper, Scalar, RhsMapper, ResScalar>(rows, cols, lhs, rhs, res, resIncr, alpha); \
2450 } \
2451 };
2452
2453#define EIGEN_POWER_GEMV_REAL_SPECIALIZE_ROW(Scalar) \
2454 template <typename Index, typename LhsMapper, bool ConjugateLhs, typename RhsMapper, bool ConjugateRhs, int Version> \
2455 struct general_matrix_vector_product<Index, Scalar, LhsMapper, RowMajor, ConjugateLhs, Scalar, RhsMapper, \
2456 ConjugateRhs, Version> { \
2457 typedef typename ScalarBinaryOpTraits<Scalar, Scalar>::ReturnType ResScalar; \
2458 \
2459 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE static void run(Index rows, Index cols, const LhsMapper& lhs, \
2460 const RhsMapper& rhs, ResScalar* res, Index resIncr, \
2461 ResScalar alpha) { \
2462 if (numext::is_exactly_zero(alpha)) return; \
2463 gemv_row<Scalar, LhsMapper, Scalar, RhsMapper, ResScalar>(rows, cols, lhs, rhs, res, resIncr, alpha); \
2464 } \
2465 };
2466
2467EIGEN_POWER_GEMV_REAL_SPECIALIZE_COL(float)
2468EIGEN_POWER_GEMV_REAL_SPECIALIZE_COL(double)
2469EIGEN_POWER_GEMV_REAL_SPECIALIZE_ROW(float)
2470EIGEN_POWER_GEMV_REAL_SPECIALIZE_ROW(double)
2471
2472#ifdef USE_GEMV_MMA
2473#define gemv_bf16_col gemvMMA_bfloat16_col
2474#define gemv_bf16_row gemvMMA_bfloat16_row
2475#else
2476#define gemv_bf16_col gemv_bfloat16_col
2477#define gemv_bf16_row gemv_bfloat16_row
2478#endif
2479
2480#define EIGEN_POWER_GEMV_REAL_SPECIALIZE_COL_BFLOAT16() \
2481 template <typename Index, typename LhsMapper, bool ConjugateLhs, typename RhsMapper, bool ConjugateRhs, int Version> \
2482 struct general_matrix_vector_product<Index, bfloat16, LhsMapper, ColMajor, ConjugateLhs, bfloat16, RhsMapper, \
2483 ConjugateRhs, Version> { \
2484 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE static void run(Index rows, Index cols, const LhsMapper& lhs, \
2485 const RhsMapper& rhs, bfloat16* res, Index resIncr, \
2486 bfloat16 alpha) { \
2487 if (numext::is_exactly_zero(alpha)) return; \
2488 gemv_bf16_col<LhsMapper, RhsMapper>(rows, cols, lhs, rhs, res, resIncr, alpha); \
2489 } \
2490 };
2491
2492#define EIGEN_POWER_GEMV_REAL_SPECIALIZE_ROW_BFLOAT16() \
2493 template <typename Index, typename LhsMapper, bool ConjugateLhs, typename RhsMapper, bool ConjugateRhs, int Version> \
2494 struct general_matrix_vector_product<Index, bfloat16, LhsMapper, RowMajor, ConjugateLhs, bfloat16, RhsMapper, \
2495 ConjugateRhs, Version> { \
2496 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE static void run(Index rows, Index cols, const LhsMapper& lhs, \
2497 const RhsMapper& rhs, bfloat16* res, Index resIncr, \
2498 bfloat16 alpha) { \
2499 if (numext::is_exactly_zero(alpha)) return; \
2500 gemv_bf16_row<LhsMapper, RhsMapper>(rows, cols, lhs, rhs, res, resIncr, alpha); \
2501 } \
2502 };
2503
2504EIGEN_POWER_GEMV_REAL_SPECIALIZE_COL_BFLOAT16()
2505EIGEN_POWER_GEMV_REAL_SPECIALIZE_ROW_BFLOAT16()
2506
2507template <typename ResScalar, typename PResPacket, typename ResPacket, typename LhsPacket, typename RhsPacket>
2508EIGEN_ALWAYS_INLINE ScalarBlock<ResScalar, 2> predux_complex(PResPacket& a0, PResPacket& b0, const ResPacket& a1,
2509 const ResPacket& b1) {
2510 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) {
2511 a0 = padd(a0, a1);
2512 b0 = padd(b0, b1);
2513 }
2514 return predux_complex<ResScalar, PResPacket>(a0, b0);
2515}
2516
2517#define GEMV_LOADPACKET_ROW_COMPLEX(iter) loadLhsPacket<Scalar, LhsScalar, LhsMapper, PLhsPacket>(lhs, i + (iter), j)
2518
2519#define GEMV_LOADPACKET_ROW_COMPLEX_DATA(iter) convertReal(GEMV_LOADPACKET_ROW_COMPLEX(iter))
2520
2521#define GEMV_PROCESS_ROW_COMPLEX_SINGLE_WORK(which, N) \
2522 j = 0; \
2523 for (; j + LhsPacketSize <= cols; j += LhsPacketSize) { \
2524 const RhsScalar& b1 = rhs2(j); \
2525 RhsScalar* b = const_cast<RhsScalar*>(&b1); \
2526 GEMV_UNROLL_ROW(which, N) \
2527 }
2528
2529#define GEMV_PROCESS_END_ROW_COMPLEX(N) \
2530 for (; j < cols; ++j) { \
2531 RhsScalar b0 = rhs2(j); \
2532 GEMV_UNROLL_ROW_HALF(GEMV_MULT_COMPLEX, (N >> 1)) \
2533 } \
2534 GEMV_UNROLL_ROW_HALF(GEMV_STORE_ROW_COMPLEX, (N >> 1))
2535
2536#ifdef USE_GEMV_MMA
2537#define GEMV_INIT_ROW_COMPLEX_MMA(iter, N) \
2538 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
2539 __builtin_mma_xxsetaccz(&e0##iter); \
2540 }
2541
2542#define GEMV_LOADPAIR_ROW_COMPLEX_MMA(iter1, iter2) \
2543 GEMV_BUILDPAIR_MMA(a##iter1, GEMV_LOADPACKET_ROW_COMPLEX_DATA(iter2), GEMV_LOADPACKET_ROW_COMPLEX_DATA((iter2) + 1));
2544
2545#define GEMV_WORK_ROW_COMPLEX_MMA(iter, N) \
2546 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
2547 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
2548 PLhsPacket a##iter = GEMV_LOADPACKET_ROW_COMPLEX(iter); \
2549 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, PLhsPacket, RhsScalar, RhsPacket, ResPacket, \
2550 ConjugateLhs, ConjugateRhs, RowMajor>(a##iter, b, &e0##iter); \
2551 } else { \
2552 __vector_pair a##iter; \
2553 GEMV_LOADPAIR_ROW_COMPLEX_MMA(iter, iter << 1) \
2554 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, __vector_pair, RhsScalar, RhsPacket, ResPacket, \
2555 ConjugateLhs, ConjugateRhs, RowMajor>(a##iter, b, &e0##iter); \
2556 } \
2557 }
2558
2559#define GEMV_PREDUX4_COMPLEX_MMA(iter1, iter2, iter3, N) \
2560 EIGEN_IF_CONSTEXPR (N > iter1) { \
2561 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_FLOAT) { \
2562 cc##iter1 = predux_complex<ResScalar, ScalarPacket, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>( \
2563 &e0##iter2, &e0##iter3); \
2564 } else { \
2565 cc##iter1 = \
2566 predux_complex<ResScalar, ScalarPacket, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>(&e0##iter1); \
2567 } \
2568 } else { \
2569 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2570 }
2571
2572#define GEMV_PROCESS_ROW_COMPLEX_SINGLE_MMA(N) \
2573 GEMV_UNROLL_ROW(GEMV_INIT_ROW_COMPLEX_MMA, N) \
2574 GEMV_PROCESS_ROW_COMPLEX_SINGLE_WORK(GEMV_WORK_ROW_COMPLEX_MMA, N)
2575
2576#define GEMV_PROCESS_ROW_COMPLEX_ONE_MMA(N) \
2577 for (; i < n##N; i += N) { \
2578 GEMV_PROCESS_ROW_COMPLEX_SINGLE_MMA(N) \
2579 GEMV_UNROLL_ROW_HALF(GEMV_PREDUX4_COMPLEX_MMA, (N >> 1)) \
2580 GEMV_PROCESS_END_ROW_COMPLEX(N); \
2581 }
2582#endif
2583
2584#define GEMV_WORK_ROW_COMPLEX(iter, N) \
2585 EIGEN_IF_CONSTEXPR (N > iter) { \
2586 PLhsPacket a##iter = GEMV_LOADPACKET_ROW_COMPLEX(iter); \
2587 gemv_mult_complex<ScalarPacket, PLhsPacket, RhsScalar, RhsPacket, PResPacket, ResPacket, ConjugateLhs, \
2588 ConjugateRhs, RowMajor>(a##iter, b, c0##iter, c1##iter); \
2589 }
2590
2591#define GEMV_PREDUX4_COMPLEX(iter1, iter2, iter3, N) \
2592 EIGEN_IF_CONSTEXPR (N > iter1) { \
2593 cc##iter1 = predux_complex<ResScalar, PResPacket, ResPacket, LhsPacket, RhsPacket>(c0##iter2, c0##iter3, \
2594 c1##iter2, c1##iter3); \
2595 } else { \
2596 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2597 }
2598
2599#define GEMV_MULT_COMPLEX(iter1, iter2, iter3, N) \
2600 EIGEN_IF_CONSTEXPR (N > iter1) { \
2601 cc##iter1.scalar[0] += cj.pmul(lhs(i + iter2, j), b0); \
2602 cc##iter1.scalar[1] += cj.pmul(lhs(i + iter3, j), b0); \
2603 }
2604
2605#define GEMV_STORE_ROW_COMPLEX(iter1, iter2, iter3, N) \
2606 EIGEN_IF_CONSTEXPR (N > iter1) { \
2607 storeMaddData<ResScalar>(res + ((i + iter2) * resIncr), alpha, cc##iter1.scalar[0]); \
2608 storeMaddData<ResScalar>(res + ((i + iter3) * resIncr), alpha, cc##iter1.scalar[1]); \
2609 }
2610
2611#define GEMV_PROCESS_ROW_COMPLEX_SINGLE_NEW(N) \
2612 GEMV_UNROLL_ROW(GEMV_INIT_COMPLEX, N) \
2613 GEMV_PROCESS_ROW_COMPLEX_SINGLE_WORK(GEMV_WORK_ROW_COMPLEX, N)
2614
2617#define GEMV_PROCESS_ROW_COMPLEX_ONE_NEW(N) \
2618 for (; i < n##N; i += N) { \
2619 GEMV_PROCESS_ROW_COMPLEX_SINGLE_NEW(N) \
2620 GEMV_UNROLL_ROW_HALF(GEMV_PREDUX4_COMPLEX, (N >> 1)) \
2621 GEMV_PROCESS_END_ROW_COMPLEX(N); \
2622 }
2623
2624#define GEMV_PROCESS_ROW_COMPLEX_PREDUX_NEW(iter) \
2625 EIGEN_IF_CONSTEXPR (GEMV_IS_COMPLEX_COMPLEX) { \
2626 c0##iter = padd(c0##iter, c1##iter); \
2627 } \
2628 dd0 = predux(c0##iter);
2629
2630#if EIGEN_COMP_LLVM
2631#define GEMV_PROCESS_ROW_COMPLEX_SINGLE(N) GEMV_PROCESS_ROW_COMPLEX_SINGLE_NEW(N)
2632
2633#define GEMV_PROCESS_ROW_COMPLEX_ONE(N) GEMV_PROCESS_ROW_COMPLEX_ONE_NEW(N)
2634
2635#define GEMV_PROCESS_ROW_COMPLEX_PREDUX(iter) GEMV_PROCESS_ROW_COMPLEX_PREDUX_NEW(iter)
2636#else
2637// gcc seems to be reading and writing registers unnecessarily to memory.
2638// Use the old way for complex double until it is fixed.
2639
2640#define GEMV_LOADPACKET_ROW_COMPLEX_OLD(iter) lhs.template load<LhsPacket, LhsAlignment>(i + (iter), j)
2641
2642#define GEMV_INIT_COMPLEX_OLD(iter, N) \
2643 EIGEN_UNUSED_VARIABLE(c0##iter); \
2644 EIGEN_IF_CONSTEXPR (N > iter) { \
2645 c1##iter = pset_zero<ResPacket>(); \
2646 } else { \
2647 EIGEN_UNUSED_VARIABLE(c1##iter); \
2648 }
2649
2650#define GEMV_WORK_ROW_COMPLEX_OLD(iter, N) \
2651 EIGEN_IF_CONSTEXPR (N > iter) { \
2652 LhsPacket a##iter = GEMV_LOADPACKET_ROW_COMPLEX_OLD(iter); \
2653 c1##iter = pcj.pmadd(a##iter, b0, c1##iter); \
2654 }
2655
2656#define GEMV_PREDUX4_COMPLEX_OLD(iter1, iter2, iter3, N) \
2657 EIGEN_IF_CONSTEXPR (N > iter1) { \
2658 cc##iter1.scalar[0] = predux(c1##iter2); \
2659 cc##iter1.scalar[1] = predux(c1##iter3); \
2660 } else { \
2661 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2662 }
2663
2664#define GEMV_PROCESS_ROW_COMPLEX_SINGLE_OLD(N) \
2665 GEMV_UNROLL_ROW(GEMV_INIT_COMPLEX_OLD, N) \
2666 j = 0; \
2667 for (; j + LhsPacketSize <= cols; j += LhsPacketSize) { \
2668 RhsPacket b0 = rhs2.template load<RhsPacket, Unaligned>(j); \
2669 GEMV_UNROLL_ROW(GEMV_WORK_ROW_COMPLEX_OLD, N) \
2670 }
2671
2672#define GEMV_PROCESS_ROW_COMPLEX_ONE_OLD(N) \
2673 for (; i < n##N; i += N) { \
2674 GEMV_PROCESS_ROW_COMPLEX_SINGLE_OLD(N) \
2675 GEMV_UNROLL_ROW_HALF(GEMV_PREDUX4_COMPLEX_OLD, (N >> 1)) \
2676 GEMV_PROCESS_END_ROW_COMPLEX(N) \
2677 }
2678
2679#define GEMV_PROCESS_ROW_COMPLEX_PREDUX_OLD(iter) dd0 = predux(c1##iter);
2680
2681#if (__GNUC__ > 10)
2682#define GEMV_PROCESS_ROW_COMPLEX_IS_NEW 1
2683#else
2684#define GEMV_PROCESS_ROW_COMPLEX_IS_NEW (sizeof(Scalar) == sizeof(float)) || GEMV_IS_COMPLEX_COMPLEX
2685#endif
2686
2687#define GEMV_PROCESS_ROW_COMPLEX_SINGLE(N) \
2688 if (GEMV_PROCESS_ROW_COMPLEX_IS_NEW) { \
2689 GEMV_PROCESS_ROW_COMPLEX_SINGLE_NEW(N) \
2690 } else { \
2691 GEMV_PROCESS_ROW_COMPLEX_SINGLE_OLD(N) \
2692 }
2693
2694#define GEMV_PROCESS_ROW_COMPLEX_ONE(N) \
2695 if (GEMV_PROCESS_ROW_COMPLEX_IS_NEW) { \
2696 GEMV_PROCESS_ROW_COMPLEX_ONE_NEW(N) \
2697 } else { \
2698 GEMV_PROCESS_ROW_COMPLEX_ONE_OLD(N) \
2699 }
2700
2701#define GEMV_PROCESS_ROW_COMPLEX_PREDUX(iter) \
2702 if (GEMV_PROCESS_ROW_COMPLEX_IS_NEW) { \
2703 GEMV_PROCESS_ROW_COMPLEX_PREDUX_NEW(iter) \
2704 } else { \
2705 GEMV_PROCESS_ROW_COMPLEX_PREDUX_OLD(iter) \
2706 }
2707#endif
2708
2709#ifdef USE_GEMV_MMA
2710#define GEMV_PROCESS_ROW_COMPLEX(N) GEMV_PROCESS_ROW_COMPLEX_ONE_MMA(N)
2711#else
2712#define GEMV_PROCESS_ROW_COMPLEX(N) GEMV_PROCESS_ROW_COMPLEX_ONE(N)
2713#endif
2714
2715template <typename Scalar, typename LhsScalar, typename LhsMapper, bool ConjugateLhs, bool LhsIsReal,
2716 typename RhsScalar, typename RhsMapper, bool ConjugateRhs, bool RhsIsReal, typename ResScalar>
2717EIGEN_STRONG_INLINE void gemv_complex_row(Index rows, Index cols, const LhsMapper& alhs, const RhsMapper& rhs,
2718 ResScalar* res, Index resIncr, ResScalar alpha) {
2719 typedef gemv_traits<LhsScalar, RhsScalar> Traits;
2720
2721 typedef typename Traits::LhsPacket LhsPacket;
2722 typedef typename Traits::RhsPacket RhsPacket;
2723 typedef typename Traits::ResPacket ResPacket;
2724
2725 typedef typename packet_traits<Scalar>::type ScalarPacket;
2726 typedef typename packet_traits<LhsScalar>::type PLhsPacket;
2727 typedef typename packet_traits<ResScalar>::type PResPacket;
2728 typedef gemv_traits<ResPacket, ResPacket> PTraits;
2729
2730 // The following copy tells the compiler that lhs's attributes are not modified outside this function
2731 // This helps GCC to generate proper code.
2732 LhsMapper lhs(alhs);
2733 typename RhsMapper::LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
2734
2735 eigen_internal_assert(rhs.stride() == 1);
2736 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
2737#if !EIGEN_COMP_LLVM
2738 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs> pcj;
2739#endif
2740
2741 // TODO: fine tune the following heuristic. The rationale is that if the matrix is very large,
2742 // processing 8 rows at once might be counter productive wrt cache.
2743#ifndef GCC_ONE_VECTORPAIR_BUG
2744 const Index n8 = lhs.stride() * sizeof(LhsScalar) > 32000 ? (rows - 7) : (rows - 7);
2745 const Index n4 = rows - 3;
2746 const Index n2 = rows - 1;
2747#endif
2748
2749 // LhsAlignment stays Unaligned; enabling aligned reads would require
2750 // propagating the Mapper's Alignment through the run() template.
2751 enum {
2752 LhsAlignment = Unaligned,
2753 ResPacketSize = PTraits::ResPacketSize,
2754 LhsPacketSize = PTraits::LhsPacketSize,
2755 RhsPacketSize = PTraits::RhsPacketSize,
2756 };
2757
2758 Index i = 0, j;
2759 PResPacket c00, c01, c02, c03, c04, c05, c06, c07;
2760 ResPacket c10, c11, c12, c13, c14, c15, c16, c17;
2761#ifdef USE_GEMV_MMA
2762 __vector_quad e00, e01, e02, e03, e04, e05, e06, e07;
2763 GEMV_UNUSED_ROW(8, e0)
2764 GEMV_UNUSED_EXTRA(1, c0)
2765 GEMV_UNUSED_EXTRA(1, c1)
2766#endif
2767 ResScalar dd0;
2768#ifndef GCC_ONE_VECTORPAIR_BUG
2769 ScalarBlock<ResScalar, 2> cc0, cc1, cc2, cc3;
2770#ifdef USE_GEMV_MMA
2771 EIGEN_IF_CONSTEXPR (!GEMV_IS_COMPLEX_COMPLEX)
2772#endif
2773 {
2774 GEMV_PROCESS_ROW_COMPLEX(8)
2775 }
2776 GEMV_PROCESS_ROW_COMPLEX(4)
2777 GEMV_PROCESS_ROW_COMPLEX(2)
2778#endif
2779 for (; i < rows; ++i) {
2780 GEMV_PROCESS_ROW_COMPLEX_SINGLE(1)
2781 GEMV_PROCESS_ROW_COMPLEX_PREDUX(0)
2782 for (; j < cols; ++j) {
2783 dd0 += cj.pmul(lhs(i, j), rhs2(j));
2784 }
2785 res[i * resIncr] += alpha * dd0;
2786 }
2787}
2788
2789#define EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(Scalar, LhsScalar, RhsScalar) \
2790 template <typename Index, typename LhsMapper, bool ConjugateLhs, typename RhsMapper, bool ConjugateRhs, int Version> \
2791 struct general_matrix_vector_product<Index, LhsScalar, LhsMapper, ColMajor, ConjugateLhs, RhsScalar, RhsMapper, \
2792 ConjugateRhs, Version> { \
2793 typedef typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType ResScalar; \
2794 \
2795 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE static void run(Index rows, Index cols, const LhsMapper& lhs, \
2796 const RhsMapper& rhs, ResScalar* res, Index resIncr, \
2797 ResScalar alpha) { \
2798 if (numext::is_exactly_zero(alpha)) return; \
2799 gemv_complex_col<Scalar, LhsScalar, LhsMapper, ConjugateLhs, sizeof(Scalar) == sizeof(LhsScalar), RhsScalar, \
2800 RhsMapper, ConjugateRhs, sizeof(Scalar) == sizeof(RhsScalar), ResScalar>(rows, cols, lhs, rhs, \
2801 res, resIncr, alpha); \
2802 } \
2803 };
2804
2805#define EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(Scalar, LhsScalar, RhsScalar) \
2806 template <typename Index, typename LhsMapper, bool ConjugateLhs, typename RhsMapper, bool ConjugateRhs, int Version> \
2807 struct general_matrix_vector_product<Index, LhsScalar, LhsMapper, RowMajor, ConjugateLhs, RhsScalar, RhsMapper, \
2808 ConjugateRhs, Version> { \
2809 typedef typename ScalarBinaryOpTraits<LhsScalar, RhsScalar>::ReturnType ResScalar; \
2810 \
2811 EIGEN_DEVICE_FUNC EIGEN_DONT_INLINE static void run(Index rows, Index cols, const LhsMapper& lhs, \
2812 const RhsMapper& rhs, ResScalar* res, Index resIncr, \
2813 ResScalar alpha) { \
2814 if (numext::is_exactly_zero(alpha)) return; \
2815 gemv_complex_row<Scalar, LhsScalar, LhsMapper, ConjugateLhs, sizeof(Scalar) == sizeof(LhsScalar), RhsScalar, \
2816 RhsMapper, ConjugateRhs, sizeof(Scalar) == sizeof(RhsScalar), ResScalar>(rows, cols, lhs, rhs, \
2817 res, resIncr, alpha); \
2818 } \
2819 };
2820
2821EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(float, float, std::complex<float>)
2822EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(float, std::complex<float>, float)
2823EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(float, std::complex<float>, std::complex<float>)
2824EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(double, double, std::complex<double>)
2825EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(double, std::complex<double>, double)
2826EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_COL(double, std::complex<double>, std::complex<double>)
2827EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(float, float, std::complex<float>)
2828EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(float, std::complex<float>, float)
2829EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(float, std::complex<float>, std::complex<float>)
2830EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(double, double, std::complex<double>)
2831EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(double, std::complex<double>, double)
2832EIGEN_POWER_GEMV_COMPLEX_SPECIALIZE_ROW(double, std::complex<double>, std::complex<double>)
2833
2834#endif // EIGEN_MATRIX_VECTOR_PRODUCT_ALTIVEC_H
@ Unaligned
Definition Constants.h:236
@ ColMajor
Definition Constants.h:319