reduce computation

This commit is contained in:
conrad
2022-02-23 14:07:44 +10:00
parent c392ff8cf7
commit fe2c7bd1b3
3 changed files with 34 additions and 16 deletions
+3 -3
View File
@@ -226,11 +226,11 @@ auxlib::inv_sympd_rcond(Mat<eT>& A, eT& out_rcond, const eT rcond_threshold)
arma_extra_debug_print("lapack::potrf()");
lapack::potrf(&uplo, &n, A.memptr(), &n, &info);
if(info != 0) { return false; }
if(info != 0) { out_rcond = eT(0); return false; }
out_rcond = auxlib::lu_rcond_sympd<T>(A, norm_val);
if( (rcond_threshold > eT(0)) && (rcond < rcond_threshold) ) { return false; }
if( (rcond_threshold > eT(0)) && (out_rcond < rcond_threshold) ) { return false; }
arma_extra_debug_print("lapack::potri()");
lapack::potri(&uplo, &n, A.memptr(), &n, &info);
@@ -287,7 +287,7 @@ auxlib::inv_sympd_rcond(Mat< std::complex<T> >& A, T& out_rcond, const T rcond_t
arma_extra_debug_print("lapack::potrf()");
lapack::potrf(&uplo, &n, A.memptr(), &n, &info);
if(info != 0) { return false; }
if(info != 0) { out_rcond = T(0); return false; }
out_rcond = auxlib::lu_rcond_sympd<T>(A, norm_val);
+30 -12
View File
@@ -56,7 +56,7 @@ op_inv_rcond::apply_direct_gen(Mat<typename T1::elem_type>& out_inv, typename T1
template<typename T1>
inline
bool
op_inv_rcond::apply_direct_spd(Mat<typename T1::elem_type>& out_inv, typename T1::pod_type& out_rcond, const Base<typename T1::elem_type,T1>& expr)
op_inv_rcond::apply_direct_spd(Mat<typename T1::elem_type>& out, typename T1::pod_type& out_rcond, const Base<typename T1::elem_type,T1>& expr)
{
arma_extra_debug_sigprint();
@@ -65,22 +65,40 @@ op_inv_rcond::apply_direct_spd(Mat<typename T1::elem_type>& out_inv, typename T1
typedef typename T1::elem_type eT;
typedef typename T1::pod_type T;
const Mat<eT> A = expr.get_ref();
out = expr.get_ref();
arma_debug_check( (A.is_square() == false), "inv_sympd(): given matrix must be square sized" );
arma_debug_check( (out.is_square() == false), "inv_sympd(): given matrix must be square sized" );
const bool status = op_inv_spd::apply_direct<T1,false>(out_inv, A, uword(0));
if(status)
if((arma_config::debug) && (auxlib::rudimentary_sym_check(out) == false))
{
out_rcond = op_cond::rcond(expr.get_ref());
}
else
{
out_rcond = T(0);
if(is_cx<eT>::no ) { arma_debug_warn_level(1, "inv_sympd(): given matrix is not symmetric"); }
if(is_cx<eT>::yes) { arma_debug_warn_level(1, "inv_sympd(): given matrix is not hermitian"); }
}
return status;
const uword N = out.n_rows;
if(is_cx<eT>::yes)
{
arma_extra_debug_print("op_inv_spd: checking imaginary components of diagonal elements");
const T tol = T(100) * std::numeric_limits<T>::epsilon(); // allow some leeway
const eT* colmem = out.memptr();
for(uword i=0; i<N; ++i)
{
const eT& out_ii = colmem[i];
const T out_ii_imag = access::tmp_imag(out_ii);
if(std::abs(out_ii_imag) > tol) { return false; }
colmem += N;
}
}
// TODO: optimisation for diagonal matrices
return auxlib::inv_sympd_rcond(out, out_rcond, T(-1));
}
+1 -1
View File
@@ -149,7 +149,7 @@ op_inv_spd::apply_direct(Mat<typename T1::elem_type>& out, const Base<typename T
{
arma_extra_debug_print("op_inv_spd: detected diagonal matrix");
const eT* colmem = out.memptr();
eT* colmem = out.memptr();
for(uword i=0; i<N; ++i)
{