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

Method CopyTrainedLayersFromHDF5

src/caffe/net.cpp:788–835  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

786
787template <typename Dtype>
788void Net<Dtype>::CopyTrainedLayersFromHDF5(const string trained_filename) {
789 hid_t file_hid = H5Fopen(trained_filename.c_str(), H5F_ACC_RDONLY,
790 H5P_DEFAULT);
791 CHECK_GE(file_hid, 0) << "Couldn't open " << trained_filename;
792 hid_t data_hid = H5Gopen2(file_hid, "data", H5P_DEFAULT);
793 CHECK_GE(data_hid, 0) << "Error reading weights from " << trained_filename;
794 int num_layers = hdf5_get_num_links(data_hid);
795 for (int i = 0; i < num_layers; ++i) {
796 string source_layer_name = hdf5_get_name_by_idx(data_hid, i);
797 if (!layer_names_index_.count(source_layer_name)) {
798 LOG(INFO) << "Ignoring source layer " << source_layer_name;
799 continue;
800 }
801 int target_layer_id = layer_names_index_[source_layer_name];
802 DLOG(INFO) << "Copying source layer " << source_layer_name;
803 vector<shared_ptr<Blob<Dtype> > >& target_blobs =
804 layers_[target_layer_id]->blobs();
805 hid_t layer_hid = H5Gopen2(data_hid, source_layer_name.c_str(),
806 H5P_DEFAULT);
807 CHECK_GE(layer_hid, 0)
808 << "Error reading weights from " << trained_filename;
809 // Check that source layer doesn't have more params than target layer
810 int num_source_params = hdf5_get_num_links(layer_hid);
811 CHECK_LE(num_source_params, target_blobs.size())
812 << "Incompatible number of blobs for layer " << source_layer_name;
813 for (int j = 0; j < target_blobs.size(); ++j) {
814 ostringstream oss;
815 oss << j;
816 string dataset_name = oss.str();
817 int target_net_param_id = param_id_vecs_[target_layer_id][j];
818 if (!H5Lexists(layer_hid, dataset_name.c_str(), H5P_DEFAULT)) {
819 // Target param doesn't exist in source weights...
820 if (param_owners_[target_net_param_id] != -1) {
821 // ...but it's weight-shared in target, so that's fine.
822 continue;
823 } else {
824 LOG(FATAL) << "Incompatible number of blobs for layer "
825 << source_layer_name;
826 }
827 }
828 hdf5_load_nd_dataset(layer_hid, dataset_name.c_str(), 0, kMaxBlobAxes,
829 target_blobs[j].get());
830 }
831 H5Gclose(layer_hid);
832 }
833 H5Gclose(data_hid);
834 H5Fclose(file_hid);
835}
836
837template <typename Dtype>
838void Net<Dtype>::ToProto(NetParameter* param, bool write_diff) const {

Callers 1

Net_LoadHDF5Function · 0.80

Calls 5

hdf5_get_num_linksFunction · 0.85
hdf5_get_name_by_idxFunction · 0.85
countMethod · 0.80
sizeMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected