| 325 | } // namespace |
| 326 | |
| 327 | CompTaskNode* GenCompTaskNode( |
| 328 | const OpNode* op_node, int64_t parallel_id, |
| 329 | const std::function<StreamId(const OpNode* op_node, int64_t parallel_id, TaskType task_type)>& |
| 330 | GetOrCreateStreamId) { |
| 331 | const ParallelDesc& parallel_desc = op_node->parallel_desc(); |
| 332 | int64_t parallel_num = parallel_desc.parallel_num(); |
| 333 | CompTaskNode* comp_task_node = NewCompTaskNode4OpNode(op_node); |
| 334 | int64_t machine_id = CHECK_JUST(parallel_desc.MachineId4ParallelId(parallel_id)); |
| 335 | comp_task_node->set_machine_id(machine_id); |
| 336 | comp_task_node->mut_parallel_ctx()->set_parallel_id(parallel_id); |
| 337 | comp_task_node->mut_parallel_ctx()->set_parallel_num(parallel_num); |
| 338 | StreamId stream_id = GetOrCreateStreamId(op_node, parallel_id, comp_task_node->GetTaskType()); |
| 339 | comp_task_node->set_thrd_id(EncodeStreamIdToInt64(stream_id)); |
| 340 | comp_task_node->set_op_node(op_node); |
| 341 | return comp_task_node; |
| 342 | } |
| 343 | |
| 344 | void GenSortedCompTaskNodes(const OpNode* op_node, std::vector<CompTaskNode*>* sorted_comp_tasks) { |
| 345 | int64_t parallel_idx = 0; |
no test coverage detected