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

Function GetMthdForBldSubTskGph

oneflow/core/graph/task_graph.cpp:370–444  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

368}
369
370BldSubTskGphMthd GetMthdForBldSubTskGph(const OpEdge* op_edge) {
371 const OpNode* src_node = op_edge->src_node();
372 const OpNode* dst_node = op_edge->dst_node();
373 const ParallelDesc& src_pd = src_node->parallel_desc();
374 const ParallelDesc& dst_pd = dst_node->parallel_desc();
375 const OperatorConf& src_op_conf = src_node->op().op_conf();
376 const OperatorConf& dst_op_conf = dst_node->op().op_conf();
377
378 // WaitAndSendIds -> Reentrantlock
379 if (src_op_conf.has_wait_and_send_ids_conf() && dst_op_conf.has_reentrant_lock_conf()) {
380 CHECK_EQ(src_pd.parallel_num(), 1);
381 CHECK_EQ(dst_pd.parallel_num(), 1);
382 return &TaskGraph::BldSubTskGphByBoxing;
383 }
384
385 // *Tick -> *Tick
386 if (IsTickOpConf(src_op_conf) || IsTickOpConf(dst_op_conf)) {
387 if (src_op_conf.has_source_tick_conf()) {
388 CHECK(dst_op_conf.has_tick_conf());
389 CHECK_EQ(src_pd.parallel_num(), 1);
390 CHECK_EQ(dst_pd.parallel_num(), 1);
391 return &TaskGraph::BldSubTskGphByBoxing;
392 } else if (dst_op_conf.has_sink_tick_conf()) {
393 CHECK(src_op_conf.has_tick_conf() || src_op_conf.has_sink_tick_conf());
394 CHECK_EQ(src_pd.parallel_num(), 1);
395 CHECK_EQ(dst_pd.parallel_num(), 1);
396 return &TaskGraph::BldSubTskGphByBoxing;
397 } else if (IsSubsetTickOpConf(src_op_conf)) {
398 return &TaskGraph::BldSubTskGphBySrcSubsetConnect;
399 } else if (IsSubsetTickOpConf(dst_op_conf)) {
400 return &TaskGraph::BldSubTskGphByDstSubsetConnect;
401 } else if (IsTickOpConf(src_op_conf) && IsTickOpConf(dst_op_conf)) {
402 if (src_pd.parallel_num() == dst_pd.parallel_num()) {
403 return &TaskGraph::BldSubTskGphByOneToOne;
404 } else {
405 CHECK_EQ(src_pd.parallel_num(), 1);
406 return &TaskGraph::BldSubTskGphByBroadcastToBroadcast;
407 }
408 }
409 }
410
411 std::shared_ptr<CompTaskNode> src_comp_task(NewCompTaskNode4OpNode(src_node));
412 std::shared_ptr<CompTaskNode> dst_comp_task(NewCompTaskNode4OpNode(dst_node));
413 // NOTE(chengcheng): MUST use TaskType instead of OpTypeCase because may
414 // Multi-op corresponding to SAME TaskType such as:
415 // DistributeConcatOpConf and DistributeAddOpConf -> TaskType::kDistributeConcat
416 // DistributeSplitOpConf and DistributeCloneOpConf -> TaskType::kDistributeSplit
417 // * -> DistributeConcat
418 if (dst_comp_task->GetTaskType() == TaskType::kDistributeConcat) {
419 return &TaskGraph::BldSubTskGphByPartialInLbiConnect;
420 }
421
422 // DistributeSplit -> *
423 if (src_comp_task->GetTaskType() == TaskType::kDistributeSplit) {
424 return &TaskGraph::BldSubTskGphByPartialOutLbiConnect;
425 }
426
427 // NormalForward -> DecodeH2D

Callers 1

InitMethod · 0.85

Calls 11

IsSubsetTickOpConfFunction · 0.85
NewCompTaskNode4OpNodeFunction · 0.85
src_nodeMethod · 0.80
dst_nodeMethod · 0.80
hierarchyMethod · 0.80
IsTickOpConfFunction · 0.70
parallel_descMethod · 0.45
opMethod · 0.45
parallel_numMethod · 0.45
GetTaskTypeMethod · 0.45

Tested by

no test coverage detected