unified runtime bug of auto moand and host device of multi graphs

This commit is contained in:
limingqi107 2021-07-23 11:07:01 +08:00
parent 34e2581dfc
commit 0b8937cd6c
4 changed files with 75 additions and 57 deletions

View File

@ -1124,10 +1124,27 @@ void KernelGraph::ReplaceInternalOutput(const AnfNodePtr &node, const AnfNodePtr
void KernelGraph::CacheInternalParameterToFrontNode(const AnfNodePtr &parameter,
const AnfWithOutIndex &front_node_with_index) {
if (parameter == nullptr) {
if ((parameter == nullptr) || (front_node_with_index.first == nullptr)) {
return;
}
internal_parameter_to_front_node_map_[parameter] = front_node_with_index;
auto front_outputs = AnfAlgo::GetAllOutputWithIndex(front_node_with_index.first);
AnfWithOutIndex new_front_node_with_index;
if (front_node_with_index.second < front_outputs.size()) {
new_front_node_with_index = front_outputs[front_node_with_index.second];
} else {
new_front_node_with_index = front_node_with_index;
}
if (new_front_node_with_index.first == nullptr) {
return;
}
MS_LOG(INFO) << "Cache internal parameter: " << parameter->DebugString()
<< " to front node: " << new_front_node_with_index.first->DebugString()
<< " with index: " << new_front_node_with_index.second
<< ", from front node: " << front_node_with_index.first->DebugString()
<< " with index: " << front_node_with_index.second;
internal_parameter_to_front_node_map_[parameter] = new_front_node_with_index;
}
AnfWithOutIndex KernelGraph::GetFrontNodeByInternalParameter(const AnfNodePtr &parameter) const {

View File

@ -667,10 +667,6 @@ void SessionBasic::GetNewCNodeInputs(const CNodePtr &cnode, KernelGraph *graph,
}
cnode_inputs->push_back(parameter_from_cnode);
(*other_graph_cnode)[anf] = parameter_from_cnode;
KernelWithIndex front_node_with_index(anf, 0);
MS_LOG(INFO) << "The " << input_idx << " input of node:" << cnode->fullname_with_scope()
<< " is from front node:" << anf->fullname_with_scope();
graph->CacheInternalParameterToFrontNode(parameter_from_cnode, front_node_with_index);
}
}
}

View File

