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: 79.4 % 223 177
Test Date: 2026-08-04 10:52:23 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(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 &paramAddr, array<CcuRep::Variable, CCU_PARAM_NUM_PER_DIE> &param)
     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
        

Generated by: LCOV version 2.0-1