11#ifndef EIGEN_MATRIX_VECTOR_PRODUCT_ALTIVEC_H
12#define EIGEN_MATRIX_VECTOR_PRODUCT_ALTIVEC_H
15#include "../../InternalHeaderCheck.h"
17#if defined(__MMA__) && !EIGEN_ALTIVEC_DISABLE_MMA
18#if EIGEN_COMP_LLVM || (__GNUC__ > 10 || __GNUC_MINOR__ >= 3)
22#if !EIGEN_COMP_LLVM && (__GNUC__ < 11)
24#define GCC_ONE_VECTORPAIR_BUG
32#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
33#define EIGEN_POWER_GEMV_PREFETCH(p) prefetch(p)
35#define EIGEN_POWER_GEMV_PREFETCH(p)
39#if !__has_builtin(__builtin_vsx_assemble_pair)
40#define __builtin_vsx_assemble_pair __builtin_mma_assemble_pair
42#if !__has_builtin(__builtin_vsx_disassemble_pair)
43#define __builtin_vsx_disassemble_pair __builtin_mma_disassemble_pair
48#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
49 __builtin_vsx_assemble_pair(&dst, (__vector unsigned char)src2, (__vector unsigned char)src1)
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)
56#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
57 __builtin_vsx_assemble_pair(&dst, (__vector unsigned char)src1, (__vector unsigned char)src2)
60#define GEMV_BUILDPAIR_MMA(dst, src1, src2) \
61 __builtin_vsx_build_pair(&dst, (__vector unsigned char)src1, (__vector unsigned char)src2)
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>)))
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)));
76template <
typename ResScalar>
77EIGEN_ALWAYS_INLINE
void storeMaddData(ResScalar* res, ResScalar& alpha, ResScalar& data) {
78 *res += (alpha * data);
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)
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)
85#define GEMV_GETN(N) (((N) * ResPacketSize) >> 2)
87#define GEMV_LOADPACKET_COL(iter) lhs.template load<LhsPacket, LhsAlignment>(i + ((iter) * LhsPacketSize), j)
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)
94#define GEMV_UNUSED_VAR(iter, N, which) \
95 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) <= iter) { \
96 EIGEN_UNUSED_VARIABLE(which##iter); \
99#define GEMV_UNUSED_EXTRA_VAR(iter, N, which) \
100 EIGEN_IF_CONSTEXPR (N <= iter) { \
101 EIGEN_UNUSED_VARIABLE(which##iter); \
104#define GEMV_UNUSED_EXTRA(N, which) GEMV_UNROLL3(GEMV_UNUSED_EXTRA_VAR, N, which)
106#define GEMV_UNUSED(N, which) GEMV_UNROLL3(GEMV_UNUSED_VAR, N, which)
108#define GEMV_INIT_MMA(iter, N) \
109 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
110 __builtin_mma_xxsetaccz(&e##iter); \
114#define GEMV_LOADPAIR_COL_MMA(iter1, iter2) \
115 GEMV_BUILDPAIR_MMA(b##iter1, GEMV_LOADPACKET_COL(iter2), GEMV_LOADPACKET_COL((iter2) + 1));
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));
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); \
128 GEMV_LOADPAIR_COL_MMA(iter, iter << 1) \
129 EIGEN_UNUSED_VARIABLE(g##iter); \
132 EIGEN_UNUSED_VARIABLE(b##iter); \
133 EIGEN_UNUSED_VARIABLE(g##iter); \
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); \
141 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter, b##iter, a0); \
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); \
151 GEMV_LOADPAIR_COL_MMA(iter2, iter2 << 1) \
152 GEMV_LOADPAIR_COL_MMA(iter3, iter3 << 1) \
155 EIGEN_UNUSED_VARIABLE(b##iter2); \
156 EIGEN_UNUSED_VARIABLE(b##iter3); \
158 EIGEN_UNUSED_VARIABLE(g##iter2); \
159 EIGEN_UNUSED_VARIABLE(g##iter3);
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) { \
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]); \
169 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter2, b##iter2, a0); \
170 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&e##iter3, b##iter3, a0); \
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)) \
179 GEMV_UNROLL(GEMV_LOAD1A_COL_MMA, N) \
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)) \
186 GEMV_UNROLL(GEMV_WORK1A_COL_MMA, N) \
189#define GEMV_LOAD_COL_MMA(N) GEMV_UNROLL(GEMV_LOAD1A_COL_MMA, N)
191#define GEMV_WORK_COL_MMA(N) GEMV_UNROLL(GEMV_WORK1A_COL_MMA, N)
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]; \
203#define GEMV_LOADPAIR2_COL_MMA(iter1, iter2) \
204 b##iter1 = *reinterpret_cast<__vector_pair*>(res + i + ((iter2) * ResPacketSize));
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); \
212 GEMV_LOADPAIR2_COL_MMA(iter2, iter2 << 1); \
213 GEMV_LOADPAIR2_COL_MMA(iter3, iter3 << 1); \
216 EIGEN_UNUSED_VARIABLE(b##iter2); \
217 EIGEN_UNUSED_VARIABLE(b##iter3); \
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]);
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" \
232 : "wa"(result##iter3.packet[0]), "wa"(result##iter2.packet[0]), "wa"(palpha)); \
234 __asm__("xvmaddadp %0,%x1,%x3\n\txvmaddadp %L0,%x2,%x3" \
236 : "wa"(result##iter2.packet[2]), "wa"(result##iter2.packet[0]), "wa"(palpha)); \
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); \
245 GEMV_WORKPAIR2_COL_MMA(iter2, iter2, iter2 << 1); \
246 GEMV_WORKPAIR2_COL_MMA(iter3, iter3, iter3 << 1); \
250#define GEMV_STOREPAIR2_COL_MMA(iter1, iter2) \
251 *reinterpret_cast<__vector_pair*>(res + i + ((iter2) * ResPacketSize)) = b##iter1;
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]); \
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) \
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); \
269 GEMV_STOREPAIR2_COL_MMA(iter2, iter2 << 1) \
270 GEMV_STOREPAIR2_COL_MMA(iter3, iter3 << 1) \
274#define GEMV_PROCESS_COL_ONE_MMA(N) \
275 GEMV_UNROLL(GEMV_INIT_MMA, N) \
277 __vector_pair b0, b1, b2, b3, b4, b5, b6, b7; \
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) \
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)) \
292 i += (ResPacketSize * N);
295#define GEMV_INIT(iter, N) \
296 EIGEN_IF_CONSTEXPR (N > iter) { \
297 c##iter = pset1<ResPacket>(ResScalar(0)); \
299 EIGEN_UNUSED_VARIABLE(c##iter); \
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); \
308#define GEMV_PREFETCH(iter, N)
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); \
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)))); \
323#define GEMV_PROCESS_COL_ONE(N) \
324 GEMV_UNROLL(GEMV_INIT, N) \
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);
335#define GEMV_PROCESS_COL(N) GEMV_PROCESS_COL_ONE_MMA(N)
337#define GEMV_PROCESS_COL(N) GEMV_PROCESS_COL_ONE(N)
342template <
typename LhsPacket,
typename RhsPacket,
bool accumulate>
343EIGEN_ALWAYS_INLINE
void pger_vecMMA_acc(__vector_quad* acc,
const RhsPacket& a,
const LhsPacket& b) {
345 __builtin_mma_xvf32gerpp(acc, (__vector
unsigned char)a, (__vector
unsigned char)b);
347 __builtin_mma_xvf32ger(acc, (__vector
unsigned char)a, (__vector
unsigned char)b);
352template <
typename LhsPacket,
typename RhsPacket,
bool accumulate>
353EIGEN_ALWAYS_INLINE
void pger_vecMMA_acc(__vector_quad* acc, __vector_pair& a,
const LhsPacket& b) {
355 __builtin_mma_xvf64gerpp(acc, a, (__vector
unsigned char)b);
357 __builtin_mma_xvf64ger(acc, a, (__vector
unsigned char)b);
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;
367 typedef typename Traits::LhsPacket LhsPacket;
368 typedef typename Traits::RhsPacket RhsPacket;
369 typedef typename Traits::ResPacket ResPacket;
371 EIGEN_UNUSED_VARIABLE(resIncr);
372 eigen_internal_assert(resIncr == 1);
379 conj_helper<LhsScalar, RhsScalar, false, false> cj;
380 conj_helper<LhsPacket, RhsPacket, false, false> pcj;
382 const Index lhsStride = lhs.stride();
387 ResPacketSize = Traits::ResPacketSize,
388 LhsPacketSize = Traits::LhsPacketSize,
389 RhsPacketSize = Traits::RhsPacketSize,
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;
397 const Index n1 = rows - 1 * ResPacketSize + 1;
398#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
399 const Index prefetch_dist = 64 * LhsPacketSize;
403 const Index block_cols = cols < 128 ? cols : (lhsStride *
sizeof(LhsScalar) < 16000 ? 16 : 8);
404 ResPacket palpha = pset1<ResPacket>(alpha);
406 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
407 Index jend = numext::mini(j2 + block_cols, cols);
409 ResPacket c0, c1, c2, c3, c4, c5, c6, c7;
411 __vector_quad e0, e1, e2, e3, e4, e5, e6, e7;
412 PacketBlock<ResPacket, 4> result0, result1, result2, result3, result4, result5, result6, result7;
414 GEMV_UNUSED(8, result)
415 GEMV_UNUSED_EXTRA(1, c)
417#ifndef GCC_ONE_VECTORPAIR_BUG
432 GEMV_PROCESS_COL_ONE(1)
434 for (; i < rows; ++i) {
438 d0 += cj.pmul(lhs(i, j), rhs2(j, 0));
439 }
while (++j < jend);
440 res[i] += alpha * d0;
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);
456template <Index num_acc,
bool extraRows, Index size>
457EIGEN_ALWAYS_INLINE
void outputVecColResults(Packet4f (&acc)[num_acc][size],
float* result,
const Packet4f& pAlpha,
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);
463 EIGEN_IF_CONSTEXPR (extraRows) {
464 outputVecCol<true>(acc[real_acc][0], result + real_acc * 4, pAlpha, extra_rows);
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};
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);
475 EIGEN_IF_CONSTEXPR (!zero) {
476 b1 = lhs.template loadPacket<Packet8bf>(k * 4, 1);
478 a0[k + 0][1] = oneConvertBF16Hi(b1.m_val);
480 a0[k + 0][0] = oneConvertBF16Hi(c0.m_val);
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);
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]);
499template <
typename RhsMapper,
bool linear>
500struct loadColData_impl {
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];
506 for (Index i = 0; i < n; i++) {
507 to[i] = rhs(j + i, 0);
509 return pload<Packet8bf>(to);
513template <
typename RhsMapper>
514struct loadColData_impl<RhsMapper, true> {
516 static EIGEN_ALWAYS_INLINE Packet8bf run(RhsMapper& rhs, Index j) {
517 return rhs.template loadPacket<Packet8bf>(j + 0, 0);
521template <
typename RhsMapper,
bool linear>
522EIGEN_ALWAYS_INLINE Packet8bf loadColData(RhsMapper& rhs, Index j) {
523 return loadColData_impl<RhsMapper, linear>::run(rhs, j);
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);
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);
536 using LhsSubMapper =
typename LhsMapper::SubMapper;
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);
543 multVecVSX<num_acc, zero>(acc, a0, b0);
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];
554#define MAX_BFLOAT16_VEC_ACC_VSX 8
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,
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);
564 Packet4f acc[num_acc][2];
566 zeroAccumulators<num_acc, 2>(acc);
568 using LhsSubMapper =
typename LhsMapper::SubMapper;
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);
575 vecColLoopVSX<num_acc, LhsSubMapper, RhsMapper, true, linear>(cend - 1, lhs2, rhs, acc);
578 addResultsVSX<num_acc>(acc);
580 outputVecColResults<num_acc, extraRows, 2>(acc, result, pAlpha, extra_rows);
583 }
while (multiIters && (step <= rows - (row += step)));
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);
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) {
600 colVSXVecColLoopBodyExtraN<7, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
603 colVSXVecColLoopBodyExtraN<6, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
606 colVSXVecColLoopBodyExtraN<5, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
609 colVSXVecColLoopBodyExtraN<4, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
612 colVSXVecColLoopBodyExtraN<3, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
615 colVSXVecColLoopBodyExtraN<2, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
618 colVSXVecColLoopBodyExtraN<1, LhsMapper, RhsMapper, extraRows, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
621 EIGEN_IF_CONSTEXPR (extraRows) {
622 colVSXVecColLoopBody<1, LhsMapper, RhsMapper, true, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
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) {
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,
638 colVSXVecColLoopBodyExtra<LhsMapper, RhsMapper, true, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
640 colVSXVecColLoopBodyExtra<LhsMapper, RhsMapper, false, linear>(row, cend, rows, lhs, rhs, pAlpha, result);
644template <const Index size,
bool inc, Index delta>
645EIGEN_ALWAYS_INLINE
void storeBF16fromResult(bfloat16* dst,
const Packet8bf& data, Index resInc, Index extra) {
647 EIGEN_IF_CONSTEXPR (size < 8) {
648 pscatter_partial(dst + delta * resInc, data, resInc, extra);
650 pscatter(dst + delta * resInc, data, resInc);
653 EIGEN_IF_CONSTEXPR (size < 8) {
654 pstoreu_partial(dst + delta, data, extra);
656 pstoreu(dst + delta, data);
661template <const Index size,
bool inc = false>
662EIGEN_ALWAYS_INLINE
void convertPointerF32toBF16VSX(Index& i,
float* result, Index rows, bfloat16*& dst,
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);
671 EIGEN_IF_CONSTEXPR (size >= 32) {
672 r32.packet[2] = convertF32toBF16VSX(result + i + 16);
673 r32.packet[3] = convertF32toBF16VSX(result + i + 24);
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);
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);
684 dst += extra * resInc;
685 EIGEN_IF_CONSTEXPR (size != 32) break;
689template <
bool inc = false>
690EIGEN_ALWAYS_INLINE
void convertArrayPointerF32toBF16VSX(
float* result, Index rows, bfloat16* dst, Index resInc = 1) {
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);
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;
704 RhsSubMapper rhs2 = rhs.getSubMapper(j2, 0);
705 calcVSXVecColLoops<LhsMapper, RhsSubMapper, false>(jend - j2, rows, lhs, rhs2, pAlpha, result);
709template <
typename RhsMapper,
typename LhsMapper>
710struct UseStride<RhsMapper, LhsMapper,
711 std::enable_if_t<std::is_member_function_pointer<decltype(&RhsMapper::stride)>::value>>
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;
717 RhsSubMapper rhs2 = rhs.getSubMapper(j2, 0);
718 if (rhs.stride() == 1) {
719 calcVSXVecColLoops<LhsMapper, RhsSubMapper, true>(jend - j2, rows, lhs, rhs2, pAlpha, result);
721 calcVSXVecColLoops<LhsMapper, RhsSubMapper, false>(jend - j2, rows, lhs, rhs2, pAlpha, result);
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);
737 const Index lhsStride = lhs.stride();
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);
744 ei_declare_aligned_stack_constructed_variable(
float, result, rows, 0);
746 convertArrayPointerBF16toF32(result, 1, rows, res);
748 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
749 Index jend = numext::mini(j2 + block_cols, cols);
751 using LhsSubMapper =
typename LhsMapper::SubMapper;
753 LhsSubMapper lhs2 = lhs.getSubMapper(0, j2);
754 UseStride<RhsMapper, LhsSubMapper>::run(j2, jend, rows, lhs2, rhs2, pAlpha, result);
757 convertArrayPointerF32toBF16VSX(result, rows, res);
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;
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);
768 if (num_acc > (k + 3)) {
769 pstoreu(result + k, d0);
772 pstoreu_partial(result + k, d0, extra);
774 memcpy((
void*)(result + k), (
void*)(&d0),
sizeof(
float) * extra);
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);
788 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 8);
790 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 12);
792 acc[k][0] += vec_sld(acc[k][0], acc[k][0], 4);
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])));
807 acc[k + 0][0] =
reinterpret_cast<Packet4f
>(vec_perm(acc[k + 0][0], acc[k + 2][0], p16uc_TRANSPOSE64_HI));
814EIGEN_ALWAYS_INLINE Packet8us loadPacketPartialZero(
const Packet8us& data, Index extra_cols) {
815 Packet16uc shift = pset1<Packet16uc>(8 * 2 * (8 - extra_cols));
817 return reinterpret_cast<Packet8us
>(vec_slo(vec_sro(
reinterpret_cast<Packet16uc
>(data), shift), shift));
819 return reinterpret_cast<Packet8us
>(vec_sro(vec_slo(
reinterpret_cast<Packet16uc
>(data), shift), shift));
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,
827 Packet4f a0[num_acc][2], b0[2];
831 b1 = rhs.template loadPacketPartial<Packet8bf>(j, extra_cols);
833 b1 = loadPacketPartialZero(b1.m_val, extra_cols);
836 b1 = rhs.template loadPacket<Packet8bf>(j);
838 b0[0] = oneConvertBF16Hi(b1.m_val);
839 b0[1] = oneConvertBF16Lo(b1.m_val);
841 const LhsMapper lhs2 = lhs.getSubMapper(0, j);
842 for (Index k = 0; k < num_acc; k++) {
844 a1 = lhs2.template loadPacketPartial<Packet8bf>(k, 0, extra_cols);
846 a1 = loadPacketPartialZero(a1.m_val, extra_cols);
849 a1 = lhs2.template loadPacket<Packet8bf>(k, 0);
851 a0[k][0] = oneConvertBF16Hi(a1.m_val);
852 a0[k][1] = oneConvertBF16Lo(a1.m_val);
855 multVecVSX<num_acc, false>(acc, a0, b0);
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],
862 for (; j + 8 <= cols; j += 8) {
863 multVSXVecLoop<num_acc, LhsMapper, RhsMapper, false>(acc, lhs, rhs, j, extra_cols);
867 multVSXVecLoop<num_acc, LhsMapper, RhsMapper, true>(acc, lhs, rhs, j, extra_cols);
871template <const Index num_acc,
typename LhsMapper,
typename RhsMapper>
872void colVSXVecLoopBody(Index& row, Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs,
const Packet4f& pAlpha,
874 constexpr bool multiIters = (num_acc == MAX_BFLOAT16_VEC_ACC_VSX);
875 const Index extra_cols = (cols & 7);
878 Packet4f acc[num_acc][2];
880 zeroAccumulators<num_acc, 2>(acc);
882 const LhsMapper lhs2 = lhs.getSubMapper(row, 0);
883 vecVSXLoop<num_acc, LhsMapper, RhsMapper>(cols, lhs2, rhs, acc, extra_cols);
885 addResultsVSX<num_acc>(acc);
887 preduxVecResultsVSX<num_acc>(acc);
889 outputVecResults<num_acc, 2>(acc, result, pAlpha);
892 }
while (multiIters && (num_acc <= rows - (row += num_acc)));
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);
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) {
908 colVSXVecLoopBodyExtraN<7, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
911 colVSXVecLoopBodyExtraN<6, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
914 colVSXVecLoopBodyExtraN<5, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
917 colVSXVecLoopBodyExtraN<4, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
920 colVSXVecLoopBodyExtraN<3, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
923 colVSXVecLoopBodyExtraN<2, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
926 colVSXVecLoopBodyExtraN<1, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
931template <
typename LhsMapper,
typename RhsMapper>
932EIGEN_ALWAYS_INLINE
void calcVSXVecLoops(Index cols, Index rows, LhsMapper& lhs, RhsMapper& rhs,
const Packet4f& pAlpha,
935 if (rows >= MAX_BFLOAT16_VEC_ACC_VSX) {
936 colVSXVecLoopBody<MAX_BFLOAT16_VEC_ACC_VSX, LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
939 colVSXVecLoopBodyExtra<LhsMapper, RhsMapper>(row, cols, rows, lhs, rhs, pAlpha, result);
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;
950 LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
952 eigen_internal_assert(rhs.stride() == 1);
954 float falpha = Eigen::bfloat16_impl::bfloat16_to_float(alpha);
955 const Packet4f pAlpha = pset1<Packet4f>(falpha);
957 ei_declare_aligned_stack_constructed_variable(
float, result, rows, 0);
959 convertArrayPointerBF16toF32(result, 1, rows, res);
961 convertArrayPointerBF16toF32<true>(result, 1, rows, res, resIncr);
963 calcVSXVecLoops<LhsMapper, LinearMapper>(cols, rows, lhs, rhs2, pAlpha, result);
965 convertArrayPointerF32toBF16VSX(result, rows, res);
967 convertArrayPointerF32toBF16VSX<true>(result, rows, res, resIncr);
971#undef MAX_BFLOAT16_VEC_ACC_VSX
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};
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};
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};
1007#define COMPLEX_DELTA 0
1009#define COMPLEX_DELTA 2
1013EIGEN_ALWAYS_INLINE Packet2cf pconj2(
const Packet2cf& a) {
1014 return Packet2cf(pxor(a.v,
reinterpret_cast<Packet4f
>(p16uc_COMPLEX32_CONJ_XOR)));
1017EIGEN_ALWAYS_INLINE Packet1cd pconj2(
const Packet1cd& a) {
1018 return Packet1cd(pxor(a.v,
reinterpret_cast<Packet2d
>(p16uc_COMPLEX64_CONJ_XOR)));
1022EIGEN_ALWAYS_INLINE Packet2cf pconjinv(
const Packet2cf& a) {
1023#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1024 return Packet2cf(Packet4f(vec_neg(Packet2d(a.v))));
1026 return Packet2cf(pxor(a.v,
reinterpret_cast<Packet4f
>(p16uc_COMPLEX32_CONJ_XOR2)));
1030EIGEN_ALWAYS_INLINE Packet1cd pconjinv(
const Packet1cd& a) {
1031 return Packet1cd(pxor(a.v,
reinterpret_cast<Packet2d
>(p16uc_COMPLEX64_CONJ_XOR2)));
1034#if defined(_ARCH_PWR8) && (!EIGEN_COMP_LLVM || __clang_major__ >= 12)
1039EIGEN_ALWAYS_INLINE Packet2cf pcplxflipconj(
const Packet2cf& a) {
1041 return Packet2cf(Packet4f(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX32_CONJ_XOR, p16uc_COMPLEX32_XORFLIP)));
1043 return pcplxflip(pconj2(a));
1047EIGEN_ALWAYS_INLINE Packet1cd pcplxflipconj(
const Packet1cd& a) {
1049 return Packet1cd(Packet2d(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX64_CONJ_XOR, p16uc_COMPLEX64_XORFLIP)));
1051 return pcplxflip(pconj2(a));
1056EIGEN_ALWAYS_INLINE Packet2cf pcplxconjflip(
const Packet2cf& a) {
1058 return Packet2cf(Packet4f(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX32_CONJ_XOR2, p16uc_COMPLEX32_XORFLIP)));
1060 return pconj2(pcplxflip(a));
1064EIGEN_ALWAYS_INLINE Packet1cd pcplxconjflip(
const Packet1cd& a) {
1066 return Packet1cd(Packet2d(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX64_CONJ_XOR2, p16uc_COMPLEX64_XORFLIP)));
1068 return pconj2(pcplxflip(a));
1073EIGEN_ALWAYS_INLINE Packet2cf pnegate2(
const Packet2cf& a) {
1074#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1075 return Packet2cf(vec_neg(a.v));
1077 return Packet2cf(pxor(a.v,
reinterpret_cast<Packet4f
>(p16uc_COMPLEX32_NEGATE)));
1081EIGEN_ALWAYS_INLINE Packet1cd pnegate2(
const Packet1cd& a) {
1082#ifdef EIGEN_VECTORIZE_POWER8_VECTOR
1083 return Packet1cd(vec_neg(a.v));
1085 return Packet1cd(pxor(a.v,
reinterpret_cast<Packet2d
>(p16uc_COMPLEX64_NEGATE)));
1090EIGEN_ALWAYS_INLINE Packet2cf pcplxflipnegate(
const Packet2cf& a) {
1092 return Packet2cf(Packet4f(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX32_NEGATE, p16uc_COMPLEX32_XORFLIP)));
1094 return pcplxflip(pnegate2(a));
1098EIGEN_ALWAYS_INLINE Packet1cd pcplxflipnegate(
const Packet1cd& a) {
1100 return Packet1cd(Packet2d(vec_permxor(Packet16uc(a.v), p16uc_COMPLEX64_NEGATE, p16uc_COMPLEX64_XORFLIP)));
1102 return pcplxflip(pnegate2(a));
1107EIGEN_ALWAYS_INLINE Packet2cf pcplxflip2(
const Packet2cf& a) {
1108 return Packet2cf(Packet4f(vec_perm(Packet16uc(a.v), Packet16uc(a.v), p16uc_COMPLEX32_XORFLIP)));
1111EIGEN_ALWAYS_INLINE Packet1cd pcplxflip2(
const Packet1cd& a) {
1112#ifdef EIGEN_VECTORIZE_VSX
1113 return Packet1cd(__builtin_vsx_xxpermdi(a.v, a.v, 2));
1115 return Packet1cd(Packet2d(vec_perm(Packet16uc(a.v), Packet16uc(a.v), p16uc_COMPLEX64_XORFLIP)));
1120EIGEN_ALWAYS_INLINE Packet4f pload_complex_half(std::complex<float>* src) {
1122#ifdef EIGEN_VECTORIZE_VSX
1124 __asm__(
"lxsdx %x0,%y1" :
"=wa"(t) :
"Z"(*src));
1126 *
reinterpret_cast<std::complex<float>*
>(
reinterpret_cast<float*
>(&t) + COMPLEX_DELTA) = *src;
1132template <
typename RhsScalar>
1133EIGEN_ALWAYS_INLINE
void pload_realimag(RhsScalar* src, Packet4f& r, Packet4f& i) {
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)));
1138 Packet4f t = pload_complex_half(src);
1139 r = vec_splat(t, COMPLEX_DELTA + 0);
1140 i = vec_splat(t, COMPLEX_DELTA + 1);
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)));
1150 Packet2d t = ploadu<Packet2d>(
reinterpret_cast<double*
>(src));
1151 r = vec_splat(t, 0);
1152 i = vec_splat(t, 1);
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};
1160const Packet16uc p16uc_MERGEO = {0x04, 0x05, 0x06, 0x07, 0x14, 0x15, 0x16, 0x17,
1161 0x0C, 0x0D, 0x0E, 0x0F, 0x1C, 0x1D, 0x1E, 0x1F};
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);
1172 r = vec_perm(t, t, p16uc_MERGEE);
1173 i = vec_perm(t, t, p16uc_MERGEO);
1177template <
typename RhsScalar>
1178EIGEN_ALWAYS_INLINE
void pload_realimag_row(RhsScalar* src, Packet2d& r, Packet2d& i) {
1179 return pload_realimag(src, r, i);
1183EIGEN_ALWAYS_INLINE Packet4f pload_realimag_combine(std::complex<float>* src) {
1184#ifdef EIGEN_VECTORIZE_VSX
1186 __asm__(
"lxvdsx %x0,%y1" :
"=wa"(ret) :
"Z"(*(
reinterpret_cast<double*
>(src) + 0)));
1189 return Packet4f(ploaddup<Packet2d>(
reinterpret_cast<double*
>(src)));
1193EIGEN_ALWAYS_INLINE Packet2d pload_realimag_combine(std::complex<double>* src) {
return ploadu<Packet1cd>(src).v; }
1196EIGEN_ALWAYS_INLINE Packet4f pload_realimag_combine_row(std::complex<float>* src) {
return ploadu<Packet2cf>(src).v; }
1198EIGEN_ALWAYS_INLINE Packet2d pload_realimag_combine_row(std::complex<double>* src) {
return ploadu<Packet1cd>(src).v; }
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);
1206 return ploadu<Packet4f>(
reinterpret_cast<float*
>(src));
1210template <
typename ResPacket>
1211EIGEN_ALWAYS_INLINE Packet2d pload_complex(std::complex<double>* src) {
1212 return ploadu<Packet2d>(
reinterpret_cast<double*
>(src));
1216template <
typename ResPacket>
1217EIGEN_ALWAYS_INLINE Packet4f pload_complex(Packet2cf* src) {
1221template <
typename ResPacket>
1222EIGEN_ALWAYS_INLINE Packet2d pload_complex(Packet1cd* src) {
1227EIGEN_ALWAYS_INLINE Packet4f pload_complex_full(std::complex<float>* src) {
1228 return Packet4f(ploaddup<Packet2d>(
reinterpret_cast<double*
>(src)));
1231EIGEN_ALWAYS_INLINE Packet2d pload_complex_full(std::complex<double>* src) {
return ploadu<Packet1cd>(src).v; }
1234EIGEN_ALWAYS_INLINE Packet4f pload_complex_full_row(std::complex<float>* src) {
return ploadu<Packet2cf>(src).v; }
1236EIGEN_ALWAYS_INLINE Packet2d pload_complex_full_row(std::complex<double>* src) {
return pload_complex_full(src); }
1239EIGEN_ALWAYS_INLINE Packet4f pload_real(
float* src) {
return pset1<Packet4f>(*src); }
1241EIGEN_ALWAYS_INLINE Packet2d pload_real(
double* src) {
return pset1<Packet2d>(*src); }
1243EIGEN_ALWAYS_INLINE Packet4f pload_real(
const Packet4f& src) {
return src; }
1245EIGEN_ALWAYS_INLINE Packet2d pload_real(
const Packet2d& src) {
return src; }
1248EIGEN_ALWAYS_INLINE Packet4f pload_real_full(
float* src) {
1249 Packet4f ret = ploadu<Packet4f>(src);
1250 return vec_mergeh(ret, ret);
1253EIGEN_ALWAYS_INLINE Packet2d pload_real_full(
double* src) {
return pload_real(src); }
1255EIGEN_ALWAYS_INLINE Packet4f pload_real_full(std::complex<float>* src) {
1256 return pload_complex_full(src);
1259EIGEN_ALWAYS_INLINE Packet2d pload_real_full(std::complex<double>* src) {
1260 return pload_complex_full(src);
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);
1269 return ploadu<Packet4f>(src);
1273template <
typename ResPacket>
1274EIGEN_ALWAYS_INLINE Packet2d pload_real_row(
double* src) {
1275 return pload_real(src);
1282EIGEN_ALWAYS_INLINE Packet2cf padd(
const Packet2cf& a,
const std::complex<float>& b) {
1283 EIGEN_UNUSED_VARIABLE(b);
1287EIGEN_ALWAYS_INLINE Packet1cd padd(
const Packet1cd& a,
const std::complex<double>& b) {
1288 EIGEN_UNUSED_VARIABLE(b);
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());
1299template <
typename Scalar,
typename ResScalar,
typename ResPacket,
int which>
1300EIGEN_ALWAYS_INLINE Packet2cf pset1_complex(std::complex<float>& alpha) {
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];
1309template <
typename Scalar,
typename ResScalar,
typename ResPacket,
int which>
1310EIGEN_ALWAYS_INLINE Packet1cd pset1_complex(std::complex<double>& alpha) {
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));
1318template <
typename Packet>
1319EIGEN_ALWAYS_INLINE Packet pset_zero() {
1320 return pset1<Packet>(__UNPACK_TYPE__(Packet)(0));
1324EIGEN_ALWAYS_INLINE Packet2cf pset_zero<Packet2cf>() {
1325 return Packet2cf(pset1<Packet4f>(
float(0)));
1329EIGEN_ALWAYS_INLINE Packet1cd pset_zero<Packet1cd>() {
1330 return Packet1cd(pset1<Packet2d>(
double(0)));
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>();
1344template <
typename PResPacket,
typename ResPacket,
typename ResScalar,
typename Scalar>
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);
1357template <
typename ScalarPacket,
typename AlphaData>
1358EIGEN_ALWAYS_INLINE ScalarPacket pmadd_complex(
const ScalarPacket& c0,
const ScalarPacket& c2,
const ScalarPacket& c4,
1360 return pmadd(c2, b0.separate.i.v, pmadd(c0, b0.separate.r.v, c4));
1364template <
typename Scalar,
typename ScalarPacket,
typename PResPacket,
typename ResPacket,
typename ResScalar,
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);
1373 ScalarPacket c4 = pload_complex<ResPacket>(res);
1374 PResPacket c3 = PResPacket(pmadd_complex<ScalarPacket, AlphaData>(c0.v, c2.v, c4, b0));
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,
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);
1393 __vector_pair a = *
reinterpret_cast<__vector_pair*
>(res + (iter2 * ResPacketSize));
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);
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));
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));
1409 *
reinterpret_cast<__vector_pair*
>(res + (iter2 * ResPacketSize)) = a;
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)));
1420 return lhs.template load<LhsPacket, Unaligned>(i + 0, j);
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);
1431 return vec_madd(a, b, c);
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);
1441 return vec_madd(a, b, c);
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;
1450 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1451 b0 = pset1<RhsPacket>(*b);
1453 b0 = ploadu<RhsPacket>(b);
1455 c0 = pcj.pmadd(a0, b0, c0);
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);
1466 pload_realimag_row<RhsScalar>(b, br, bi);
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);
1473 c0 = PResPacket(cr);
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) {
1481 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1482 b0 = pload_complex_full(b);
1484 b0 = pload_complex_full_row(b);
1486 ScalarPacket cri = pmadd_complex_real<PResPacket, ScalarPacket, ConjugateRhs>(a0, b0, c0.v);
1487 c0 = PResPacket(cri);
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);
1496 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1499 b0 = pload_real_row<ResPacket>(b);
1501 ScalarPacket cri = pmadd_complex_real<PResPacket, ScalarPacket, ConjugateLhs>(a1, b0, c0.v);
1502 c0 = PResPacket(cri);
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); \
1513GEMV_MULT_COMPLEX_COMPLEX(Packet2cf, std::complex<float>, Packet2cf)
1514GEMV_MULT_COMPLEX_COMPLEX(Packet1cd, std::complex<double>, Packet1cd)
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); \
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)
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); \
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>)
1544template <
typename T>
1545EIGEN_ALWAYS_INLINE T convertReal(T a) {
1549EIGEN_ALWAYS_INLINE Packet4f convertReal(
const Packet2cf& a) {
return a.v; }
1551EIGEN_ALWAYS_INLINE Packet2d convertReal(
const Packet1cd& a) {
return a.v; }
1554template <
typename T>
1555EIGEN_ALWAYS_INLINE T convertComplex(T a) {
1559EIGEN_ALWAYS_INLINE Packet2cf convertComplex(
const Packet4f& a) {
return Packet2cf(a); }
1561EIGEN_ALWAYS_INLINE Packet1cd convertComplex(
const Packet2d& a) {
return Packet1cd(a); }
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));
1569template <
typename ScalarPacket,
typename LhsPacket,
typename SLhsPacket,
typename ResPacket>
1570EIGEN_ALWAYS_INLINE
void pload_complex_MMA(__vector_pair&) {
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);
1580 __builtin_mma_xvf32gerpp(acc, (__vector
unsigned char)a, (__vector
unsigned char)b);
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);
1590 __builtin_mma_xvf64gerpp(acc, (__vector_pair)a, (__vector
unsigned char)b);
1594template <
typename LhsPacket,
typename RhsPacket,
bool NegativeAccumulate>
1595EIGEN_ALWAYS_INLINE
void pger_vecMMA(__vector_quad*, __vector_pair&,
const Packet4f&) {
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);
1608 return pger_vecMMA<RealPacket, RealPacket, false>(c, b, a.v);
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);
1620 return pger_vecMMA<RealPacket, __vector_pair, false>(c, a, b);
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);
1633 return pger_vecMMA<RealPacket, RealPacket, false>(c, a2, b2);
1636 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1637 return pger_vecMMA<RealPacket, RealPacket, false>(c, b, a2);
1639 return pger_vecMMA<RealPacket, RealPacket, false>(c, a2, b);
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);
1651 return pger_vecMMA<RealPacket, __vector_pair, false>(c, a, b);
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) {
1660 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1661 b0 = pload_realimag_combine(b);
1663 b0 = pload_realimag_combine_row(b);
1665 pmadd_complex_complex_MMA<ScalarPacket, LhsPacket, ConjugateLhs, ConjugateRhs, false>(a0, b0, c0);
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);
1674 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1677 b0 = pload_real_row<ResPacket>(b);
1679 pmadd_complex_real_MMA<ScalarPacket, LhsPacket, ConjugateLhs, ColMajor>(a0, b0, c0);
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) {
1687 EIGEN_IF_CONSTEXPR (StorageOrder == ColMajor) {
1688 b0 = pload_complex_full(b);
1690 b0 = pload_complex_full_row(b);
1692 pmadd_complex_real_MMA<ScalarPacket, LhsPacket, ConjugateRhs,
1693 (
sizeof(RhsScalar) ==
sizeof(std::complex<float>)) ? StorageOrder :
ColMajor>(a0, b0, c0);
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); \
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>)
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);
1716 gemv_mult_real_complex_MMA<ScalarPacket, LhsPacket, SLhsPacket, RhsScalar, ResPacket, ConjugateLhs, ConjugateRhs,
1717 StorageOrder>(a0, b, c0);
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); \
1729GEMV_MULT_REAL_COMPLEX_MMA(Packet4f, std::complex<float>)
1730GEMV_MULT_REAL_COMPLEX_MMA(Packet2d, std::complex<double>)
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); \
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)
1746template <
typename Scalar,
typename ScalarPacket,
typename LhsPacket,
typename RhsPacket,
bool ConjugateLhs,
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;
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;
1767 result0.packet[1] = pconjinv(convertComplex(result0.packet[1])).v;
1768 result0.packet[3] = pconjinv(convertComplex(result0.packet[3])).v;
1770 result0.packet[0] = vec_add(result0.packet[0], result0.packet[1]);
1771 result0.packet[2] = vec_add(result0.packet[2], result0.packet[3]);
1773 result0.packet[0][1] = result0.packet[1][1];
1774 result0.packet[2][1] = result0.packet[3][1];
1779template <
typename Scalar,
typename ScalarPacket,
typename LhsPacket,
typename RhsPacket,
bool ConjugateLhs,
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;
1788 EIGEN_IF_CONSTEXPR (ConjugateRhs) {
1789 result0.packet[1] = pcplxconjflip(convertComplex(result0.packet[1])).v;
1791 result0.packet[1] = pcplxflipconj(convertComplex(result0.packet[1])).v;
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;
1800 result0.packet[0] = vec_mergee(result0.packet[0], result0.packet[1]);
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);
1810 disassembleResults4<Scalar, ScalarPacket, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>(c0, result0);
1815#define GEMV_GETN_COMPLEX(N) (((N) * ResPacketSize) >> 1)
1817#define GEMV_LOADPACKET_COL_COMPLEX(iter) \
1818 loadLhsPacket<Scalar, LhsScalar, LhsMapper, PLhsPacket>(lhs, i + ((iter) * ResPacketSize), j)
1820#define GEMV_LOADPACKET_COL_COMPLEX_DATA(iter) convertReal(GEMV_LOADPACKET_COL_COMPLEX(iter))
1823#define GEMV_INIT_COL_COMPLEX_MMA(iter, N) \
1824 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
1825 __builtin_mma_xxsetaccz(&e0##iter); \
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);
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); \
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)); \
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); \
1851 GEMV_LOADPAIR_COL_COMPLEX_MMA(iter, iter << 1) \
1854 EIGEN_UNUSED_VARIABLE(a##iter); \
1855 EIGEN_UNUSED_VARIABLE(f##iter); \
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); \
1864 gemv_mult_complex_MMA<ScalarPacket, LhsScalar, PLhsPacket, __vector_pair, RhsScalar, RhsPacket, ResPacket, \
1865 ConjugateLhs, ConjugateRhs, ColMajor>(a##iter, b, &e0##iter); \
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));
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); \
1878 GEMV_LOADPAIR2_COL_COMPLEX_MMA(iter2, iter2 << 1); \
1879 GEMV_LOADPAIR2_COL_COMPLEX_MMA(iter3, iter3 << 1); \
1882 EIGEN_UNUSED_VARIABLE(a##iter2); \
1883 EIGEN_UNUSED_VARIABLE(a##iter3); \
1885 EIGEN_UNUSED_VARIABLE(f##iter2); \
1886 EIGEN_UNUSED_VARIABLE(f##iter3);
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) { \
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); \
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); \
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)) \
1910 GEMV_UNROLL(GEMV_LOAD1_COL_COMPLEX_MMA, N) \
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)) \
1917 GEMV_UNROLL(GEMV_WORK1_COL_COMPLEX_MMA, N) \
1920#define GEMV_LOAD_COL_COMPLEX_MMA(N) GEMV_UNROLL(GEMV_LOAD1_COL_COMPLEX_MMA, N)
1922#define GEMV_WORK_COL_COMPLEX_MMA(N) GEMV_UNROLL(GEMV_WORK1_COL_COMPLEX_MMA, N)
1925#define GEMV_DISASSEMBLE_COMPLEX_MMA(iter) \
1926 disassembleResults<Scalar, ScalarPacket, ResPacketSize, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>( \
1927 &e0##iter, result0##iter);
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)); \
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)); \
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); \
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); \
1965#define GEMV_PROCESS_COL_COMPLEX_ONE_MMA(N) \
1966 GEMV_UNROLL(GEMV_INIT_COL_COMPLEX_MMA, N) \
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) \
1978 GEMV_UNROLL_HALF(GEMV_STORE2_COL_COMPLEX_MMA, (N >> 1)) \
1980 i += (ResPacketSize * N);
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); \
1988 EIGEN_UNUSED_VARIABLE(c0##iter); \
1989 EIGEN_UNUSED_VARIABLE(c1##iter); \
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); \
1998 EIGEN_UNUSED_VARIABLE(f##iter); \
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); \
2006 pstoreu_pmadd_complex<Scalar, ScalarPacket, PResPacket, ResPacket, ResScalar, AlphaData>( \
2007 c0##iter, alpha_data, res + i + (iter * ResPacketSize)); \
2011#define GEMV_PROCESS_COL_COMPLEX_ONE(N) \
2012 GEMV_UNROLL(GEMV_INIT_COMPLEX, N) \
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);
2023#if defined(USE_GEMV_MMA) && (EIGEN_COMP_LLVM || defined(USE_SLOWER_GEMV_MMA))
2024#define USE_GEMV_COL_COMPLEX_MMA
2027#ifdef USE_GEMV_COL_COMPLEX_MMA
2028#define GEMV_PROCESS_COL_COMPLEX(N) GEMV_PROCESS_COL_COMPLEX_ONE_MMA(N)
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) \
2035 GEMV_PROCESS_COL_COMPLEX_ONE(N) \
2038#define GEMV_PROCESS_COL_COMPLEX(N) GEMV_PROCESS_COL_COMPLEX_ONE(N)
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;
2048 typedef typename Traits::LhsPacket LhsPacket;
2049 typedef typename Traits::RhsPacket RhsPacket;
2050 typedef typename Traits::ResPacket ResPacket;
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;
2057 EIGEN_UNUSED_VARIABLE(resIncr);
2058 eigen_internal_assert(resIncr == 1);
2062 LhsMapper lhs(alhs);
2063 RhsMapper rhs2(rhs);
2065 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
2067 const Index lhsStride = lhs.stride();
2072 ResPacketSize = PTraits::ResPacketSize,
2073 LhsPacketSize = PTraits::LhsPacketSize,
2074 RhsPacketSize = PTraits::RhsPacketSize,
2076#ifdef EIGEN_POWER_USE_GEMV_PREFETCH
2077 const Index prefetch_dist = 64 * LhsPacketSize;
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;
2085 const Index n1 = rows - 1 * ResPacketSize + 1;
2088 const Index block_cols = cols < 128 ? cols : (lhsStride *
sizeof(LhsScalar) < 16000 ? 16 : 8);
2090 typedef alpha_store<PResPacket, ResPacket, ResScalar, Scalar> AlphaData;
2091 AlphaData alpha_data(alpha);
2093 for (Index j2 = 0; j2 < cols; j2 += block_cols) {
2094 Index jend = numext::mini(j2 + block_cols, cols);
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;
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;
2104 GEMV_UNUSED(8, result0)
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)
2111#ifndef GCC_ONE_VECTORPAIR_BUG
2114 GEMV_PROCESS_COL_COMPLEX(8)
2118 GEMV_PROCESS_COL_COMPLEX(4)
2121 GEMV_PROCESS_COL_COMPLEX(2)
2128 GEMV_PROCESS_COL_COMPLEX_ONE(1)
2130 for (; i < rows; ++i) {
2134 d0 += cj.pmul(lhs(i, j), rhs2(j, 0));
2135 }
while (++j < jend);
2136 res[i] += alpha * d0;
2141template <
typename Scalar,
int N>
2147static Packet16uc p16uc_ELEMENT_3 = {0x0c, 0x0d, 0x0e, 0x0f, 0x1c, 0x1d, 0x1e, 0x1f,
2148 0x0c, 0x0d, 0x0e, 0x0f, 0x1c, 0x1d, 0x1e, 0x1f};
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);
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]);
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);
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]);
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;
2197 result0.packet[1] = pcplxflipconj(convertComplex(result0.packet[1])).v;
2199 result0.packet[0] = vec_add(result0.packet[0], result0.packet[1]);
2201 EIGEN_IF_CONSTEXPR (ConjugateLhs && (
sizeof(LhsPacket) ==
sizeof(std::complex<float>))) {
2202 result0.packet[0] = pconj2(convertComplex(result0.packet[0])).v;
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]);
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);
2221template <
typename ResScalar,
typename ResPacket,
typename LhsPacket,
typename RhsPacket,
bool ConjugateLhs,
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);
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);
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]);
2239template <
typename ResScalar,
typename ResPacket,
typename LhsPacket,
typename RhsPacket,
bool ConjugateLhs,
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;
2253 result0.packet[1] = pconj2(convertComplex(result0.packet[1])).v;
2254 result0.packet[3] = pconj2(convertComplex(result0.packet[3])).v;
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));
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);
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]);
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);
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);
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)
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)
2287#define GEMV_LOADPACKET_ROW(iter) lhs.template load<LhsPacket, Unaligned>(i + (iter), j)
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)
2294#define GEMV_UNUSED_ROW(N, which) GEMV_UNROLL3_ROW(GEMV_UNUSED_VAR, N, which)
2296#define GEMV_INIT_ROW(iter, N) \
2297 EIGEN_IF_CONSTEXPR (GEMV_GETN(N) > iter) { \
2298 __builtin_mma_xxsetaccz(&c##iter); \
2301#define GEMV_LOADPAIR_ROW(iter1, iter2) \
2302 GEMV_BUILDPAIR_MMA(b##iter1, GEMV_LOADPACKET_ROW(iter2), GEMV_LOADPACKET_ROW((iter2) + 1));
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)); \
2309 __vector_pair b##iter; \
2310 GEMV_LOADPAIR_ROW(iter, iter << 1) \
2311 pger_vecMMA_acc<LhsPacket, RhsPacket, true>(&c##iter, b##iter, a0); \
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); \
2320 cc##iter1 = predux_real<ResScalar, ResPacket>(&c##iter1); \
2323 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2326#define GEMV_INIT_ROW(iter, N) \
2327 EIGEN_IF_CONSTEXPR (N > iter) { \
2328 c##iter = pset1<ResPacket>(ResScalar(0)); \
2330 EIGEN_UNUSED_VARIABLE(c##iter); \
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); \
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); \
2342 EIGEN_UNUSED_VARIABLE(cc##iter1); \
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); \
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]); \
2359#define GEMV_PROCESS_ROW(N) \
2360 for (; i < n##N; i += N) { \
2361 GEMV_UNROLL_ROW(GEMV_INIT_ROW, N) \
2363 for (; j + LhsPacketSize <= cols; j += LhsPacketSize) { \
2364 RhsPacket a0 = rhs2.template load<RhsPacket, Unaligned>(j); \
2365 GEMV_UNROLL_ROW(GEMV_WORK_ROW, N) \
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)) \
2372 GEMV_UNROLL_ROW_HALF(GEMV_STORE_ROW, (N >> 1)) \
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;
2380 typedef typename Traits::LhsPacket LhsPacket;
2381 typedef typename Traits::RhsPacket RhsPacket;
2382 typedef typename Traits::ResPacket ResPacket;
2386 LhsMapper lhs(alhs);
2387 typename RhsMapper::LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
2389 eigen_internal_assert(rhs.stride() == 1);
2390 conj_helper<LhsScalar, RhsScalar, false, false> cj;
2391 conj_helper<LhsPacket, RhsPacket, false, false> pcj;
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;
2405 ResPacketSize = Traits::ResPacketSize,
2406 LhsPacketSize = Traits::LhsPacketSize,
2407 RhsPacketSize = Traits::RhsPacketSize,
2412 __vector_quad c0, c1, c2, c3, c4, c5, c6, c7;
2413 GEMV_UNUSED_ROW(8, c)
2415 ResPacket c0, c1, c2, c3, c4, c5, c6, c7;
2417#ifndef GCC_ONE_VECTORPAIR_BUG
2418 ScalarBlock<ResScalar, 2> cc0, cc1, cc2, cc3;
2423 for (; i < rows; ++i) {
2424 ResPacket d0 = pset1<ResPacket>(ResScalar(0));
2426 for (; j + LhsPacketSize <= cols; j += LhsPacketSize) {
2427 RhsPacket b0 = rhs2.template load<RhsPacket, Unaligned>(j);
2429 d0 = pcj.pmadd(lhs.template load<LhsPacket, LhsAlignment>(i + 0, j), b0, d0);
2431 ResScalar dd0 = predux(d0);
2432 for (; j < cols; ++j) {
2433 dd0 += cj.pmul(lhs(i, j), rhs2(j));
2435 res[i * resIncr] += alpha * dd0;
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; \
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); \
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; \
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); \
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)
2473#define gemv_bf16_col gemvMMA_bfloat16_col
2474#define gemv_bf16_row gemvMMA_bfloat16_row
2476#define gemv_bf16_col gemv_bfloat16_col
2477#define gemv_bf16_row gemv_bfloat16_row
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, \
2487 if (numext::is_exactly_zero(alpha)) return; \
2488 gemv_bf16_col<LhsMapper, RhsMapper>(rows, cols, lhs, rhs, res, resIncr, alpha); \
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, \
2499 if (numext::is_exactly_zero(alpha)) return; \
2500 gemv_bf16_row<LhsMapper, RhsMapper>(rows, cols, lhs, rhs, res, resIncr, alpha); \
2504EIGEN_POWER_GEMV_REAL_SPECIALIZE_COL_BFLOAT16()
2505EIGEN_POWER_GEMV_REAL_SPECIALIZE_ROW_BFLOAT16()
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) {
2514 return predux_complex<ResScalar, PResPacket>(a0, b0);
2517#define GEMV_LOADPACKET_ROW_COMPLEX(iter) loadLhsPacket<Scalar, LhsScalar, LhsMapper, PLhsPacket>(lhs, i + (iter), j)
2519#define GEMV_LOADPACKET_ROW_COMPLEX_DATA(iter) convertReal(GEMV_LOADPACKET_ROW_COMPLEX(iter))
2521#define GEMV_PROCESS_ROW_COMPLEX_SINGLE_WORK(which, N) \
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) \
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)) \
2534 GEMV_UNROLL_ROW_HALF(GEMV_STORE_ROW_COMPLEX, (N >> 1))
2537#define GEMV_INIT_ROW_COMPLEX_MMA(iter, N) \
2538 EIGEN_IF_CONSTEXPR (GEMV_GETN_COMPLEX(N) > iter) { \
2539 __builtin_mma_xxsetaccz(&e0##iter); \
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));
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); \
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); \
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); \
2566 predux_complex<ResScalar, ScalarPacket, LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs>(&e0##iter1); \
2569 EIGEN_UNUSED_VARIABLE(cc##iter1); \
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)
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); \
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); \
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); \
2596 EIGEN_UNUSED_VARIABLE(cc##iter1); \
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); \
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]); \
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)
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); \
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); \
2628 dd0 = predux(c0##iter);
2631#define GEMV_PROCESS_ROW_COMPLEX_SINGLE(N) GEMV_PROCESS_ROW_COMPLEX_SINGLE_NEW(N)
2633#define GEMV_PROCESS_ROW_COMPLEX_ONE(N) GEMV_PROCESS_ROW_COMPLEX_ONE_NEW(N)
2635#define GEMV_PROCESS_ROW_COMPLEX_PREDUX(iter) GEMV_PROCESS_ROW_COMPLEX_PREDUX_NEW(iter)
2640#define GEMV_LOADPACKET_ROW_COMPLEX_OLD(iter) lhs.template load<LhsPacket, LhsAlignment>(i + (iter), j)
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>(); \
2647 EIGEN_UNUSED_VARIABLE(c1##iter); \
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); \
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); \
2661 EIGEN_UNUSED_VARIABLE(cc##iter1); \
2664#define GEMV_PROCESS_ROW_COMPLEX_SINGLE_OLD(N) \
2665 GEMV_UNROLL_ROW(GEMV_INIT_COMPLEX_OLD, N) \
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) \
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) \
2679#define GEMV_PROCESS_ROW_COMPLEX_PREDUX_OLD(iter) dd0 = predux(c1##iter);
2682#define GEMV_PROCESS_ROW_COMPLEX_IS_NEW 1
2684#define GEMV_PROCESS_ROW_COMPLEX_IS_NEW (sizeof(Scalar) == sizeof(float)) || GEMV_IS_COMPLEX_COMPLEX
2687#define GEMV_PROCESS_ROW_COMPLEX_SINGLE(N) \
2688 if (GEMV_PROCESS_ROW_COMPLEX_IS_NEW) { \
2689 GEMV_PROCESS_ROW_COMPLEX_SINGLE_NEW(N) \
2691 GEMV_PROCESS_ROW_COMPLEX_SINGLE_OLD(N) \
2694#define GEMV_PROCESS_ROW_COMPLEX_ONE(N) \
2695 if (GEMV_PROCESS_ROW_COMPLEX_IS_NEW) { \
2696 GEMV_PROCESS_ROW_COMPLEX_ONE_NEW(N) \
2698 GEMV_PROCESS_ROW_COMPLEX_ONE_OLD(N) \
2701#define GEMV_PROCESS_ROW_COMPLEX_PREDUX(iter) \
2702 if (GEMV_PROCESS_ROW_COMPLEX_IS_NEW) { \
2703 GEMV_PROCESS_ROW_COMPLEX_PREDUX_NEW(iter) \
2705 GEMV_PROCESS_ROW_COMPLEX_PREDUX_OLD(iter) \
2710#define GEMV_PROCESS_ROW_COMPLEX(N) GEMV_PROCESS_ROW_COMPLEX_ONE_MMA(N)
2712#define GEMV_PROCESS_ROW_COMPLEX(N) GEMV_PROCESS_ROW_COMPLEX_ONE(N)
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;
2721 typedef typename Traits::LhsPacket LhsPacket;
2722 typedef typename Traits::RhsPacket RhsPacket;
2723 typedef typename Traits::ResPacket ResPacket;
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;
2732 LhsMapper lhs(alhs);
2733 typename RhsMapper::LinearMapper rhs2 = rhs.getLinearMapper(0, 0);
2735 eigen_internal_assert(rhs.stride() == 1);
2736 conj_helper<LhsScalar, RhsScalar, ConjugateLhs, ConjugateRhs> cj;
2738 conj_helper<LhsPacket, RhsPacket, ConjugateLhs, ConjugateRhs> pcj;
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;
2753 ResPacketSize = PTraits::ResPacketSize,
2754 LhsPacketSize = PTraits::LhsPacketSize,
2755 RhsPacketSize = PTraits::RhsPacketSize,
2759 PResPacket c00, c01, c02, c03, c04, c05, c06, c07;
2760 ResPacket c10, c11, c12, c13, c14, c15, c16, c17;
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)
2768#ifndef GCC_ONE_VECTORPAIR_BUG
2769 ScalarBlock<ResScalar, 2> cc0, cc1, cc2, cc3;
2771 EIGEN_IF_CONSTEXPR (!GEMV_IS_COMPLEX_COMPLEX)
2774 GEMV_PROCESS_ROW_COMPLEX(8)
2776 GEMV_PROCESS_ROW_COMPLEX(4)
2777 GEMV_PROCESS_ROW_COMPLEX(2)
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));
2785 res[i * resIncr] += alpha * dd0;
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; \
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); \
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; \
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); \
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>)
@ Unaligned
Definition Constants.h:236
@ ColMajor
Definition Constants.h:319