@ -893,7 +893,7 @@ void GraphScheduler::Link(ActorSet *actor_set, const GraphCompilerInfo &graph_co
auto input_node = AnfAlgo::GetInputNode(kernel, i);
// Link the control arrows of kernel actor by the auto monad, the inputs include monad node.
if (AnfAlgo::IsOneOfPrimitiveCNode(input_node, auto_monad_prims)) {
LinkControlArrowByAutoMonad(kernel_actor, input_node);
LinkControlArrowByAutoMonad(kernel_actor, input_node, graph);
}
if (HasAbstractMonad(input_node)) {
auto_monad_actors.emplace_back(kernel_actor);
@ -1401,42 +1401,39 @@ void GraphScheduler::LinkDataArrowForInternalParameter(const AnfNodePtr &interna
MS_EXCEPTION_IF_NULL(graph);
MS_EXCEPTION_IF_NULL(to_actor);
// Parameter ---> front node ---> actor.
auto front_node_with_index = graph->GetFrontNodeByInternalParameter(internal_parameter);
MS_EXCEPTION_IF_NULL(front_node_with_index.first);
const auto &front_output_with_index =
AnfAlgo::VisitKernelWithReturnType(front_node_with_index.first, front_node_with_index.second, false);
// Parameter ---> front node.
auto front_output_with_index = graph->GetFrontNodeByInternalParameter(internal_parameter);
auto front_output_node = front_output_with_index.first;
MS_EXCEPTION_IF_NULL(front_output_node);
if (IsSwitchActor(front_output_node)) {
auto switch_actor = dynamic_cast<SwitchActor *>(FetchActor(front_output_node->DebugString()));
MS_EXCEPTION_IF_NULL(switch_actor);
LinkDataArrowForSwitchActor(switch_actor, 0, to_actor, to_kernel_with_input_idx.second);
to_actor->input_datas_num_++;
return;
}
if (IsPersistentDeviceTensor(front_output_node)) {
to_actor->device_tensor_store_keys_.emplace_back(to_kernel_with_input_idx.second, front_output_node.get());
return;
}
if (graph_output_to_actor_.count(front_output_with_index) == 0 && (!IsSwitchActor(front_output_node))) {
// front node ---> actor.
if (graph_output_to_actor_.count(front_output_with_index) == 0) {
MS_LOG(EXCEPTION) << "Can't find actor by front node:" << AnfAlgo::GetNodeDebugString(front_output_node)
<< ", internal parameter:" << AnfAlgo::GetNodeDebugString(internal_parameter);
}
auto actor_pair = graph_output_to_actor_[front_output_with_index];
if (actor_pair.first != nullptr) {
MS_LOG(INFO) << "Graph " << graph->graph_id() << " internal parameter:" << internal_parameter->DebugString()
<< ", corresponding front node:" << front_output_node->fullname_with_scope()
<< " with index:" << front_output_with_index.second
<< ", from actor:" << actor_pair.first->GetAID().Name() << " with index:" << actor_pair.second
<< ", to actor:" << to_actor->GetAID().Name() << " with index:" << to_kernel_with_input_idx.second;
}
MS_EXCEPTION_IF_NULL(actor_pair.first);
MS_LOG(INFO) << "Graph " << graph->graph_id() << " internal parameter:" << internal_parameter->DebugString()
<< ", corresponding front node:" << front_output_node->fullname_with_scope()
<< " with index:" << front_output_with_index.second
<< ", from actor:" << actor_pair.first->GetAID().Name() << " with index:" << actor_pair.second
<< ", to actor:" << to_actor->GetAID().Name() << " with index:" << to_kernel_with_input_idx.second;
if (IsDeviceQueueDSActor(front_output_node)) {
auto from_actor = dynamic_cast<DeviceQueueDataSourceActor *>(actor_pair.first);
auto from_kernel_with_output_idx = KernelWithIndex(from_actor->data_kernel_, actor_pair.second);
LinkDataArrowForDeviceDSActor(from_actor, to_actor, from_kernel_with_output_idx, to_kernel_with_input_idx);
} else if (IsSwitchActor(front_output_node)) {
const auto &actor_name = front_output_node->DebugString();
const auto &actor = FetchActor(actor_name);
MS_EXCEPTION_IF_NULL(actor);
auto switch_actor = dynamic_cast<SwitchActor *>(actor);
LinkDataArrowForSwitchActor(switch_actor, 0, to_actor, to_kernel_with_input_idx.second);
to_actor->input_datas_num_++;
} else if (IsKernelActor(front_output_node)) {
auto from_actor = dynamic_cast<KernelActor *>(actor_pair.first);
auto from_kernel_with_output_idx = KernelWithIndex(from_actor->kernel_, actor_pair.second);
@ -1446,7 +1443,7 @@ void GraphScheduler::LinkDataArrowForInternalParameter(const AnfNodePtr &interna
auto from_kernel_with_output_idx = KernelWithIndex(from_actor->data_nodes_[actor_pair.second], 0);
LinkDataArrowForHostDSActor(from_actor, to_actor, from_kernel_with_output_idx, to_kernel_with_input_idx);
} else {
MS_LOG(EXCEPTION) << "Invalid internal parameter: " << internal_parameter->fullname_with_scope();
MS_LOG(EXCEPTION) << "Invalid internal parameter: " << internal_parameter->DebugString();
}
}
@ -1615,26 +1612,25 @@ void GraphScheduler::LinkDataArrowForCopyActor(OpActor<DeviceTensor> *from_actor
UpdateRefCount(copy_actor->output_.get());
}
void GraphScheduler::LinkControlArrowByAutoMonad(KernelActor *to_actor, const AnfNodePtr &from_node) {
void GraphScheduler::LinkControlArrowByAutoMonad(KernelActor *to_actor, const AnfNodePtr &from_node,
const KernelGraphPtr &graph) {
MS_EXCEPTION_IF_NULL(to_actor);
MS_EXCEPTION_IF_NULL(from_node);
if (!from_node->isa<CNode>()) {
return;
}
// Find the real input node, include the monad node and make tuple node.
const std::vector<PrimitivePtr> return_types = {prim::kPrimDepend, prim::kPrimUpdateState, prim::kPrimLoad,
prim::kPrimMakeTuple};
const auto &input_kernel_with_output_idx = AnfAlgo::VisitKernelWithReturnType(from_node, 0, false, return_types);
MS_EXCEPTION_IF_NULL(input_kernel_with_output_idx.first);
if (!input_kernel_with_output_idx.first->isa<CNode>()) {
return;
auto input_anfnode = input_kernel_with_output_idx.first;
CNodePtr input_cnode = nullptr;
if (input_anfnode->isa<CNode>()) {
input_cnode = input_anfnode->cast<CNodePtr>();
}
const auto &input_cnode = input_kernel_with_output_idx.first->cast<CNodePtr>();
MS_EXCEPTION_IF_NULL(input_cnode);
// Make tuple node needs to be expanded.
if (AnfAlgo::CheckPrimitiveType(input_cnode, prim::kPrimMakeTuple)) {
if (AnfAlgo::CheckPrimitiveType(input_anfnode, prim::kPrimMakeTuple)) {
MS_EXCEPTION_IF_NULL(input_cnode);
for (size_t i = 1; i < input_cnode->inputs().size(); ++i) {
LinkControlArrowByAutoMonad(to_actor, input_cnode->input(i));
LinkControlArrowByAutoMonad(to_actor, input_cnode->input(i), graph);
}
return;
}
@ -1643,20 +1639,22 @@ void GraphScheduler::LinkControlArrowByAutoMonad(KernelActor *to_actor, const An
prim::kPrimDepend, prim::kPrimUpdateState, prim::kPrimLoad, prim::kPrimMakeTuple};
// Get the real depend input by monad node which needs to link the control arrow.
std::vector<AnfNodePtr> real_depend_inputs;
if (AnfAlgo::CheckPrimitiveType(input_cnode, prim::kPrimDepend) ||
AnfAlgo::CheckPrimitiveType(input_cnode, prim::kPrimLoad)) {
if (AnfAlgo::CheckPrimitiveType(input_anfnode, prim::kPrimDepend) ||
AnfAlgo::CheckPrimitiveType(input_anfnode, prim::kPrimLoad)) {
MS_EXCEPTION_IF_NULL(input_cnode);
real_depend_inputs.push_back(input_cnode->input(kDependAttachNodeIndex));
// The real input may be this scene: depend/load --> load/depend, so need add the control arrow for real input
// node in this scene.
if (AnfAlgo::IsOneOfPrimitiveCNode(input_cnode->input(kRealInputIndexInDepend), recursion_prims)) {
real_depend_inputs.push_back(input_cnode->input(kRealInputIndexInDepend));
}
} else if (AnfAlgo::CheckPrimitiveType(input_cnode, prim::kPrimUpdateState)) {
} else if (AnfAlgo::CheckPrimitiveType(input_anfnode, prim::kPrimUpdateState)) {
MS_EXCEPTION_IF_NULL(input_cnode);
for (size_t i = kUpdateStateRealInput; i < input_cnode->inputs().size(); ++i) {
real_depend_inputs.push_back(input_cnode->input(i));
}
} else {
real_depend_inputs.push_back(input_cnode);
real_depend_inputs.push_back(input_anfnode);
}
for (const auto &real_depend_input : real_depend_inputs) {
@ -1664,19 +1662,27 @@ void GraphScheduler::LinkControlArrowByAutoMonad(KernelActor *to_actor, const An
auto real_depend_kernel = real_depend_input_with_idx.first;
// The monad node and make tuple node need recursion.
if (AnfAlgo::IsOneOfPrimitiveCNode(real_depend_kernel, recursion_prims)) {
LinkControlArrowByAutoMonad(to_actor, real_depend_kernel);
LinkControlArrowByAutoMonad(to_actor, real_depend_kernel, graph);
continue;
}
if (!IsKernelActor(real_depend_kernel)) {
continue;
KernelActor *from_actor = nullptr;
if (IsKernelActor(real_depend_kernel)) {
from_actor = dynamic_cast<KernelActor *>(FetchActor(real_depend_kernel->fullname_with_scope()));
} else if (IsInternalParameter(real_depend_kernel, graph)) {
auto front_output_with_index = graph->GetFrontNodeByInternalParameter(real_depend_kernel);
MS_EXCEPTION_IF_NULL(front_output_with_index.first);
if (IsKernelActor(front_output_with_index.first)) {
if (graph_output_to_actor_.count(front_output_with_index) == 0) {
MS_LOG(EXCEPTION) << "Can't find actor by front node:" << front_output_with_index.first->DebugString();
}
from_actor = dynamic_cast<KernelActor *>(graph_output_to_actor_[front_output_with_index].first);
}
}
// Link the control arrow between the kernel actors.
const auto &from_actor = dynamic_cast<KernelActor *>(FetchActor(real_depend_kernel->fullname_with_scope()));
if (from_actor == nullptr) {
continue;
}
MS_LOG(INFO) << "Link control arrow by auto monad, from actor: " << real_depend_kernel->fullname_with_scope()
MS_LOG(INFO) << "Link control arrow by auto monad, from actor: " << from_actor->GetAID().Name()
<< ", to actor: " << to_actor->GetAID().Name();
from_actor->output_control_arrows_.emplace_back(to_actor->GetAID());
to_actor->input_controls_num_++;
@ -2666,10 +2672,7 @@ void GraphScheduler::PersistDeviceTensor(const GraphCompilerInfo &graph_compiler
MS_EXCEPTION_IF_NULL(input_node);
AnfNodePtr sub_front_node = nullptr;
if (IsInternalParameter(input_node, graph)) {
auto front_node_with_index = graph->GetFrontNodeByInternalParameter(input_node);
MS_EXCEPTION_IF_NULL(front_node_with_index.first);
const auto &front_output_with_index =
AnfAlgo::VisitKernelWithReturnType(front_node_with_index.first, front_node_with_index.second, false);
auto front_output_with_index = graph->GetFrontNodeByInternalParameter(input_node);
sub_front_node = front_output_with_index.first;
} else if (IsPersistentDeviceTensor(input_node) || HasAbstractRef(input_node)) {
sub_front_node = FetchFrontNodeByBackendNode(input_node, graph);
@ -3031,10 +3034,11 @@ void GraphScheduler::DumpDeviceTensorStore(const GraphCompilerInfo &graph_compil
continue;
}
const auto &front_node = FetchFrontNodeByBackendNode(value_node, graph);
MS_EXCEPTION_IF_NULL(front_node);
const auto device_tensors = DeviceTensorStore::GetInstance().Fetch(front_node.get());
ofs << "\t\tdevcie tensor key:" << front_node->fullname_with_scope() << "\tvalue size:" << device_tensors.size()
<< "\n";
ofs << "\t\tdevcie tensor key:" << front_node->DebugString() << "\tvalue size:" << device_tensors.size() << "\n";
for (const auto &device_tensor : device_tensors) {
MS_EXCEPTION_IF_NULL(device_tensor);
ofs << "\t\t\tdevcie tensor value:" << device_tensor << "\tptr:" << device_tensor->GetPtr()
<< "\tsize:" << device_tensor->GetSize() << "\toriginal_ref_count:" << device_tensor->original_ref_count()
<< "\tdevice_type:" << device_tensor->DeviceType() << "\n ";
@ -3053,9 +3057,10 @@ void GraphScheduler::DumpDeviceTensorStore(const GraphCompilerInfo &graph_compil
front_node = graph_compiler_info.control_node_parser_->FetchRootGraphFrontNodeBySubFrontNode(sub_front_node);
}
const auto device_tensors = DeviceTensorStore::GetInstance().Fetch(front_node.get());
ofs << "\t\tdevcie tensor key:" << front_node->fullname_with_scope() << "\tvalue size:" << device_tensors.size()
<< "\n";
MS_EXCEPTION_IF_NULL(front_node);
ofs << "\t\tdevcie tensor key:" << front_node->DebugString() << "\tvalue size:" << device_tensors.size() << "\n";
for (const auto &device_tensor : device_tensors) {
MS_EXCEPTION_IF_NULL(device_tensor);
ofs << "\t\t\tdevcie tensor value:" << device_tensor << "\tptr:" << device_tensor->GetPtr()
<< "\tsize:" << device_tensor->GetSize() << "\toriginal_ref_count:" << device_tensor->original_ref_count()
<< "\tdevice_type:" << device_tensor->DeviceType() << "\n ";

View File

@ -218,7 +218,7 @@ class GraphScheduler {
// 2. The processing of linking control arrows.
void LinkControlArrowForLoopCountActor(LoopCountActor *loop_count_actor, const ActorSet *actor_set,
const ControlNodeParserPtr &parser);
void LinkControlArrowByAutoMonad(KernelActor *to_actor, const AnfNodePtr &from_node);
void LinkControlArrowByAutoMonad(KernelActor *to_actor, const AnfNodePtr &from_node, const KernelGraphPtr &graph);
// The skipped node doesn't run, so need link the control arrow between the inputs and user of skipped node.
void LinkControlArrowBySkippedNode(KernelActor *to_actor, const AnfNodePtr &skipped_node);
// Link the control arrows for allreduce kernel by the send/recv nodes in the kernel graph.