diff --git a/src/mlpack/bindings/cli/add_to_cli11.hpp b/src/mlpack/bindings/cli/add_to_cli11.hpp index f4ae48a608..ceb03c64e2 100644 --- a/src/mlpack/bindings/cli/add_to_cli11.hpp +++ b/src/mlpack/bindings/cli/add_to_cli11.hpp @@ -48,7 +48,7 @@ void AddToCLI11(const std::string& cliName, { using TupleType = std::tuple::type>; TupleType& tuple = *boost::any_cast(¶m.value); - std::get<1>(tuple) = boost::any_cast(value); + std::get<0>(std::get<1>(tuple)) = boost::any_cast(value); param.wasPassed = true; }, param.desc.c_str()); @@ -110,7 +110,7 @@ void AddToCLI11(const std::string& cliName, { using TupleType = std::tuple::type>; TupleType& tuple = *boost::any_cast(¶m.value); - std::get<1>(tuple) = boost::any_cast(value); + std::get<0>(std::get<1>(tuple)) = boost::any_cast(value); param.wasPassed = true; }, param.desc.c_str()); diff --git a/src/mlpack/bindings/cli/get_param.hpp b/src/mlpack/bindings/cli/get_param.hpp index eaaa813d61..d401e0e554 100644 --- a/src/mlpack/bindings/cli/get_param.hpp +++ b/src/mlpack/bindings/cli/get_param.hpp @@ -53,8 +53,10 @@ T& GetParam( // happens. typedef std::tuple::type> TupleType; TupleType& tuple = *boost::any_cast(&d.value); - const std::string& value = std::get<1>(tuple); + const std::string& value = std::get<0>(std::get<1>(tuple)); T& matrix = std::get<0>(tuple); + size_t& n_rows = std::get<1>(std::get<1>(tuple)); + size_t& n_cols = std::get<2>(std::get<1>(tuple)); if (d.input && !d.loaded) { // Call correct data::Load() function. @@ -62,6 +64,8 @@ T& GetParam( data::Load(value, matrix, true); else data::Load(value, matrix, true, !d.noTranspose); + n_rows = matrix.n_rows; + n_cols = matrix.n_cols; d.loaded = true; } @@ -81,13 +85,17 @@ T& GetParam( { // If this is an input parameter, we need to load both the matrix and the // dataset info. - typedef std::tuple TupleType; + typedef std::tuple> TupleType; TupleType* tuple = boost::any_cast(&d.value); - const std::string& value = std::get<1>(*tuple); + const std::string& value = std::get<0>(std::get<1>(*tuple)); T& t = std::get<0>(*tuple); + size_t& n_rows = std::get<1>(std::get<1>(*tuple)); + size_t& n_cols = std::get<2>(std::get<1>(*tuple)); if (d.input && !d.loaded) { data::Load(value, std::get<1>(t), std::get<0>(t), true, !d.noTranspose); + n_rows = std::get<1>(t).n_rows; + n_cols = std::get<1>(t).n_cols; d.loaded = true; } diff --git a/src/mlpack/bindings/cli/get_printable_param_impl.hpp b/src/mlpack/bindings/cli/get_printable_param_impl.hpp index 9e2584f6b5..8a3cb211cb 100644 --- a/src/mlpack/bindings/cli/get_printable_param_impl.hpp +++ b/src/mlpack/bindings/cli/get_printable_param_impl.hpp @@ -83,13 +83,14 @@ std::string GetPrintableParam( const TupleType* tuple = boost::any_cast(&data.value); std::ostringstream oss; - oss << "'" << std::get<1>(*tuple) << "'"; + oss << "'" << std::get<0>(std::get<1>(*tuple)) << "'"; - if (std::get<1>(*tuple) != "") + if (std::get<0>(std::get<1>(*tuple)) != "") { // Make sure the matrix is loaded so that we can print its size. - T& mat = GetParam(const_cast(data)); - std::string matDescription = GetMatrixSize(mat); + GetParam(const_cast(data)); + std::string matDescription = std::to_string(std::get<1>(std::get<1>(*tuple))) + "x"; + matDescription += std::to_string(std::get<2>(std::get<1>(*tuple))) + " matrix"; oss << " (" << matDescription << ")"; } diff --git a/src/mlpack/bindings/cli/get_raw_param.hpp b/src/mlpack/bindings/cli/get_raw_param.hpp index 35d544f08b..46c3956a1f 100644 --- a/src/mlpack/bindings/cli/get_raw_param.hpp +++ b/src/mlpack/bindings/cli/get_raw_param.hpp @@ -48,7 +48,7 @@ T& GetRawParam( arma::mat>>::value>::type* = 0) { // Don't load the matrix. - typedef std::tuple TupleType; + typedef std::tuple> TupleType; T& value = std::get<0>(*boost::any_cast(&d.value)); return value; } diff --git a/src/mlpack/bindings/cli/in_place_copy.hpp b/src/mlpack/bindings/cli/in_place_copy.hpp index ca0d7667f5..e2094e53d4 100644 --- a/src/mlpack/bindings/cli/in_place_copy.hpp +++ b/src/mlpack/bindings/cli/in_place_copy.hpp @@ -41,7 +41,7 @@ void InPlaceCopyInternal( /** * Modify the filename for any type that needs to be loaded from disk to match - * the filename of the input parameter. + * the filename of the input parameter. For matrix/datasetinfo parameter. * * @param d ParamData object we want to make into an in-place copy. * @param input ParamData object whose filename we should copy. @@ -53,12 +53,34 @@ void InPlaceCopyInternal( const typename std::enable_if< arma::is_arma_type::value || std::is_same>::value || - data::HasSerialize::value>::type* = 0) + std::tuple>::value>::type* = 0) { // Make the output filename the same as the input filename. typedef std::tuple::type> TupleType; TupleType& tuple = *boost::any_cast(&d.value); + std::string& value = std::get<0>(std::get<1>(tuple)); + + const TupleType& inputTuple = *boost::any_cast(&input.value); + value = std::get<0>(std::get<1>(inputTuple)); +} + +/** + * Modify the filename for any type that needs to be loaded from disk to match + * the filename of the input parameter. For Serializable object. + * + * @param d ParamData object we want to make into an in-place copy. + * @param input ParamData object whose filename we should copy. + */ +template +void InPlaceCopyInternal( + util::ParamData& d, + util::ParamData& input, + const typename std::enable_if< + data::HasSerialize::value>::type* = 0) +{ + // Make the output filename the same as the input filename. + typedef std::tuple::type> TupleType; + TupleType& tuple = *boost::any_cast(&d.value); std::string& value = std::get<1>(tuple); const TupleType& inputTuple = *boost::any_cast(&input.value); diff --git a/src/mlpack/bindings/cli/output_param_impl.hpp b/src/mlpack/bindings/cli/output_param_impl.hpp index c4dcfddbf4..ab2f1e8822 100644 --- a/src/mlpack/bindings/cli/output_param_impl.hpp +++ b/src/mlpack/bindings/cli/output_param_impl.hpp @@ -53,10 +53,10 @@ void OutputParamImpl( util::ParamData& data, const typename boost::enable_if>::type* /* junk */) { - typedef std::tuple TupleType; + typedef std::tuple> TupleType; const T& output = std::get<0>(*boost::any_cast(&data.value)); const std::string& filename = - std::get<1>(*boost::any_cast(&data.value)); + std::get<0>(std::get<1>(*boost::any_cast(&data.value))); if (output.n_elem > 0 && filename != "") { @@ -95,10 +95,10 @@ void OutputParamImpl( std::tuple>>::type* /* junk */) { // Output the matrix with the mappings. - typedef std::tuple TupleType; + typedef std::tuple> TupleType; const T& tuple = std::get<0>(*boost::any_cast(&data.value)); const std::string& filename = - std::get<1>(*boost::any_cast(&data.value)); + std::get<0>(std::get<1>(*boost::any_cast(&data.value))); const arma::mat& matrix = std::get<1>(tuple); // The mapping isn't taken into account. We should write a data::Save() diff --git a/src/mlpack/bindings/cli/parameter_type.hpp b/src/mlpack/bindings/cli/parameter_type.hpp index ddc289ae32..036240b375 100644 --- a/src/mlpack/bindings/cli/parameter_type.hpp +++ b/src/mlpack/bindings/cli/parameter_type.hpp @@ -53,7 +53,7 @@ struct ParameterType template struct ParameterType> { - typedef std::string type; + typedef std::tuple type; }; /** @@ -65,7 +65,7 @@ struct ParameterType> template struct ParameterType> { - typedef std::string type; + typedef std::tuple type; }; /** @@ -76,7 +76,7 @@ struct ParameterType> template struct ParameterType> { - typedef std::string type; + typedef std::tuple type; }; /** @@ -86,7 +86,7 @@ template struct ParameterType, arma::Mat>> { - typedef std::string type; + typedef std::tuple type; }; } // namespace cli diff --git a/src/mlpack/bindings/cli/set_param.hpp b/src/mlpack/bindings/cli/set_param.hpp index 8295b51242..f800fe0553 100644 --- a/src/mlpack/bindings/cli/set_param.hpp +++ b/src/mlpack/bindings/cli/set_param.hpp @@ -51,8 +51,8 @@ void SetParam( } /** - * Set a matrix parameter, a matrix/dataset info parameter, or a serializable - * object. These set the filename referring to the parameter. + * Set a matrix parameter, a matrix/dataset info parameter. + * These set the filename referring to the parameter. */ template void SetParam( @@ -65,7 +65,7 @@ void SetParam( // We're setting the string filename. typedef std::tuple::type> TupleType; TupleType& tuple = *boost::any_cast(&d.value); - std::get<1>(tuple) = boost::any_cast(value); + std::get<0>(std::get<1>(tuple)) = boost::any_cast(value); } /**