From 55059a4e308b6ef364ec223afd57539afa24c032 Mon Sep 17 00:00:00 2001 From: conrad Date: Mon, 28 Oct 2024 15:44:25 +1000 Subject: [PATCH] refactor inv_sym() to handle complex hermitian matrices --- include/armadillo_bits/auxlib_bones.hpp | 3 + include/armadillo_bits/auxlib_meat.hpp | 68 ++++++++++++++++++++- include/armadillo_bits/def_lapack.hpp | 54 +++++++++------- include/armadillo_bits/translate_lapack.hpp | 61 +++++++++++++----- 4 files changed, 147 insertions(+), 39 deletions(-) diff --git a/include/armadillo_bits/auxlib_bones.hpp b/include/armadillo_bits/auxlib_bones.hpp index 336867c7..e8399e97 100644 --- a/include/armadillo_bits/auxlib_bones.hpp +++ b/include/armadillo_bits/auxlib_bones.hpp @@ -46,6 +46,9 @@ class auxlib template inline static bool inv_sym(Mat& A); + template + inline static bool inv_sym(Mat< std::complex >& A); + template inline static bool inv_sympd(Mat& A, bool& out_sympd_state); diff --git a/include/armadillo_bits/auxlib_meat.hpp b/include/armadillo_bits/auxlib_meat.hpp index 9e992516..6a955214 100644 --- a/include/armadillo_bits/auxlib_meat.hpp +++ b/include/armadillo_bits/auxlib_meat.hpp @@ -242,7 +242,6 @@ auxlib::inv_tr_rcond(Mat& A, typename get_pod_type::result& out_rcond, c -// TODO: create specialisation for complex hermitian matrices, which replaces sytrf/sytri with hetrf/hetri template inline bool @@ -306,6 +305,73 @@ auxlib::inv_sym(Mat& A) +template +inline +bool +auxlib::inv_sym(Mat< std::complex >& A) + { + arma_debug_sigprint(); + + // NOTE: the function name is required for overloading, but is a misnomer: it processes hermitian complex matrices + + if(A.is_empty()) { return true; } + + #if defined(ARMA_USE_LAPACK) + { + typedef typename std::complex eT; + + arma_conform_assert_blas_size(A); + + char uplo = 'L'; + blas_int n = blas_int(A.n_rows); + blas_int lda = blas_int(A.n_rows); + blas_int lwork = (std::max)(blas_int(podarray_prealloc_n_elem::val), n); + blas_int info = 0; + + podarray ipiv(A.n_rows); + + if(n > 16) + { + eT work_query[2] = {}; + blas_int lwork_query = -1; + + arma_debug_print("lapack::hetrf()"); + lapack::hetrf(&uplo, &n, A.memptr(), &lda, ipiv.memptr(), &work_query[0], &lwork_query, &info); + + if(info != 0) { return false; } + + blas_int lwork_proposed = static_cast( access::tmp_real(work_query[0]) ); + + lwork = (std::max)(lwork_proposed, lwork); + } + + podarray work( static_cast(lwork) ); + + arma_debug_print("lapack::hetrf()"); + lapack::hetrf(&uplo, &n, A.memptr(), &lda, ipiv.memptr(), work.memptr(), &lwork, &info); + + if(info != 0) { return false; } + + arma_debug_print("lapack::hetri()"); + lapack::hetri(&uplo, &n, A.memptr(), &lda, ipiv.memptr(), work.memptr(), &info); + + if(info != 0) { return false; } + + A = symmatl(A); + + return true; + } + #else + { + arma_ignore(A); + arma_stop_logic_error("inv_sym(): use of LAPACK must be enabled"); + return false; + } + #endif + } + + + template inline bool diff --git a/include/armadillo_bits/def_lapack.hpp b/include/armadillo_bits/def_lapack.hpp index b9dad379..d9456816 100644 --- a/include/armadillo_bits/def_lapack.hpp +++ b/include/armadillo_bits/def_lapack.hpp @@ -271,13 +271,15 @@ #define arma_ssytrf ssytrf #define arma_dsytrf dsytrf - #define arma_csytrf csytrf - #define arma_zsytrf zsytrf + + #define arma_chetrf chetrf + #define arma_zhetrf zhetrf #define arma_ssytri ssytri #define arma_dsytri dsytri - #define arma_csytri csytri - #define arma_zsytri zsytri + + #define arma_chetri chetri + #define arma_zhetri zhetri #else @@ -517,13 +519,15 @@ #define arma_ssytrf SSYTRF #define arma_dsytrf DSYTRF - #define arma_csytrf CSYTRF - #define arma_zsytrf ZSYTRF + + #define arma_chetrf CHETRF + #define arma_zhetrf ZHETRF #define arma_ssytri SSYTRI #define arma_dsytri DSYTRI - #define arma_csytri CSYTRI - #define arma_zsytri ZSYTRI + + #define arma_chetri CHETRI + #define arma_zhetri ZHETRI #endif @@ -866,19 +870,21 @@ extern "C" void arma_fortran(arma_cpstrf)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* piv, blas_int* rank, const float* tol, float* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; void arma_fortran(arma_zpstrf)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* piv, blas_int* rank, const double* tol, double* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - // factorisation of symmetric matrix + // factorisation of symmetric matrix (real) void arma_fortran(arma_ssytrf)(const char* uplo, const blas_int* n, float* a, const blas_int* lda, blas_int* ipiv, float* work, const blas_int* lwork, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; void arma_fortran(arma_dsytrf)(const char* uplo, const blas_int* n, double* a, const blas_int* lda, blas_int* ipiv, double* work, const blas_int* lwork, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - void arma_fortran(arma_csytrf)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, const blas_int* lwork, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - void arma_fortran(arma_zsytrf)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, const blas_int* lwork, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - // TODO: replace csytrf/zsytrf with chetrf/zhetrf - // inverse of symmetric matrix (using pre-computed factorisation) + // factorisation of hermitian matrix (complex) + void arma_fortran(arma_chetrf)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, const blas_int* lwork, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; + void arma_fortran(arma_zhetrf)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, const blas_int* lwork, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; + + // inverse of symmetric matrix using pre-computed factorisation (real) void arma_fortran(arma_ssytri)(const char* uplo, const blas_int* n, float* a, const blas_int* lda, blas_int* ipiv, float* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; void arma_fortran(arma_dsytri)(const char* uplo, const blas_int* n, double* a, const blas_int* lda, blas_int* ipiv, double* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - void arma_fortran(arma_csytri)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - void arma_fortran(arma_zsytri)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; - // TODO: replace csytri/zsytri with chetri/zhetri + + // inverse of hermitian matrix using pre-computed factorisation (complex) + void arma_fortran(arma_chetri)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; + void arma_fortran(arma_zhetri)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, blas_int* info, blas_len uplo_len) ARMA_NOEXCEPT; #else @@ -1204,17 +1210,21 @@ extern "C" void arma_fortran(arma_cpstrf)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* piv, blas_int* rank, const float* tol, float* work, blas_int* info) ARMA_NOEXCEPT; void arma_fortran(arma_zpstrf)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* piv, blas_int* rank, const double* tol, double* work, blas_int* info) ARMA_NOEXCEPT; - // factorisation of symmetric matrix + // factorisation of symmetric matrix (real) void arma_fortran(arma_ssytrf)(const char* uplo, const blas_int* n, float* a, const blas_int* lda, blas_int* ipiv, float* work, const blas_int* lwork, blas_int* info) ARMA_NOEXCEPT; void arma_fortran(arma_dsytrf)(const char* uplo, const blas_int* n, double* a, const blas_int* lda, blas_int* ipiv, double* work, const blas_int* lwork, blas_int* info) ARMA_NOEXCEPT; - void arma_fortran(arma_csytrf)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, const blas_int* lwork, blas_int* info) ARMA_NOEXCEPT; - void arma_fortran(arma_zsytrf)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, const blas_int* lwork, blas_int* info) ARMA_NOEXCEPT; - // inverse of symmetric matrix (using pre-computed factorisation) + // factorisation of hermitian matrix (complex) + void arma_fortran(arma_chetrf)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, const blas_int* lwork, blas_int* info) ARMA_NOEXCEPT; + void arma_fortran(arma_zhetrf)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, const blas_int* lwork, blas_int* info) ARMA_NOEXCEPT; + + // inverse of symmetric matrix using pre-computed factorisation (real) void arma_fortran(arma_ssytri)(const char* uplo, const blas_int* n, float* a, const blas_int* lda, blas_int* ipiv, float* work, blas_int* info) ARMA_NOEXCEPT; void arma_fortran(arma_dsytri)(const char* uplo, const blas_int* n, double* a, const blas_int* lda, blas_int* ipiv, double* work, blas_int* info) ARMA_NOEXCEPT; - void arma_fortran(arma_csytri)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, blas_int* info) ARMA_NOEXCEPT; - void arma_fortran(arma_zsytri)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, blas_int* info) ARMA_NOEXCEPT; + + // inverse of hermitian matrix using pre-computed factorisation (complex) + void arma_fortran(arma_chetri)(const char* uplo, const blas_int* n, blas_cxf* a, const blas_int* lda, blas_int* ipiv, blas_cxf* work, blas_int* info) ARMA_NOEXCEPT; + void arma_fortran(arma_zhetri)(const char* uplo, const blas_int* n, blas_cxd* a, const blas_int* lda, blas_int* ipiv, blas_cxd* work, blas_int* info) ARMA_NOEXCEPT; #endif } diff --git a/include/armadillo_bits/translate_lapack.hpp b/include/armadillo_bits/translate_lapack.hpp index 6eaffc43..b3797063 100644 --- a/include/armadillo_bits/translate_lapack.hpp +++ b/include/armadillo_bits/translate_lapack.hpp @@ -1350,19 +1350,34 @@ namespace lapack arma_type_check(( is_supported_blas_type::value == false )); #if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS) - if( is_float::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } - else if( is_double::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } - else if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_csytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } - else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } + if( is_float::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } + else if(is_double::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } #else - if( is_float::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } - else if( is_double::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } - else if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_csytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } - else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } + if( is_float::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } + else if(is_double::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } #endif } + + template + inline + void + hetrf(const char* uplo, const blas_int* n, eT* a, const blas_int* lda, blas_int* ipiv, eT* work, blas_int* lwork, blas_int* info) + { + arma_type_check(( is_supported_blas_type::value == false )); + + #if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS) + if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_chetrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } + else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zhetrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); } + #else + if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_chetrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } + else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zhetrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); } + #endif + } + + + template inline void @@ -1371,15 +1386,29 @@ namespace lapack arma_type_check(( is_supported_blas_type::value == false )); #if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS) - if( is_float::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } - else if( is_double::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } - else if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_csytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } - else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } + if( is_float::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } + else if(is_double::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } #else - if( is_float::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } - else if( is_double::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } - else if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_csytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } - else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } + if( is_float::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } + else if(is_double::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } + #endif + } + + + + template + inline + void + hetri(const char* uplo, const blas_int* n, eT* a, const blas_int* lda, blas_int* ipiv, eT* work, blas_int* info) + { + arma_type_check(( is_supported_blas_type::value == false )); + + #if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS) + if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_chetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } + else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zhetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); } + #else + if( is_cx_float::value) { typedef blas_cxf T; arma_fortran(arma_chetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } + else if(is_cx_double::value) { typedef blas_cxd T; arma_fortran(arma_zhetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); } #endif }