| 61 | |
| 62 | template <typename Dtype> |
| 63 | void Solver<Dtype>::InitTrainNet() { |
| 64 | const int num_train_nets = param_.has_net() + param_.has_net_param() + |
| 65 | param_.has_train_net() + param_.has_train_net_param(); |
| 66 | const string& field_names = "net, net_param, train_net, train_net_param"; |
| 67 | CHECK_GE(num_train_nets, 1) << "SolverParameter must specify a train net " |
| 68 | << "using one of these fields: " << field_names; |
| 69 | CHECK_LE(num_train_nets, 1) << "SolverParameter must not contain more than " |
| 70 | << "one of these fields specifying a train_net: " << field_names; |
| 71 | NetParameter net_param; |
| 72 | if (param_.has_train_net_param()) { |
| 73 | LOG_IF(INFO, Caffe::root_solver()) |
| 74 | << "Creating training net specified in train_net_param."; |
| 75 | net_param.CopyFrom(param_.train_net_param()); |
| 76 | } else if (param_.has_train_net()) { |
| 77 | LOG_IF(INFO, Caffe::root_solver()) |
| 78 | << "Creating training net from train_net file: " << param_.train_net(); |
| 79 | ReadNetParamsFromTextFileOrDie(param_.train_net(), &net_param); |
| 80 | } |
| 81 | if (param_.has_net_param()) { |
| 82 | LOG_IF(INFO, Caffe::root_solver()) |
| 83 | << "Creating training net specified in net_param."; |
| 84 | net_param.CopyFrom(param_.net_param()); |
| 85 | } |
| 86 | if (param_.has_net()) { |
| 87 | LOG_IF(INFO, Caffe::root_solver()) |
| 88 | << "Creating training net from net file: " << param_.net(); |
| 89 | ReadNetParamsFromTextFileOrDie(param_.net(), &net_param); |
| 90 | } |
| 91 | // Set the correct NetState. We start with the solver defaults (lowest |
| 92 | // precedence); then, merge in any NetState specified by the net_param itself; |
| 93 | // finally, merge in any NetState specified by the train_state (highest |
| 94 | // precedence). |
| 95 | NetState net_state; |
| 96 | net_state.set_phase(TRAIN); |
| 97 | net_state.MergeFrom(net_param.state()); |
| 98 | net_state.MergeFrom(param_.train_state()); |
| 99 | net_param.mutable_state()->CopyFrom(net_state); |
| 100 | net_.reset(new Net<Dtype>(net_param)); |
| 101 | } |
| 102 | |
| 103 | template <typename Dtype> |
| 104 | void Solver<Dtype>::InitTestNets() { |
nothing calls this directly
no test coverage detected