MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / EvalOpenBlas

Function EvalOpenBlas

tensorflow/lite/experimental/ruy/test.h:1008–1098  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1006
1007template <typename Scalar, typename Spec>
1008void 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

Callers 1

RunMethod · 0.85

Calls 4

infinityFunction · 0.85
TransposeFunction · 0.70
RunFunction · 0.50
getMethod · 0.45

Tested by

no test coverage detected