34#ifndef EIGEN_GENERAL_MATRIX_VECTOR_BLAS_H
35#define EIGEN_GENERAL_MATRIX_VECTOR_BLAS_H
38#include "../InternalHeaderCheck.h"
53template <
typename Index,
typename LhsScalar,
int StorageOrder,
bool ConjugateLhs,
typename RhsScalar,
55struct general_matrix_vector_product_gemv;
57#define EIGEN_BLAS_GEMV_SPECIALIZE(Scalar) \
58 template <typename Index, bool ConjugateLhs, bool ConjugateRhs> \
59 struct general_matrix_vector_product<Index, Scalar, const_blas_data_mapper<Scalar, Index, ColMajor>, ColMajor, \
60 ConjugateLhs, Scalar, const_blas_data_mapper<Scalar, Index, RowMajor>, \
61 ConjugateRhs, Specialized> { \
62 static void run(Index rows, Index cols, const const_blas_data_mapper<Scalar, Index, ColMajor>& lhs, \
63 const const_blas_data_mapper<Scalar, Index, RowMajor>& rhs, Scalar* res, Index resIncr, \
65 EIGEN_IF_CONSTEXPR (ConjugateLhs) { \
66 general_matrix_vector_product<Index, Scalar, const_blas_data_mapper<Scalar, Index, ColMajor>, ColMajor, \
67 ConjugateLhs, Scalar, const_blas_data_mapper<Scalar, Index, RowMajor>, \
68 ConjugateRhs, BuiltIn>::run(rows, cols, lhs, rhs, res, resIncr, alpha); \
70 general_matrix_vector_product_gemv<Index, Scalar, ColMajor, ConjugateLhs, Scalar, ConjugateRhs>::run( \
71 rows, cols, lhs.data(), lhs.stride(), rhs.data(), rhs.stride(), res, resIncr, alpha); \
75 template <typename Index, bool ConjugateLhs, bool ConjugateRhs> \
76 struct general_matrix_vector_product<Index, Scalar, const_blas_data_mapper<Scalar, Index, RowMajor>, RowMajor, \
77 ConjugateLhs, Scalar, const_blas_data_mapper<Scalar, Index, ColMajor>, \
78 ConjugateRhs, Specialized> { \
79 static void run(Index rows, Index cols, const const_blas_data_mapper<Scalar, Index, RowMajor>& lhs, \
80 const const_blas_data_mapper<Scalar, Index, ColMajor>& rhs, Scalar* res, Index resIncr, \
82 general_matrix_vector_product_gemv<Index, Scalar, RowMajor, ConjugateLhs, Scalar, ConjugateRhs>::run( \
83 rows, cols, lhs.data(), lhs.stride(), rhs.data(), rhs.stride(), res, resIncr, alpha); \
87EIGEN_BLAS_GEMV_SPECIALIZE(
double)
88EIGEN_BLAS_GEMV_SPECIALIZE(
float)
89EIGEN_BLAS_GEMV_SPECIALIZE(dcomplex)
90EIGEN_BLAS_GEMV_SPECIALIZE(scomplex)
92#define EIGEN_BLAS_GEMV_SPECIALIZATION(EIGTYPE, BLASTYPE, BLASFUNC) \
93 template <typename Index, int LhsStorageOrder, bool ConjugateLhs, bool ConjugateRhs> \
94 struct general_matrix_vector_product_gemv<Index, EIGTYPE, LhsStorageOrder, ConjugateLhs, EIGTYPE, ConjugateRhs> { \
95 typedef Matrix<EIGTYPE, Dynamic, 1, ColMajor> GEMVVector; \
97 static void run(Index rows, Index cols, const EIGTYPE* lhs, Index lhsStride, const EIGTYPE* rhs, Index rhsIncr, \
98 EIGTYPE* res, Index resIncr, EIGTYPE alpha) { \
99 if (rows == 0 || cols == 0) return; \
100 BlasIndex m = convert_index<BlasIndex>(rows), n = convert_index<BlasIndex>(cols), \
101 lda = convert_index<BlasIndex>(lhsStride), incx = convert_index<BlasIndex>(rhsIncr), \
102 incy = convert_index<BlasIndex>(resIncr); \
103 const EIGTYPE beta(1); \
104 const EIGTYPE* x_ptr; \
105 char trans = (LhsStorageOrder == ColMajor) ? 'N' : (ConjugateLhs) ? 'C' : 'T'; \
106 EIGEN_IF_CONSTEXPR (LhsStorageOrder == RowMajor) { \
107 m = convert_index<BlasIndex>(cols); \
108 n = convert_index<BlasIndex>(rows); \
111 EIGEN_IF_CONSTEXPR (ConjugateRhs) { \
112 Map<const GEMVVector, 0, InnerStride<> > map_x(rhs, cols, 1, InnerStride<>(incx)); \
113 x_tmp = map_x.conjugate(); \
114 x_ptr = x_tmp.data(); \
119 BLASFUNC(&trans, &m, &n, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)lhs, &lda, \
120 (const BLASTYPE*)x_ptr, &incx, (const BLASTYPE*)&numext::real_ref(beta), (BLASTYPE*)res, &incy); \
125EIGEN_BLAS_GEMV_SPECIALIZATION(
double,
double, dgemv)
126EIGEN_BLAS_GEMV_SPECIALIZATION(
float,
float, sgemv)
127EIGEN_BLAS_GEMV_SPECIALIZATION(dcomplex, MKL_Complex16, zgemv)
128EIGEN_BLAS_GEMV_SPECIALIZATION(scomplex, MKL_Complex8, cgemv)
130EIGEN_BLAS_GEMV_SPECIALIZATION(
double,
double, EIGEN_BLAS_SYM(dgemv))
131EIGEN_BLAS_GEMV_SPECIALIZATION(
float,
float, EIGEN_BLAS_SYM(sgemv))
132EIGEN_BLAS_GEMV_SPECIALIZATION(dcomplex,
double, EIGEN_BLAS_SYM(zgemv))
133EIGEN_BLAS_GEMV_SPECIALIZATION(scomplex,
float, EIGEN_BLAS_SYM(cgemv))
136#undef EIGEN_BLAS_GEMV_SPECIALIZE
137#undef EIGEN_BLAS_GEMV_SPECIALIZATION