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 "mc2_context.h"
12 : #include "ccu_task_arg_mc2.h"
13 : #include "const_val.h"
14 :
15 : namespace Hccl {
16 :
17 : using namespace std;
18 :
19 : const string OP_SELECTOR_LABEL = "OpSelector";
20 : // HBM参数index
21 : const uint32_t HBM_PARAM_IDX_0 = 0;
22 : const uint32_t HBM_PARAM_IDX_1 = 1;
23 : const uint32_t HBM_PARAM_IDX_2 = 2;
24 : const uint32_t HBM_PARAM_IDX_3 = 3;
25 :
26 : const uint32_t SINGLE_DIE = 1; // 单Die数量
27 : const uint32_t DOUBLE_DIE = 2; // 双Die数量
28 : const uint32_t DIE0_ID = 0; // Die0 ID
29 : const uint32_t DIE1_ID = 1; // Die1 ID
30 : const string DIE1_START_SIG = "Die1StartSig"; // 双Die场景,Die0通知Die1开始执行的信号
31 : const string DIE1_END_SIG = "Die1EndSig"; // 双Die场景,Die1通知Die0执行完成的信号
32 :
33 9 : void Mc2ContextBase::SetAlgoTemplateInfo(const map<uint64_t, uint32_t>& algoTemplateInfo)
34 : {
35 9 : algoTemplateInfo_ = algoTemplateInfo;
36 25 : for (const auto& pair : algoTemplateInfo_) {
37 48 : HCCL_INFO("[Mc2Context::SetAlgoTemplateInfo] algoSignature[%llu] startInstr[%u]", pair.first, pair.second);
38 : }
39 9 : }
40 :
41 5 : void Mc2ContextBase::SetMissionNumAndId(uint32_t miNum, uint32_t miIndex)
42 : {
43 5 : this->missionNum = miNum;
44 5 : this->missionIndex = miIndex;
45 5 : if (missionIndex >= miNum) {
46 0 : THROW<InvalidParamsException>("MC2 High Level API SetMissionNumAndId Failed: Invalid Mission Config");
47 : }
48 5 : if (miNum > 1) { // 多Mission场景才需要导入导出
49 4 : if (miIndex == 0) {
50 : // missionIndex = 0 为Master,需要missionNum - 1个导入导出信号以及missionNum - 1个导入变量
51 4 : for (uint32_t i = 0; i < miNum - 1; ++i) {
52 2 : exportMissoinSig.push_back(CreateMaskSignal());
53 2 : ExportMaskSignal(
54 4 : exportMissoinSig[i], "master_sig_" + std::to_string(GetDieId()) + "_" + std::to_string(i + 1));
55 2 : importMissionSig.push_back(
56 4 : ImportMaskSignal("slave_sig_" + std::to_string(GetDieId()) + "_" + std::to_string(i + 1)));
57 2 : importMissionVar.push_back(
58 4 : ImportVariable("slave_var_" + std::to_string(GetDieId()) + "_" + std::to_string(i + 1)));
59 : }
60 : } else {
61 : // missionIndex > 0 为Slave,需要1个导入导出信号以及1个导出变量
62 2 : importMissionSig.push_back(
63 4 : ImportMaskSignal("master_sig_" + std::to_string(GetDieId()) + "_" + std::to_string(miIndex)));
64 2 : exportMissoinSig.push_back(CreateMaskSignal());
65 2 : ExportMaskSignal(
66 4 : exportMissoinSig[0], "slave_sig_" + std::to_string(GetDieId()) + "_" + std::to_string(miIndex));
67 2 : exportMissionVar.push_back(CreateVariable());
68 2 : ExportVariable(
69 4 : exportMissionVar[0], "slave_var_" + std::to_string(GetDieId()) + "_" + std::to_string(miIndex));
70 : }
71 : }
72 5 : }
73 :
74 2 : void Mc2ContextBase::MissionPreSync(CcuRep::Variable& func)
75 : {
76 2 : if (missionNum == 1) {
77 2 : return;
78 : }
79 0 : if (missionIndex == 0) {
80 0 : for (uint32_t i = 0; i < missionNum - 1; ++i) {
81 0 : LocalCtxPostVar(func, importMissionVar[i], importMissionSig[i]);
82 : }
83 : } else {
84 0 : LocalWait(exportMissoinSig[0]);
85 0 : func = exportMissionVar[0];
86 : }
87 : }
88 :
89 2 : void Mc2ContextBase::MissionPostSync()
90 : {
91 2 : if (missionNum == 1) {
92 2 : return;
93 : }
94 0 : if (missionIndex == 0) {
95 0 : for (uint32_t i = 0; i < missionNum - 1; ++i) {
96 0 : LocalWait(exportMissoinSig[i]);
97 : }
98 : } else {
99 0 : LocalCtxPost(importMissionSig[0]);
100 : }
101 : }
102 :
103 3 : void Mc2ContextBase::GenOpSelector()
104 : {
105 3 : if (algoTemplateInfo_.empty()) {
106 1 : THROW<InvalidParamsException>("MC2 High Level API GenOpSelector Failed: Empty AlgoTemplateInfo");
107 : }
108 :
109 : {
110 : std::string funcName
111 2 : = OP_SELECTOR_LABEL + "_" + std::to_string(GetDieId()) + "_" + std::to_string(missionIndex);
112 2 : CcuRep::FuncBlock selectorFunc(this, funcName);
113 :
114 2 : CcuRep::Variable opCode = CreateVariable(); // 函数入参,算子FuncBlock的signature
115 2 : selectorFunc.DefineInArg(opCode);
116 :
117 2 : CcuRep::Variable opAddr = CreateVariable(); // 函数出参,命中算子FuncBlock的函数地址
118 2 : selectorFunc.DefineOutArg(opAddr);
119 2 : opAddr = INVALID_U64; // opAddr初值为非法值,如果命中算子则会被改为对应的函数地址
120 :
121 6 : for (auto entry : algoTemplateInfo_) {
122 8 : CCU_IF(opCode == entry.first) { opAddr = entry.second; }
123 : }
124 2 : }
125 2 : }
126 :
127 3 : void Mc2ContextBase::Algorithm()
128 : {
129 3 : GenOpSelector();
130 2 : GenCircularQueue();
131 2 : }
132 :
133 6 : void Mc2Context::SetCommAddr(uint64_t syncAddr, uint64_t paramAddr)
134 : {
135 6 : waitAddr_ = syncAddr;
136 6 : if (syncAddr > (UINT64_MAX - CCU_TASK_NUM_MAX * CCU_ONE_PARAM_SIZE)) {
137 0 : THROW<InvalidParamsException>("MC2 High Level API SetDieNum Failed: integer overflow occurs");
138 : }
139 6 : recordAddr_ = syncAddr + CCU_TASK_NUM_MAX * CCU_ONE_PARAM_SIZE; // 偏移8轮的总宽度
140 6 : paramAddr_ = paramAddr;
141 6 : }
142 :
143 7 : void Mc2Context::SetDieNum(uint32_t dieNum)
144 : {
145 7 : dieNum_ = dieNum;
146 : // 参数合法值判断
147 7 : bool isDieNumValid = (dieNum_ == SINGLE_DIE || dieNum_ == DOUBLE_DIE);
148 7 : bool isDieIdValid = (GetDieId() == DIE0_ID || GetDieId() == DIE1_ID);
149 7 : if (!(isDieNumValid && isDieIdValid)) {
150 1 : THROW<InvalidParamsException>("MC2 High Level API SetDieNum Failed: Invalid Die Config");
151 : }
152 :
153 6 : if (dieNum_ == DOUBLE_DIE) { // 双Die场景才需要导入导出
154 : // 导出信号
155 4 : exportDieSig = CreateMaskSignal();
156 : // Die0: export完成信号给Die1,Die1: export开始信号给Die0
157 4 : const string& exportSigLabel = (GetDieId() == DIE1_ID) ? DIE1_START_SIG : DIE1_END_SIG;
158 4 : ExportMaskSignal(exportDieSig, exportSigLabel);
159 :
160 : // 导入信号
161 4 : const string& importSigLabel = (GetDieId() == DIE1_ID) ? DIE1_END_SIG : DIE1_START_SIG;
162 4 : importDieSig = ImportMaskSignal(importSigLabel);
163 : }
164 6 : }
165 :
166 2 : void Mc2Context::GenCircularQueue()
167 : {
168 : // 存放Token的寄存器
169 2 : CcuRep::Variable token = CreateVariable();
170 : // 从SQE中载入Token
171 2 : Load(token);
172 :
173 : // 存放《选择函数返回的FuncCall地址》的寄存器,选择函数的出参,循环队列内部使用
174 2 : CcuRep::Variable opAddr = CreateVariable();
175 :
176 : // 存放《控制repeat循环执行的条件》的寄存器
177 2 : CcuRep::Variable repeatCond = CreateVariable();
178 2 : repeatCond = 0;
179 :
180 : // 存放《轮次执行开始信号》的寄存器,初值为 0
181 2 : CcuRep::Variable turnStartSig = CreateVariable();
182 2 : turnStartSig = 0;
183 : // 存放《轮次执行完成信号》的寄存器,在循环中固定为 1
184 2 : CcuRep::Variable turnEndSig = CreateVariable();
185 2 : turnEndSig = 1;
186 :
187 2 : CcuRep::Variable waitStartAddr = CreateVariable();
188 2 : waitStartAddr = waitAddr_;
189 2 : CcuRep::Variable recordStartAddr = CreateVariable();
190 2 : recordStartAddr = recordAddr_;
191 2 : CcuRep::Variable paramStartAddr = CreateVariable();
192 2 : paramStartAddr = paramAddr_;
193 2 : CcuRep::Variable waitAddr = CreateVariable();
194 2 : waitAddr = waitAddr_;
195 2 : CcuRep::Variable recordAddr = CreateVariable();
196 2 : recordAddr = recordAddr_;
197 2 : CcuRep::Variable paramAddr = CreateVariable();
198 2 : paramAddr = paramAddr_;
199 :
200 2 : CcuRep::Variable ckeSize = CreateVariable();
201 2 : ckeSize = CCU_ONE_PARAM_SIZE;
202 2 : CcuRep::Variable paramSize = CreateVariable();
203 2 : paramSize = CCU_PARAM_NUM_MAX * CCU_ONE_PARAM_SIZE;
204 :
205 2 : CcuRep::Variable queueIdx = CreateVariable();
206 2 : queueIdx = 0;
207 2 : CcuRep::Variable queueEnd = CreateVariable();
208 2 : queueEnd = CCU_TASK_NUM_MAX;
209 2 : CcuRep::Variable one = CreateVariable();
210 2 : one = 1;
211 : // 存放《每轮算子参数》的寄存器
212 2 : array<CcuRep::Variable, CCU_PARAM_NUM_PER_DIE> param;
213 66 : for (uint32_t i = 0; i < CCU_PARAM_NUM_PER_DIE; ++i) {
214 64 : param[i] = CreateContinuousVariable();
215 : }
216 :
217 6 : CCU_WHILE(repeatCond == 0)
218 : {
219 : // 在context中依次加入8轮指令
220 2 : if (waitAddr_ > (UINT64_MAX - (CCU_TASK_NUM_MAX - 1) * CCU_ONE_PARAM_SIZE)
221 2 : || recordAddr_ > (UINT64_MAX - (CCU_TASK_NUM_MAX - 1) * CCU_ONE_PARAM_SIZE)
222 2 : || paramAddr_ > (UINT64_MAX - (CCU_TASK_NUM_MAX - 1) * CCU_PARAM_NUM_MAX * CCU_ONE_PARAM_SIZE)) {
223 0 : THROW<InvalidParamsException>("MC2 High Level API SetDieNum Failed: integer overflow occurs");
224 : }
225 : // 等待本轮开始信号
226 2 : WaitTurnStartSig(waitAddr, turnStartSig);
227 :
228 : // 读取本轮参数
229 2 : LoadFuncParamFromMemory(paramAddr, param);
230 :
231 2 : MissionPreSync(param[HBM_PARAM_IDX_0]);
232 :
233 : // 第一个参数为opCode, 如果参数中opCode非法则跳出循环队列
234 4 : CCU_IF(param[HBM_PARAM_IDX_0] == INVALID_U64) { CCU_BREAK; }
235 :
236 : // 调用OpSelector
237 : std::string funcName
238 2 : = OP_SELECTOR_LABEL + "_" + std::to_string(GetDieId()) + "_" + std::to_string(missionIndex);
239 2 : auto selectFunc = Func(funcName);
240 2 : selectFunc.SetInArg(param[HBM_PARAM_IDX_0]);
241 2 : selectFunc.SetOutArg(opAddr);
242 2 : selectFunc.AppendToContext();
243 :
244 : // 检查OpSelector是否命中算子,如果没命中则跳出循环队列
245 4 : CCU_IF(opAddr == INVALID_U64) { CCU_BREAK; }
246 :
247 : // 调用算子Func
248 2 : auto opFunc = Func(opAddr);
249 : // 传入参 param[1-31] + token,token需要放在第三个
250 2 : opFunc.SetInArg(param[HBM_PARAM_IDX_1]);
251 2 : opFunc.SetInArg(param[HBM_PARAM_IDX_2]);
252 2 : opFunc.SetInArg(token);
253 60 : for (uint32_t i = HBM_PARAM_IDX_3; i < CCU_PARAM_NUM_PER_DIE; ++i) {
254 58 : opFunc.SetInArg(param[i]);
255 : }
256 2 : opFunc.AppendToContext();
257 :
258 2 : MissionPostSync();
259 :
260 : // Set本轮完成信号
261 2 : SetTurnEndSig(recordAddr, turnEndSig);
262 2 : waitAddr += ckeSize;
263 2 : recordAddr += ckeSize;
264 2 : paramAddr += paramSize;
265 2 : queueIdx += one;
266 6 : CCU_IF(queueIdx == static_cast<u64>(CCU_TASK_NUM_MAX))
267 : {
268 2 : waitAddr = waitStartAddr;
269 2 : recordAddr = recordStartAddr;
270 2 : paramAddr = paramStartAddr;
271 2 : queueIdx = 0;
272 2 : }
273 4 : }
274 2 : }
275 :
276 2 : void Mc2Context::WaitTurnStartSig(const CcuRep::Variable& hbmSigAddr, CcuRep::Variable& turnStartSig)
277 : {
278 2 : if (dieNum_ == SINGLE_DIE) {
279 : // 单Die场景: 等待HBM中的信号
280 3 : CCU_WHILE(turnStartSig != 1)
281 : {
282 : // 循环读HBM对应地址的信号到Xn,直到Xn中的信号值为1
283 1 : LoadVariable(hbmSigAddr, turnStartSig);
284 1 : }
285 1 : turnStartSig = 0; // reset Xn
286 1 : StoreVariable(turnStartSig, hbmSigAddr); // reset HBM
287 : } else {
288 : // 双Die场景
289 1 : if (GetDieId() == DIE0_ID) {
290 : // 双Die场景Die0: 等待HBM中的信号,收到HBM信号之后再给Die1发信号,通知Die1开始
291 3 : CCU_WHILE(turnStartSig != 1)
292 : {
293 : // 循环读HBM对应地址的信号到Xn,直到Xn中的信号值为1
294 1 : LoadVariable(hbmSigAddr, turnStartSig);
295 1 : }
296 1 : turnStartSig = 0; // reset Xn
297 1 : StoreVariable(turnStartSig, hbmSigAddr); // reset HBM
298 : // 给Die1发开始信号
299 1 : LocalCtxPost(importDieSig, 1);
300 0 : } else if (GetDieId() == DIE1_ID) {
301 : // 双Die场景Die1: 等待Die0的信号
302 0 : LocalWait(exportDieSig, 1); // LocalWait会自动reset CKE
303 : }
304 : }
305 2 : }
306 :
307 2 : void Mc2Context::SetTurnEndSig(const CcuRep::Variable& hbmSigAddr, const CcuRep::Variable& turnEndSig)
308 : {
309 2 : if (dieNum_ == SINGLE_DIE) {
310 : // 单Die场景: Set本轮完成信号到HBM
311 1 : StoreVariable(turnEndSig, hbmSigAddr);
312 : } else {
313 : // 双Die场景
314 1 : if (GetDieId() == DIE0_ID) {
315 : // 双Die场景Die0: 等待Die1执行完成信号,然后Set本轮完成信号到HBM
316 1 : LocalWait(exportDieSig, 1); // LocalWait会自动reset CKE
317 1 : StoreVariable(turnEndSig, hbmSigAddr);
318 0 : } else if (GetDieId() == DIE1_ID) {
319 : // 双Die场景Die1: 通知Die0执行完成
320 0 : LocalCtxPost(importDieSig, 1);
321 : }
322 : }
323 2 : }
324 :
325 2 : void Mc2Context::LoadFuncParamFromMemory(
326 : CcuRep::Variable& paramAddr, array<CcuRep::Variable, CCU_PARAM_NUM_PER_DIE>& param)
327 : {
328 : // 双Die场景Die1需要读后32个参数,其他场景都是读前32个参数
329 2 : CcuRep::Variable doubleDie = CreateVariable();
330 2 : doubleDie = CCU_PARAM_NUM_PER_DIE * CCU_ONE_PARAM_SIZE;
331 2 : CcuRep::Variable addr = CreateVariable();
332 2 : addr = paramAddr;
333 2 : if (dieNum_ == DOUBLE_DIE && GetDieId() == DIE1_ID) {
334 0 : addr += doubleDie;
335 : }
336 :
337 : // 一次性读取本轮32个参数
338 2 : LoadVariable(addr, param[0], CCU_PARAM_NUM_PER_DIE);
339 2 : }
340 :
341 0 : vector<uint64_t> Mc2Context::GeneArgs(const CcuTaskArg& arg)
342 : {
343 0 : const CcuTaskArgMc2* taskArg = dynamic_cast<const CcuTaskArgMc2*>(&arg);
344 0 : uint64_t tokenInfo = taskArg->token_;
345 0 : return {tokenInfo};
346 : }
347 :
348 0 : void Mc2SlaveContext::GenCircularQueue()
349 : {
350 : // 算子签名
351 0 : CcuRep::Variable signature = CreateVariable();
352 : // 存放《选择函数返回的FuncCall地址》的寄存器,选择函数的出参,循环队列内部使用
353 0 : CcuRep::Variable opAddr = CreateVariable();
354 :
355 : // 存放《控制repeat循环执行的条件》的寄存器
356 0 : CcuRep::Variable repeatCond = CreateVariable();
357 0 : repeatCond = 0;
358 :
359 0 : CCU_WHILE(repeatCond == 0)
360 : {
361 : // 在context中依次加入8轮指令
362 0 : MissionPreSync(signature);
363 :
364 : // 第一个参数为opCode, 如果参数中opCode非法则跳出循环队列
365 0 : CCU_IF(signature == INVALID_U64) { CCU_BREAK; }
366 :
367 : // 调用OpSelector
368 : std::string funcName
369 0 : = OP_SELECTOR_LABEL + "_" + std::to_string(GetDieId()) + "_" + std::to_string(missionIndex);
370 0 : auto selectFunc = Func(funcName);
371 0 : selectFunc.SetInArg(signature);
372 0 : selectFunc.SetOutArg(opAddr);
373 0 : selectFunc.AppendToContext();
374 :
375 : // 检查OpSelector是否命中算子,如果没命中则跳出循环队列
376 0 : CCU_IF(opAddr == INVALID_U64) { CCU_BREAK; }
377 :
378 : // 调用算子Func
379 0 : auto opFunc = Func(opAddr);
380 :
381 0 : opFunc.AppendToContext();
382 :
383 0 : MissionPostSync();
384 0 : }
385 0 : }
386 :
387 0 : vector<uint64_t> Mc2SlaveContext::GeneArgs([[maybe_unused]] const CcuTaskArg& arg) { return {}; }
388 :
389 : } // namespace Hccl
|