Eigen  5.0.1
 
Loading...
Searching...
No Matches
MatrixProductCommon.h
1// #define EIGEN_POWER_USE_PREFETCH // Use prefetching in gemm routines
2// SPDX-FileCopyrightText: The Eigen Authors
3// SPDX-License-Identifier: MPL-2.0
4
5#ifndef EIGEN_MATRIX_PRODUCT_COMMON_ALTIVEC_H
6#define EIGEN_MATRIX_PRODUCT_COMMON_ALTIVEC_H
7
8#ifdef EIGEN_POWER_USE_PREFETCH
9#define EIGEN_POWER_PREFETCH(p) prefetch(p)
10#else
11#define EIGEN_POWER_PREFETCH(p)
12#endif
13
14// IWYU pragma: private
15#include "../../InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20
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);
25
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,
30 const Packet& pMask);
31
32template <typename Packet>
33EIGEN_ALWAYS_INLINE Packet bmask(const Index remaining_rows);
34
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);
41
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,
48 const Packet& pMask);
49
50template <typename DataMapper>
51EIGEN_ALWAYS_INLINE void convertArrayBF16toF32(float* result, Index cols, Index rows, const DataMapper& src);
52
53template <const Index size, bool non_unit_stride, Index delta>
54EIGEN_ALWAYS_INLINE void storeBF16fromResult(bfloat16* dst, const Packet8bf& data, Index resInc, Index extra = 0);
55
56template <bool non_unit_stride = false>
57EIGEN_ALWAYS_INLINE void convertArrayPointerBF16toF32(float* result, Index cols, Index rows, bfloat16* src,
58 Index resInc = 1);
59
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);
63
64template <Index num_acc, bool extraRows, Index size = 4>
65EIGEN_ALWAYS_INLINE void outputVecColResults(Packet4f (&acc)[num_acc][size], float* result, const Packet4f& pAlpha,
66 Index extra_rows);
67
68template <Index num_acc, Index size = 4>
69EIGEN_ALWAYS_INLINE void outputVecResults(Packet4f (&acc)[num_acc][size], float* result, const Packet4f& pAlpha);
70
71template <typename RhsMapper, bool linear>
72EIGEN_ALWAYS_INLINE Packet8bf loadColData(RhsMapper& rhs, Index j);
73
74template <typename Packet>
75EIGEN_ALWAYS_INLINE Packet ploadLhs(const __UNPACK_TYPE__(Packet) * lhs);
76
77template <typename DataMapper, typename Packet, const Index accCols, int StorageOrder, bool Complex, int N,
78 bool full = true>
79EIGEN_ALWAYS_INLINE void bload(PacketBlock<Packet, N*(Complex ? 2 : 1)>& acc, const DataMapper& res, Index row,
80 Index col);
81
82template <typename DataMapper, typename Packet, int N>
83EIGEN_ALWAYS_INLINE void bstore(PacketBlock<Packet, N>& acc, const DataMapper& res, Index row);
84
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,
87 Index elements);
88
89template <typename DataMapper, typename Packet, Index N>
90EIGEN_ALWAYS_INLINE void bstore_partial(PacketBlock<Packet, N>& acc, const DataMapper& res, Index row, Index elements);
91
92template <typename Packet, int N>
93EIGEN_ALWAYS_INLINE void bscale(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ, const Packet& pAlpha);
94
95template <typename Packet, int N, bool mask>
96EIGEN_ALWAYS_INLINE void bscale(PacketBlock<Packet, N>& acc, PacketBlock<Packet, N>& accZ, const Packet& pAlpha,
97 const Packet& pMask);
98
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);
103
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);
108
109#define MICRO_NORMAL(iter) (accCols == accCols2) || (unroll_factor != (iter + 1))
110
111#define MICRO_UNROLL_ITER1(func, N) \
112 switch (remaining_rows) { \
113 default: \
114 func(N, 0) break; \
115 case 1: \
116 func(N, 1) break; \
117 case 2: \
118 if (sizeof(Scalar) == sizeof(float)) { \
119 func(N, 2) \
120 } \
121 break; \
122 case 3: \
123 if (sizeof(Scalar) == sizeof(float)) { \
124 func(N, 3) \
125 } \
126 break; \
127 }
128
129#define MICRO_UNROLL_ITER(func, N) \
130 if (remaining_rows) { \
131 func(N, true); \
132 } else { \
133 func(N, false); \
134 }
135
136#define MICRO_NORMAL_PARTIAL(iter) full || (unroll_factor != (iter + 1))
137
138#define MICRO_COMPLEX_UNROLL_ITER(func, N) MICRO_UNROLL_ITER1(func, N)
139
140#define MICRO_NORMAL_COLS(iter, a, b) ((MICRO_NORMAL(iter)) ? a : b)
141
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); \
146 } else { \
147 EIGEN_UNUSED_VARIABLE(lhsV##iter); \
148 }
149
150#define MICRO_LOAD_ONE(iter) MICRO_LOAD1(lhs_ptr, iter)
151
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)); \
155 } else { \
156 EIGEN_UNUSED_VARIABLE(lhsVi##iter); \
157 } \
158 MICRO_LOAD1(lhs_ptr_real, iter)
159
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; \
165 } else { \
166 lhsV##iter = ploadu_partial<Packet>(lhs_ptr##iter, accCols2); \
167 lhs_ptr##iter += accCols2; \
168 } \
169 } else { \
170 EIGEN_UNUSED_VARIABLE(lhsV##iter); \
171 }
172
173#define MICRO_LOAD_PARTIAL_ONE(iter) MICRO_LOAD1_PARTIAL(lhs_ptr, iter)
174
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); \
179 } else { \
180 lhsVi##iter = ploadu_partial<Packet>(lhs_ptr_real##iter + imag_delta2, accCols2); \
181 } \
182 } else { \
183 EIGEN_UNUSED_VARIABLE(lhsVi##iter); \
184 } \
185 MICRO_LOAD1_PARTIAL(lhs_ptr_real, iter)
186
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); \
191 } else { \
192 EIGEN_UNUSED_VARIABLE(lhs_ptr##iter); \
193 }
194
195#define MICRO_SRC_PTR_ONE(iter) MICRO_SRC_PTR1(lhs_ptr, 1, iter)
196
197#define MICRO_COMPLEX_SRC_PTR_ONE(iter) MICRO_SRC_PTR1(lhs_ptr_real, advanceRows, iter)
198
199#define MICRO_PREFETCH1(lhs_ptr, iter) \
200 if (unroll_factor > iter) { \
201 EIGEN_POWER_PREFETCH(lhs_ptr##iter); \
202 }
203
204#define MICRO_PREFETCH_ONE(iter) MICRO_PREFETCH1(lhs_ptr, iter)
205
206#define MICRO_COMPLEX_PREFETCH_ONE(iter) MICRO_PREFETCH1(lhs_ptr_real, iter)
207
208#define MICRO_UPDATE \
209 if (accCols == accCols2) { \
210 EIGEN_UNUSED_VARIABLE(offsetA); \
211 row += unroll_factor * accCols; \
212 }
213
214#define MICRO_COMPLEX_UPDATE \
215 MICRO_UPDATE \
216 if (LhsIsReal || (accCols == accCols2)) { \
217 EIGEN_UNUSED_VARIABLE(imag_delta2); \
218 }
219
220} // end namespace internal
221} // end namespace Eigen
222
223#endif // EIGEN_MATRIX_PRODUCT_COMMON_ALTIVEC_H