5#ifndef EIGEN_MATRIX_PRODUCT_COMMON_ALTIVEC_H
6#define EIGEN_MATRIX_PRODUCT_COMMON_ALTIVEC_H
8#ifdef EIGEN_POWER_USE_PREFETCH
9#define EIGEN_POWER_PREFETCH(p) prefetch(p)
11#define EIGEN_POWER_PREFETCH(p)
15#include "../../InternalHeaderCheck.h"
21template <
typename Scalar,
typename Packet,
typename DataMapper, const Index accRows, const Index accCols>
22EIGEN_ALWAYS_INLINE
void gemm_extra_row(
const DataMapper& res,
const Scalar* lhs_base,
const Scalar* rhs_base,
23 Index depth, Index strideA, Index offsetA, Index strideB, Index row, Index rows,
24 Index remaining_rows,
const Packet& pAlpha,
const Packet& pMask);
26template <
typename Scalar,
typename Packet,
typename DataMapper, const Index accCols>
27EIGEN_ALWAYS_INLINE
void gemm_extra_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB, Index depth,
28 Index strideA, Index offsetA, Index strideB, Index offsetB, Index col,
29 Index rows, Index cols, Index remaining_rows,
const Packet& pAlpha,
32template <
typename Packet>
33EIGEN_ALWAYS_INLINE Packet bmask(
const Index remaining_rows);
35template <
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
const Index accRows,
36 const Index accCols,
bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
37EIGEN_ALWAYS_INLINE
void gemm_complex_extra_row(
const DataMapper& res,
const Scalar* lhs_base,
const Scalar* rhs_base,
38 Index depth, Index strideA, Index offsetA, Index strideB, Index row,
39 Index rows, Index remaining_rows,
const Packet& pAlphaReal,
40 const Packet& pAlphaImag,
const Packet& pMask);
42template <
typename Scalar,
typename Packet,
typename Packetc,
typename DataMapper,
const Index accCols,
43 bool ConjugateLhs,
bool ConjugateRhs,
bool LhsIsReal,
bool RhsIsReal>
44EIGEN_ALWAYS_INLINE
void gemm_complex_extra_cols(
const DataMapper& res,
const Scalar* blockA,
const Scalar* blockB,
45 Index depth, Index strideA, Index offsetA, Index strideB,
46 Index offsetB, Index col, Index rows, Index cols, Index remaining_rows,
47 const Packet& pAlphaReal,
const Packet& pAlphaImag,
50template <
typename DataMapper>
51EIGEN_ALWAYS_INLINE
void convertArrayBF16toF32(
float* result, Index cols, Index rows,
const DataMapper& src);
53template <const Index size,
bool non_unit_str
ide, Index delta>
54EIGEN_ALWAYS_INLINE
void storeBF16fromResult(bfloat16* dst,
const Packet8bf& data, Index resInc, Index extra = 0);
56template <
bool non_unit_str
ide = false>
57EIGEN_ALWAYS_INLINE
void convertArrayPointerBF16toF32(
float* result, Index cols, Index rows, bfloat16* src,
60template <
bool rhsExtraCols,
bool lhsExtraRows>
61EIGEN_ALWAYS_INLINE
void storeResults(Packet4f (&acc)[4], Index rows,
const Packet4f& pAlpha,
float* result,
62 Index extra_cols, Index extra_rows);
64template <Index num_acc,
bool extraRows, Index size = 4>
65EIGEN_ALWAYS_INLINE
void outputVecColResults(Packet4f (&acc)[num_acc][size],
float* result,
const Packet4f& pAlpha,
68template <Index num_acc, Index size = 4>
69EIGEN_ALWAYS_INLINE
void outputVecResults(Packet4f (&acc)[num_acc][size],
float* result,
const Packet4f& pAlpha);
71template <
typename RhsMapper,
bool linear>
72EIGEN_ALWAYS_INLINE Packet8bf loadColData(RhsMapper& rhs, Index j);
74template <
typename Packet>
75EIGEN_ALWAYS_INLINE Packet ploadLhs(
const __UNPACK_TYPE__(Packet) * lhs);
77template <
typename DataMapper,
typename Packet,
const Index accCols,
int StorageOrder,
bool Complex,
int N,
79EIGEN_ALWAYS_INLINE
void bload(PacketBlock<Packet, N*(Complex ? 2 : 1)>& acc,
const DataMapper& res, Index row,
82template <
typename DataMapper,
typename Packet,
int N>
83EIGEN_ALWAYS_INLINE
void bstore(PacketBlock<Packet, N>& acc,
const DataMapper& res, Index row);
85template <
typename DataMapper,
typename Packet, const Index accCols,
bool Complex, Index N,
bool full = true>
86EIGEN_ALWAYS_INLINE
void bload_partial(PacketBlock<Packet, N*(Complex ? 2 : 1)>& acc,
const DataMapper& res, Index row,
89template <
typename DataMapper,
typename Packet, Index N>
90EIGEN_ALWAYS_INLINE
void bstore_partial(PacketBlock<Packet, N>& acc,
const DataMapper& res, Index row, Index elements);
92template <
typename Packet,
int N>
93EIGEN_ALWAYS_INLINE
void bscale(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ,
const Packet& pAlpha);
95template <
typename Packet,
int N,
bool mask>
96EIGEN_ALWAYS_INLINE
void bscale(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ,
const Packet& pAlpha,
99template <
typename Packet,
int N,
bool mask>
100EIGEN_ALWAYS_INLINE
void bscalec(PacketBlock<Packet, N>& aReal, PacketBlock<Packet, N>& aImag,
const Packet& bReal,
101 const Packet& bImag, PacketBlock<Packet, N>& cReal, PacketBlock<Packet, N>& cImag,
102 const Packet& pMask);
104template <
typename Packet,
typename Packetc,
int N,
bool full>
105EIGEN_ALWAYS_INLINE
void bcouple(PacketBlock<Packet, N>& taccReal, PacketBlock<Packet, N>& taccImag,
106 PacketBlock<Packetc, N * 2>& tRes, PacketBlock<Packetc, N>& acc1,
107 PacketBlock<Packetc, N>& acc2);
109#define MICRO_NORMAL(iter) (accCols == accCols2) || (unroll_factor != (iter + 1))
111#define MICRO_UNROLL_ITER1(func, N) \
112 switch (remaining_rows) { \
118 if (sizeof(Scalar) == sizeof(float)) { \
123 if (sizeof(Scalar) == sizeof(float)) { \
129#define MICRO_UNROLL_ITER(func, N) \
130 if (remaining_rows) { \
136#define MICRO_NORMAL_PARTIAL(iter) full || (unroll_factor != (iter + 1))
138#define MICRO_COMPLEX_UNROLL_ITER(func, N) MICRO_UNROLL_ITER1(func, N)
140#define MICRO_NORMAL_COLS(iter, a, b) ((MICRO_NORMAL(iter)) ? a : b)
142#define MICRO_LOAD1(lhs_ptr, iter) \
143 if (unroll_factor > iter) { \
144 lhsV##iter = ploadLhs<Packet>(lhs_ptr##iter); \
145 lhs_ptr##iter += MICRO_NORMAL_COLS(iter, accCols, accCols2); \
147 EIGEN_UNUSED_VARIABLE(lhsV##iter); \
150#define MICRO_LOAD_ONE(iter) MICRO_LOAD1(lhs_ptr, iter)
152#define MICRO_COMPLEX_LOAD_ONE(iter) \
153 if (!LhsIsReal && (unroll_factor > iter)) { \
154 lhsVi##iter = ploadLhs<Packet>(lhs_ptr_real##iter + MICRO_NORMAL_COLS(iter, imag_delta, imag_delta2)); \
156 EIGEN_UNUSED_VARIABLE(lhsVi##iter); \
158 MICRO_LOAD1(lhs_ptr_real, iter)
160#define MICRO_LOAD1_PARTIAL(lhs_ptr, iter) \
161 if (unroll_factor > iter) { \
162 if (MICRO_NORMAL(iter)) { \
163 lhsV##iter = ploadLhs<Packet>(lhs_ptr##iter); \
164 lhs_ptr##iter += accCols; \
166 lhsV##iter = ploadu_partial<Packet>(lhs_ptr##iter, accCols2); \
167 lhs_ptr##iter += accCols2; \
170 EIGEN_UNUSED_VARIABLE(lhsV##iter); \
173#define MICRO_LOAD_PARTIAL_ONE(iter) MICRO_LOAD1_PARTIAL(lhs_ptr, iter)
175#define MICRO_COMPLEX_LOAD_PARTIAL_ONE(iter) \
176 if (!LhsIsReal && (unroll_factor > iter)) { \
177 if (MICRO_NORMAL(iter)) { \
178 lhsVi##iter = ploadLhs<Packet>(lhs_ptr_real##iter + imag_delta); \
180 lhsVi##iter = ploadu_partial<Packet>(lhs_ptr_real##iter + imag_delta2, accCols2); \
183 EIGEN_UNUSED_VARIABLE(lhsVi##iter); \
185 MICRO_LOAD1_PARTIAL(lhs_ptr_real, iter)
187#define MICRO_SRC_PTR1(lhs_ptr, advRows, iter) \
188 if (unroll_factor > iter) { \
189 lhs_ptr##iter = lhs_base + (row + (iter * accCols)) * strideA * advRows - \
190 MICRO_NORMAL_COLS(iter, 0, (accCols - accCols2) * offsetA); \
192 EIGEN_UNUSED_VARIABLE(lhs_ptr##iter); \
195#define MICRO_SRC_PTR_ONE(iter) MICRO_SRC_PTR1(lhs_ptr, 1, iter)
197#define MICRO_COMPLEX_SRC_PTR_ONE(iter) MICRO_SRC_PTR1(lhs_ptr_real, advanceRows, iter)
199#define MICRO_PREFETCH1(lhs_ptr, iter) \
200 if (unroll_factor > iter) { \
201 EIGEN_POWER_PREFETCH(lhs_ptr##iter); \
204#define MICRO_PREFETCH_ONE(iter) MICRO_PREFETCH1(lhs_ptr, iter)
206#define MICRO_COMPLEX_PREFETCH_ONE(iter) MICRO_PREFETCH1(lhs_ptr_real, iter)
208#define MICRO_UPDATE \
209 if (accCols == accCols2) { \
210 EIGEN_UNUSED_VARIABLE(offsetA); \
211 row += unroll_factor * accCols; \
214#define MICRO_COMPLEX_UPDATE \
216 if (LhsIsReal || (accCols == accCols2)) { \
217 EIGEN_UNUSED_VARIABLE(imag_delta2); \