| 30 | extern Barrier g_barrier; |
| 31 | |
| 32 | void MultiTrainer::Initialize(const TrainerDesc& trainer_desc, |
| 33 | Dataset* dataset) { |
| 34 | thread_num_ = trainer_desc.thread_num(); |
| 35 | SetDataset(dataset); |
| 36 | |
| 37 | ParseDumpConfig(trainer_desc); |
| 38 | mpi_rank_ = trainer_desc.mpi_rank(); |
| 39 | mpi_size_ = trainer_desc.mpi_size(); |
| 40 | dump_file_num_ = trainer_desc.dump_file_num(); |
| 41 | for (int i = 0; i < trainer_desc.downpour_param().stat_var_names_size(); |
| 42 | i++) { |
| 43 | need_merge_var_names_.push_back( |
| 44 | trainer_desc.downpour_param().stat_var_names(i)); |
| 45 | } |
| 46 | use_ps_gpu_ = trainer_desc.use_ps_gpu(); |
| 47 | use_gpu_graph_ = trainer_desc.use_gpu_graph(); |
| 48 | VLOG(3) << "Initialize use_ps_gpu_:" << use_ps_gpu_ |
| 49 | << "; use_gpu_graph_:" << use_gpu_graph_; |
| 50 | user_define_dump_filename_ = trainer_desc.user_define_dump_filename(); |
| 51 | // get filelist from trainer_desc here |
| 52 | const std::vector<paddle::framework::DataFeed*> readers = |
| 53 | dataset->GetReaders(); |
| 54 | VLOG(3) << "readers num: " << readers.size(); |
| 55 | // change thread num to readers num |
| 56 | thread_num_ = static_cast<int>(readers.size()); |
| 57 | VLOG(3) << "worker thread num: " << thread_num_; |
| 58 | workers_.resize(thread_num_); |
| 59 | |
| 60 | g_barrier.reset(thread_num_); |
| 61 | for (int i = 0; i < thread_num_; ++i) { |
| 62 | workers_[i] = DeviceWorkerFactory::CreateDeviceWorker( |
| 63 | trainer_desc.device_worker_name()); |
| 64 | workers_[i]->SetNeedDumpField(need_dump_field_); |
| 65 | workers_[i]->SetNeedDumpParam(need_dump_param_); |
| 66 | workers_[i]->SetDumpFieldVector(dump_fields_); |
| 67 | workers_[i]->SetDumpParamVector(dump_param_); |
| 68 | workers_[i]->InitRandomDumpConfig(trainer_desc); |
| 69 | workers_[i]->Initialize(trainer_desc); |
| 70 | workers_[i]->SetDeviceIndex(i); |
| 71 | workers_[i]->SetDataFeed(readers[i]); |
| 72 | workers_[i]->SetThreadNum(thread_num_); |
| 73 | } |
| 74 | |
| 75 | // set debug here |
| 76 | SetDebug(trainer_desc.debug()); |
| 77 | } |
| 78 | |
| 79 | std::string MultiTrainer::GetDumpPath(int tid) { |
| 80 | if (!user_define_dump_filename_.empty()) { |
nothing calls this directly
no test coverage detected