faster handling of submatrix rows

This commit is contained in:
conrad
2025-10-26 20:07:38 +10:00
parent 485b705757
commit 65c42d01cd
3 changed files with 265 additions and 44 deletions
+28 -5
View File
@@ -602,7 +602,30 @@ op_accu_mat::apply(const subview<eT>& X)
const uword X_n_rows = X.n_rows;
const uword X_n_cols = X.n_cols;
if(X_n_rows == 1) { return op_accu_mat::apply( static_cast< const subview_row<eT>& >(X) ); }
if(X_n_rows == 1)
{
const uword X_m_n_rows = X.m.n_rows;
const eT* mem_ptr = X.colptr(0);
eT val1 = eT(0);
eT val2 = eT(0);
uword j;
for(j=1; j < X_n_cols; j+=2)
{
val1 += (*mem_ptr); mem_ptr += X_m_n_rows;
val2 += (*mem_ptr); mem_ptr += X_m_n_rows;
}
if((j-1) < X_n_cols)
{
val1 += (*mem_ptr);
}
return val1 + val2;
}
if(X_n_cols == 1) { return arrayops::accumulate( X.colptr(0), X_n_rows ); }
@@ -640,7 +663,7 @@ op_accu_mat::apply(const subview_row<eT>& X)
const uword X_m_n_rows = X.m.n_rows;
const uword X_n_cols = X.n_cols;
const eT* row_mem = &(X.m.at(X.aux_row1,X.aux_col1));
const eT* mem_ptr = X.rowmem;
eT val1 = eT(0);
eT val2 = eT(0);
@@ -649,13 +672,13 @@ op_accu_mat::apply(const subview_row<eT>& X)
for(j=1; j < X_n_cols; j+=2)
{
val1 += (*row_mem); row_mem += X_m_n_rows;
val2 += (*row_mem); row_mem += X_m_n_rows;
val1 += (*mem_ptr); mem_ptr += X_m_n_rows;
val2 += (*mem_ptr); mem_ptr += X_m_n_rows;
}
if((j-1) < X_n_cols)
{
val1 += (*row_mem);
val1 += (*mem_ptr);
}
return val1 + val2;
+17 -1
View File
@@ -85,7 +85,7 @@ class subview : public Base< eT, subview<eT> >
template<typename T1> inline void operator-= (const SpBase<eT,T1>& x);
template<typename T1> inline void operator%= (const SpBase<eT,T1>& x);
template<typename T1> inline void operator/= (const SpBase<eT,T1>& x);
template<typename T1, typename gen_type>
inline typename enable_if2< is_same_type<typename T1::elem_type, eT>::value, void>::result operator=(const Gen<T1,gen_type>& x);
@@ -396,6 +396,11 @@ class subview_col : public subview<eT>
inline void zeros();
inline void ones();
arma_warn_unused inline bool is_finite() const;
arma_warn_unused inline bool has_inf() const;
arma_warn_unused inline bool has_nan() const;
arma_inline eT at_alt (const uword i) const;
arma_inline eT& operator[](const uword i);
@@ -528,6 +533,8 @@ class subview_row : public subview<eT>
static constexpr bool is_col = false;
static constexpr bool is_xvec = false;
const eT* rowmem;
inline void operator= (const subview<eT>& x);
inline void operator= (const subview_row& x);
inline void operator= (const eT val);
@@ -545,6 +552,15 @@ class subview_row : public subview<eT>
arma_warn_unused arma_inline const Op<subview_row<eT>,op_strans> as_col() const;
inline void fill(const eT val);
inline void zeros();
inline void ones();
arma_warn_unused inline bool is_finite() const;
arma_warn_unused inline bool has_inf() const;
arma_warn_unused inline bool has_nan() const;
inline eT at_alt (const uword i) const;
inline eT& operator[](const uword i);
+220 -38
View File
@@ -1079,16 +1079,9 @@ subview<eT>::fill(const eT val)
eT* Aptr = &(A.at(s.aux_row1,s.aux_col1));
uword jj;
for(jj=1; jj < s_n_cols; jj+=2)
for(uword ii=0; ii < s_n_cols; ++ii)
{
(*Aptr) = val; Aptr += A_n_rows;
(*Aptr) = val; Aptr += A_n_rows;
}
if((jj-1) < s_n_cols)
{
(*Aptr) = val;
}
}
else
@@ -3380,7 +3373,7 @@ subview_col<eT>::operator=(const Base<eT,T1>& expr)
if(is_Mat<T1>::value)
{
const unwrap<T1> U(expr.get_ref());
const unwrap<T1> U(expr.get_ref()); // deliberately not using quasi_unwrap
arma_conform_assert_same_size(subview<eT>::n_rows, uword(1), U.M.n_rows, U.M.n_cols, "copy into submatrix");
@@ -3498,6 +3491,48 @@ subview_col<eT>::ones()
template<typename eT>
inline
bool
subview_col<eT>::is_finite() const
{
arma_debug_sigprint();
if(arma_config::fast_math_warn) { arma_warn(1, "is_finite(): detection of non-finite values is not reliable in fast math mode"); }
return arrayops::is_finite(colmem, subview<eT>::n_rows);
}
template<typename eT>
inline
bool
subview_col<eT>::has_inf() const
{
arma_debug_sigprint();
if(arma_config::fast_math_warn) { arma_warn(1, "has_inf(): detection of non-finite values is not reliable in fast math mode"); }
return arrayops::has_inf(colmem, subview<eT>::n_rows);
}
template<typename eT>
inline
bool
subview_col<eT>::has_nan() const
{
arma_debug_sigprint();
if(arma_config::fast_math_warn) { arma_warn(1, "has_nan(): detection of non-finite values is not reliable in fast math mode"); }
return arrayops::has_nan(colmem, subview<eT>::n_rows);
}
template<typename eT>
arma_inline
eT
@@ -4233,6 +4268,7 @@ template<typename eT>
inline
subview_row<eT>::subview_row(const Mat<eT>& in_m, const uword in_row)
: subview<eT>(in_m, in_row, 0, 1, in_m.n_cols)
, rowmem(subview<eT>::colptr(0))
{
arma_debug_sigprint();
}
@@ -4243,6 +4279,7 @@ template<typename eT>
inline
subview_row<eT>::subview_row(const Mat<eT>& in_m, const uword in_row, const uword in_col1, const uword in_n_cols)
: subview<eT>(in_m, in_row, in_col1, 1, in_n_cols)
, rowmem(subview<eT>::colptr(0))
{
arma_debug_sigprint();
}
@@ -4253,6 +4290,7 @@ template<typename eT>
inline
subview_row<eT>::subview_row(const subview_row<eT>& in)
: subview<eT>(in) // interprets 'subview_row' as 'subview'
, rowmem(in.rowmem)
{
arma_debug_sigprint();
}
@@ -4263,8 +4301,11 @@ template<typename eT>
inline
subview_row<eT>::subview_row(subview_row<eT>&& in)
: subview<eT>(std::move(in)) // interprets 'subview_row' as 'subview'
, rowmem(in.rowmem)
{
arma_debug_sigprint();
access::rw(in.rowmem) = nullptr;
}
@@ -4300,7 +4341,12 @@ subview_row<eT>::operator=(const eT val)
{
arma_debug_sigprint();
subview<eT>::operator=(val); // interprets 'subview_row' as 'subview'
if(subview<eT>::n_elem != 1)
{
arma_conform_assert_same_size(subview<eT>::n_rows, subview<eT>::n_cols, 1, 1, "copy into submatrix");
}
access::rw( rowmem[0] ) = val;
}
@@ -4335,7 +4381,39 @@ subview_row<eT>::operator=(const Base<eT,T1>& X)
{
arma_debug_sigprint();
subview<eT>::operator=(X);
if(is_Mat<T1>::value)
{
const unwrap<T1> U(X.get_ref()); // deliberately not using quasi_unwrap
arma_conform_assert_same_size(uword(1), subview<eT>::n_cols, U.M.n_rows, U.M.n_cols, "copy into submatrix");
const eT* UM_mem = U.M.memptr();
eT* mem_ptr = access::rwp(rowmem);
const uword local_s_n_cols = subview<eT>::n_cols;
const uword local_m_n_rows = subview<eT>::m.n_rows;
uword j;
for(j=1; j < local_s_n_cols; j+=2)
{
const eT val_i = (*UM_mem); UM_mem++;
const eT val_j = (*UM_mem); UM_mem++;
(*mem_ptr) = val_i; mem_ptr += local_m_n_rows;
(*mem_ptr) = val_j; mem_ptr += local_m_n_rows;
}
if((j-1) < local_s_n_cols)
{
(*mem_ptr) = (*UM_mem);
}
}
else
{
subview<eT>::operator=(X);
}
}
@@ -4408,14 +4486,134 @@ subview_row<eT>::as_col() const
template<typename eT>
inline
void
subview_row<eT>::fill(const eT val)
{
arma_debug_sigprint();
eT* mem_ptr = access::rwp(rowmem);
const uword local_s_n_cols = subview<eT>::n_cols;
const uword local_m_n_rows = subview<eT>::m.n_rows;
for(uword ii=0; ii < local_s_n_cols; ++ii)
{
(*mem_ptr) = val; mem_ptr += local_m_n_rows;
}
}
template<typename eT>
inline
void
subview_row<eT>::zeros()
{
arma_debug_sigprint();
(*this).fill(eT(0));
}
template<typename eT>
inline
void
subview_row<eT>::ones()
{
arma_debug_sigprint();
(*this).fill(eT(1));
}
template<typename eT>
inline
bool
subview_row<eT>::is_finite() const
{
arma_debug_sigprint();
if(arma_config::fast_math_warn) { arma_warn(1, "is_finite(): detection of non-finite values is not reliable in fast math mode"); }
const eT* mem_ptr = rowmem;
const uword local_s_n_cols = subview<eT>::n_cols;
const uword local_m_n_rows = subview<eT>::m.n_rows;
for(uword ii=0; ii < local_s_n_cols; ++ii)
{
const eT val = (*mem_ptr); mem_ptr += local_m_n_rows;
if(arma_isnonfinite(val)) { return false; }
}
return true;
}
template<typename eT>
inline
bool
subview_row<eT>::has_inf() const
{
arma_debug_sigprint();
if(arma_config::fast_math_warn) { arma_warn(1, "has_inf(): detection of non-finite values is not reliable in fast math mode"); }
const eT* mem_ptr = rowmem;
const uword local_s_n_cols = subview<eT>::n_cols;
const uword local_m_n_rows = subview<eT>::m.n_rows;
for(uword ii=0; ii < local_s_n_cols; ++ii)
{
const eT val = (*mem_ptr); mem_ptr += local_m_n_rows;
if(arma_isinf(val)) { return true; }
}
return false;
}
template<typename eT>
inline
bool
subview_row<eT>::has_nan() const
{
arma_debug_sigprint();
if(arma_config::fast_math_warn) { arma_warn(1, "has_nan(): detection of non-finite values is not reliable in fast math mode"); }
const eT* mem_ptr = rowmem;
const uword local_s_n_cols = subview<eT>::n_cols;
const uword local_m_n_rows = subview<eT>::m.n_rows;
for(uword ii=0; ii < local_s_n_cols; ++ii)
{
const eT val = (*mem_ptr); mem_ptr += local_m_n_rows;
if(arma_isnan(val)) { return true; }
}
return false;
}
template<typename eT>
inline
eT
subview_row<eT>::at_alt(const uword ii) const
{
const uword index = (ii + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return subview<eT>::m.mem[index];
return rowmem[ii * subview<eT>::m.n_rows];
}
@@ -4425,9 +4623,7 @@ inline
eT&
subview_row<eT>::operator[](const uword ii)
{
const uword index = (ii + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return access::rw( (const_cast< Mat<eT>& >(subview<eT>::m)).mem[index] );
return access::rw( rowmem[ii * subview<eT>::m.n_rows] );
}
@@ -4437,9 +4633,7 @@ inline
eT
subview_row<eT>::operator[](const uword ii) const
{
const uword index = (ii + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return subview<eT>::m.mem[index];
return rowmem[ii * subview<eT>::m.n_rows];
}
@@ -4450,10 +4644,8 @@ eT&
subview_row<eT>::operator()(const uword ii)
{
arma_conform_check_bounds( (ii >= subview<eT>::n_elem), "subview::operator(): index out of bounds" );
const uword index = (ii + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return access::rw( (const_cast< Mat<eT>& >(subview<eT>::m)).mem[index] );
return access::rw( rowmem[ii * subview<eT>::m.n_rows] );
}
@@ -4465,9 +4657,7 @@ subview_row<eT>::operator()(const uword ii) const
{
arma_conform_check_bounds( (ii >= subview<eT>::n_elem), "subview::operator(): index out of bounds" );
const uword index = (ii + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return subview<eT>::m.mem[index];
return rowmem[ii * subview<eT>::m.n_rows];
}
@@ -4479,9 +4669,7 @@ subview_row<eT>::operator()(const uword in_row, const uword in_col)
{
arma_conform_check_bounds( ((in_row > 0) || (in_col >= subview<eT>::n_cols)), "subview::operator(): index out of bounds" );
const uword index = (in_col + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return access::rw( (const_cast< Mat<eT>& >(subview<eT>::m)).mem[index] );
return access::rw( rowmem[in_col * subview<eT>::m.n_rows] );
}
@@ -4493,9 +4681,7 @@ subview_row<eT>::operator()(const uword in_row, const uword in_col) const
{
arma_conform_check_bounds( ((in_row > 0) || (in_col >= subview<eT>::n_cols)), "subview::operator(): index out of bounds" );
const uword index = (in_col + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return subview<eT>::m.mem[index];
return rowmem[in_col * subview<eT>::m.n_rows];
}
@@ -4505,9 +4691,7 @@ inline
eT&
subview_row<eT>::at(const uword, const uword in_col)
{
const uword index = (in_col + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return access::rw( (const_cast< Mat<eT>& >(subview<eT>::m)).mem[index] );
return access::rw( rowmem[in_col * subview<eT>::m.n_rows] );
}
@@ -4517,9 +4701,7 @@ inline
eT
subview_row<eT>::at(const uword, const uword in_col) const
{
const uword index = (in_col + (subview<eT>::aux_col1))*(subview<eT>::m).n_rows + (subview<eT>::aux_row1);
return subview<eT>::m.mem[index];
return rowmem[in_col * subview<eT>::m.n_rows];
}