speedup via inplace processing

This commit is contained in:
conrad
2023-04-02 23:40:05 +10:00
parent 9cc1964fce
commit 3f02dc034c
+54 -1
View File
@@ -354,7 +354,60 @@ SpSubview<eT>::operator%=(const Base<eT, T1>& x)
{
arma_extra_debug_sigprint();
return (*this).operator=( (*this) % x.get_ref() );
SpSubview<eT>& sv = (*this);
const quasi_unwrap<T1> U(x.get_ref());
const Mat<eT>& B = U.M;
arma_debug_assert_same_size(sv.n_rows, sv.n_cols, B.n_rows, B.n_cols, "element-wise multiplication");
SpMat<eT>& sv_m = access::rw(sv.m);
sv_m.sync_csc();
sv_m.invalidate_cache();
const uword m_row_start = sv.aux_row1;
const uword m_row_end = sv.aux_row1 + sv.n_rows - 1;
const uword m_col_start = sv.aux_col1;
const uword m_col_end = sv.aux_col1 + sv.n_cols - 1;
constexpr eT zero = eT(0);
bool has_zero = false;
uword count = 0;
for(uword m_col = m_col_start; m_col <= m_col_end; ++m_col)
{
const uword sv_col = m_col - m_col_start;
const uword index_start = sv_m.col_ptrs[m_col ];
const uword index_end = sv_m.col_ptrs[m_col + 1];
for(uword i=index_start; i < index_end; ++i)
{
const uword m_row = sv_m.row_indices[i];
if(m_row < m_row_start) { continue; }
if(m_row > m_row_end ) { break; }
const uword sv_row = m_row - m_row_start;
eT& m_val = access::rw(sv_m.values[i]);
const eT result = m_val * B.at(sv_row, sv_col);
m_val = result;
if(result == zero) { has_zero = true; } else { ++count; }
}
}
if(has_zero) { sv_m.remove_zeros(); }
access::rw(sv.n_nonzero) = count;
return (*this);
}