From 89e77038755375f213bdbc9eeaa56fc5a032d40f Mon Sep 17 00:00:00 2001 From: conrad Date: Mon, 28 Jul 2025 12:56:45 +1000 Subject: [PATCH] specialisation for fp16 --- include/armadillo_bits/op_mean_bones.hpp | 9 +++ include/armadillo_bits/op_mean_meat.hpp | 73 ++++++++++++++++++++++++ 2 files changed, 82 insertions(+) diff --git a/include/armadillo_bits/op_mean_bones.hpp b/include/armadillo_bits/op_mean_bones.hpp index a07ee979..bbaeee64 100644 --- a/include/armadillo_bits/op_mean_bones.hpp +++ b/include/armadillo_bits/op_mean_bones.hpp @@ -32,6 +32,10 @@ struct op_mean template inline static void apply_noalias(Mat& out, const Mat& X, const uword dim); + #if defined(ARMA_HAVE_FP16) + // template<> + inline static void apply_noalias(Mat& out, const Mat& X, const uword dim); + #endif // cubes @@ -46,6 +50,11 @@ struct op_mean template inline static eT direct_mean(const eT* X_mem, const uword N); + #if defined(ARMA_HAVE_FP16) + // template<> + inline static fp16 direct_mean(const fp16* X_mem, const uword N); + #endif + template inline static eT direct_mean_robust(const eT old_mean, const eT* X_mem, const uword N); diff --git a/include/armadillo_bits/op_mean_meat.hpp b/include/armadillo_bits/op_mean_meat.hpp index 6f7937a2..67d28e2c 100644 --- a/include/armadillo_bits/op_mean_meat.hpp +++ b/include/armadillo_bits/op_mean_meat.hpp @@ -114,6 +114,52 @@ op_mean::apply_noalias(Mat& out, const Mat& X, const uword dim) +#if defined(ARMA_HAVE_FP16) +inline +void +op_mean::apply_noalias(Mat& out, const Mat& X, const uword dim) + { + arma_debug_sigprint(); + + const uword X_n_rows = X.n_rows; + const uword X_n_cols = X.n_cols; + + if(dim == 0) + { + out.set_size((X_n_rows > 0) ? 1 : 0, X_n_cols); + + if(X_n_rows == 0) { return; } + + fp16* out_mem = out.memptr(); + + for(uword col=0; col < X_n_cols; ++col) + { + out_mem[col] = op_mean::direct_mean( X.colptr(col), X_n_rows ); + } + } + else + if(dim == 1) + { + out.set_size(X_n_rows, (X_n_cols > 0) ? 1 : 0); + + if(X_n_cols == 0) { return; } + + fp16* out_mem = out.memptr(); + + podarray tmp; + + for(uword row=0; row < X_n_rows; ++row) + { + tmp.copy_row(X, row); + + out_mem[row] = op_mean::direct_mean( tmp.memptr(), tmp.n_elem ); + } + } + } +#endif + + + // @@ -274,6 +320,33 @@ op_mean::direct_mean(const eT* X_mem, const uword N) +#if defined(ARMA_HAVE_FP16) +inline +fp16 +op_mean::direct_mean(const fp16* X_mem, const uword N) + { + arma_debug_sigprint(); + + float acc = float(0); + + for(uword i=0; i float(std::numeric_limits::max() )) { return std::numeric_limits::infinity(); } + if(mean < float(std::numeric_limits::lowest())) { return fp16(-1) * std::numeric_limits::infinity(); } + + return fp16(mean); + } +#endif + + + template inline eT