From 17e2d0709cac516bfe1f20f17f5d75e9638e9c82 Mon Sep 17 00:00:00 2001 From: ZPaC Date: Wed, 6 Apr 2022 11:33:28 +0800 Subject: [PATCH] Do rpc nodes fusion. --- .../parallel/graph_util/graph_splitter.cc | 126 +++++++++++++++++- .../parallel/graph_util/graph_splitter.h | 16 ++- 2 files changed, 139 insertions(+), 3 deletions(-) diff --git a/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.cc b/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.cc index 3e6481b9ea2..45ccda57039 100644 --- a/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.cc +++ b/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.cc @@ -205,7 +205,7 @@ CNodePtr CreateRecvNode(const FuncGraphPtr &func_graph, const InterProcessOpEdge void ParameterServerMode::PreBuildDistributedGraph() { MS_LOG(INFO) << "Start pre-building distribtued graph in Parameter Server mode."; MS_EXCEPTION_IF_NULL(node_labels_); - ProcessForSplittedOptimizer(); + ProcessForSplitOptimizer(); MS_LOG(INFO) << "End pre-building distribtued graph in Parameter Server mode."; } @@ -257,7 +257,48 @@ void ParameterServerMode::PostBuildDistributedGraph(const InterProcessOpEdgesInf MS_LOG(INFO) << "End post-building distribtued graph in Parameter Server mode."; } -void ParameterServerMode::ProcessForSplittedOptimizer() { +void ParameterServerMode::DoRpcNodeFusion() { + MS_EXCEPTION_IF_NULL(func_graph_); + std::vector all_nodes = DeepScopedGraphSearch(func_graph_->get_return()); + // Only the rpc nodes whose peer is the same process(with same OperatorLabel) can be fused. + std::map, std::vector> rpc_nodes_list_need_to_be_fused; + for (const auto &node : all_nodes) { + MS_EXCEPTION_IF_NULL(node); + if (!node->isa()) { + continue; + } + const auto &cnode = node->cast(); + std::string cnode_name = common::AnfAlgo::GetCNodeName(cnode); + if (cnode_name != kRpcSendOpName && cnode_name != kRpcRecvOpName) { + continue; + } + const auto &peer_ranks = (cnode_name == kRpcSendOpName) + ? common::AnfAlgo::GetNodeAttr>(cnode, kAttrSendDstRanks) + : common::AnfAlgo::GetNodeAttr>(cnode, kAttrRecvSrcRanks); + const auto &peer_roles = (cnode_name == kRpcSendOpName) + ? common::AnfAlgo::GetNodeAttr>(cnode, kAttrSendDstRoles) + : common::AnfAlgo::GetNodeAttr>(cnode, kAttrRecvSrcRoles); + OperatorLabel peer_label = {peer_ranks[0], peer_roles[0]}; + rpc_nodes_list_need_to_be_fused[std::make_pair(peer_label, cnode_name)].emplace_back(cnode); + } + + for (auto &rpc_nodes_fuse_info : rpc_nodes_list_need_to_be_fused) { + // Reorder the rpc nodes list according to the inter-process edge name so the inputs order of send/recv nodes can + // correspond. + std::sort(rpc_nodes_fuse_info.second.begin(), rpc_nodes_fuse_info.second.end(), + [](const CNodePtr &a, const CNodePtr &b) { + return common::AnfAlgo::GetNodeAttr(a, kAttrInterProcessEdgeName) < + common::AnfAlgo::GetNodeAttr(b, kAttrInterProcessEdgeName); + }); + if (rpc_nodes_fuse_info.first.second == kRpcSendOpName) { + FuseRpcSendNodes(rpc_nodes_fuse_info.second); + } else { + FuseRpcRecvNodes(rpc_nodes_fuse_info.second); + } + } +} + +void ParameterServerMode::ProcessForSplitOptimizer() { // Judge the node role number validation. uint32_t worker_num = ClusterContext::instance()->node_num(distributed::kEnvRoleOfWorker); if (worker_num == 0) { @@ -294,6 +335,7 @@ void ParameterServerMode::ProcessForSplittedOptimizer() { MS_LOG(EXCEPTION) << "The gradient type " << gradient_type << " is invalid."; } + const std::string &opt_device_target = GetCNodeTarget(ps_optimizer); for (size_t i = 0; i < common::AnfAlgo::GetInputNum(ps_optimizer); i++) { auto input = common::AnfAlgo::GetInputNode(ps_optimizer, i); // If the input is not a cnode, no inter-process edge is added so no node with multiple inputs should be created. @@ -311,6 +353,8 @@ void ParameterServerMode::ProcessForSplittedOptimizer() { func_graph_->manager()->SetEdge(ps_optimizer, i + 1, real_div_node); node_labels_->insert(std::make_pair(accum_node, node_labels_->at(ps_optimizer))); node_labels_->insert(std::make_pair(real_div_node, node_labels_->at(ps_optimizer))); + common::AnfAlgo::SetNodeAttr(kAttrPrimitiveTarget, MakeValue(opt_device_target), accum_node); + common::AnfAlgo::SetNodeAttr(kAttrPrimitiveTarget, MakeValue(opt_device_target), real_div_node); } else if (i == indices_index) { // Create the node to replace origin indices. AnfNodePtr new_indices_input = CreateNodeWithInterProcessEdgeOnPServer( @@ -318,6 +362,7 @@ void ParameterServerMode::ProcessForSplittedOptimizer() { func_graph_->manager()->SetEdge(ps_optimizer, i + 1, new_indices_input); node_labels_->insert(std::make_pair(new_indices_input, node_labels_->at(ps_optimizer))); + common::AnfAlgo::SetNodeAttr(kAttrPrimitiveTarget, MakeValue(opt_device_target), new_indices_input); } else { std::pair make_tuple_get_item_nodes = CreateNodesForMakeTuple(input, (role_ == distributed::kEnvRoleOfWorker) ? rank_id_ : 0, worker_num); @@ -327,6 +372,8 @@ void ParameterServerMode::ProcessForSplittedOptimizer() { func_graph_->manager()->SetEdge(ps_optimizer, i + 1, tuple_get_item_node); node_labels_->insert(std::make_pair(make_tuple_node, node_labels_->at(ps_optimizer))); node_labels_->insert(std::make_pair(tuple_get_item_node, node_labels_->at(ps_optimizer))); + common::AnfAlgo::SetNodeAttr(kAttrPrimitiveTarget, MakeValue(opt_device_target), make_tuple_node); + common::AnfAlgo::SetNodeAttr(kAttrPrimitiveTarget, MakeValue(opt_device_target), tuple_get_item_node); } } } @@ -473,6 +520,78 @@ CNodePtr ParameterServerMode::CreateNodeWithInterProcessEdgeOnPServer(const std: return new_node; } +bool ParameterServerMode::FuseRpcSendNodes(const std::vector &rpc_send_nodes) { + std::vector send_inputs = {NewValueNode(std::make_shared(kRpcSendOpName))}; + AbstractBasePtrList abstract_list = {}; + std::string fused_inter_process_edge_name = ""; + for (const auto &send_node : rpc_send_nodes) { + MS_EXCEPTION_IF_NULL(send_node); + for (size_t i = 1; i < send_node->inputs().size(); i++) { + auto input_i = send_node->inputs()[i]; + MS_EXCEPTION_IF_NULL(input_i); + // If the input of send is monad, do not pass it to fused send node. + if (HasAbstractMonad(input_i)) { + continue; + } + send_inputs.emplace_back(input_i); + } + abstract_list.emplace_back(send_node->abstract()); + fused_inter_process_edge_name.append( + common::AnfAlgo::GetNodeAttr(send_node, kAttrInterProcessEdgeName)); + } + + CNodePtr fused_send_node = func_graph_->NewCNode(send_inputs); + MS_EXCEPTION_IF_NULL(fused_send_node); + fused_send_node->set_abstract(std::make_shared(abstract_list)); + common::AnfAlgo::SetNodeAttr(kAttrInterProcessEdgeName, MakeValue(fused_inter_process_edge_name), fused_send_node); + + for (size_t j = 0; j < rpc_send_nodes.size(); j++) { + auto index_node = NewValueNode(MakeValue(SizeToLong(j))); + MS_EXCEPTION_IF_NULL(index_node); + std::vector tuple_get_item_inputs = {NewValueNode(std::make_shared(prim::kTupleGetItem)), + fused_send_node, index_node}; + CNodePtr tuple_get_item_node = func_graph_->NewCNode(tuple_get_item_inputs); + MS_EXCEPTION_IF_NULL(tuple_get_item_node); + tuple_get_item_node->set_abstract(abstract_list[j]); + func_graph_->manager()->Replace(rpc_send_nodes[j], tuple_get_item_node); + } + return true; +} + +bool ParameterServerMode::FuseRpcRecvNodes(const std::vector &rpc_recv_nodes) { + std::vector recv_inputs = {NewValueNode(std::make_shared(kRpcRecvOpName))}; + AbstractBasePtrList abstract_list = {}; + std::string fused_inter_process_edge_name = ""; + for (const auto &recv_node : rpc_recv_nodes) { + MS_EXCEPTION_IF_NULL(recv_node); + for (size_t i = 1; i < recv_node->inputs().size(); i++) { + auto input_i = recv_node->inputs()[i]; + MS_EXCEPTION_IF_NULL(input_i); + recv_inputs.emplace_back(input_i); + } + abstract_list.emplace_back(recv_node->abstract()); + fused_inter_process_edge_name.append( + common::AnfAlgo::GetNodeAttr(recv_node, kAttrInterProcessEdgeName)); + } + + CNodePtr fused_recv_node = func_graph_->NewCNode(recv_inputs); + MS_EXCEPTION_IF_NULL(fused_recv_node); + fused_recv_node->set_abstract(std::make_shared(abstract_list)); + common::AnfAlgo::SetNodeAttr(kAttrInterProcessEdgeName, MakeValue(fused_inter_process_edge_name), fused_recv_node); + + for (size_t j = 0; j < rpc_recv_nodes.size(); j++) { + auto index_node = NewValueNode(MakeValue(SizeToLong(j))); + MS_EXCEPTION_IF_NULL(index_node); + std::vector tuple_get_item_inputs = {NewValueNode(std::make_shared(prim::kTupleGetItem)), + fused_recv_node, index_node}; + CNodePtr tuple_get_item_node = func_graph_->NewCNode(tuple_get_item_inputs); + MS_EXCEPTION_IF_NULL(tuple_get_item_node); + tuple_get_item_node->set_abstract(abstract_list[j]); + func_graph_->manager()->Replace(rpc_recv_nodes[j], tuple_get_item_node); + } + return true; +} + GraphSplitter::GraphSplitter(const FuncGraphPtr &func_graph, uint32_t rank_id, const std::string &role) : func_graph_(func_graph), rank_id_(rank_id), @@ -520,6 +639,9 @@ void GraphSplitter::Run() { // Step 7: Postbuild the graph after splitting. exec_mode_->PostBuildDistributedGraph(comm_edges); + + // Step 8: Fuse the rpc nodes to improve performance. + exec_mode_->DoRpcNodeFusion(); } void GraphSplitter::DyeGraph() { diff --git a/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.h b/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.h index bd88465723b..64b0772bfd0 100644 --- a/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.h +++ b/mindspore/ccsrc/frontend/parallel/graph_util/graph_splitter.h @@ -172,6 +172,9 @@ class DistributedExecutionMode { // Input 'comm_edges' represents the inter-process edges generated after splitting the graph. virtual void PostBuildDistributedGraph(const InterProcessOpEdgesInfo &comm_edges) {} + // After building the distributed graph, do rpc node fusion to decrease the overhead of network communication. + virtual void DoRpcNodeFusion() {} + protected: FuncGraphPtr func_graph_; @@ -197,10 +200,11 @@ class ParameterServerMode : public DistributedExecutionMode { void PreBuildDistributedGraph() override; void PostBuildDistributedGraph(const InterProcessOpEdgesInfo &comm_edges) override; + void DoRpcNodeFusion() override; private: // Process optimizers split to the parameter server. - void ProcessForSplittedOptimizer(); + void ProcessForSplitOptimizer(); // Filter out all optimizer nodes which are set on parameter server from the graph. std::vector FilterServerAwareOptimizerList(const std::vector &nodes); @@ -228,6 +232,16 @@ class ParameterServerMode : public DistributedExecutionMode { CNodePtr CreateNodeWithInterProcessEdgeOnPServer(const std::string &many_to_one_node_name, const AnfNodePtr &real_input, size_t index_of_real_input, uint32_t total_inputs_number); + + // Fuse RpcSend and RpcRecv nodes for Parameter Server optimizers. Only one fused send node should be corresponding to + // one fused recv node, vice versa. + void FuseRpcNodesForSplitOptimizer(); + + // Fuse the given rpc send nodes list. Only nodes which send data to the same peer can be fused. + bool FuseRpcSendNodes(const std::vector &rpc_send_nodes); + + // Fuse the given rpc recv nodes list. Only nodes which recv data from the same peer can be fused. + bool FuseRpcRecvNodes(const std::vector &rpc_recv_nodes); }; // The class is used as an action in pipeline. It will process the graph and split the nodes to each process in the