| 543 | } |
| 544 | |
| 545 | TaskNode* 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 | |
| 593 | TaskNode* TaskGraph::GetProxyNode(TaskNode* src_node, const LogicalBlobId& lbi, |
| 594 | const ParallelDesc& dst_parallel_desc, int64_t dst_parallel_id) { |
no test coverage detected