faster handling of submatrix rows
This commit is contained in:
@@ -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;
|
||||
|
||||
@@ -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);
|
||||
|
||||
@@ -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];
|
||||
}
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user