refactor inv_sym() to handle complex hermitian matrices

This commit is contained in:
conrad
2024-10-28 15:44:25 +10:00
parent 98c7800566
commit 55059a4e30
4 changed files with 147 additions and 39 deletions
+3
View File
@@ -46,6 +46,9 @@ class auxlib
template<typename eT>
inline static bool inv_sym(Mat<eT>& A);
template<typename T>
inline static bool inv_sym(Mat< std::complex<T> >& A);
template<typename eT>
inline static bool inv_sympd(Mat<eT>& A, bool& out_sympd_state);
+67 -1
View File
@@ -242,7 +242,6 @@ auxlib::inv_tr_rcond(Mat<eT>& A, typename get_pod_type<eT>::result& out_rcond, c
// TODO: create specialisation for complex hermitian matrices, which replaces sytrf/sytri with hetrf/hetri
template<typename eT>
inline
bool
@@ -306,6 +305,73 @@ auxlib::inv_sym(Mat<eT>& A)
template<typename T>
inline
bool
auxlib::inv_sym(Mat< std::complex<T> >& 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<T> 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<blas_int> 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<blas_int>( access::tmp_real(work_query[0]) );
lwork = (std::max)(lwork_proposed, lwork);
}
podarray<eT> work( static_cast<uword>(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<typename eT>
inline
bool
+32 -22
View File
@@ -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
}
+45 -16
View File
@@ -1350,19 +1350,34 @@ namespace lapack
arma_type_check(( is_supported_blas_type<eT>::value == false ));
#if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS)
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); }
else if( is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); }
else if( is_cx_float<eT>::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<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); }
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); }
else if(is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info, 1); }
#else
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
else if( is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
else if( is_cx_float<eT>::value) { typedef blas_cxf T; arma_fortran(arma_csytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
else if(is_cx_double<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
else if(is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
#endif
}
template<typename eT>
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<eT>::value == false ));
#if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS)
if( is_cx_float<eT>::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<eT>::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<eT>::value) { typedef blas_cxf T; arma_fortran(arma_chetrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
else if(is_cx_double<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zhetrf)(uplo, n, (T*)a, lda, ipiv, (T*)work, lwork, info); }
#endif
}
template<typename eT>
inline
void
@@ -1371,15 +1386,29 @@ namespace lapack
arma_type_check(( is_supported_blas_type<eT>::value == false ));
#if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS)
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
else if( is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
else if( is_cx_float<eT>::value) { typedef blas_cxf T; arma_fortran(arma_csytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
else if(is_cx_double<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
else if(is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
#else
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
else if( is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
else if( is_cx_float<eT>::value) { typedef blas_cxf T; arma_fortran(arma_csytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
else if(is_cx_double<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
if( is_float<eT>::value) { typedef float T; arma_fortran(arma_ssytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
else if(is_double<eT>::value) { typedef double T; arma_fortran(arma_dsytri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
#endif
}
template<typename eT>
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<eT>::value == false ));
#if defined(ARMA_USE_FORTRAN_HIDDEN_ARGS)
if( is_cx_float<eT>::value) { typedef blas_cxf T; arma_fortran(arma_chetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
else if(is_cx_double<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zhetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info, 1); }
#else
if( is_cx_float<eT>::value) { typedef blas_cxf T; arma_fortran(arma_chetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
else if(is_cx_double<eT>::value) { typedef blas_cxd T; arma_fortran(arma_zhetri)(uplo, n, (T*)a, lda, ipiv, (T*)work, info); }
#endif
}