MCPcopy Create free account
hub / github.com/BVLC/caffe / ShareTrainedLayersWith

Method ShareTrainedLayersWith

src/caffe/net.cpp:665–694  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

663
664template <typename Dtype>
665void Net<Dtype>::ShareTrainedLayersWith(const Net* other) {
666 int num_source_layers = other->layers().size();
667 for (int i = 0; i < num_source_layers; ++i) {
668 Layer<Dtype>* source_layer = other->layers()[i].get();
669 const string& source_layer_name = other->layer_names()[i];
670 int target_layer_id = 0;
671 while (target_layer_id != layer_names_.size() &&
672 layer_names_[target_layer_id] != source_layer_name) {
673 ++target_layer_id;
674 }
675 if (target_layer_id == layer_names_.size()) {
676 LOG(INFO) << "Ignoring source layer " << source_layer_name;
677 continue;
678 }
679 DLOG(INFO) << "Copying source layer " << source_layer_name;
680 vector<shared_ptr<Blob<Dtype> > >& target_blobs =
681 layers_[target_layer_id]->blobs();
682 CHECK_EQ(target_blobs.size(), source_layer->blobs().size())
683 << "Incompatible number of blobs for layer " << source_layer_name;
684 for (int j = 0; j < target_blobs.size(); ++j) {
685 Blob<Dtype>* source_blob = source_layer->blobs()[j].get();
686 CHECK(target_blobs[j]->shape() == source_blob->shape())
687 << "Cannot share param " << j << " weights from layer '"
688 << source_layer_name << "'; shape mismatch. Source param shape is "
689 << source_blob->shape_string() << "; target param shape is "
690 << target_blobs[j]->shape_string();
691 target_blobs[j]->ShareData(*source_blob);
692 }
693 }
694}
695
696template <typename Dtype>
697void Net<Dtype>::BackwardFrom(int start) {

Callers 2

share_weightsFunction · 0.80
TestMethod · 0.80

Calls 5

shapeMethod · 0.80
shape_stringMethod · 0.80
ShareDataMethod · 0.80
sizeMethod · 0.45
getMethod · 0.45

Tested by 1

TestMethod · 0.64