LCOV - code coverage report
Current view: top level - legacy/ascend950/framework/ccu/ccu_mc2 - mc2_context.cpp (source / functions) Coverage Total Hit
Test: coverage.info Lines: 80.7 % 212 171
Test Date: 2026-08-18 17:47:01 Functions: 80.0 % 15 12

            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
        

Generated by: LCOV version 2.0-1