Add function to get the input size of a given network.

This commit is contained in:
Marcus Edel
2016-04-09 13:31:20 +02:00
parent ba826b1959
commit d0708a5f6d
2 changed files with 71 additions and 1 deletions
+32
View File
@@ -140,6 +140,38 @@ template<typename T, typename P>
typename std::enable_if<
!HasGradientCheck<T, P&(T::*)()>::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<size_t I = 0, typename... Tp>
typename std::enable_if<I < sizeof...(Tp), size_t>::type
NetworkInputSize(std::tuple<Tp...>& network);
template<size_t I, typename... Tp>
typename std::enable_if<I == sizeof...(Tp), size_t>::type
NetworkInputSize(std::tuple<Tp...>& 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 T, typename P>
typename std::enable_if<
!HasWeightsCheck<T, P&(T::*)()>::value, size_t>::type
LayerInputSize(T& layer, P& output);
template<typename T, typename P>
typename std::enable_if<
HasWeightsCheck<T, P&(T::*)()>::value, size_t>::type
LayerInputSize(T& layer, P& output);
} // namespace ann
} // namespace mlpack
+39 -1
View File
@@ -34,7 +34,7 @@ typename std::enable_if<
HasWeightsCheck<T, P&(T::*)()>::value, size_t>::type
LayerSize(T& layer, P& /* unused */)
{
return layer.Weights().n_elem;
return layer.Weights().n_elem;
}
template<typename T, typename P>
@@ -166,6 +166,44 @@ LayerGradients(T& /* unused */,
return 0;
}
template<size_t I, typename... Tp>
typename std::enable_if<I == sizeof...(Tp), size_t>::type
NetworkInputSize(std::tuple<Tp...>& /* unused */)
{
return 0;
}
template<size_t I, typename... Tp>
typename std::enable_if<I < sizeof...(Tp), size_t>::type
NetworkInputSize(std::tuple<Tp...>& network)
{
const size_t inputSize = LayerInputSize(std::get<I>(network), std::get<I>(
network).OutputParameter());
if (inputSize)
{
return inputSize;
}
return NetworkInputSize<I + 1, Tp...>(network);
}
template<typename T, typename P>
typename std::enable_if<
HasWeightsCheck<T, P&(T::*)()>::value, size_t>::type
LayerInputSize(T& layer, P& /* unused */)
{
return layer.Weights().n_cols;
}
template<typename T, typename P>
typename std::enable_if<
!HasWeightsCheck<T, P&(T::*)()>::value, size_t>::type
LayerInputSize(T& /* unused */, P& /* unused */)
{
return 0;
}
} // namespace ann
} // namespace mlpack