From 3e6578d7bf77795a31cd9cc199c21cf1384a872a Mon Sep 17 00:00:00 2001 From: "S. Zayd Enam" Date: Tue, 15 Jul 2014 02:45:43 +0000 Subject: [PATCH 1/2] Train a network with the top N layers from a saved state removed and replaced with different layers --- include/caffe/net.hpp | 2 +- include/caffe/solver.hpp | 8 ++++---- src/caffe/net.cpp | 4 ++-- src/caffe/solver.cpp | 20 +++++++++++++------- tools/train_net.cpp | 12 +++++++++--- 5 files changed, 29 insertions(+), 17 deletions(-) diff --git a/include/caffe/net.hpp b/include/caffe/net.hpp index cbd5becde2d..1b08f5b0239 100644 --- a/include/caffe/net.hpp +++ b/include/caffe/net.hpp @@ -64,7 +64,7 @@ class Net { void ShareTrainedLayersWith(Net* other); // For an already initialized net, CopyTrainedLayersFrom() copies the already // trained layers from another net parameter instance. - void CopyTrainedLayersFrom(const NetParameter& param); + void CopyTrainedLayersFrom(const NetParameter& param, const int remove_from_top = 0); void CopyTrainedLayersFrom(const string trained_filename); // Writes the net to a proto. void ToProto(NetParameter* param, bool write_diff = false); diff --git a/include/caffe/solver.hpp b/include/caffe/solver.hpp index 3112c59e0fc..a5cd46fa240 100644 --- a/include/caffe/solver.hpp +++ b/include/caffe/solver.hpp @@ -16,7 +16,7 @@ class Solver { void Init(const SolverParameter& param); // The main entry of the solver function. In default, iter will be zero. Pass // in a non-zero iter number to resume training for a pre-trained net. - virtual void Solve(const char* resume_file = NULL); + virtual void Solve(const char* resume_file = NULL, const int remove_from_top = 0); inline void Solve(const string resume_file) { Solve(resume_file.c_str()); } virtual ~Solver() {} inline shared_ptr > net() { return net_; } @@ -39,8 +39,8 @@ class Solver { // The Restore function implements how one should restore the solver to a // previously snapshotted state. You should implement the RestoreSolverState() // function that restores the state from a SolverState protocol buffer. - void Restore(const char* resume_file); - virtual void RestoreSolverState(const SolverState& state) = 0; + void Restore(const char* resume_file, const int remove_from_top = 0); + virtual void RestoreSolverState(const SolverState& state, const int remove_from_top = 0) = 0; SolverParameter param_; int iter_; @@ -64,7 +64,7 @@ class SGDSolver : public Solver { Dtype GetLearningRate(); virtual void ComputeUpdateValue(); virtual void SnapshotSolverState(SolverState * state); - virtual void RestoreSolverState(const SolverState& state); + virtual void RestoreSolverState(const SolverState& state, const int remove_from_top = 0); // history maintains the historical momentum data. vector > > history_; diff --git a/src/caffe/net.cpp b/src/caffe/net.cpp index c0f124b0ea7..c44cc27d55b 100644 --- a/src/caffe/net.cpp +++ b/src/caffe/net.cpp @@ -414,9 +414,9 @@ void Net::ShareTrainedLayersWith(Net* other) { } template -void Net::CopyTrainedLayersFrom(const NetParameter& param) { +void Net::CopyTrainedLayersFrom(const NetParameter& param, const int remove_from_top) { int num_source_layers = param.layers_size(); - for (int i = 0; i < num_source_layers; ++i) { + for (int i = 0; i < num_source_layers - remove_from_top; ++i) { const LayerParameter& source_layer = param.layers(i); const string& source_layer_name = source_layer.name(); int target_layer_id = 0; diff --git a/src/caffe/solver.cpp b/src/caffe/solver.cpp index 769618175ac..08c9d6db918 100644 --- a/src/caffe/solver.cpp +++ b/src/caffe/solver.cpp @@ -80,13 +80,19 @@ void Solver::Init(const SolverParameter& param) { } template -void Solver::Solve(const char* resume_file) { +void Solver::Solve(const char* resume_file, const int remove_from_top) { Caffe::set_phase(Caffe::TRAIN); LOG(INFO) << "Solving " << net_->name(); PreSolve(); iter_ = 0; - if (resume_file) { + if (resume_file && remove_from_top) { + LOG(INFO) << "Restoring previous solver status from " << resume_file; + LOG(INFO) << "Removing " << remove_from_top << " layers from top of network"; + Restore(resume_file, remove_from_top); + } + + else if (remove_from_top) { LOG(INFO) << "Restoring previous solver status from " << resume_file; Restore(resume_file); } @@ -202,16 +208,16 @@ void Solver::Snapshot() { } template -void Solver::Restore(const char* state_file) { +void Solver::Restore(const char* state_file, const int remove_from_top) { SolverState state; NetParameter net_param; ReadProtoFromBinaryFile(state_file, &state); if (state.has_learned_net()) { ReadProtoFromBinaryFile(state.learned_net().c_str(), &net_param); - net_->CopyTrainedLayersFrom(net_param); + net_->CopyTrainedLayersFrom(net_param, remove_from_top); } iter_ = state.iter(); - RestoreSolverState(state); + RestoreSolverState(state, remove_from_top); } @@ -331,8 +337,8 @@ void SGDSolver::SnapshotSolverState(SolverState* state) { } template -void SGDSolver::RestoreSolverState(const SolverState& state) { - CHECK_EQ(state.history_size(), history_.size()) +void SGDSolver::RestoreSolverState(const SolverState& state, const int remove_from_top) { + CHECK_EQ(state.history_size(), history_.size() + remove_from_top) << "Incorrect length of history blobs."; LOG(INFO) << "SGDSolver: restoring history"; for (int i = 0; i < history_.size(); ++i) { diff --git a/tools/train_net.cpp b/tools/train_net.cpp index 7c6f23e6240..25d9a0c29b4 100644 --- a/tools/train_net.cpp +++ b/tools/train_net.cpp @@ -15,8 +15,8 @@ using namespace caffe; // NOLINT(build/namespaces) int main(int argc, char** argv) { ::google::InitGoogleLogging(argv[0]); - if (argc < 2 || argc > 3) { - LOG(ERROR) << "Usage: train_net solver_proto_file [resume_point_file]"; + if (argc < 2 || argc > 4) { + LOG(ERROR) << "Usage: train_net solver_proto_file [resume_point_file] [num_remove_top_layers]"; return 1; } @@ -28,7 +28,13 @@ int main(int argc, char** argv) { if (argc == 3) { LOG(INFO) << "Resuming from " << argv[2]; solver.Solve(argv[2]); - } else { + } + else if (argc == 4) { + LOG(INFO) << "Resuming from " << argv[2]; + LOG(INFO) << "Removing top " << argv[3] << " layers"; + solver.Solve(argv[2], atoi(argv[3])); + } + else { solver.Solve(); } LOG(INFO) << "Optimization Done."; From c1e8f6b0bd0e1df2d5c9dbc5a31a8023cc88e4fd Mon Sep 17 00:00:00 2001 From: "S. Zayd Enam" Date: Tue, 15 Jul 2014 19:03:09 +0000 Subject: [PATCH 2/2] Revert Train a network with the top N layers from a saved state removed and replaced with different layers. Functionality already exists in Caffe --- include/caffe/net.hpp | 2 +- include/caffe/solver.hpp | 8 ++++---- src/caffe/net.cpp | 4 ++-- src/caffe/solver.cpp | 20 +++++++------------- tools/train_net.cpp | 12 +++--------- 5 files changed, 17 insertions(+), 29 deletions(-) diff --git a/include/caffe/net.hpp b/include/caffe/net.hpp index 1b08f5b0239..cbd5becde2d 100644 --- a/include/caffe/net.hpp +++ b/include/caffe/net.hpp @@ -64,7 +64,7 @@ class Net { void ShareTrainedLayersWith(Net* other); // For an already initialized net, CopyTrainedLayersFrom() copies the already // trained layers from another net parameter instance. - void CopyTrainedLayersFrom(const NetParameter& param, const int remove_from_top = 0); + void CopyTrainedLayersFrom(const NetParameter& param); void CopyTrainedLayersFrom(const string trained_filename); // Writes the net to a proto. void ToProto(NetParameter* param, bool write_diff = false); diff --git a/include/caffe/solver.hpp b/include/caffe/solver.hpp index a5cd46fa240..3112c59e0fc 100644 --- a/include/caffe/solver.hpp +++ b/include/caffe/solver.hpp @@ -16,7 +16,7 @@ class Solver { void Init(const SolverParameter& param); // The main entry of the solver function. In default, iter will be zero. Pass // in a non-zero iter number to resume training for a pre-trained net. - virtual void Solve(const char* resume_file = NULL, const int remove_from_top = 0); + virtual void Solve(const char* resume_file = NULL); inline void Solve(const string resume_file) { Solve(resume_file.c_str()); } virtual ~Solver() {} inline shared_ptr > net() { return net_; } @@ -39,8 +39,8 @@ class Solver { // The Restore function implements how one should restore the solver to a // previously snapshotted state. You should implement the RestoreSolverState() // function that restores the state from a SolverState protocol buffer. - void Restore(const char* resume_file, const int remove_from_top = 0); - virtual void RestoreSolverState(const SolverState& state, const int remove_from_top = 0) = 0; + void Restore(const char* resume_file); + virtual void RestoreSolverState(const SolverState& state) = 0; SolverParameter param_; int iter_; @@ -64,7 +64,7 @@ class SGDSolver : public Solver { Dtype GetLearningRate(); virtual void ComputeUpdateValue(); virtual void SnapshotSolverState(SolverState * state); - virtual void RestoreSolverState(const SolverState& state, const int remove_from_top = 0); + virtual void RestoreSolverState(const SolverState& state); // history maintains the historical momentum data. vector > > history_; diff --git a/src/caffe/net.cpp b/src/caffe/net.cpp index c44cc27d55b..c0f124b0ea7 100644 --- a/src/caffe/net.cpp +++ b/src/caffe/net.cpp @@ -414,9 +414,9 @@ void Net::ShareTrainedLayersWith(Net* other) { } template -void Net::CopyTrainedLayersFrom(const NetParameter& param, const int remove_from_top) { +void Net::CopyTrainedLayersFrom(const NetParameter& param) { int num_source_layers = param.layers_size(); - for (int i = 0; i < num_source_layers - remove_from_top; ++i) { + for (int i = 0; i < num_source_layers; ++i) { const LayerParameter& source_layer = param.layers(i); const string& source_layer_name = source_layer.name(); int target_layer_id = 0; diff --git a/src/caffe/solver.cpp b/src/caffe/solver.cpp index 08c9d6db918..769618175ac 100644 --- a/src/caffe/solver.cpp +++ b/src/caffe/solver.cpp @@ -80,19 +80,13 @@ void Solver::Init(const SolverParameter& param) { } template -void Solver::Solve(const char* resume_file, const int remove_from_top) { +void Solver::Solve(const char* resume_file) { Caffe::set_phase(Caffe::TRAIN); LOG(INFO) << "Solving " << net_->name(); PreSolve(); iter_ = 0; - if (resume_file && remove_from_top) { - LOG(INFO) << "Restoring previous solver status from " << resume_file; - LOG(INFO) << "Removing " << remove_from_top << " layers from top of network"; - Restore(resume_file, remove_from_top); - } - - else if (remove_from_top) { + if (resume_file) { LOG(INFO) << "Restoring previous solver status from " << resume_file; Restore(resume_file); } @@ -208,16 +202,16 @@ void Solver::Snapshot() { } template -void Solver::Restore(const char* state_file, const int remove_from_top) { +void Solver::Restore(const char* state_file) { SolverState state; NetParameter net_param; ReadProtoFromBinaryFile(state_file, &state); if (state.has_learned_net()) { ReadProtoFromBinaryFile(state.learned_net().c_str(), &net_param); - net_->CopyTrainedLayersFrom(net_param, remove_from_top); + net_->CopyTrainedLayersFrom(net_param); } iter_ = state.iter(); - RestoreSolverState(state, remove_from_top); + RestoreSolverState(state); } @@ -337,8 +331,8 @@ void SGDSolver::SnapshotSolverState(SolverState* state) { } template -void SGDSolver::RestoreSolverState(const SolverState& state, const int remove_from_top) { - CHECK_EQ(state.history_size(), history_.size() + remove_from_top) +void SGDSolver::RestoreSolverState(const SolverState& state) { + CHECK_EQ(state.history_size(), history_.size()) << "Incorrect length of history blobs."; LOG(INFO) << "SGDSolver: restoring history"; for (int i = 0; i < history_.size(); ++i) { diff --git a/tools/train_net.cpp b/tools/train_net.cpp index 25d9a0c29b4..7c6f23e6240 100644 --- a/tools/train_net.cpp +++ b/tools/train_net.cpp @@ -15,8 +15,8 @@ using namespace caffe; // NOLINT(build/namespaces) int main(int argc, char** argv) { ::google::InitGoogleLogging(argv[0]); - if (argc < 2 || argc > 4) { - LOG(ERROR) << "Usage: train_net solver_proto_file [resume_point_file] [num_remove_top_layers]"; + if (argc < 2 || argc > 3) { + LOG(ERROR) << "Usage: train_net solver_proto_file [resume_point_file]"; return 1; } @@ -28,13 +28,7 @@ int main(int argc, char** argv) { if (argc == 3) { LOG(INFO) << "Resuming from " << argv[2]; solver.Solve(argv[2]); - } - else if (argc == 4) { - LOG(INFO) << "Resuming from " << argv[2]; - LOG(INFO) << "Removing top " << argv[3] << " layers"; - solver.Solve(argv[2], atoi(argv[3])); - } - else { + } else { solver.Solve(); } LOG(INFO) << "Optimization Done.";