diff --git a/src/NativeWrappers/ACML/blas.c b/src/NativeWrappers/ACML/blas.c index 8086bfbf..7ed9b151 100644 --- a/src/NativeWrappers/ACML/blas.c +++ b/src/NativeWrappers/ACML/blas.c @@ -2,6 +2,7 @@ #include "wrapper_common.h" enum CBLAS_TRANSPOSE {CblasNoTrans=111, CblasTrans=112, CblasConjTrans=113, CblasConjNoTrans=114}; +char getTransChar(TRANSPOSE); DLLEXPORT void s_axpy(const int n, const float alpha, float x[], float y[]){ saxpy(n, alpha, x, 1, y, 1); @@ -54,27 +55,44 @@ DLLEXPORT doublecomplex z_dot_product(int n, doublecomplex x[], doublecomplex y[ DLLEXPORT void s_matrix_multiply(const enum TRANSPOSE transA, const enum TRANSPOSE transB, const int m, const int n, const int k, float alpha, float x[], float y[], float beta, float c[]){ int lda = transA == CblasNoTrans ? m : k; int ldb = transB == CblasNoTrans ? k : n; - - sgemm(transA, transB, m, n, k, alpha, x, lda, y, ldb, beta, c, m); + char transAchar = getTransChar(transA); + char transBchar = getTransChar(transB); + sgemm(transAchar, transBchar, m, n, k, alpha, x, lda, y, ldb, beta, c, m); } DLLEXPORT void d_matrix_multiply(const enum TRANSPOSE transA, const enum TRANSPOSE transB, const int m, const int n, const int k, double alpha, double x[], double y[], double beta, double c[]){ int lda = transA == CblasNoTrans ? m : k; int ldb = transB == CblasNoTrans ? k : n; - - dgemm(transA, transB, m, n, k, alpha, x, lda, y, ldb, beta, c, m); + char transAchar = getTransChar(transA); + char transBchar = getTransChar(transB); + dgemm(transAchar, transBchar, m, n, k, alpha, x, lda, y, ldb, beta, c, m); } DLLEXPORT void c_matrix_multiply(const enum TRANSPOSE transA, const enum TRANSPOSE transB, const int m, const int n, const int k, complex alpha, complex x[], complex y[], complex beta, complex c[]){ int lda = transA == CblasNoTrans ? m : k; int ldb = transB == CblasNoTrans ? k : n; - - cgemm(transA, transB, m, n, k, &alpha, x, lda, y, ldb, &beta, c, m); + char transAchar = getTransChar(transA); + char transBchar = getTransChar(transB); + cgemm(transAchar, transBchar, m, n, k, &alpha, x, lda, y, ldb, &beta, c, m); } DLLEXPORT void z_matrix_multiply(const enum TRANSPOSE transA, const enum TRANSPOSE transB, const int m, const int n, const int k, doublecomplex alpha, doublecomplex x[], doublecomplex y[], doublecomplex beta, doublecomplex c[]){ int lda = transA == CblasNoTrans ? m : k; int ldb = transB == CblasNoTrans ? k : n; - - zgemm(transA, transB, m, n, k, &alpha, x, lda, y, ldb, &beta, c, m); + char transAchar = getTransChar(transA); + char transBchar = getTransChar(transB); + zgemm(transAchar, transBchar, m, n, k, &alpha, x, lda, y, ldb, &beta, c, m); +} + +char getTransChar(enum TRANSPOSE trans){ + char cTrans; + switch( trans ){ + case CblasNoTrans : cTrans = 'N'; + break; + case CblasTrans : cTrans = 'T'; + break; + case CblasConjTrans : cTrans = 'C'; + break; + } + return cTrans; }