Eigen  5.0.1
 
Loading...
Searching...
No Matches
TriangularMatrixMatrix_BLAS.h
1/*
2 Copyright (c) 2011, Intel Corporation. All rights reserved.
3
4 Redistribution and use in source and binary forms, with or without modification,
5 are permitted provided that the following conditions are met:
6
7 * Redistributions of source code must retain the above copyright notice, this
8 list of conditions and the following disclaimer.
9 * Redistributions in binary form must reproduce the above copyright notice,
10 this list of conditions and the following disclaimer in the documentation
11 and/or other materials provided with the distribution.
12 * Neither the name of Intel Corporation nor the names of its contributors may
13 be used to endorse or promote products derived from this software without
14 specific prior written permission.
15
16 THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND
17 ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED
18 WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
19 DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT OWNER OR CONTRIBUTORS BE LIABLE FOR
20 ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES
21 (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES;
22 LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON
23 ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT
24 (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS
25 SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
26
27 ********************************************************************************
28 * Content : Eigen bindings to BLAS F77
29 * Triangular matrix * matrix product functionality based on ?TRMM.
30 ********************************************************************************
31*/
32// SPDX-License-Identifier: BSD-3-Clause
33
34#ifndef EIGEN_TRIANGULAR_MATRIX_MATRIX_BLAS_H
35#define EIGEN_TRIANGULAR_MATRIX_MATRIX_BLAS_H
36
37// IWYU pragma: private
38#include "../InternalHeaderCheck.h"
39
40namespace Eigen {
41
42namespace internal {
43
44template <typename Scalar, typename Index, int Mode, bool LhsIsTriangular, int LhsStorageOrder, bool ConjugateLhs,
45 int RhsStorageOrder, bool ConjugateRhs, int ResStorageOrder>
46struct product_triangular_matrix_matrix_trmm
47 : product_triangular_matrix_matrix<Scalar, Index, Mode, LhsIsTriangular, LhsStorageOrder, ConjugateLhs,
48 RhsStorageOrder, ConjugateRhs, ResStorageOrder, 1, BuiltIn> {};
49
50// try to go to BLAS specialization
51#define EIGEN_BLAS_TRMM_SPECIALIZE(Scalar, LhsIsTriangular) \
52 template <typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder, bool ConjugateRhs> \
53 struct product_triangular_matrix_matrix<Scalar, Index, Mode, LhsIsTriangular, LhsStorageOrder, ConjugateLhs, \
54 RhsStorageOrder, ConjugateRhs, ColMajor, 1, Specialized> { \
55 static inline void run(Index _rows, Index _cols, Index _depth, const Scalar* _lhs, Index lhsStride, \
56 const Scalar* _rhs, Index rhsStride, Scalar* res, Index resIncr, Index resStride, \
57 Scalar alpha, level3_blocking<Scalar, Scalar>& blocking) { \
58 EIGEN_ONLY_USED_FOR_DEBUG(resIncr); \
59 eigen_assert(resIncr == 1); \
60 product_triangular_matrix_matrix_trmm<Scalar, Index, Mode, LhsIsTriangular, LhsStorageOrder, ConjugateLhs, \
61 RhsStorageOrder, ConjugateRhs, ColMajor>::run(_rows, _cols, _depth, _lhs, \
62 lhsStride, _rhs, rhsStride, \
63 res, resStride, alpha, \
64 blocking); \
65 } \
66 };
67
68EIGEN_BLAS_TRMM_SPECIALIZE(double, true)
69EIGEN_BLAS_TRMM_SPECIALIZE(double, false)
70EIGEN_BLAS_TRMM_SPECIALIZE(dcomplex, true)
71EIGEN_BLAS_TRMM_SPECIALIZE(dcomplex, false)
72EIGEN_BLAS_TRMM_SPECIALIZE(float, true)
73EIGEN_BLAS_TRMM_SPECIALIZE(float, false)
74EIGEN_BLAS_TRMM_SPECIALIZE(scomplex, true)
75EIGEN_BLAS_TRMM_SPECIALIZE(scomplex, false)
76
77// implements col-major += alpha * op(triangular) * op(general)
78#define EIGEN_BLAS_TRMM_L(EIGTYPE, BLASTYPE, EIGPREFIX, BLASFUNC) \
79 template <typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder, bool ConjugateRhs> \
80 struct product_triangular_matrix_matrix_trmm<EIGTYPE, Index, Mode, true, LhsStorageOrder, ConjugateLhs, \
81 RhsStorageOrder, ConjugateRhs, ColMajor> { \
82 enum { \
83 IsLower = (Mode & Lower) == Lower, \
84 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1, \
85 IsUnitDiag = (Mode & UnitDiag) ? 1 : 0, \
86 IsZeroDiag = (Mode & ZeroDiag) ? 1 : 0, \
87 LowUp = IsLower ? Lower : Upper, \
88 conjA = ((LhsStorageOrder == ColMajor) && ConjugateLhs) ? 1 : 0 \
89 }; \
90 \
91 static void run(Index _rows, Index _cols, Index _depth, const EIGTYPE* _lhs, Index lhsStride, const EIGTYPE* _rhs, \
92 Index rhsStride, EIGTYPE* res, Index resStride, EIGTYPE alpha, \
93 level3_blocking<EIGTYPE, EIGTYPE>& blocking) { \
94 if (_rows == 0 || _cols == 0 || _depth == 0) return; \
95 Index diagSize = (std::min)(_rows, _depth); \
96 Index rows = IsLower ? _rows : diagSize; \
97 Index depth = IsLower ? diagSize : _depth; \
98 Index cols = _cols; \
99 \
100 typedef Matrix<EIGTYPE, Dynamic, Dynamic, LhsStorageOrder> MatrixLhs; \
101 typedef Matrix<EIGTYPE, Dynamic, Dynamic, RhsStorageOrder> MatrixRhs; \
102 \
103 /* Non-square case - doesn't fit to BLAS ?TRMM. Fall to default triangular product or call BLAS ?GEMM*/ \
104 if (rows != depth) { \
105 /* FIXME handle mkl_domain_get_max_threads */ \
106 /*int nthr = mkl_domain_get_max_threads(EIGEN_BLAS_DOMAIN_BLAS);*/ int nthr = 1; \
107 \
108 if (((nthr == 1) && (((std::max)(rows, depth) - diagSize) / (double)diagSize < 0.5))) { \
109 /* Most likely no benefit to call TRMM or GEMM from BLAS */ \
110 product_triangular_matrix_matrix<EIGTYPE, Index, Mode, true, LhsStorageOrder, ConjugateLhs, RhsStorageOrder, \
111 ConjugateRhs, ColMajor, 1, BuiltIn>::run(_rows, _cols, _depth, _lhs, \
112 lhsStride, _rhs, rhsStride, res, \
113 1, resStride, alpha, blocking); \
114 /*std::cout << "TRMM_L: A is not square! Go to Eigen TRMM implementation!\n";*/ \
115 } else { \
116 /* Make sense to call GEMM */ \
117 Map<const MatrixLhs, 0, OuterStride<> > lhsMap(_lhs, rows, depth, OuterStride<>(lhsStride)); \
118 MatrixLhs aa_tmp = lhsMap.template triangularView<Mode>(); \
119 BlasIndex aStride = convert_index<BlasIndex>(aa_tmp.outerStride()); \
120 gemm_blocking_space<ColMajor, EIGTYPE, EIGTYPE, Dynamic, Dynamic, Dynamic> gemm_blocking(_rows, _cols, \
121 _depth, 1, true); \
122 general_matrix_matrix_product<Index, EIGTYPE, LhsStorageOrder, ConjugateLhs, EIGTYPE, RhsStorageOrder, \
123 ConjugateRhs, ColMajor, 1>::run(rows, cols, depth, aa_tmp.data(), aStride, \
124 _rhs, rhsStride, res, 1, resStride, alpha, \
125 gemm_blocking, 0); \
126 \
127 /*std::cout << "TRMM_L: A is not square! Go to BLAS GEMM implementation! " << nthr<<" \n";*/ \
128 } \
129 return; \
130 } \
131 char side = 'L', transa, uplo, diag = 'N'; \
132 EIGTYPE* b; \
133 const EIGTYPE* a; \
134 BlasIndex m, n, lda, ldb; \
135 \
136 /* Set m, n */ \
137 m = convert_index<BlasIndex>(diagSize); \
138 n = convert_index<BlasIndex>(cols); \
139 \
140 /* Set trans */ \
141 transa = (LhsStorageOrder == RowMajor) ? ((ConjugateLhs) ? 'C' : 'T') : 'N'; \
142 \
143 /* Set b, ldb */ \
144 Map<const MatrixRhs, 0, OuterStride<> > rhs(_rhs, depth, cols, OuterStride<>(rhsStride)); \
145 MatrixX##EIGPREFIX b_tmp; \
146 \
147 EIGEN_IF_CONSTEXPR (ConjugateRhs) \
148 b_tmp = rhs.conjugate(); \
149 else \
150 b_tmp = rhs; \
151 b = b_tmp.data(); \
152 ldb = convert_index<BlasIndex>(b_tmp.outerStride()); \
153 \
154 /* Set uplo */ \
155 uplo = IsLower ? 'L' : 'U'; \
156 EIGEN_IF_CONSTEXPR (LhsStorageOrder == RowMajor) { \
157 uplo = (uplo == 'L') ? 'U' : 'L'; \
158 } \
159 /* Set a, lda */ \
160 Map<const MatrixLhs, 0, OuterStride<> > lhs(_lhs, rows, depth, OuterStride<>(lhsStride)); \
161 MatrixLhs a_tmp; \
162 \
163 EIGEN_IF_CONSTEXPR ((conjA != 0) || (SetDiag == 0)) { \
164 EIGEN_IF_CONSTEXPR (conjA) \
165 a_tmp = lhs.conjugate(); \
166 else \
167 a_tmp = lhs; \
168 EIGEN_IF_CONSTEXPR (IsZeroDiag) \
169 a_tmp.diagonal().setZero(); \
170 else EIGEN_IF_CONSTEXPR (IsUnitDiag) \
171 a_tmp.diagonal().setOnes(); \
172 a = a_tmp.data(); \
173 lda = convert_index<BlasIndex>(a_tmp.outerStride()); \
174 } else { \
175 a = _lhs; \
176 lda = convert_index<BlasIndex>(lhsStride); \
177 } \
178 /*std::cout << "TRMM_L: A is square! Go to BLAS TRMM implementation! \n";*/ \
179 /* call ?trmm*/ \
180 BLASFUNC(&side, &uplo, &transa, &diag, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, \
181 &lda, (BLASTYPE*)b, &ldb); \
182 \
183 /* Add op(a_triangular)*b into res*/ \
184 Map<MatrixX##EIGPREFIX, 0, OuterStride<> > res_tmp(res, rows, cols, OuterStride<>(resStride)); \
185 res_tmp = res_tmp + b_tmp; \
186 } \
187 };
188
189#ifdef EIGEN_USE_MKL
190EIGEN_BLAS_TRMM_L(double, double, d, dtrmm)
191EIGEN_BLAS_TRMM_L(dcomplex, MKL_Complex16, cd, ztrmm)
192EIGEN_BLAS_TRMM_L(float, float, f, strmm)
193EIGEN_BLAS_TRMM_L(scomplex, MKL_Complex8, cf, ctrmm)
194#else
195EIGEN_BLAS_TRMM_L(double, double, d, EIGEN_BLAS_SYM(dtrmm))
196EIGEN_BLAS_TRMM_L(dcomplex, double, cd, EIGEN_BLAS_SYM(ztrmm))
197EIGEN_BLAS_TRMM_L(float, float, f, EIGEN_BLAS_SYM(strmm))
198EIGEN_BLAS_TRMM_L(scomplex, float, cf, EIGEN_BLAS_SYM(ctrmm))
199#endif
200
201// implements col-major += alpha * op(general) * op(triangular)
202#define EIGEN_BLAS_TRMM_R(EIGTYPE, BLASTYPE, EIGPREFIX, BLASFUNC) \
203 template <typename Index, int Mode, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder, bool ConjugateRhs> \
204 struct product_triangular_matrix_matrix_trmm<EIGTYPE, Index, Mode, false, LhsStorageOrder, ConjugateLhs, \
205 RhsStorageOrder, ConjugateRhs, ColMajor> { \
206 enum { \
207 IsLower = (Mode & Lower) == Lower, \
208 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1, \
209 IsUnitDiag = (Mode & UnitDiag) ? 1 : 0, \
210 IsZeroDiag = (Mode & ZeroDiag) ? 1 : 0, \
211 LowUp = IsLower ? Lower : Upper, \
212 conjA = ((RhsStorageOrder == ColMajor) && ConjugateRhs) ? 1 : 0 \
213 }; \
214 \
215 static void run(Index _rows, Index _cols, Index _depth, const EIGTYPE* _lhs, Index lhsStride, const EIGTYPE* _rhs, \
216 Index rhsStride, EIGTYPE* res, Index resStride, EIGTYPE alpha, \
217 level3_blocking<EIGTYPE, EIGTYPE>& blocking) { \
218 if (_rows == 0 || _cols == 0 || _depth == 0) return; \
219 Index diagSize = (std::min)(_cols, _depth); \
220 Index rows = _rows; \
221 Index depth = IsLower ? _depth : diagSize; \
222 Index cols = IsLower ? diagSize : _cols; \
223 \
224 typedef Matrix<EIGTYPE, Dynamic, Dynamic, LhsStorageOrder> MatrixLhs; \
225 typedef Matrix<EIGTYPE, Dynamic, Dynamic, RhsStorageOrder> MatrixRhs; \
226 \
227 /* Non-square case - doesn't fit to BLAS ?TRMM. Fall to default triangular product or call BLAS ?GEMM*/ \
228 if (cols != depth) { \
229 int nthr = 1 /*mkl_domain_get_max_threads(EIGEN_BLAS_DOMAIN_BLAS)*/; \
230 \
231 if ((nthr == 1) && (((std::max)(cols, depth) - diagSize) / (double)diagSize < 0.5)) { \
232 /* Most likely no benefit to call TRMM or GEMM from BLAS*/ \
233 product_triangular_matrix_matrix<EIGTYPE, Index, Mode, false, LhsStorageOrder, ConjugateLhs, \
234 RhsStorageOrder, ConjugateRhs, ColMajor, 1, BuiltIn>::run(_rows, _cols, \
235 _depth, _lhs, \
236 lhsStride, _rhs, \
237 rhsStride, res, \
238 1, resStride, \
239 alpha, blocking); \
240 /*std::cout << "TRMM_R: A is not square! Go to Eigen TRMM implementation!\n";*/ \
241 } else { \
242 /* Make sense to call GEMM */ \
243 Map<const MatrixRhs, 0, OuterStride<> > rhsMap(_rhs, depth, cols, OuterStride<>(rhsStride)); \
244 MatrixRhs aa_tmp = rhsMap.template triangularView<Mode>(); \
245 BlasIndex aStride = convert_index<BlasIndex>(aa_tmp.outerStride()); \
246 gemm_blocking_space<ColMajor, EIGTYPE, EIGTYPE, Dynamic, Dynamic, Dynamic> gemm_blocking(_rows, _cols, \
247 _depth, 1, true); \
248 general_matrix_matrix_product<Index, EIGTYPE, LhsStorageOrder, ConjugateLhs, EIGTYPE, RhsStorageOrder, \
249 ConjugateRhs, ColMajor, 1>::run(rows, cols, depth, _lhs, lhsStride, \
250 aa_tmp.data(), aStride, res, 1, resStride, \
251 alpha, gemm_blocking, 0); \
252 \
253 /*std::cout << "TRMM_R: A is not square! Go to BLAS GEMM implementation! " << nthr<<" \n";*/ \
254 } \
255 return; \
256 } \
257 char side = 'R', transa, uplo, diag = 'N'; \
258 EIGTYPE* b; \
259 const EIGTYPE* a; \
260 BlasIndex m, n, lda, ldb; \
261 \
262 /* Set m, n */ \
263 m = convert_index<BlasIndex>(rows); \
264 n = convert_index<BlasIndex>(diagSize); \
265 \
266 /* Set trans */ \
267 transa = (RhsStorageOrder == RowMajor) ? ((ConjugateRhs) ? 'C' : 'T') : 'N'; \
268 \
269 /* Set b, ldb */ \
270 Map<const MatrixLhs, 0, OuterStride<> > lhs(_lhs, rows, depth, OuterStride<>(lhsStride)); \
271 MatrixX##EIGPREFIX b_tmp; \
272 \
273 EIGEN_IF_CONSTEXPR (ConjugateLhs) \
274 b_tmp = lhs.conjugate(); \
275 else \
276 b_tmp = lhs; \
277 b = b_tmp.data(); \
278 ldb = convert_index<BlasIndex>(b_tmp.outerStride()); \
279 \
280 /* Set uplo */ \
281 uplo = IsLower ? 'L' : 'U'; \
282 EIGEN_IF_CONSTEXPR (RhsStorageOrder == RowMajor) { \
283 uplo = (uplo == 'L') ? 'U' : 'L'; \
284 } \
285 /* Set a, lda */ \
286 Map<const MatrixRhs, 0, OuterStride<> > rhs(_rhs, depth, cols, OuterStride<>(rhsStride)); \
287 MatrixRhs a_tmp; \
288 \
289 EIGEN_IF_CONSTEXPR ((conjA != 0) || (SetDiag == 0)) { \
290 EIGEN_IF_CONSTEXPR (conjA) \
291 a_tmp = rhs.conjugate(); \
292 else \
293 a_tmp = rhs; \
294 EIGEN_IF_CONSTEXPR (IsZeroDiag) \
295 a_tmp.diagonal().setZero(); \
296 else EIGEN_IF_CONSTEXPR (IsUnitDiag) \
297 a_tmp.diagonal().setOnes(); \
298 a = a_tmp.data(); \
299 lda = convert_index<BlasIndex>(a_tmp.outerStride()); \
300 } else { \
301 a = _rhs; \
302 lda = convert_index<BlasIndex>(rhsStride); \
303 } \
304 /*std::cout << "TRMM_R: A is square! Go to BLAS TRMM implementation! \n";*/ \
305 /* call ?trmm*/ \
306 BLASFUNC(&side, &uplo, &transa, &diag, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, \
307 &lda, (BLASTYPE*)b, &ldb); \
308 \
309 /* Add op(a_triangular)*b into res*/ \
310 Map<MatrixX##EIGPREFIX, 0, OuterStride<> > res_tmp(res, rows, cols, OuterStride<>(resStride)); \
311 res_tmp = res_tmp + b_tmp; \
312 } \
313 };
314
315#ifdef EIGEN_USE_MKL
316EIGEN_BLAS_TRMM_R(double, double, d, dtrmm)
317EIGEN_BLAS_TRMM_R(dcomplex, MKL_Complex16, cd, ztrmm)
318EIGEN_BLAS_TRMM_R(float, float, f, strmm)
319EIGEN_BLAS_TRMM_R(scomplex, MKL_Complex8, cf, ctrmm)
320#else
321EIGEN_BLAS_TRMM_R(double, double, d, EIGEN_BLAS_SYM(dtrmm))
322EIGEN_BLAS_TRMM_R(dcomplex, double, cd, EIGEN_BLAS_SYM(ztrmm))
323EIGEN_BLAS_TRMM_R(float, float, f, EIGEN_BLAS_SYM(strmm))
324EIGEN_BLAS_TRMM_R(scomplex, float, cf, EIGEN_BLAS_SYM(ctrmm))
325#endif
326
327#undef EIGEN_BLAS_TRMM_SPECIALIZE
328#undef EIGEN_BLAS_TRMM_L
329#undef EIGEN_BLAS_TRMM_R
330} // end namespace internal
331
332} // end namespace Eigen
333
334#endif // EIGEN_TRIANGULAR_MATRIX_MATRIX_BLAS_H