MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / Initialize

Method Initialize

paddle/fluid/framework/multi_trainer.cc:32–77  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

30extern Barrier g_barrier;
31
32void 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
79std::string MultiTrainer::GetDumpPath(int tid) {
80 if (!user_define_dump_filename_.empty()) {

Callers

nothing calls this directly

Calls 14

GetReadersMethod · 0.80
SetNeedDumpFieldMethod · 0.80
SetNeedDumpParamMethod · 0.80
SetDumpFieldVectorMethod · 0.80
SetDumpParamVectorMethod · 0.80
InitRandomDumpConfigMethod · 0.80
SetDataFeedMethod · 0.80
push_backMethod · 0.45
sizeMethod · 0.45
resizeMethod · 0.45
resetMethod · 0.45
SetDeviceIndexMethod · 0.45

Tested by

no test coverage detected