From 7faec6fdcefa7980db6c3c4fc1b400cd9577ec94 Mon Sep 17 00:00:00 2001 From: conrad Date: Tue, 10 Jun 2025 19:30:19 +1000 Subject: [PATCH] simplifications --- include/armadillo_bits/op_mean_bones.hpp | 2 +- include/armadillo_bits/op_mean_meat.hpp | 32 +++++++++++++----------- 2 files changed, 19 insertions(+), 15 deletions(-) diff --git a/include/armadillo_bits/op_mean_bones.hpp b/include/armadillo_bits/op_mean_bones.hpp index 3d60b44b..6b9307d8 100644 --- a/include/armadillo_bits/op_mean_bones.hpp +++ b/include/armadillo_bits/op_mean_bones.hpp @@ -85,7 +85,7 @@ class op_mean_omit 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); + inline static eT direct_mean(const eT* X_mem, const uword N, functor is_omitted, podarray& work); template inline static typename T1::elem_type mean_all(const T1& 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 c4de13d0..7a076a74 100644 --- a/include/armadillo_bits/op_mean_meat.hpp +++ b/include/armadillo_bits/op_mean_meat.hpp @@ -97,7 +97,7 @@ op_mean::apply_noalias(Mat& out, const Mat& X, const uword dim) if(out.internal_has_nonfinite()) { - podarray tmp(X_n_cols, arma_nozeros_indicator()); + podarray tmp; for(uword row=0; row < X_n_rows; ++row) { @@ -107,7 +107,7 @@ op_mean::apply_noalias(Mat& out, const Mat& X, const uword dim) { tmp.copy_row(X, row); - out_mem[row] = op_mean::direct_mean_robust(old_mean, tmp.memptr(), X_n_cols); + out_mem[row] = op_mean::direct_mean_robust(old_mean, tmp.memptr(), tmp.n_elem); } } } @@ -203,7 +203,7 @@ op_mean::apply_noalias(Cube& out, const Cube& X, const uword dim) { const Mat tmp_mat('j', X.slice_memptr(slice), X_n_rows, X_n_cols); - podarray tmp_vec(X_n_cols, arma_nozeros_indicator()); + podarray tmp_vec; for(uword row=0; row < X_n_rows; ++row) { @@ -213,7 +213,7 @@ op_mean::apply_noalias(Cube& out, const Cube& X, const uword dim) { tmp_vec.copy(tmp_mat, row); - out_mem[row] = op_mean::direct_mean_robust(old_mean, tmp_vec, X_n_cols); + out_mem[row] = op_mean::direct_mean_robust(old_mean, tmp_vec.memptr(), tmp_vec.n_elem); } } } @@ -248,7 +248,7 @@ op_mean::apply_noalias(Cube& out, const Cube& X, const uword dim) { for(uword slice=0; slice < X_n_slices; ++slice) { tmp[slice] = X.at(row,col,slice); } - out.at(row,col,0) = op_mean::direct_mean_robust(mean, tmp.memptr(), X_n_slices); + out.at(row,col,0) = op_mean::direct_mean_robust(mean, tmp.memptr(), tmp.n_elem); } } } @@ -409,6 +409,8 @@ op_mean_omit::apply_noalias(Mat& out, const Mat& X, const uword dim, fun const uword X_n_rows = X.n_rows; const uword X_n_cols = X.n_cols; + podarray work; + if(dim == 0) { out.set_size((X_n_rows > 0) ? 1 : 0, X_n_cols); @@ -419,7 +421,7 @@ op_mean_omit::apply_noalias(Mat& out, const Mat& X, const uword dim, fun 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 ); + out_mem[col] = op_mean_omit::direct_mean(X.colptr(col), X_n_rows, is_omitted, work); } } else @@ -431,13 +433,13 @@ op_mean_omit::apply_noalias(Mat& out, const Mat& X, const uword dim, fun eT* out_mem = out.memptr(); - podarray tmp(X_n_cols, arma_nozeros_indicator()); + podarray tmp; 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); + out_mem[row] = op_mean_omit::direct_mean(tmp.memptr(), X_n_cols, is_omitted, work); } } } @@ -447,7 +449,7 @@ op_mean_omit::apply_noalias(Mat& out, const Mat& X, const uword dim, fun template inline eT -op_mean_omit::direct_mean(const eT* X_mem, const uword N, functor is_omitted) +op_mean_omit::direct_mean(const eT* X_mem, const uword N, functor is_omitted, podarray& work) { arma_debug_sigprint(); @@ -470,9 +472,9 @@ op_mean_omit::direct_mean(const eT* X_mem, const uword N, functor is_omitted) arma_debug_print("op_mean_omit::direct_mean(): possible overflow; fallback to robust mean calculation"); - podarray Y(N, arma_nozeros_indicator()); // TODO: it may be more efficient to declare Y outside of this function; amortise mem allocation penalty + work.set_size(N); - eT* Y_mem = Y.memptr(); + eT* work_mem = work.memptr(); count = 0; @@ -480,10 +482,10 @@ op_mean_omit::direct_mean(const eT* X_mem, const uword N, functor is_omitted) { const eT tmp = X_mem[i]; - if(is_omitted(tmp) == false) { Y_mem[count] = tmp; ++count; } + if(is_omitted(tmp) == false) { work_mem[count] = tmp; ++count; } } - return op_mean::direct_mean_robust(val, Y_mem, count); + return op_mean::direct_mean_robust(val, work_mem, count); } @@ -515,7 +517,9 @@ op_mean_omit::mean_all(const T1& X, const elem_opts::omit_indicator&) if(omit_mode == 2) { return arma_isnonfinite(x); } }; - return op_mean_omit::direct_mean(A.memptr(), A_n_elem, is_omitted); + podarray work; + + return op_mean_omit::direct_mean(A.memptr(), A_n_elem, is_omitted, work); }