From d0708a5f6d437edbbb574cdac599b58eccc259ba Mon Sep 17 00:00:00 2001 From: Marcus Edel Date: Fri, 8 Apr 2016 15:43:32 +0200 Subject: [PATCH] Add function to get the input size of a given network. --- src/mlpack/methods/ann/network_util.hpp | 32 ++++++++++++++++ src/mlpack/methods/ann/network_util_impl.hpp | 40 +++++++++++++++++++- 2 files changed, 71 insertions(+), 1 deletion(-) diff --git a/src/mlpack/methods/ann/network_util.hpp b/src/mlpack/methods/ann/network_util.hpp index c0b222788e..193d5a8d3f 100644 --- a/src/mlpack/methods/ann/network_util.hpp +++ b/src/mlpack/methods/ann/network_util.hpp @@ -140,6 +140,38 @@ template typename std::enable_if< !HasGradientCheck::value, size_t>::type LayerGradients(T& layer, arma::mat& gradients, size_t offset, P& output); + +/** + * Auxiliary function to get the input size of the specified network. + * + * @param network The network used for specifying the input size. + * @return The input size. + */ +template +typename std::enable_if::type +NetworkInputSize(std::tuple& network); + +template +typename std::enable_if::type +NetworkInputSize(std::tuple& network); + +/** + * Auxiliary function to get the input size of the specified layer. + * + * @param layer The layer used for specifying the input size. + * @param output The layer output parameter. + * @return The input size. + */ +template +typename std::enable_if< + !HasWeightsCheck::value, size_t>::type +LayerInputSize(T& layer, P& output); + +template +typename std::enable_if< + HasWeightsCheck::value, size_t>::type +LayerInputSize(T& layer, P& output); + } // namespace ann } // namespace mlpack diff --git a/src/mlpack/methods/ann/network_util_impl.hpp b/src/mlpack/methods/ann/network_util_impl.hpp index fbb404147e..34cc84de40 100644 --- a/src/mlpack/methods/ann/network_util_impl.hpp +++ b/src/mlpack/methods/ann/network_util_impl.hpp @@ -34,7 +34,7 @@ typename std::enable_if< HasWeightsCheck::value, size_t>::type LayerSize(T& layer, P& /* unused */) { - return layer.Weights().n_elem; + return layer.Weights().n_elem; } template @@ -166,6 +166,44 @@ LayerGradients(T& /* unused */, return 0; } +template +typename std::enable_if::type +NetworkInputSize(std::tuple& /* unused */) +{ + return 0; +} + +template +typename std::enable_if::type +NetworkInputSize(std::tuple& network) +{ + const size_t inputSize = LayerInputSize(std::get(network), std::get( + network).OutputParameter()); + + if (inputSize) + { + return inputSize; + } + + return NetworkInputSize(network); +} + +template +typename std::enable_if< + HasWeightsCheck::value, size_t>::type +LayerInputSize(T& layer, P& /* unused */) +{ + return layer.Weights().n_cols; +} + +template +typename std::enable_if< + !HasWeightsCheck::value, size_t>::type +LayerInputSize(T& /* unused */, P& /* unused */) +{ + return 0; +} + } // namespace ann } // namespace mlpack