avoid inserting zero when diagonal element is equal to sigma

This commit is contained in:
conrad
2021-01-01 15:58:08 +10:00
parent 6525c8e4c2
commit 66a9bddbf7
+131 -93
View File
@@ -293,6 +293,8 @@ sp_auxlib::eigs_sym_newarp(Col<eT>& eigval, Mat<eT>& eigvec, const SpMat<eT>& X,
{
const newarp::SparseGenRealShiftSolve<eT> op(X, sigma);
if(op.valid == false) { return false; }
arma_debug_check( (op.n_rows != op.n_cols), "eigs_sym(): given matrix must be square sized" );
arma_debug_check( (n_eigvals >= op.n_rows), "eigs_sym(): n_eigvals must be less than the number of rows in the matrix" );
@@ -1596,25 +1598,10 @@ sp_auxlib::spsolve_refine(Mat<typename T1::elem_type>& X, typename T1::pod_type&
// We store in column-major CSC.
out.Stype = superlu::SLU_NC;
if(is_float<eT>::value)
{
out.Dtype = superlu::SLU_S;
}
else
if(is_double<eT>::value)
{
out.Dtype = superlu::SLU_D;
}
else
if(is_cx_float<eT>::value)
{
out.Dtype = superlu::SLU_C;
}
else
if(is_cx_double<eT>::value)
{
out.Dtype = superlu::SLU_Z;
}
if( is_float<eT>::value) { out.Dtype = superlu::SLU_S; }
else if( is_double<eT>::value) { out.Dtype = superlu::SLU_D; }
else if( is_cx_float<eT>::value) { out.Dtype = superlu::SLU_C; }
else if(is_cx_double<eT>::value) { out.Dtype = superlu::SLU_Z; }
out.Mtype = superlu::SLU_GE; // Just a general matrix. We don't know more now.
@@ -1662,116 +1649,167 @@ sp_auxlib::spsolve_refine(Mat<typename T1::elem_type>& X, typename T1::pod_type&
sp_auxlib::copy_to_supermatrix_with_shift(superlu::SuperMatrix& out, const SpMat<eT>& A, const eT shift)
{
arma_extra_debug_sigprint();
arma_debug_check( (A.is_square() == false), "sp_auxlib::copy_to_supermatrix_with_shift(): given matrix must be square sized" );
if(shift == eT(0)) { return sp_auxlib::copy_to_supermatrix(out, A); }
// We store in column-major CSC.
out.Stype = superlu::SLU_NC;
if(is_float<eT>::value)
{
out.Dtype = superlu::SLU_S;
}
else
if(is_double<eT>::value)
{
out.Dtype = superlu::SLU_D;
}
else
if(is_cx_float<eT>::value)
{
out.Dtype = superlu::SLU_C;
}
else
if(is_cx_double<eT>::value)
{
out.Dtype = superlu::SLU_Z;
}
if( is_float<eT>::value) { out.Dtype = superlu::SLU_S; }
else if( is_double<eT>::value) { out.Dtype = superlu::SLU_D; }
else if( is_cx_float<eT>::value) { out.Dtype = superlu::SLU_C; }
else if(is_cx_double<eT>::value) { out.Dtype = superlu::SLU_Z; }
out.Mtype = superlu::SLU_GE; // Just a general matrix. We don't know more now.
// We have to actually create the object which stores the data.
// This gets cleaned by destroy_supermatrix().
superlu::NCformat* nc = (superlu::NCformat*)superlu::malloc(sizeof(superlu::NCformat));
if(nc == nullptr) { return false; }
A.sync();
// Since A needs to be subtracted by shift*I, we treat the diagonal elements of out
// as non-zero elements, even if some of them may be numerically zero.
// As a result, we need to recompute the total number of non-zero elements.
const uword search_cols = std::min(A.n_rows, A.n_cols);
uword nnz_new = A.n_nonzero + search_cols;
for(uword j = 0; j < search_cols; j++)
uword n_nonzero_diag_old = 0;
uword n_nonzero_diag_new = 0;
const uword n_search_cols = (std::min)(A.n_rows, A.n_cols);
for(uword j=0; j < n_search_cols; ++j)
{
// A.col_ptrs[j+1]-A.col_ptrs[j] is the number of non-zero elements in the j-th column of A
const uword idx_start = A.col_ptrs[j];
const uword idx_end = A.col_ptrs[j + 1];
// Test whether A.row_indices[start:end] contains j
// If yes, it means A[j,j] is already there
for(uword i = idx_start; (i < idx_end) && (A.row_indices[i] <= j); i++)
const uword col_offset = A.col_ptrs[j ];
const uword next_col_offset = A.col_ptrs[j + 1];
const uword* start_ptr = &(A.row_indices[ col_offset]);
const uword* end_ptr = &(A.row_indices[next_col_offset]);
const uword wanted_row = j;
const uword* pos_ptr = std::lower_bound(start_ptr, end_ptr, wanted_row); // binary search
if( (pos_ptr != end_ptr) && ((*pos_ptr) == wanted_row) )
{
if(A.row_indices[i] == j)
{
// Decrease nnz_new by one if the diagonal element already exists
nnz_new--;
break;
}
// element on the main diagonal is non-zero
++n_nonzero_diag_old;
const uword offset = uword(pos_ptr - start_ptr);
const uword index = offset + col_offset;
const eT new_val = A.values[index] - shift;
if(new_val != eT(0)) { ++n_nonzero_diag_new; }
}
else
{
// element on the main diagonal is zero, but sigma is non-zero,
// so the number of new non-zero elments on the diagonal is increased
++n_nonzero_diag_new;
}
}
nc->nnz = nnz_new;
nc->nzval = (void*) superlu::malloc(sizeof(eT) * nnz_new );
const uword out_n_nonzero = A.n_nonzero - n_nonzero_diag_old + n_nonzero_diag_new;
arma_extra_debug_print( arma_str::format("A.n_nonzero: %d") % A.n_nonzero );
arma_extra_debug_print( arma_str::format("n_nonzero_diag_old: %d") % n_nonzero_diag_old );
arma_extra_debug_print( arma_str::format("n_nonzero_diag_new: %d") % n_nonzero_diag_new );
arma_extra_debug_print( arma_str::format("out_n_nonzero: %d") % out_n_nonzero );
nc->nnz = out_n_nonzero;
nc->nzval = (void*) superlu::malloc(sizeof(eT) * out_n_nonzero );
nc->colptr = (superlu::int_t*)superlu::malloc(sizeof(superlu::int_t) * (A.n_cols + 1));
nc->rowind = (superlu::int_t*)superlu::malloc(sizeof(superlu::int_t) * nnz_new );
nc->rowind = (superlu::int_t*)superlu::malloc(sizeof(superlu::int_t) * out_n_nonzero );
if( (nc->nzval == nullptr) || (nc->colptr == nullptr) || (nc->rowind == nullptr) ) { return false; }
// fill the matrix column by column, and insert diagonal elements when necessary
nc->colptr[0] = 0;
// Fill the matrix column by column, and insert diagonal elements when necessary
eT* values_current = (eT*) nc->nzval;
superlu::int_t* rowind_current = nc->rowind;
for(uword j = 0; j < A.n_cols; j++)
uword count = 0;
for(uword j=0; j < A.n_cols; ++j)
{
const uword idx_start = A.col_ptrs[j];
const uword idx_end = A.col_ptrs[j + 1];
const uword idx_start = A.col_ptrs[j ];
const uword idx_end = A.col_ptrs[j + 1];
const eT* values_start = values_current;
uword i;
// Segment 1: elements in the strictly upper triangular part
for(i = idx_start; (i < idx_end) && A.row_indices[i] < j; i++, values_current++, rowind_current++)
uword i = idx_start;
// elements in the upper triangular part, excluding the main diagonal
for(; (i < idx_end) && (A.row_indices[i] < j); ++i)
{
*values_current = A.values[i];
*rowind_current = A.row_indices[i];
(*values_current) = A.values[i];
(*rowind_current) = A.row_indices[i];
++values_current;
++rowind_current;
++count;
}
// Segment 2-a: A[j,j] exists
if( (i < idx_end) && A.row_indices[i] == j )
// elements on the main diagonal
if( (i < idx_end) && (A.row_indices[i] == j) )
{
*values_current = A.values[i] - shift;
*rowind_current = j;
values_current++;
rowind_current++;
i++;
// A(j,j) is non-zero
const eT new_diag_val = A.values[i] - shift;
if(new_diag_val != eT(0))
{
(*values_current) = new_diag_val;
(*rowind_current) = j;
++values_current;
++rowind_current;
++count;
}
++i;
}
// Segment 2-b: A[j,j] does not exist, so we insert a new element
else if(j < search_cols)
else
{
*values_current = -shift;
*rowind_current = j;
values_current++;
rowind_current++;
// A(j,j) is zero, so insert a new element
if(j < n_search_cols)
{
(*values_current) = -shift;
(*rowind_current) = j;
++values_current;
++rowind_current;
++count;
}
}
// Segment 3: elements in the strictly lower triangular part
for(; i < idx_end; i++, values_current++, rowind_current++)
// elements in the lower triangular part, excluding the main diagonal
for(; i < idx_end; ++i)
{
*values_current = A.values[i];
*rowind_current = A.row_indices[i];
(*values_current) = A.values[i];
(*rowind_current) = A.row_indices[i];
++values_current;
++rowind_current;
++count;
}
// Number of non-zero elements in the j-th column of out
// number of non-zero elements in the j-th column of out
const uword nnz_col = values_current - values_start;
nc->colptr[j + 1] = nc->colptr[j] + nnz_col;
}
arma_extra_debug_print( arma_str::format("count: %d") % count );
arma_check( (count != out_n_nonzero), "internal error: sp_auxlib::copy_to_supermatrix_with_shift(): count != out_n_nonzero" );
out.nrow = A.n_rows;
out.ncol = A.n_cols;
out.Store = (void*) nc;
@@ -1965,12 +2003,12 @@ sp_auxlib::run_aupd_plain
// Call saupd() or naupd() with the current parameters.
if(sym)
{
arma_extra_debug_print("arpack::saupd");
arma_extra_debug_print("arpack::saupd()");
arpack::saupd(&ido, &bmat, &n, which, &nev, &tol, resid.memptr(), &ncv, v.memptr(), &ldv, iparam.memptr(), ipntr.memptr(), workd.memptr(), workl.memptr(), &lworkl, &info);
}
else
{
arma_extra_debug_print("arpack::naupd");
arma_extra_debug_print("arpack::naupd()");
arpack::naupd(&ido, &bmat, &n, which, &nev, &tol, resid.memptr(), &ncv, v.memptr(), &ldv, iparam.memptr(), ipntr.memptr(), workd.memptr(), workl.memptr(), &lworkl, rwork.memptr(), &info);
}
@@ -2212,12 +2250,12 @@ sp_auxlib::run_aupd_shiftinvert
// Call saupd() or naupd() with the current parameters.
if(sym)
{
arma_extra_debug_print("arpack::saupd");
arma_extra_debug_print("arpack::saupd()");
arpack::saupd(&ido, &bmat, &n, which, &nev, &tol, resid.memptr(), &ncv, v.memptr(), &ldv, iparam.memptr(), ipntr.memptr(), workd.memptr(), workl.memptr(), &lworkl, &info);
}
else
{
arma_extra_debug_print("arpack::naupd");
arma_extra_debug_print("arpack::naupd()");
arpack::naupd(&ido, &bmat, &n, which, &nev, &tol, resid.memptr(), &ncv, v.memptr(), &ldv, iparam.memptr(), ipntr.memptr(), workd.memptr(), workl.memptr(), &lworkl, rwork.memptr(), &info);
}