| 786 | |
| 787 | template <typename Dtype> |
| 788 | void 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 | |
| 837 | template <typename Dtype> |
| 838 | void Net<Dtype>::ToProto(NetParameter* param, bool write_diff) const { |
no test coverage detected