| 1006 | |
| 1007 | template <typename Scalar, typename Spec> |
| 1008 | void EvalOpenBlas(const Matrix<Scalar>& lhs, const Matrix<Scalar>& rhs, |
| 1009 | const Spec& spec, int max_num_threads, Matrix<Scalar>* dst) { |
| 1010 | RUY_CHECK_EQ(lhs.zero_point, 0); |
| 1011 | RUY_CHECK_EQ(rhs.zero_point, 0); |
| 1012 | RUY_CHECK_EQ(dst->zero_point, 0); |
| 1013 | RUY_CHECK_EQ(spec.multiplier_fixedpoint, 0); |
| 1014 | RUY_CHECK_EQ(spec.multiplier_exponent, 0); |
| 1015 | |
| 1016 | Matrix<Scalar> gemm_lhs; |
| 1017 | Matrix<Scalar> gemm_rhs; |
| 1018 | Matrix<Scalar> gemm_dst; |
| 1019 | gemm_dst = *dst; |
| 1020 | |
| 1021 | // Use Transpose to reduce to the all-column-major case. |
| 1022 | // Notice that ruy::Matrix merely holds a pointer, does not own data, |
| 1023 | // so Transpose is cheap -- no actual matrix data is being transposed here. |
| 1024 | if (dst->layout.order == Order::kColMajor) { |
| 1025 | gemm_lhs = lhs; |
| 1026 | gemm_rhs = rhs; |
| 1027 | } else { |
| 1028 | gemm_lhs = rhs; |
| 1029 | gemm_rhs = lhs; |
| 1030 | Transpose(&gemm_lhs); |
| 1031 | Transpose(&gemm_rhs); |
| 1032 | Transpose(&gemm_dst); |
| 1033 | } |
| 1034 | bool transposed_lhs = false; |
| 1035 | bool transposed_rhs = false; |
| 1036 | |
| 1037 | if (gemm_lhs.layout.order == Order::kRowMajor) { |
| 1038 | Transpose(&gemm_lhs); |
| 1039 | transposed_lhs = true; |
| 1040 | } |
| 1041 | if (gemm_rhs.layout.order == Order::kRowMajor) { |
| 1042 | Transpose(&gemm_rhs); |
| 1043 | transposed_rhs = true; |
| 1044 | } |
| 1045 | |
| 1046 | RUY_CHECK(gemm_lhs.layout.order == Order::kColMajor); |
| 1047 | RUY_CHECK(gemm_rhs.layout.order == Order::kColMajor); |
| 1048 | RUY_CHECK(gemm_dst.layout.order == Order::kColMajor); |
| 1049 | |
| 1050 | char transa = transposed_lhs ? 'T' : 'N'; |
| 1051 | char transb = transposed_rhs ? 'T' : 'N'; |
| 1052 | int m = gemm_lhs.layout.rows; |
| 1053 | int n = gemm_rhs.layout.cols; |
| 1054 | int k = gemm_lhs.layout.cols; |
| 1055 | float alpha = 1; |
| 1056 | Scalar* a = gemm_lhs.data.get(); |
| 1057 | int lda = gemm_lhs.layout.stride; |
| 1058 | Scalar* b = gemm_rhs.data.get(); |
| 1059 | int ldb = gemm_rhs.layout.stride; |
| 1060 | float beta = 0; |
| 1061 | Scalar* c = gemm_dst.data.get(); |
| 1062 | int ldc = gemm_dst.layout.stride; |
| 1063 | GenericBlasGemm<Scalar>::Run(&transa, &transb, &m, &n, &k, &alpha, a, &lda, b, |
| 1064 | &ldb, &beta, c, &ldc); |
| 1065 | |