directly avoid Proxy if possible

This commit is contained in:
conrad
2025-01-25 22:29:44 +10:00
parent bce19ebcbd
commit 426ce1372a
2 changed files with 104 additions and 63 deletions
+4 -2
View File
@@ -68,9 +68,11 @@ class op_all
static inline bool all_vec(T1& X);
template<typename T1>
static inline void apply_helper(Mat<uword>& out, const Proxy<T1>& P, const uword dim);
template<typename eT>
static inline void apply_mat_noalias(Mat<uword>& out, const Mat<eT>& X, const uword dim);
template<typename T1>
static inline void apply_proxy_noalias(Mat<uword>& out, const Proxy<T1>& P, const uword dim);
template<typename T1>
static inline void apply(Mat<uword>& out, const mtOp<uword, T1, op_all>& X);
+100 -61
View File
@@ -277,17 +277,15 @@ op_all::all_vec(T1& X)
template<typename T1>
template<typename eT>
inline
void
op_all::apply_helper(Mat<uword>& out, const Proxy<T1>& P, const uword dim)
op_all::apply_mat_noalias(Mat<uword>& out, const Mat<eT>& X, const uword dim)
{
arma_debug_sigprint();
const uword n_rows = P.get_n_rows();
const uword n_cols = P.get_n_cols();
typedef typename Proxy<T1>::elem_type eT;
const uword n_rows = X.n_rows;
const uword n_cols = X.n_cols;
if(dim == 0) // traverse rows (ie. process each column)
{
@@ -297,37 +295,18 @@ op_all::apply_helper(Mat<uword>& out, const Proxy<T1>& P, const uword dim)
uword* out_mem = out.memptr();
if(is_Mat<typename Proxy<T1>::stored_type>::value)
for(uword col=0; col < n_cols; ++col)
{
const unwrap<typename Proxy<T1>::stored_type> U(P.Q);
const eT* colmem = X.colptr(col);
for(uword col=0; col < n_cols; ++col)
uword count = 0;
for(uword row=0; row < n_rows; ++row)
{
const eT* colmem = U.M.colptr(col);
uword count = 0;
for(uword row=0; row < n_rows; ++row)
{
count += (colmem[row] != eT(0)) ? uword(1) : uword(0);
}
out_mem[col] = (n_rows == count) ? uword(1) : uword(0);
}
}
else
{
for(uword col=0; col < n_cols; ++col)
{
uword count = 0;
for(uword row=0; row < n_rows; ++row)
{
if(P.at(row,col) != eT(0)) { ++count; }
}
out_mem[col] = (n_rows == count) ? uword(1) : uword(0);
count += (colmem[row] != eT(0)) ? uword(1) : uword(0);
}
out_mem[col] = (n_rows == count) ? uword(1) : uword(0);
}
}
else
@@ -338,31 +317,15 @@ op_all::apply_helper(Mat<uword>& out, const Proxy<T1>& P, const uword dim)
// internal dual use of 'out': keep the counts for each row
if(is_Mat<typename Proxy<T1>::stored_type>::value)
for(uword col=0; col < n_cols; ++col)
{
const unwrap<typename Proxy<T1>::stored_type> U(P.Q);
const eT* colmem = X.colptr(col);
for(uword col=0; col < n_cols; ++col)
for(uword row=0; row < n_rows; ++row)
{
const eT* colmem = U.M.colptr(col);
for(uword row=0; row < n_rows; ++row)
{
out_mem[row] += (colmem[row] != eT(0)) ? uword(1) : uword(0);
}
out_mem[row] += (colmem[row] != eT(0)) ? uword(1) : uword(0);
}
}
else
{
for(uword col=0; col < n_cols; ++col)
{
for(uword row=0; row < n_rows; ++row)
{
if(P.at(row,col) != eT(0)) { ++out_mem[row]; }
}
}
}
// see what the counts tell us
@@ -370,7 +333,63 @@ op_all::apply_helper(Mat<uword>& out, const Proxy<T1>& P, const uword dim)
{
out_mem[row] = (n_cols == out_mem[row]) ? uword(1) : uword(0);
}
}
}
template<typename T1>
inline
void
op_all::apply_proxy_noalias(Mat<uword>& out, const Proxy<T1>& P, const uword dim)
{
arma_debug_sigprint();
typedef typename Proxy<T1>::elem_type eT;
const uword n_rows = P.get_n_rows();
const uword n_cols = P.get_n_cols();
if(dim == 0) // traverse rows (ie. process each column)
{
out.zeros(1, n_cols);
if(out.n_elem == 0) { return; }
uword* out_mem = out.memptr();
for(uword col=0; col < n_cols; ++col)
{
uword count = 0;
for(uword row=0; row < n_rows; ++row)
{
if(P.at(row,col) != eT(0)) { ++count; }
}
out_mem[col] = (n_rows == count) ? uword(1) : uword(0);
}
}
else
{
out.zeros(n_rows, 1);
uword* out_mem = out.memptr();
// internal dual use of 'out': keep the counts for each row
for(uword col=0; col < n_cols; ++col)
for(uword row=0; row < n_rows; ++row)
{
if(P.at(row,col) != eT(0)) { ++out_mem[row]; }
}
// see what the counts tell us
for(uword row=0; row < n_rows; ++row)
{
out_mem[row] = (n_cols == out_mem[row]) ? uword(1) : uword(0);
}
}
}
@@ -385,19 +404,39 @@ op_all::apply(Mat<uword>& out, const mtOp<uword, T1, op_all>& X)
const uword dim = X.aux_uword_a;
const Proxy<T1> P(X.m);
if(P.is_alias(out) == false)
if( (is_Mat<T1>::value) || (is_Mat<typename Proxy<T1>::stored_type>::value) || (arma_config::openmp && Proxy<T1>::use_mp) )
{
op_all::apply_helper(out, P, dim);
const quasi_unwrap<T1> U(X.m);
if(U.is_alias(out) == false)
{
op_all::apply_mat_noalias(out, U.M, dim);
}
else
{
Mat<uword> tmp;
op_all::apply_mat_noalias(tmp, U.M, dim);
out.steal_mem(tmp);
}
}
else
{
Mat<uword> out2;
const Proxy<T1> P(X.m);
op_all::apply_helper(out2, P, dim);
out.steal_mem(out2);
if(P.is_alias(out) == false)
{
op_all::apply_proxy_noalias(out, P, dim);
}
else
{
Mat<uword> tmp;
op_all::apply_proxy_noalias(tmp, P, dim);
out.steal_mem(tmp);
}
}
}