34#ifndef EIGEN_TRIANGULAR_MATRIX_MATRIX_BLAS_H
35#define EIGEN_TRIANGULAR_MATRIX_MATRIX_BLAS_H
38#include "../InternalHeaderCheck.h"
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> {};
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, \
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)
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> { \
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 \
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; \
100 typedef Matrix<EIGTYPE, Dynamic, Dynamic, LhsStorageOrder> MatrixLhs; \
101 typedef Matrix<EIGTYPE, Dynamic, Dynamic, RhsStorageOrder> MatrixRhs; \
104 if (rows != depth) { \
108 if (((nthr == 1) && (((std::max)(rows, depth) - diagSize) / (double)diagSize < 0.5))) { \
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); \
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, \
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, \
131 char side = 'L', transa, uplo, diag = 'N'; \
134 BlasIndex m, n, lda, ldb; \
137 m = convert_index<BlasIndex>(diagSize); \
138 n = convert_index<BlasIndex>(cols); \
141 transa = (LhsStorageOrder == RowMajor) ? ((ConjugateLhs) ? 'C' : 'T') : 'N'; \
144 Map<const MatrixRhs, 0, OuterStride<> > rhs(_rhs, depth, cols, OuterStride<>(rhsStride)); \
145 MatrixX##EIGPREFIX b_tmp; \
147 EIGEN_IF_CONSTEXPR (ConjugateRhs) \
148 b_tmp = rhs.conjugate(); \
152 ldb = convert_index<BlasIndex>(b_tmp.outerStride()); \
155 uplo = IsLower ? 'L' : 'U'; \
156 EIGEN_IF_CONSTEXPR (LhsStorageOrder == RowMajor) { \
157 uplo = (uplo == 'L') ? 'U' : 'L'; \
160 Map<const MatrixLhs, 0, OuterStride<> > lhs(_lhs, rows, depth, OuterStride<>(lhsStride)); \
163 EIGEN_IF_CONSTEXPR ((conjA != 0) || (SetDiag == 0)) { \
164 EIGEN_IF_CONSTEXPR (conjA) \
165 a_tmp = lhs.conjugate(); \
168 EIGEN_IF_CONSTEXPR (IsZeroDiag) \
169 a_tmp.diagonal().setZero(); \
170 else EIGEN_IF_CONSTEXPR (IsUnitDiag) \
171 a_tmp.diagonal().setOnes(); \
173 lda = convert_index<BlasIndex>(a_tmp.outerStride()); \
176 lda = convert_index<BlasIndex>(lhsStride); \
180 BLASFUNC(&side, &uplo, &transa, &diag, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, \
181 &lda, (BLASTYPE*)b, &ldb); \
184 Map<MatrixX##EIGPREFIX, 0, OuterStride<> > res_tmp(res, rows, cols, OuterStride<>(resStride)); \
185 res_tmp = res_tmp + b_tmp; \
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)
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))
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> { \
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 \
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; \
224 typedef Matrix<EIGTYPE, Dynamic, Dynamic, LhsStorageOrder> MatrixLhs; \
225 typedef Matrix<EIGTYPE, Dynamic, Dynamic, RhsStorageOrder> MatrixRhs; \
228 if (cols != depth) { \
231 if ((nthr == 1) && (((std::max)(cols, depth) - diagSize) / (double)diagSize < 0.5)) { \
233 product_triangular_matrix_matrix<EIGTYPE, Index, Mode, false, LhsStorageOrder, ConjugateLhs, \
234 RhsStorageOrder, ConjugateRhs, ColMajor, 1, BuiltIn>::run(_rows, _cols, \
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, \
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); \
257 char side = 'R', transa, uplo, diag = 'N'; \
260 BlasIndex m, n, lda, ldb; \
263 m = convert_index<BlasIndex>(rows); \
264 n = convert_index<BlasIndex>(diagSize); \
267 transa = (RhsStorageOrder == RowMajor) ? ((ConjugateRhs) ? 'C' : 'T') : 'N'; \
270 Map<const MatrixLhs, 0, OuterStride<> > lhs(_lhs, rows, depth, OuterStride<>(lhsStride)); \
271 MatrixX##EIGPREFIX b_tmp; \
273 EIGEN_IF_CONSTEXPR (ConjugateLhs) \
274 b_tmp = lhs.conjugate(); \
278 ldb = convert_index<BlasIndex>(b_tmp.outerStride()); \
281 uplo = IsLower ? 'L' : 'U'; \
282 EIGEN_IF_CONSTEXPR (RhsStorageOrder == RowMajor) { \
283 uplo = (uplo == 'L') ? 'U' : 'L'; \
286 Map<const MatrixRhs, 0, OuterStride<> > rhs(_rhs, depth, cols, OuterStride<>(rhsStride)); \
289 EIGEN_IF_CONSTEXPR ((conjA != 0) || (SetDiag == 0)) { \
290 EIGEN_IF_CONSTEXPR (conjA) \
291 a_tmp = rhs.conjugate(); \
294 EIGEN_IF_CONSTEXPR (IsZeroDiag) \
295 a_tmp.diagonal().setZero(); \
296 else EIGEN_IF_CONSTEXPR (IsUnitDiag) \
297 a_tmp.diagonal().setOnes(); \
299 lda = convert_index<BlasIndex>(a_tmp.outerStride()); \
302 lda = convert_index<BlasIndex>(rhsStride); \
306 BLASFUNC(&side, &uplo, &transa, &diag, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, \
307 &lda, (BLASTYPE*)b, &ldb); \
310 Map<MatrixX##EIGPREFIX, 0, OuterStride<> > res_tmp(res, rows, cols, OuterStride<>(resStride)); \
311 res_tmp = res_tmp + b_tmp; \
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)
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))
327#undef EIGEN_BLAS_TRMM_SPECIALIZE
328#undef EIGEN_BLAS_TRMM_L
329#undef EIGEN_BLAS_TRMM_R