34#ifndef EIGEN_GENERAL_MATRIX_MATRIX_BLAS_H
35#define EIGEN_GENERAL_MATRIX_MATRIX_BLAS_H
38#include "../InternalHeaderCheck.h"
53#define EIGEN_BLAS_GEMM_SPECIALIZATION(EIGTYPE, EIGPREFIX, BLASTYPE, BLASFUNC) \
54 template <typename Index, int LhsStorageOrder, bool ConjugateLhs, int RhsStorageOrder, bool ConjugateRhs> \
55 struct general_matrix_matrix_product<Index, EIGTYPE, LhsStorageOrder, ConjugateLhs, EIGTYPE, RhsStorageOrder, \
56 ConjugateRhs, ColMajor, 1> { \
57 typedef gebp_traits<EIGTYPE, EIGTYPE> Traits; \
59 static void run(Index rows, Index cols, Index depth, const EIGTYPE* lhs_, Index lhsStride, const EIGTYPE* rhs_, \
60 Index rhsStride, EIGTYPE* res, Index resIncr, Index resStride, EIGTYPE alpha, \
61 level3_blocking<EIGTYPE, EIGTYPE>& , GemmParallelInfo<Index>* ) { \
63 if (rows == 0 || cols == 0 || depth == 0) return; \
64 EIGEN_ONLY_USED_FOR_DEBUG(resIncr); \
65 eigen_assert(resIncr == 1); \
66 char transa, transb; \
67 BlasIndex m, n, k, lda, ldb, ldc; \
68 const EIGTYPE *a, *b; \
70 MatrixX##EIGPREFIX a_tmp, b_tmp; \
73 transa = (LhsStorageOrder == RowMajor) ? ((ConjugateLhs) ? 'C' : 'T') : 'N'; \
74 transb = (RhsStorageOrder == RowMajor) ? ((ConjugateRhs) ? 'C' : 'T') : 'N'; \
77 m = convert_index<BlasIndex>(rows); \
78 n = convert_index<BlasIndex>(cols); \
79 k = convert_index<BlasIndex>(depth); \
82 lda = convert_index<BlasIndex>(lhsStride); \
83 ldb = convert_index<BlasIndex>(rhsStride); \
84 ldc = convert_index<BlasIndex>(resStride); \
87 EIGEN_IF_CONSTEXPR ((LhsStorageOrder == ColMajor) && (ConjugateLhs)) { \
88 Map<const MatrixX##EIGPREFIX, 0, OuterStride<> > lhs(lhs_, m, k, OuterStride<>(lhsStride)); \
89 a_tmp = lhs.conjugate(); \
91 lda = convert_index<BlasIndex>(a_tmp.outerStride()); \
95 EIGEN_IF_CONSTEXPR ((RhsStorageOrder == ColMajor) && (ConjugateRhs)) { \
96 Map<const MatrixX##EIGPREFIX, 0, OuterStride<> > rhs(rhs_, k, n, OuterStride<>(rhsStride)); \
97 b_tmp = rhs.conjugate(); \
99 ldb = convert_index<BlasIndex>(b_tmp.outerStride()); \
103 BLASFUNC(&transa, &transb, &m, &n, &k, (const BLASTYPE*)&numext::real_ref(alpha), (const BLASTYPE*)a, &lda, \
104 (const BLASTYPE*)b, &ldb, (const BLASTYPE*)&numext::real_ref(beta), (BLASTYPE*)res, &ldc); \
109EIGEN_BLAS_GEMM_SPECIALIZATION(
double, d,
double, dgemm)
110EIGEN_BLAS_GEMM_SPECIALIZATION(
float, f,
float, sgemm)
111EIGEN_BLAS_GEMM_SPECIALIZATION(dcomplex, cd, MKL_Complex16, zgemm)
112EIGEN_BLAS_GEMM_SPECIALIZATION(scomplex, cf, MKL_Complex8, cgemm)
114EIGEN_BLAS_GEMM_SPECIALIZATION(
double, d,
double, EIGEN_BLAS_SYM(dgemm))
115EIGEN_BLAS_GEMM_SPECIALIZATION(
float, f,
float, EIGEN_BLAS_SYM(sgemm))
116EIGEN_BLAS_GEMM_SPECIALIZATION(dcomplex, cd,
double, EIGEN_BLAS_SYM(zgemm))
117EIGEN_BLAS_GEMM_SPECIALIZATION(scomplex, cf,
float, EIGEN_BLAS_SYM(cgemm))
122#if EIGEN_USE_OPENBLAS_BFLOAT16
126void EIGEN_BLAS_SYM(sbgemm)(
const char* trans_a,
const char* trans_b,
const EIGEN_BLAS_INT* M,
const EIGEN_BLAS_INT* N,
127 const EIGEN_BLAS_INT* K,
const float* alpha,
const Eigen::bfloat16* A,
128 const EIGEN_BLAS_INT* lda,
const Eigen::bfloat16* B,
const EIGEN_BLAS_INT* ldb,
129 const float* beta,
float* C,
const EIGEN_BLAS_INT* ldc);
132template <
typename Index,
int LhsStorageOrder,
bool ConjugateLhs,
int RhsStorageOrder,
bool ConjugateRhs>
133struct general_matrix_matrix_product<Index, Eigen::bfloat16, LhsStorageOrder, ConjugateLhs, Eigen::bfloat16,
134 RhsStorageOrder, ConjugateRhs,
ColMajor, 1> {
135 typedef gebp_traits<Eigen::bfloat16, Eigen::bfloat16> Traits;
137 static void run(Index rows, Index cols, Index depth,
const Eigen::bfloat16* lhs_, Index lhsStride,
138 const Eigen::bfloat16* rhs_, Index rhsStride, Eigen::bfloat16* res, Index resIncr, Index resStride,
139 Eigen::bfloat16 alpha, level3_blocking<Eigen::bfloat16, Eigen::bfloat16>& ,
140 GemmParallelInfo<Index>* ) {
142 if (rows == 0 || cols == 0 || depth == 0)
return;
143 EIGEN_ONLY_USED_FOR_DEBUG(resIncr);
144 eigen_assert(resIncr == 1);
146 BlasIndex m, n, k, lda, ldb, ldc;
147 const Eigen::bfloat16 *a, *b;
149 float falpha =
static_cast<float>(alpha);
150 float fbeta = float(1.0);
152 using MatrixXbf16 = Matrix<Eigen::bfloat16, Dynamic, Dynamic>;
153 MatrixXbf16 a_tmp, b_tmp;
157 transa = (LhsStorageOrder ==
RowMajor) ? ((ConjugateLhs) ?
'C' :
'T') :
'N';
158 transb = (RhsStorageOrder ==
RowMajor) ? ((ConjugateRhs) ?
'C' :
'T') :
'N';
161 m = convert_index<BlasIndex>(rows);
162 n = convert_index<BlasIndex>(cols);
163 k = convert_index<BlasIndex>(depth);
166 lda = convert_index<BlasIndex>(lhsStride);
167 ldb = convert_index<BlasIndex>(rhsStride);
168 ldc = convert_index<BlasIndex>(m);
171 EIGEN_IF_CONSTEXPR ((LhsStorageOrder ==
ColMajor) && (ConjugateLhs)) {
172 Map<const MatrixXbf16, 0, OuterStride<> > lhs(lhs_, m, k, OuterStride<>(lhsStride));
173 a_tmp = lhs.conjugate();
175 lda = convert_index<BlasIndex>(a_tmp.outerStride());
180 EIGEN_IF_CONSTEXPR ((RhsStorageOrder ==
ColMajor) && (ConjugateRhs)) {
181 Map<const MatrixXbf16, 0, OuterStride<> > rhs(rhs_, k, n, OuterStride<>(rhsStride));
182 b_tmp = rhs.conjugate();
184 ldb = convert_index<BlasIndex>(b_tmp.outerStride());
192 EIGEN_BLAS_SYM(sbgemm)
193 (&transa, &transb, &m, &n, &k, (
const float*)&numext::real_ref(falpha), a, &lda, b, &ldb,
194 (
const float*)&numext::real_ref(fbeta), r_tmp.data(), &ldc);
197 Map<MatrixXbf16, 0, OuterStride<> > result(res, m, n, OuterStride<>(resStride));
198 result = r_tmp.cast<Eigen::bfloat16>();
204#undef EIGEN_BLAS_GEMM_SPECIALIZATION
@ ColMajor
Definition Constants.h:319
@ RowMajor
Definition Constants.h:321
Matrix< float, Dynamic, Dynamic > MatrixXf
Dynamic×Dynamic matrix of type float.
Definition Matrix.h:488