34#ifndef EIGEN_TRIANGULAR_MATRIX_VECTOR_BLAS_H
35#define EIGEN_TRIANGULAR_MATRIX_VECTOR_BLAS_H
38#include "../InternalHeaderCheck.h"
50template <
typename Index,
int Mode,
typename LhsScalar,
bool ConjLhs,
typename RhsScalar,
bool ConjRhs,
52struct triangular_matrix_vector_product_trmv
53 : triangular_matrix_vector_product<Index, Mode, LhsScalar, ConjLhs, RhsScalar, ConjRhs, StorageOrder, BuiltIn> {};
55#define EIGEN_BLAS_TRMV_SPECIALIZE(Scalar) \
56 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
57 struct triangular_matrix_vector_product<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, ColMajor, Specialized> { \
58 static void run(Index rows_, Index cols_, const Scalar* lhs_, Index lhsStride, const Scalar* rhs_, Index rhsIncr, \
59 Scalar* res_, Index resIncr, Scalar alpha) { \
60 triangular_matrix_vector_product_trmv<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, ColMajor>::run( \
61 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
64 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
65 struct triangular_matrix_vector_product<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, RowMajor, Specialized> { \
66 static void run(Index rows_, Index cols_, const Scalar* lhs_, Index lhsStride, const Scalar* rhs_, Index rhsIncr, \
67 Scalar* res_, Index resIncr, Scalar alpha) { \
68 triangular_matrix_vector_product_trmv<Index, Mode, Scalar, ConjLhs, Scalar, ConjRhs, RowMajor>::run( \
69 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
73EIGEN_BLAS_TRMV_SPECIALIZE(
double)
74EIGEN_BLAS_TRMV_SPECIALIZE(
float)
75EIGEN_BLAS_TRMV_SPECIALIZE(dcomplex)
76EIGEN_BLAS_TRMV_SPECIALIZE(scomplex)
79#define EIGEN_BLAS_TRMV_CM(EIGTYPE, BLASTYPE, EIGPREFIX, BLASPREFIX, BLASPOSTFIX) \
80 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
81 struct triangular_matrix_vector_product_trmv<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, 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 \
89 static void run(Index rows_, Index cols_, const EIGTYPE* lhs_, Index lhsStride, const EIGTYPE* rhs_, \
90 Index rhsIncr, EIGTYPE* res_, Index resIncr, EIGTYPE alpha) { \
91 if (rows_ == 0 || cols_ == 0) return; \
92 EIGEN_IF_CONSTEXPR (ConjLhs || IsZeroDiag) { \
93 triangular_matrix_vector_product<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, ColMajor, BuiltIn>::run( \
94 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
97 Index size = (std::min)(rows_, cols_); \
98 Index rows = IsLower ? rows_ : size; \
99 Index cols = IsLower ? size : cols_; \
101 typedef VectorX##EIGPREFIX VectorRhs; \
105 Map<const VectorRhs, 0, InnerStride<> > rhs(rhs_, cols, InnerStride<>(rhsIncr)); \
107 EIGEN_IF_CONSTEXPR (ConjRhs) \
108 x_tmp = rhs.conjugate(); \
115 char trans, uplo, diag; \
116 BlasIndex m, n, lda, incx, incy; \
121 n = convert_index<BlasIndex>(size); \
122 lda = convert_index<BlasIndex>(lhsStride); \
124 incy = convert_index<BlasIndex>(resIncr); \
128 uplo = IsLower ? 'L' : 'U'; \
129 diag = IsUnitDiag ? 'U' : 'N'; \
132 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, trmv), BLASPOSTFIX) \
133 (&uplo, &trans, &diag, &n, (const BLASTYPE*)lhs_, &lda, (BLASTYPE*)x, &incx); \
136 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, axpy), BLASPOSTFIX) \
137 (&n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)x, &incx, (BLASTYPE*)res_, &incy); \
139 if (size < (std::max)(rows, cols)) { \
140 EIGEN_IF_CONSTEXPR (ConjRhs) \
141 x_tmp = rhs.conjugate(); \
146 y = res_ + size * resIncr; \
148 m = convert_index<BlasIndex>(rows - size); \
149 n = convert_index<BlasIndex>(size); \
153 a = lhs_ + size * lda; \
154 m = convert_index<BlasIndex>(size); \
155 n = convert_index<BlasIndex>(cols - size); \
157 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, gemv), BLASPOSTFIX) \
158 (&trans, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, &lda, (const BLASTYPE*)x, \
159 &incx, (const BLASTYPE*)&numext::real_ref(beta), (BLASTYPE*)y, &incy); \
165EIGEN_BLAS_TRMV_CM(
double,
double, d, d, )
166EIGEN_BLAS_TRMV_CM(dcomplex, MKL_Complex16, cd, z, )
167EIGEN_BLAS_TRMV_CM(
float,
float, f, s, )
168EIGEN_BLAS_TRMV_CM(scomplex, MKL_Complex8, cf, c, )
170EIGEN_BLAS_TRMV_CM(
double,
double, d, d, EIGEN_BLAS_POSTFIX)
171EIGEN_BLAS_TRMV_CM(dcomplex,
double, cd, z, EIGEN_BLAS_POSTFIX)
172EIGEN_BLAS_TRMV_CM(
float,
float, f, s, EIGEN_BLAS_POSTFIX)
173EIGEN_BLAS_TRMV_CM(scomplex,
float, cf, c, EIGEN_BLAS_POSTFIX)
177#define EIGEN_BLAS_TRMV_RM(EIGTYPE, BLASTYPE, EIGPREFIX, BLASPREFIX, BLASPOSTFIX) \
178 template <typename Index, int Mode, bool ConjLhs, bool ConjRhs> \
179 struct triangular_matrix_vector_product_trmv<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, RowMajor> { \
181 IsLower = (Mode & Lower) == Lower, \
182 SetDiag = (Mode & (ZeroDiag | UnitDiag)) ? 0 : 1, \
183 IsUnitDiag = (Mode & UnitDiag) ? 1 : 0, \
184 IsZeroDiag = (Mode & ZeroDiag) ? 1 : 0, \
185 LowUp = IsLower ? Lower : Upper \
187 static void run(Index rows_, Index cols_, const EIGTYPE* lhs_, Index lhsStride, const EIGTYPE* rhs_, \
188 Index rhsIncr, EIGTYPE* res_, Index resIncr, EIGTYPE alpha) { \
189 if (rows_ == 0 || cols_ == 0) return; \
190 EIGEN_IF_CONSTEXPR (IsZeroDiag) { \
191 triangular_matrix_vector_product<Index, Mode, EIGTYPE, ConjLhs, EIGTYPE, ConjRhs, RowMajor, BuiltIn>::run( \
192 rows_, cols_, lhs_, lhsStride, rhs_, rhsIncr, res_, resIncr, alpha); \
195 Index size = (std::min)(rows_, cols_); \
196 Index rows = IsLower ? rows_ : size; \
197 Index cols = IsLower ? size : cols_; \
199 typedef VectorX##EIGPREFIX VectorRhs; \
203 Map<const VectorRhs, 0, InnerStride<> > rhs(rhs_, cols, InnerStride<>(rhsIncr)); \
205 EIGEN_IF_CONSTEXPR (ConjRhs) \
206 x_tmp = rhs.conjugate(); \
213 char trans, uplo, diag; \
214 BlasIndex m, n, lda, incx, incy; \
219 n = convert_index<BlasIndex>(size); \
220 lda = convert_index<BlasIndex>(lhsStride); \
222 incy = convert_index<BlasIndex>(resIncr); \
225 trans = ConjLhs ? 'C' : 'T'; \
226 uplo = IsLower ? 'U' : 'L'; \
227 diag = IsUnitDiag ? 'U' : 'N'; \
230 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, trmv), BLASPOSTFIX) \
231 (&uplo, &trans, &diag, &n, (const BLASTYPE*)lhs_, &lda, (BLASTYPE*)x, &incx); \
234 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, axpy), BLASPOSTFIX) \
235 (&n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)x, &incx, (BLASTYPE*)res_, &incy); \
237 if (size < (std::max)(rows, cols)) { \
238 EIGEN_IF_CONSTEXPR (ConjRhs) \
239 x_tmp = rhs.conjugate(); \
244 y = res_ + size * resIncr; \
245 a = lhs_ + size * lda; \
246 m = convert_index<BlasIndex>(rows - size); \
247 n = convert_index<BlasIndex>(size); \
252 m = convert_index<BlasIndex>(size); \
253 n = convert_index<BlasIndex>(cols - size); \
255 EIGEN_CAT(EIGEN_CAT(BLASPREFIX, gemv), BLASPOSTFIX) \
256 (&trans, &n, &m, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, &lda, (const BLASTYPE*)x, \
257 &incx, (const BLASTYPE*)&numext::real_ref(beta), (BLASTYPE*)y, &incy); \
263EIGEN_BLAS_TRMV_RM(
double,
double, d, d, )
264EIGEN_BLAS_TRMV_RM(dcomplex, MKL_Complex16, cd, z, )
265EIGEN_BLAS_TRMV_RM(
float,
float, f, s, )
266EIGEN_BLAS_TRMV_RM(scomplex, MKL_Complex8, cf, c, )
268EIGEN_BLAS_TRMV_RM(
double,
double, d, d, EIGEN_BLAS_POSTFIX)
269EIGEN_BLAS_TRMV_RM(dcomplex,
double, cd, z, EIGEN_BLAS_POSTFIX)
270EIGEN_BLAS_TRMV_RM(
float,
float, f, s, EIGEN_BLAS_POSTFIX)
271EIGEN_BLAS_TRMV_RM(scomplex,
float, cf, c, EIGEN_BLAS_POSTFIX)
274#undef EIGEN_BLAS_TRMV_RM
275#undef EIGEN_BLAS_TRMV_SPECIALIZE
276#undef EIGEN_BLAS_TRMV_CM