refactor inv_sym() to handle complex hermitian matrices
This commit is contained in:
@@ -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);
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user