From f599f7cf8b012f7edee6e498cf2e6b6ba57ffa83 Mon Sep 17 00:00:00 2001 From: akhandait Date: Sat, 7 Jul 2018 16:19:02 +0530 Subject: [PATCH] Overload Evaluate() function. --- src/mlpack/methods/ann/ffn.hpp | 9 +++++++++ src/mlpack/methods/ann/ffn_impl.hpp | 27 +++++++++++++++++++++++++++ 2 files changed, 36 insertions(+) diff --git a/src/mlpack/methods/ann/ffn.hpp b/src/mlpack/methods/ann/ffn.hpp index 7c826ae65f..8c9a8356e4 100644 --- a/src/mlpack/methods/ann/ffn.hpp +++ b/src/mlpack/methods/ann/ffn.hpp @@ -134,6 +134,15 @@ class FFN */ void Predict(arma::mat predictors, arma::mat& results); + /** + * Evaluate the feedforward network with the given ppredictors and responses. + * This functions is usually used to monitor progress while training. + * + * @param predictors Input variables. + * @param responses Target outputs for input variables. + */ + double Evaluate(arma::mat predictors, arma::mat responses); + /** * Evaluate the feedforward network with the given parameters. This function * is usually called by the optimizer to train the model. diff --git a/src/mlpack/methods/ann/ffn_impl.hpp b/src/mlpack/methods/ann/ffn_impl.hpp index b180b9b0ec..03cfbae1c1 100644 --- a/src/mlpack/methods/ann/ffn_impl.hpp +++ b/src/mlpack/methods/ann/ffn_impl.hpp @@ -202,6 +202,33 @@ void FFN::Predict( results.col(i) = resultsTemp.col(0); } } +template +double FFN::Evaluate( + arma::mat predictors, arma::mat responses) +{ + if (parameter.is_empty()) + ResetParameters(); + + if (!deterministic) + { + deterministic = true; + ResetDeterministic(); + } + + Forward(std::move(predictors)); + + double res = outputLayer.Forward( + std::move(boost::apply_visitor(outputParameterVisitor, network.back())), + std::move(responses)); + + for (size_t i = 0; i < network.size(); ++i) + { + res += boost::apply_visitor(lossVisitor, network[i]); + } + + return res; +} template