expand sum() with omit_mode

This commit is contained in:
conrad
2025-06-05 13:55:41 +10:00
parent 65012ac724
commit 74720f366d
5 changed files with 193 additions and 3 deletions
-1
View File
@@ -366,7 +366,6 @@ namespace elem_opts
{
template<int omit_mode> struct omit_indicator {};
static constexpr omit_indicator<0> omit_none;
static constexpr omit_indicator<1> omit_nan;
static constexpr omit_indicator<2> omit_nonfinite;
}
-2
View File
@@ -1019,7 +1019,6 @@ accu(const T1& X, const elem_opts::omit_indicator<omit_mode>&)
auto is_omitted = [](const eT& x) -> bool
{
if(omit_mode == 0) { return false; }
if(omit_mode == 1) { return arma_isnan(x); }
if(omit_mode == 2) { return (arma_isfinite(x) == false); }
};
@@ -1490,7 +1489,6 @@ accu(const SpBase<typename T1::elem_type, T1>& expr, const elem_opts::omit_indic
auto is_omitted = [](const eT& x) -> bool
{
if(omit_mode == 0) { return false; }
if(omit_mode == 1) { return arma_isnan(x); }
if(omit_mode == 2) { return (arma_isfinite(x) == false); }
};
+13
View File
@@ -72,6 +72,19 @@ sum(const T1& X, const uword dim)
template<typename T1, int omit_mode>
arma_warn_unused
arma_inline
typename enable_if2< is_arma_type<T1>::value, const Op<T1, op_sum_omit> >::result
sum(const T1& X, const uword dim, const elem_opts::omit_indicator<omit_mode>& indicator)
{
arma_debug_sigprint();
return Op<T1, op_sum_omit>(X, dim, uword(omit_mode));
}
template<typename T>
arma_warn_unused
arma_inline
+14
View File
@@ -59,4 +59,18 @@ class op_sum
};
class op_sum_omit
: public traits_op_xvec
{
public:
template<typename T1>
inline static void apply(Mat<typename T1::elem_type>& out, const Op<T1, op_sum_omit>& in);
template<typename T1, typename functor>
inline static void apply_proxy_noalias(Mat<typename T1::elem_type>& out, const Proxy<T1>& P, const uword dim, functor is_omitted);
};
//! @}
+166
View File
@@ -415,6 +415,172 @@ op_sum::apply_proxy_noalias(Mat<typename T1::elem_type>& out, const Proxy<T1>& P
//
template<typename T1>
inline
void
op_sum_omit::apply(Mat<typename T1::elem_type>& out, const Op<T1,op_sum_omit>& in)
{
arma_debug_sigprint();
typedef typename T1::elem_type eT;
const uword dim = in.aux_uword_a;
const uword omit_mode = in.aux_uword_b;
arma_conform_check( (dim > 1), "sum(): parameter 'dim' must be 0 or 1" );
auto is_omitted_1 = [](const eT& x) -> bool { return arma_isnan(x); };
auto is_omitted_2 = [](const eT& x) -> bool { return (arma_isfinite(x) == false); };
const Proxy<T1> P(in.m);
if(P.is_alias(out))
{
Mat<eT> tmp;
if(omit_mode == 1) { op_sum_omit::apply_proxy_noalias(tmp, P, dim, is_omitted_1); }
if(omit_mode == 2) { op_sum_omit::apply_proxy_noalias(tmp, P, dim, is_omitted_2); }
out.steal_mem(tmp);
}
else
{
if(omit_mode == 1) { op_sum_omit::apply_proxy_noalias(out, P, dim, is_omitted_1); }
if(omit_mode == 2) { op_sum_omit::apply_proxy_noalias(out, P, dim, is_omitted_2); }
}
}
template<typename T1, typename functor>
inline
void
op_sum_omit::apply_proxy_noalias(Mat<typename T1::elem_type>& out, const Proxy<T1>& P, const uword dim, functor is_omitted)
{
arma_debug_sigprint();
typedef typename T1::elem_type eT;
constexpr eT eT_zero = eT(0);
const uword P_n_rows = P.get_n_rows();
const uword P_n_cols = P.get_n_cols();
const uword out_n_rows = (dim == 0) ? uword(1) : P_n_rows;
const uword out_n_cols = (dim == 0) ? P_n_cols : uword(1);
out.set_size(out_n_rows, out_n_cols);
if(P.get_n_elem() == 0) { out.zeros(); return; }
eT* out_mem = out.memptr();
if(Proxy<T1>::use_at == false)
{
if(dim == 0)
{
uword count = 0;
for(uword col=0; col < P_n_cols; ++col)
{
eT val1 = eT(0);
eT val2 = eT(0);
uword j;
for(j=1; j < P_n_rows; j+=2)
{
const eT tmp1 = P[count]; ++count;
const eT tmp2 = P[count]; ++count;
val1 += is_omitted(tmp1) ? eT_zero : tmp1;
val2 += is_omitted(tmp2) ? eT_zero : tmp2;
}
if((j-1) < P_n_rows)
{
const eT tmp1 = P[count]; ++count;
val1 += is_omitted(tmp1) ? eT_zero : tmp1;
}
out_mem[col] = (val1 + val2);
}
}
else
{
uword count = 0;
for(uword row=0; row < P_n_rows; ++row)
{
const eT tmp = P[count]; ++count;
out_mem[row] = is_omitted(tmp) ? eT_zero : tmp;
}
for(uword col=1; col < P_n_cols; ++col)
for(uword row=0; row < P_n_rows; ++row)
{
const eT tmp = P[count]; ++count;
out_mem[row] += is_omitted(tmp) ? eT_zero : tmp;
}
}
}
else
{
if(dim == 0)
{
for(uword col=0; col < P_n_cols; ++col)
{
eT val1 = eT(0);
eT val2 = eT(0);
uword i,j;
for(i=0, j=1; j < P_n_rows; i+=2, j+=2)
{
const eT tmp1 = P.at(i,col);
const eT tmp2 = P.at(j,col);
val1 += is_omitted(tmp1) ? eT_zero : tmp1;
val2 += is_omitted(tmp2) ? eT_zero : tmp2;
}
if(i < P_n_rows)
{
const eT tmp1 = P.at(i,col);
val1 += is_omitted(tmp1) ? eT_zero : tmp1;
}
out_mem[col] = (val1 + val2);
}
}
else
{
for(uword row=0; row < P_n_rows; ++row)
{
const eT tmp = P.at(row,0);
out_mem[row] = is_omitted(tmp) ? eT_zero : tmp;
}
for(uword col=1; col < P_n_cols; ++col)
for(uword row=0; row < P_n_rows; ++row)
{
const eT tmp = P.at(row,col);
out_mem[row] += is_omitted(tmp) ? eT_zero : tmp;
}
}
}
}
//
// cubes