MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / GetProxyNode

Method GetProxyNode

oneflow/core/graph/task_graph.cpp:545–591  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

543}
544
545TaskNode* TaskGraph::GetProxyNode(TaskNode* src_node, const LogicalBlobId& lbi,
546 const MemZoneId& dst_mem_zone_id) {
547 const auto& src_mem_zone_id = src_node->MemZoneId121();
548 const ProxyKey key(src_node, lbi, dst_mem_zone_id);
549 auto it = proxy2node.find(key);
550 if (it != proxy2node.cend()) {
551 // hit cache
552 return it->second;
553 } else {
554 if (src_mem_zone_id == dst_mem_zone_id) {
555 // in the same memory zone
556 proxy2node[key] = src_node;
557 return src_node;
558 } else if (dst_mem_zone_id.device_type() == DeviceType::kCPU) {
559 if (src_mem_zone_id.rank() == dst_mem_zone_id.rank()) {
560 // on the same node, not on the same device
561 // src must be not on the cpu mem zone, copy d2h first
562 CHECK(IsMemcpyDtoHSupported(src_mem_zone_id.device_type()));
563 CopyHdTaskNode* copy_task = NewNode<CopyHdTaskNode>();
564 copy_task->Init(CopyHdType::D2H, src_mem_zone_id, lbi);
565 Connect<TaskNode>(src_node, NewTaskEdgeWithLbi(lbi), copy_task);
566 proxy2node[key] = copy_task;
567 return copy_task;
568 } else {
569 // not on the same node, need CopyCommNet from src to dst
570 // build src cpu proxy first
571 TaskNode* proxy_on_src_host =
572 GetProxyNode(src_node, lbi, GetNodeCPUMemZoneId(src_mem_zone_id.rank()));
573 CopyCommNetTaskNode* copy_comm_net_task = NewNode<CopyCommNetTaskNode>();
574 copy_comm_net_task->Init(dst_mem_zone_id.rank(), lbi);
575 Connect<TaskNode>(proxy_on_src_host, NewTaskEdgeWithLbi(lbi), copy_comm_net_task);
576 proxy2node[key] = copy_comm_net_task;
577 return copy_comm_net_task;
578 }
579 } else {
580 TaskNode* proxy_on_dst_host =
581 GetProxyNode(src_node, lbi, GetNodeCPUMemZoneId(dst_mem_zone_id.rank()));
582 CHECK(IsMemcpyHtoDSupported(dst_mem_zone_id.device_type()));
583 CopyHdTaskNode* copy_task = NewNode<CopyHdTaskNode>();
584 copy_task->Init(CopyHdType::H2D, dst_mem_zone_id, lbi);
585 Connect<TaskNode>(proxy_on_dst_host, NewTaskEdgeWithLbi(lbi), copy_task);
586 proxy2node[key] = copy_task;
587 return copy_task;
588 }
589 }
590 return nullptr;
591}
592
593TaskNode* TaskGraph::GetProxyNode(TaskNode* src_node, const LogicalBlobId& lbi,
594 const ParallelDesc& dst_parallel_desc, int64_t dst_parallel_id) {

Callers 9

FOR_RANGEFunction · 0.80
BuildMethod · 0.80
FOR_RANGEFunction · 0.80
FOR_RANGEFunction · 0.80
FOR_RANGEFunction · 0.80
BuildMethod · 0.80
FOR_RANGEFunction · 0.80
FOR_RANGEFunction · 0.80
BuildMethod · 0.80

Calls 11

IsMemcpyDtoHSupportedFunction · 0.85
GetNodeCPUMemZoneIdFunction · 0.85
IsMemcpyHtoDSupportedFunction · 0.85
findMethod · 0.80
cendMethod · 0.80
MachineId4ParallelIdMethod · 0.80
DeviceId4ParallelIdMethod · 0.80
MemZoneId121Method · 0.45
device_typeMethod · 0.45
rankMethod · 0.45
InitMethod · 0.45

Tested by

no test coverage detected