diff --git a/hal/armpl/include/armpl_hal_core.hpp b/hal/armpl/include/armpl_hal_core.hpp index 4546a29db7..1d1e002b1e 100644 --- a/hal/armpl/include/armpl_hal_core.hpp +++ b/hal/armpl/include/armpl_hal_core.hpp @@ -108,6 +108,14 @@ int armpl_hal_gemm64fc(const double* src1, size_t src1_step, const double* src2, #undef cv_hal_gemm64fc #define cv_hal_gemm64fc armpl_hal_gemm64fc +int armpl_hal_LU32f(float* a, size_t a_step, int m, float* b, size_t b_step, int n, int* info); +int armpl_hal_LU64f(double* a, size_t a_step, int m, double* b, size_t b_step, int n, int* info); + +#undef cv_hal_LU32f +#define cv_hal_LU32f armpl_hal_LU32f +#undef cv_hal_LU64f +#define cv_hal_LU64f armpl_hal_LU64f + #endif // HAVE_ARMPL #endif // OPENCV_ARMPL_HAL_CORE_HPP diff --git a/hal/armpl/src/armpl_hal_core.cpp b/hal/armpl/src/armpl_hal_core.cpp index e4fc339645..80fc0e7e79 100644 --- a/hal/armpl/src/armpl_hal_core.cpp +++ b/hal/armpl/src/armpl_hal_core.cpp @@ -4,6 +4,7 @@ #include #include +#include #include #include #include @@ -11,6 +12,7 @@ #include #define ARMPL_GEMM_MIN_WORK_VOLUME 10000 +#define ARMPL_LU_SMALL_MATRIX_THRESH 100 namespace { @@ -172,6 +174,66 @@ int armpl_hal_gemm64fc(const double *src1, size_t src1_step, const double *src2, reinterpret_cast(src3), src3_step, beta, reinterpret_cast(dst), dst_step, m, n, k, flags); } +static inline armpl_int_t armpl_lapacke_gesv(int m, int n, float *a, int lda, armpl_int_t *ipiv, float *b, int ldb){ + return LAPACKE_sgesv(LAPACK_ROW_MAJOR, m, n, a, lda, ipiv, b, ldb); +} +static inline armpl_int_t armpl_lapacke_gesv(int m, int n, double *a, int lda, armpl_int_t *ipiv, double *b, int ldb){ + return LAPACKE_dgesv(LAPACK_ROW_MAJOR, m, n, a, lda, ipiv, b, ldb); +} +static inline armpl_int_t armpl_lapacke_getrf(int m, float *a, int lda, armpl_int_t *ipiv){ + return LAPACKE_sgetrf(LAPACK_ROW_MAJOR, m, m, a, lda, ipiv); +} +static inline armpl_int_t armpl_lapacke_getrf(int m, double *a, int lda, armpl_int_t *ipiv){ + return LAPACKE_dgetrf(LAPACK_ROW_MAJOR, m, m, a, lda, ipiv); +} + +template static inline int +armpl_lu(fptype *a, size_t a_step, int m, fptype *b, size_t b_step, int n, int *info) +{ + if(!info) + return CV_HAL_ERROR_NOT_IMPLEMENTED; + + int lda = (int)(a_step / sizeof(fptype)); + std::vector ipiv(m); + armpl_int_t linfo; + + if (b) + { + int ldb = (int)(b_step / sizeof(fptype)); + linfo = armpl_lapacke_gesv(m, n, a, lda, ipiv.data(), b, ldb); + } + else + { + linfo = armpl_lapacke_getrf(m, a, lda, ipiv.data()); + } + + if (linfo < 0) + return CV_HAL_ERROR_NOT_IMPLEMENTED; + + if (linfo == 0) + { + int sign = 0; + for (int i = 0; i < m; i++) + sign ^= (ipiv[i] != i + 1); + *info = sign ? -1 : 1; + } + else + *info = 0; + + return CV_HAL_ERROR_OK; +} + +int armpl_hal_LU32f(float *a, size_t a_step, int m, float *b, size_t b_step, int n, int *info) { + if (m < ARMPL_LU_SMALL_MATRIX_THRESH) + return CV_HAL_ERROR_NOT_IMPLEMENTED; + return armpl_lu(a, a_step, m, b, b_step, n, info); +} +int armpl_hal_LU64f(double *a, size_t a_step, int m, double *b, size_t b_step, int n, int *info) { + if (m < ARMPL_LU_SMALL_MATRIX_THRESH) + return CV_HAL_ERROR_NOT_IMPLEMENTED; + return armpl_lu(a, a_step, m, b, b_step, n, info); +} + enum ArmPLDFTMode { ARMPL_DFT_C2C,