From c8bff330c2d4032fbb41c2a7290064315bdad8eb Mon Sep 17 00:00:00 2001 From: Ryan Curtin Date: Fri, 16 Dec 2011 14:32:05 +0000 Subject: [PATCH] Give a few specializations for vectors, and remove using directives. --- src/mlpack/core/util/save_restore_utility.cpp | 48 +++++++++++++++---- src/mlpack/core/util/save_restore_utility.hpp | 37 +++++++++----- .../core/util/save_restore_utility_impl.hpp | 19 +++++--- 3 files changed, 77 insertions(+), 27 deletions(-) diff --git a/src/mlpack/core/util/save_restore_utility.cpp b/src/mlpack/core/util/save_restore_utility.cpp index d0090e005c..75f60febb1 100644 --- a/src/mlpack/core/util/save_restore_utility.cpp +++ b/src/mlpack/core/util/save_restore_utility.cpp @@ -4,15 +4,13 @@ * * The SaveRestoreUtility provides helper functions in saving and * restoring models. The current output file type is XML. - * - * @experimental */ #include "save_restore_utility.hpp" using namespace mlpack; using namespace utilities; -bool SaveRestoreUtility::ReadFile(std::string filename) +bool SaveRestoreUtility::ReadFile(const std::string filename) { xmlDocPtr xmlDocTree = NULL; if (NULL == (xmlDocTree = xmlReadFile(filename.c_str(), NULL, 0))) @@ -26,6 +24,7 @@ bool SaveRestoreUtility::ReadFile(std::string filename) xmlFreeDoc(xmlDocTree); return true; } + void SaveRestoreUtility::RecurseOnNodes(xmlNode* n) { xmlNodePtr current = NULL; @@ -40,7 +39,8 @@ void SaveRestoreUtility::RecurseOnNodes(xmlNode* n) RecurseOnNodes(current->children); } } -bool SaveRestoreUtility::WriteFile(std::string filename) + +bool SaveRestoreUtility::WriteFile(const std::string filename) { bool success = false; xmlDocPtr xmlDocTree = xmlNewDoc(BAD_CAST "1.0"); @@ -64,7 +64,9 @@ bool SaveRestoreUtility::WriteFile(std::string filename) xmlFreeDoc(xmlDocTree); return success; } -arma::mat& SaveRestoreUtility::LoadParameter(arma::mat& matrix, std::string name) + +arma::mat& SaveRestoreUtility::LoadParameter(arma::mat& matrix, + const std::string name) { std::map::iterator it = parameters.find(name); if (it != parameters.end()) @@ -121,7 +123,9 @@ arma::mat& SaveRestoreUtility::LoadParameter(arma::mat& matrix, std::string name } return matrix; } -std::string SaveRestoreUtility::LoadParameter(std::string str, std::string name) + +std::string SaveRestoreUtility::LoadParameter(std::string str, + const std::string name) { std::map::iterator it = parameters.find(name); if (it != parameters.end()) @@ -134,7 +138,8 @@ std::string SaveRestoreUtility::LoadParameter(std::string str, std::string name) } return ""; } -char SaveRestoreUtility::LoadParameter(char c, std::string name) + +char SaveRestoreUtility::LoadParameter(char c, const std::string name) { int temp; std::map::iterator it = parameters.find(name); @@ -151,14 +156,17 @@ char SaveRestoreUtility::LoadParameter(char c, std::string name) } return 0; } -void SaveRestoreUtility::SaveParameter(char c, std::string name) + +void SaveRestoreUtility::SaveParameter(const char c, const std::string name) { int temp = (int) c; std::ostringstream output; output << temp; parameters[name] = output.str(); } -void SaveRestoreUtility::SaveParameter(arma::mat& mat, std::string name) + +void SaveRestoreUtility::SaveParameter(const arma::mat& mat, + const std::string name) { std::ostringstream output; size_t columns = mat.n_cols; @@ -173,3 +181,25 @@ void SaveRestoreUtility::SaveParameter(arma::mat& mat, std::string name) } parameters[name] = output.str(); } + +// Special template specializations for vectors. + +namespace mlpack { +namespace utilities { + +template<> +arma::vec& SaveRestoreUtility::LoadParameter(arma::vec& t, + const std::string name) +{ + return (arma::vec&) LoadParameter((arma::mat&) t, name); +} + +template<> +void SaveRestoreUtility::SaveParameter(const arma::vec& t, + const std::string name) +{ + SaveParameter((const arma::mat&) t, name); +} + +}; // namespace utilities +}; // namespace mlpack diff --git a/src/mlpack/core/util/save_restore_utility.hpp b/src/mlpack/core/util/save_restore_utility.hpp index 5ee0aadfe9..f851799eca 100644 --- a/src/mlpack/core/util/save_restore_utility.hpp +++ b/src/mlpack/core/util/save_restore_utility.hpp @@ -7,7 +7,6 @@ * * @experimental */ - #ifndef SAVE_RESTORE_MODEL_HPP #define SAVE_RESTORE_MODEL_HPP @@ -45,66 +44,80 @@ class SaveRestoreUtility /** * ReadFile reads an XML tree from a file. */ - bool ReadFile(std::string filename); + bool ReadFile(const std::string filename); /** * WriteFile writes the XML tree to a file. */ - bool WriteFile(std::string filename); + bool WriteFile(const std::string filename); /** * LoadParameter loads a parameter from the parameters map. */ template - T& LoadParameter(T& t, std::string name); + T& LoadParameter(T& t, const std::string name); /** * LoadParameter loads a parameter from the parameters map. */ template - std::vector& LoadParameter(std::vector& v, std::string name); + std::vector& LoadParameter(std::vector& v, const std::string name); /** * LoadParameter loads a character from the parameters map. */ - char LoadParameter(char c, std::string name); + char LoadParameter(char c, const std::string name); /** * LoadParameter loads a string from the parameters map. */ - std::string LoadParameter(std::string str, std::string name); + std::string LoadParameter(std::string str, const std::string name); /** * LoadParameter loads an arma::mat from the parameters map. */ - arma::mat& LoadParameter(arma::mat& matrix, std::string name); + arma::mat& LoadParameter(arma::mat& matrix, const std::string name); /** * SaveParameter saves a parameter to the parameters map. */ template - void SaveParameter(T& t, std::string name); + void SaveParameter(const T& t, const std::string name); + + /** * SaveParameter saves a parameter to the parameters map. */ template - void SaveParameter(std::vector& v, std::string name); + void SaveParameter(const std::vector& v, const std::string name); /** * SaveParameter saves a character to the parameters map. */ - void SaveParameter(char c, std::string name); + void SaveParameter(const char c, const std::string name); /** * SaveParameter saves an arma::mat to the parameters map. */ - void SaveParameter(arma::mat& mat, std::string name); + void SaveParameter(const arma::mat& mat, const std::string name); }; +//! Specialization for arma::vec. +template<> +arma::vec& SaveRestoreUtility::LoadParameter(arma::vec& t, + const std::string name); + +//! Specialization for arma::vec. +template<> +void SaveRestoreUtility::SaveParameter(const arma::vec& t, + const std::string name); + } /* namespace utilities */ } /* namespace mlpack */ #include "save_restore_utility_impl.hpp" + + #endif diff --git a/src/mlpack/core/util/save_restore_utility_impl.hpp b/src/mlpack/core/util/save_restore_utility_impl.hpp index 98ded27195..105abb429b 100644 --- a/src/mlpack/core/util/save_restore_utility_impl.hpp +++ b/src/mlpack/core/util/save_restore_utility_impl.hpp @@ -11,11 +11,11 @@ // In case it hasn't been included already. #include "save_restore_utility.hpp" -using namespace mlpack; -using namespace mlpack::utilities; +namespace mlpack { +namespace utilities { template -T& SaveRestoreUtility::LoadParameter(T& t, std::string name) +T& SaveRestoreUtility::LoadParameter(T& t, const std::string name) { std::map::iterator it = parameters.find(name); if (it != parameters.end()) @@ -31,8 +31,10 @@ T& SaveRestoreUtility::LoadParameter(T& t, std::string name) } return t; } + template -std::vector& SaveRestoreUtility::LoadParameter(std::vector& v, std::string name) +std::vector& SaveRestoreUtility::LoadParameter(std::vector& v, + const std::string name) { std::map::iterator it = parameters.find(name); if (it != parameters.end()) @@ -59,8 +61,9 @@ std::vector& SaveRestoreUtility::LoadParameter(std::vector& v, std::string } return v; } + template -void SaveRestoreUtility::SaveParameter(T& t, std::string name) +void SaveRestoreUtility::SaveParameter(const T& t, const std::string name) { std::ostringstream output; output << t; @@ -68,7 +71,8 @@ void SaveRestoreUtility::SaveParameter(T& t, std::string name) } template -void SaveRestoreUtility::SaveParameter(std::vector& t, std::string name) +void SaveRestoreUtility::SaveParameter(const std::vector& t, + const std::string name) { std::ostringstream output; for (size_t index = 0; index < t.size(); ++index) @@ -80,4 +84,7 @@ void SaveRestoreUtility::SaveParameter(std::vector& t, std::string name) parameters[name] = vectorAsStr; } +}; // namespace utilities +}; // namespace mlpack + #endif