Line data Source code
1 : /**
2 : * Copyright (c) 2025 Huawei Technologies Co., Ltd.
3 : * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
4 : * CANN Open Software License Agreement Version 2.0 (the "License").
5 : * Please refer to the License for details. You may not use this file except in compliance with the License.
6 : * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
7 : * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
8 : * See LICENSE in the root of the software repository for the full text of the License.
9 : */
10 :
11 : #include <iostream>
12 : #include <memory>
13 :
14 : #include "interpreter.h"
15 : #include "ins_rules.h"
16 : #include "aiv_ins.h"
17 :
18 : namespace Hccl {
19 10 : Interpreter::Interpreter(CommunicatorImpl& communicator)
20 10 : : comm(communicator),
21 240 : insRuleMap(
22 10 : {{InstructionType::LOCAL_COPY, GetInsRule<InsLocalCopy>()},
23 10 : {InstructionType::LOCAL_REDUCE, GetInsRule<InsLocalReduce>()},
24 10 : {InstructionType::LOCAL_POST_TO, GetInsRule<InsLocalPostTo>()},
25 10 : {InstructionType::LOCAL_WAIT_FROM, GetInsRule<InsLocalWaitFrom>()},
26 10 : {InstructionType::LOCAL_WAIT_GROUP, GetInsRule<InsLocalWaitGroup>()},
27 10 : {InstructionType::LOCAL_BCAST_POST, GetInsRule<InsLocalBcastPost>()},
28 10 : {InstructionType::POST_READY, GetInsRule<InsPostReady>()},
29 10 : {InstructionType::WAIT_READY, GetInsRule<InsWaitReady>()},
30 10 : {InstructionType::POST_FIN, GetInsRule<InsPostFin>()},
31 10 : {InstructionType::WAIT_FIN, GetInsRule<InsWaitFin>()},
32 10 : {InstructionType::WAIT_GROUP_FIN, GetInsRule<InsWaitGroupFin>()},
33 10 : {InstructionType::POST_FIN_ACK, GetInsRule<InsPostFinAck>()},
34 10 : {InstructionType::WAIT_FIN_ACK, GetInsRule<InsWaitFinAck>()},
35 10 : {InstructionType::WRITE_WITH_FIN, GetInsRule<InsWriteWithFin>()},
36 10 : {InstructionType::WRITE_REDUCE_WITH_FIN, GetInsRule<InsWriteReduceWithFin>()},
37 10 : {InstructionType::WRITE, GetInsRule<InsWrite>()},
38 10 : {InstructionType::WRITE_REDUCE, GetInsRule<InsWriteReduce>()},
39 10 : {InstructionType::READ, GetInsRule<InsRead>()},
40 10 : {InstructionType::READ_REDUCE, GetInsRule<InsReadReduce>()},
41 10 : {InstructionType::CCU_INS, GetInsRule<CcuInstruction>()},
42 10 : {InstructionType::AICPU_INS, GetInsRule<AicpuInstruction>()},
43 20 : {InstructionType::AIV_INS, GetInsRule<AivInstruction>()}})
44 : {
45 10 : if (communicator.GetCurrentCollOperator()->opType == OpType::BARRIER) {
46 0 : taskConfig.SetNotifyWaitTime(communicator.GetNotifyTimeoutCfg().GetBarrierTimeout());
47 : } else {
48 10 : taskConfig.SetNotifyWaitTime(communicator.GetNotifyTimeoutCfg().GetNotifyTimeout());
49 : }
50 20 : }
51 :
52 7 : void Interpreter::Submit(const InsQueue& insQueue)
53 : {
54 7 : list<InsQueue::Iterator> slaveQueueIters;
55 8 : for (auto slaveQueueIter = insQueue.IterSlaves(); slaveQueueIter.HasNext(); ++slaveQueueIter) {
56 1 : slaveQueueIters.emplace_back((*slaveQueueIter).Iter());
57 7 : }
58 :
59 7 : std::set<u32> slaveStreamIndexSet;
60 8 : for (u32 slaveStreamIndex = 0; slaveStreamIndex < slaveQueueIters.size(); ++slaveStreamIndex) {
61 1 : slaveStreamIndexSet.insert(slaveStreamIndex);
62 : }
63 :
64 : // 获取指令规模,填充桶宽及初始化
65 7 : auto& streamMgr = comm.GetStreamManager();
66 7 : auto masterStream = streamMgr.GetMaster();
67 7 : streamMgr.InitBucket(UINT32_MAX);
68 7 : streamMgr.RecordStreamIdToIndex(masterStream->GetId(), UINT32_MAX);
69 7 : u32 index = 0;
70 8 : for (auto slaveIter = insQueue.IterSlaves(); slaveIter.HasNext(); ++slaveIter) {
71 1 : streamMgr.InitBucket(streamMgr.GetSlaveIndex());
72 1 : auto slaveStream = streamMgr.GetSlave();
73 1 : streamMgr.RecordStreamIdToIndex(slaveStream->GetId(), index);
74 1 : streamMgr.CaptureSlaveStream(masterStream, slaveStream);
75 7 : }
76 7 : InsQueue::Iterator masterQueueIter = insQueue.Iter();
77 :
78 : // 交替流下Task,直至全部下完
79 11 : while (!slaveQueueIters.empty() || masterQueueIter.HasNext()) {
80 4 : SubmitSlaveQueueAlternatively(slaveQueueIters, slaveStreamIndexSet);
81 4 : SubmitMasterQueueAlternatively(masterQueueIter);
82 : }
83 :
84 : // 销毁桶
85 7 : streamMgr.DestroyRecords();
86 : // 销毁流的占用状态
87 7 : streamMgr.ResetSlaveIndex(0);
88 7 : }
89 :
90 4 : void Interpreter::SubmitSlaveQueueAlternatively(
91 : list<InsQueue::Iterator>& slaveQueueIters, std::set<u32>& slaveStreamIndexSet)
92 : {
93 4 : auto& streamMgr = comm.GetStreamManager();
94 4 : auto slaveStreamIndexIter = slaveStreamIndexSet.begin();
95 8 : for (auto slaveQueueIter = slaveQueueIters.begin(); slaveQueueIter != slaveQueueIters.end();) {
96 4 : if (!slaveQueueIter->HasNext()) {
97 : // 销毁已经完成下发的流迭代器
98 3 : HCCL_INFO(
99 : "[SubmitSlaveQueueAlternatively] slave stream index(%u) interpret finish", (*slaveStreamIndexIter));
100 1 : slaveQueueIter = slaveQueueIters.erase(slaveQueueIter);
101 1 : slaveStreamIndexIter = slaveStreamIndexSet.erase(slaveStreamIndexIter);
102 1 : continue;
103 1 : }
104 3 : auto& rule = insRuleMap.at((*slaveQueueIter)->GetType());
105 3 : auto stream = streamMgr.GetSlaveByIndex(*slaveStreamIndexIter);
106 3 : rule(**slaveQueueIter, comm, *stream, taskConfig);
107 :
108 9 : HCCL_INFO(
109 : "[SubmitSlaveQueueAlternatively] slave stream index[%u], stream id[%u]. Instruction start %s",
110 : (*slaveStreamIndexIter), stream->GetId(), (*slaveQueueIter)->Describe().c_str());
111 : // 当前Task下载完毕,跳转这条流上的InsQueue里下一个task
112 3 : ++(*slaveQueueIter);
113 : // 切换至下一条流
114 3 : ++slaveQueueIter;
115 3 : ++slaveStreamIndexIter;
116 : }
117 4 : }
118 :
119 4 : void Interpreter::SubmitMasterQueueAlternatively(InsQueue::Iterator& masterQueueIter)
120 : {
121 4 : if (!masterQueueIter.HasNext()) {
122 3 : HCCL_INFO("[SubmitMasterQueueAlternatively] main stream interpret finish");
123 1 : return;
124 : }
125 3 : auto& streamMgr = comm.GetStreamManager();
126 3 : auto stream = streamMgr.GetMaster();
127 3 : auto& rule = insRuleMap.at(masterQueueIter->GetType());
128 3 : rule(*masterQueueIter, comm, *stream, taskConfig);
129 9 : HCCL_INFO(
130 : "[SubmitMasterQueueAlternatively] master stream id[%u]. Instruction start %s", stream->GetId(),
131 : masterQueueIter->Describe().c_str());
132 : // 切换至下一个task
133 3 : ++masterQueueIter;
134 : }
135 :
136 : } // namespace Hccl
|