Add function to get the input size of a given network.
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user