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(exportMissoinSig[i],
54 4 : "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(exportMissoinSig[0],
66 4 : "slave_sig_" + std::to_string(GetDieId()) + "_" + std::to_string(miIndex));
67 2 : exportMissionVar.push_back(CreateVariable());
68 2 : ExportVariable(exportMissionVar[0],
69 4 : "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 12 : CCU_IF(opCode == entry.first)
123 : {
124 4 : opAddr = entry.second;
125 4 : }
126 : }
127 2 : }
128 2 : }
129 :
130 3 : void Mc2ContextBase::Algorithm()
131 : {
132 3 : GenOpSelector();
133 2 : GenCircularQueue();
134 2 : }
135 :
136 6 : void Mc2Context::SetCommAddr(uint64_t syncAddr, uint64_t paramAddr)
137 : {
138 6 : waitAddr_ = syncAddr;
139 6 : if (syncAddr > (UINT64_MAX - CCU_TASK_NUM_MAX * CCU_ONE_PARAM_SIZE)) {
140 0 : THROW<InvalidParamsException>("MC2 High Level API SetDieNum Failed: integer overflow occurs");
141 : }
142 6 : recordAddr_ = syncAddr + CCU_TASK_NUM_MAX * CCU_ONE_PARAM_SIZE; // 偏移8轮的总宽度
143 6 : paramAddr_ = paramAddr;
144 6 : }
145 :
146 7 : void Mc2Context::SetDieNum(uint32_t dieNum)
147 : {
148 7 : dieNum_ = dieNum;
149 : // 参数合法值判断
150 7 : bool isDieNumValid = (dieNum_ == SINGLE_DIE || dieNum_ == DOUBLE_DIE);
151 7 : bool isDieIdValid = (GetDieId() == DIE0_ID || GetDieId() == DIE1_ID);
152 7 : if (!(isDieNumValid && isDieIdValid)) {
153 1 : THROW<InvalidParamsException>("MC2 High Level API SetDieNum Failed: Invalid Die Config");
154 : }
155 :
156 6 : if (dieNum_ == DOUBLE_DIE) { // 双Die场景才需要导入导出
157 : // 导出信号
158 4 : exportDieSig = CreateMaskSignal();
159 : // Die0: export完成信号给Die1,Die1: export开始信号给Die0
160 4 : const string &exportSigLabel = (GetDieId() == DIE1_ID) ? DIE1_START_SIG : DIE1_END_SIG;
161 4 : ExportMaskSignal(exportDieSig, exportSigLabel);
162 :
163 : // 导入信号
164 4 : const string &importSigLabel = (GetDieId() == DIE1_ID) ? DIE1_END_SIG : DIE1_START_SIG;
165 4 : importDieSig = ImportMaskSignal(importSigLabel);
166 : }
167 6 : }
168 :
169 2 : void Mc2Context::GenCircularQueue()
170 : {
171 : // 存放Token的寄存器
172 2 : CcuRep::Variable token = CreateVariable();
173 : // 从SQE中载入Token
174 2 : Load(token);
175 :
176 : // 存放《选择函数返回的FuncCall地址》的寄存器,选择函数的出参,循环队列内部使用
177 2 : CcuRep::Variable opAddr = CreateVariable();
178 :
179 : // 存放《控制repeat循环执行的条件》的寄存器
180 2 : CcuRep::Variable repeatCond = CreateVariable();
181 2 : repeatCond = 0;
182 :
183 : // 存放《轮次执行开始信号》的寄存器,初值为 0
184 2 : CcuRep::Variable turnStartSig = CreateVariable();
185 2 : turnStartSig = 0;
186 : // 存放《轮次执行完成信号》的寄存器,在循环中固定为 1
187 2 : CcuRep::Variable turnEndSig = CreateVariable();
188 2 : turnEndSig = 1;
189 :
190 2 : CcuRep::Variable waitStartAddr = CreateVariable();
191 2 : waitStartAddr = waitAddr_;
192 2 : CcuRep::Variable recordStartAddr = CreateVariable();
193 2 : recordStartAddr = recordAddr_;
194 2 : CcuRep::Variable paramStartAddr = CreateVariable();
195 2 : paramStartAddr = paramAddr_;
196 2 : CcuRep::Variable waitAddr = CreateVariable();
197 2 : waitAddr = waitAddr_;
198 2 : CcuRep::Variable recordAddr = CreateVariable();
199 2 : recordAddr = recordAddr_;
200 2 : CcuRep::Variable paramAddr = CreateVariable();
201 2 : paramAddr = paramAddr_;
202 :
203 2 : CcuRep::Variable ckeSize = CreateVariable();
204 2 : ckeSize = CCU_ONE_PARAM_SIZE;
205 2 : CcuRep::Variable paramSize = CreateVariable();
206 2 : paramSize = CCU_PARAM_NUM_MAX * CCU_ONE_PARAM_SIZE;
207 :
208 2 : CcuRep::Variable queueIdx = CreateVariable();
209 2 : queueIdx = 0;
210 2 : CcuRep::Variable queueEnd = CreateVariable();
211 2 : queueEnd = CCU_TASK_NUM_MAX;
212 2 : CcuRep::Variable one = CreateVariable();
213 2 : one = 1;
214 : // 存放《每轮算子参数》的寄存器
215 2 : array<CcuRep::Variable, CCU_PARAM_NUM_PER_DIE> param;
216 66 : for (uint32_t i = 0; i < CCU_PARAM_NUM_PER_DIE; ++i) {
217 64 : param[i] = CreateContinuousVariable();
218 : }
219 :
220 6 : CCU_WHILE(repeatCond == 0)
221 : {
222 : // 在context中依次加入8轮指令
223 2 : if (waitAddr_ > (UINT64_MAX - (CCU_TASK_NUM_MAX - 1) * CCU_ONE_PARAM_SIZE)
224 2 : || recordAddr_ > (UINT64_MAX - (CCU_TASK_NUM_MAX - 1) * CCU_ONE_PARAM_SIZE)
225 2 : || paramAddr_ > (UINT64_MAX - (CCU_TASK_NUM_MAX - 1) * CCU_PARAM_NUM_MAX * CCU_ONE_PARAM_SIZE)) {
226 0 : THROW<InvalidParamsException>("MC2 High Level API SetDieNum Failed: integer overflow occurs");
227 : }
228 : // 等待本轮开始信号
229 2 : WaitTurnStartSig(waitAddr, turnStartSig);
230 :
231 : // 读取本轮参数
232 2 : LoadFuncParamFromMemory(paramAddr, param);
233 :
234 2 : MissionPreSync(param[HBM_PARAM_IDX_0]);
235 :
236 : // 第一个参数为opCode, 如果参数中opCode非法则跳出循环队列
237 6 : CCU_IF(param[HBM_PARAM_IDX_0] == INVALID_U64)
238 : {
239 2 : CCU_BREAK;
240 2 : }
241 :
242 : // 调用OpSelector
243 : std::string funcName
244 2 : = OP_SELECTOR_LABEL + "_" + std::to_string(GetDieId()) + "_" + std::to_string(missionIndex);
245 2 : auto selectFunc = Func(funcName);
246 2 : selectFunc.SetInArg(param[HBM_PARAM_IDX_0]);
247 2 : selectFunc.SetOutArg(opAddr);
248 2 : selectFunc.AppendToContext();
249 :
250 : // 检查OpSelector是否命中算子,如果没命中则跳出循环队列
251 6 : CCU_IF(opAddr == INVALID_U64)
252 : {
253 2 : CCU_BREAK;
254 2 : }
255 :
256 : // 调用算子Func
257 2 : auto opFunc = Func(opAddr);
258 : // 传入参 param[1-31] + token,token需要放在第三个
259 2 : opFunc.SetInArg(param[HBM_PARAM_IDX_1]);
260 2 : opFunc.SetInArg(param[HBM_PARAM_IDX_2]);
261 2 : opFunc.SetInArg(token);
262 60 : for (uint32_t i = HBM_PARAM_IDX_3; i < CCU_PARAM_NUM_PER_DIE; ++i) {
263 58 : opFunc.SetInArg(param[i]);
264 : }
265 2 : opFunc.AppendToContext();
266 :
267 2 : MissionPostSync();
268 :
269 : // Set本轮完成信号
270 2 : SetTurnEndSig(recordAddr, turnEndSig);
271 2 : waitAddr += ckeSize;
272 2 : recordAddr += ckeSize;
273 2 : paramAddr += paramSize;
274 2 : queueIdx += one;
275 6 : CCU_IF (queueIdx == static_cast<u64>(CCU_TASK_NUM_MAX)) {
276 2 : waitAddr = waitStartAddr;
277 2 : recordAddr = recordStartAddr;
278 2 : paramAddr = paramStartAddr;
279 2 : queueIdx = 0;
280 2 : }
281 4 : }
282 2 : }
283 :
284 2 : void Mc2Context::WaitTurnStartSig(const CcuRep::Variable &hbmSigAddr, CcuRep::Variable &turnStartSig)
285 : {
286 2 : if (dieNum_ == SINGLE_DIE) {
287 : // 单Die场景: 等待HBM中的信号
288 3 : CCU_WHILE(turnStartSig != 1)
289 : {
290 : // 循环读HBM对应地址的信号到Xn,直到Xn中的信号值为1
291 1 : LoadVariable(hbmSigAddr, turnStartSig);
292 1 : }
293 1 : turnStartSig = 0; // reset Xn
294 1 : StoreVariable(turnStartSig, hbmSigAddr); // reset HBM
295 : } else {
296 : // 双Die场景
297 1 : if (GetDieId() == DIE0_ID) {
298 : // 双Die场景Die0: 等待HBM中的信号,收到HBM信号之后再给Die1发信号,通知Die1开始
299 3 : CCU_WHILE(turnStartSig != 1)
300 : {
301 : // 循环读HBM对应地址的信号到Xn,直到Xn中的信号值为1
302 1 : LoadVariable(hbmSigAddr, turnStartSig);
303 1 : }
304 1 : turnStartSig = 0; // reset Xn
305 1 : StoreVariable(turnStartSig, hbmSigAddr); // reset HBM
306 : // 给Die1发开始信号
307 1 : LocalCtxPost(importDieSig, 1);
308 0 : } else if (GetDieId() == DIE1_ID) {
309 : // 双Die场景Die1: 等待Die0的信号
310 0 : LocalWait(exportDieSig, 1); // LocalWait会自动reset CKE
311 : }
312 : }
313 2 : }
314 :
315 2 : void Mc2Context::SetTurnEndSig(const CcuRep::Variable &hbmSigAddr, const CcuRep::Variable &turnEndSig)
316 : {
317 2 : if (dieNum_ == SINGLE_DIE) {
318 : // 单Die场景: Set本轮完成信号到HBM
319 1 : StoreVariable(turnEndSig, hbmSigAddr);
320 : } else {
321 : // 双Die场景
322 1 : if (GetDieId() == DIE0_ID) {
323 : // 双Die场景Die0: 等待Die1执行完成信号,然后Set本轮完成信号到HBM
324 1 : LocalWait(exportDieSig, 1); // LocalWait会自动reset CKE
325 1 : StoreVariable(turnEndSig, hbmSigAddr);
326 0 : } else if (GetDieId() == DIE1_ID) {
327 : // 双Die场景Die1: 通知Die0执行完成
328 0 : LocalCtxPost(importDieSig, 1);
329 : }
330 : }
331 2 : }
332 :
333 2 : void Mc2Context::LoadFuncParamFromMemory(CcuRep::Variable ¶mAddr, array<CcuRep::Variable, CCU_PARAM_NUM_PER_DIE> ¶m)
334 : {
335 : // 双Die场景Die1需要读后32个参数,其他场景都是读前32个参数
336 2 : CcuRep::Variable doubleDie = CreateVariable();
337 2 : doubleDie = CCU_PARAM_NUM_PER_DIE * CCU_ONE_PARAM_SIZE;
338 2 : CcuRep::Variable addr = CreateVariable();
339 2 : addr = paramAddr;
340 2 : if (dieNum_ == DOUBLE_DIE && GetDieId() == DIE1_ID) {
341 0 : addr += doubleDie;
342 : }
343 :
344 : // 一次性读取本轮32个参数
345 2 : LoadVariable(addr, param[0], CCU_PARAM_NUM_PER_DIE);
346 2 : }
347 :
348 0 : vector<uint64_t> Mc2Context::GeneArgs(const CcuTaskArg &arg)
349 : {
350 0 : const CcuTaskArgMc2 *taskArg = dynamic_cast<const CcuTaskArgMc2 *>(&arg);
351 0 : uint64_t tokenInfo = taskArg->token_;
352 0 : return {tokenInfo};
353 : }
354 :
355 0 : void Mc2SlaveContext::GenCircularQueue()
356 : {
357 : // 算子签名
358 0 : CcuRep::Variable signature = CreateVariable();
359 : // 存放《选择函数返回的FuncCall地址》的寄存器,选择函数的出参,循环队列内部使用
360 0 : CcuRep::Variable opAddr = CreateVariable();
361 :
362 : // 存放《控制repeat循环执行的条件》的寄存器
363 0 : CcuRep::Variable repeatCond = CreateVariable();
364 0 : repeatCond = 0;
365 :
366 0 : CCU_WHILE(repeatCond == 0)
367 : {
368 : // 在context中依次加入8轮指令
369 0 : MissionPreSync(signature);
370 :
371 : // 第一个参数为opCode, 如果参数中opCode非法则跳出循环队列
372 0 : CCU_IF(signature == INVALID_U64)
373 : {
374 0 : CCU_BREAK;
375 0 : }
376 :
377 : // 调用OpSelector
378 : std::string funcName
379 0 : = OP_SELECTOR_LABEL + "_" + std::to_string(GetDieId()) + "_" + std::to_string(missionIndex);
380 0 : auto selectFunc = Func(funcName);
381 0 : selectFunc.SetInArg(signature);
382 0 : selectFunc.SetOutArg(opAddr);
383 0 : selectFunc.AppendToContext();
384 :
385 : // 检查OpSelector是否命中算子,如果没命中则跳出循环队列
386 0 : CCU_IF(opAddr == INVALID_U64)
387 : {
388 0 : CCU_BREAK;
389 0 : }
390 :
391 : // 调用算子Func
392 0 : auto opFunc = Func(opAddr);
393 :
394 0 : opFunc.AppendToContext();
395 :
396 0 : MissionPostSync();
397 0 : }
398 0 : }
399 :
400 0 : vector<uint64_t> Mc2SlaveContext::GeneArgs(const CcuTaskArg &arg)
401 : {
402 0 : return {};
403 : }
404 :
405 : } // namespace Hccl
|