From 03e92c43ba4bb1b19f736c7b01d1e8fcc778de1b Mon Sep 17 00:00:00 2001 From: conrad Date: Sun, 8 Jun 2025 01:23:11 +1000 Subject: [PATCH] expand op_mean_omit to general matrices --- include/armadillo_bits/fn_mean.hpp | 13 +++ include/armadillo_bits/op_mean_bones.hpp | 12 ++- include/armadillo_bits/op_mean_meat.hpp | 106 ++++++++++++++++++++--- 3 files changed, 116 insertions(+), 15 deletions(-) diff --git a/include/armadillo_bits/fn_mean.hpp b/include/armadillo_bits/fn_mean.hpp index 54028ecb..f3a6e927 100644 --- a/include/armadillo_bits/fn_mean.hpp +++ b/include/armadillo_bits/fn_mean.hpp @@ -73,6 +73,19 @@ mean(const T1& X, const uword dim) +template +arma_warn_unused +arma_inline +typename enable_if2< is_arma_type::value, const Op >::result +mean(const T1& X, const uword dim, const elem_opts::omit_indicator&) + { + arma_debug_sigprint(); + + return Op(X, dim, uword(omit_mode)); + } + + + template arma_warn_unused arma_inline diff --git a/include/armadillo_bits/op_mean_bones.hpp b/include/armadillo_bits/op_mean_bones.hpp index 3977879c..4c334395 100644 --- a/include/armadillo_bits/op_mean_bones.hpp +++ b/include/armadillo_bits/op_mean_bones.hpp @@ -81,11 +81,17 @@ class op_mean_omit { public: - template - inline static eT direct_mean(const eT* X_mem, const uword N, const elem_opts::omit_indicator&); + template + inline static void apply(Mat& out, const Op& in); + + template + inline static void apply_noalias(Mat& out, const Mat& X, const uword dim, functor is_omitted); + + template + inline static eT direct_mean(const eT* X_mem, const uword N, functor is_omitted); template - inline static typename T1::elem_type mean_all(const Base& X, const elem_opts::omit_indicator& indicator); + inline static typename T1::elem_type mean_all(const Base& X, const elem_opts::omit_indicator&); }; diff --git a/include/armadillo_bits/op_mean_meat.hpp b/include/armadillo_bits/op_mean_meat.hpp index 035f7924..7df97230 100644 --- a/include/armadillo_bits/op_mean_meat.hpp +++ b/include/armadillo_bits/op_mean_meat.hpp @@ -280,7 +280,7 @@ op_mean::direct_mean(const eT* X_mem, const uword N) template inline eT -op_mean::direct_mean_robust(const eT old_mean, const eT* const X_mem, const uword N) +op_mean::direct_mean_robust(const eT old_mean, const eT* X_mem, const uword N) { arma_debug_sigprint(); @@ -374,20 +374,98 @@ op_mean::robust_mean(const std::complex& A, const std::complex& B) -template +template inline -eT -op_mean_omit::direct_mean(const eT* X_mem, const uword N, const elem_opts::omit_indicator&) +void +op_mean_omit::apply(Mat& out, const Op& 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), "mean(): 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_isnonfinite(x); }; + + const quasi_unwrap U(in.m); + + if(U.is_alias(out)) + { + Mat tmp; + + if(omit_mode == 1) { op_mean_omit::apply_noalias(tmp, U.M, dim, is_omitted_1); } + if(omit_mode == 2) { op_mean_omit::apply_noalias(tmp, U.M, dim, is_omitted_2); } + + out.steal_mem(tmp); + } + else + { + if(omit_mode == 1) { op_mean_omit::apply_noalias(out, U.M, dim, is_omitted_1); } + if(omit_mode == 2) { op_mean_omit::apply_noalias(out, U.M, dim, is_omitted_2); } + } + } + + + +template +inline +void +op_mean_omit::apply_noalias(Mat& out, const Mat& X, const uword dim, functor is_omitted) { arma_debug_sigprint(); typedef typename get_pod_type::result T; - auto is_omitted = [](const eT& x) -> bool + const uword X_n_rows = X.n_rows; + const uword X_n_cols = X.n_cols; + + if(dim == 0) { - if(omit_mode == 1) { return arma_isnan(x); } - if(omit_mode == 2) { return arma_isnonfinite(x); } - }; + out.set_size((X_n_rows > 0) ? 1 : 0, X_n_cols); + + if(X_n_rows == 0) { return; } + + eT* out_mem = out.memptr(); + + for(uword col=0; col < X_n_cols; ++col) + { + out_mem[col] = op_mean_omit::direct_mean( X.colptr(col), X_n_rows, is_omitted ); + } + } + else + if(dim == 1) + { + out.set_size(X_n_rows, (X_n_cols > 0) ? 1 : 0); + + if(X_n_cols == 0) { return; } + + eT* out_mem = out.memptr(); + + podarray tmp(X_n_cols, arma_nozeros_indicator()); + + for(uword row=0; row < X_n_rows; ++row) + { + tmp.copy_row(X, row); + + out_mem[row] = op_mean_omit::direct_mean(tmp.memptr(), X_n_cols, is_omitted); + } + } + } + + + +template +inline +eT +op_mean_omit::direct_mean(const eT* X_mem, const uword N, functor is_omitted) + { + arma_debug_sigprint(); + + typedef typename get_pod_type::result T; uword count = 0; @@ -404,8 +482,6 @@ op_mean_omit::direct_mean(const eT* X_mem, const uword N, const elem_opts::omit_ if( arma_isfinite(val) || (count == 0) ) { return val; } - if( (omit_mode == 1) && arrayops::has_inf(X_mem, N) ) { return val; } - arma_debug_print("op_mean_omit::direct_mean(): possible overflow; fallback to robust mean calculation"); podarray Y(N, arma_nozeros_indicator()); @@ -429,7 +505,7 @@ op_mean_omit::direct_mean(const eT* X_mem, const uword N, const elem_opts::omit_ template inline typename T1::elem_type -op_mean_omit::mean_all(const Base& X, const elem_opts::omit_indicator& indicator) +op_mean_omit::mean_all(const Base& X, const elem_opts::omit_indicator&) { arma_debug_sigprint(); @@ -447,7 +523,13 @@ op_mean_omit::mean_all(const Base& X, const elem_opt return Datum::nan; } - return op_mean_omit::direct_mean(A.memptr(), A_n_elem, indicator); + auto is_omitted = [](const eT& x) -> bool + { + if(omit_mode == 1) { return arma_isnan(x); } + if(omit_mode == 2) { return arma_isnonfinite(x); } + }; + + return op_mean_omit::direct_mean(A.memptr(), A_n_elem, is_omitted); }