LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/all_to_all - ccu_context_all_to_all_mesh1d_2Die.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 231 0
Test Date: 2026-07-28 12:11:00 Functions: 0.0 % 11 0

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2026 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 "ccu_context_all_to_all_mesh1d_2Die.h"
      12              : #include "ccu_instruction_all_to_all_mesh1d_2Die.h"
      13              : 
      14              : namespace Hccl {
      15              : 
      16              : constexpr int CKE_IDX_0   = 0;
      17              : constexpr int CKE_IDX_1   = 1;
      18              : constexpr int CKE_IDX_2   = 2;
      19              : constexpr int INPUT_XN_ID = 0;
      20              : constexpr int OUPUT_XN_ID = 1;
      21              : constexpr int TOKEN_XN_ID = 2;
      22              : 
      23              : constexpr uint64_t CCU_MS_SIZE   = 4096;
      24              : constexpr uint64_t LOCAL_COPY_MS = 8;
      25              : 
      26            0 : CcuContextAllToAllMesh1D2Die::CcuContextAllToAllMesh1D2Die(const CcuCtxArg                   &arg,
      27              :                                                            const std::vector<CcuTransport *> &transports,
      28            0 :                                                            const CcuTransportGroup           &group)
      29            0 :     : CcuContextAlgBase(arg, transports, group)
      30              : {
      31            0 :     const CcuCtxArgAllToAllMesh1D2Die *ctxArg = dynamic_cast<const CcuCtxArgAllToAllMesh1D2Die *>(&arg);
      32            0 :     if (ctxArg == nullptr) {
      33            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::ctxArg ptr is null"));
      34              :     }
      35              : 
      36            0 :     rankId_     = ctxArg->rankId_;
      37            0 :     withMyRank_ = ctxArg->withMyRank_;
      38            0 :     rankGroup_  = ctxArg->rankGroup;
      39            0 :     if (ctxArg->dimSize_.size() > 0) {
      40            0 :         rankSize_ = ctxArg->dimSize_[0];
      41              :     }
      42            0 :     bitNumPerCKE_ = ctxArg->bitNum_;
      43            0 : }
      44              : 
      45            0 : void CcuContextAllToAllMesh1D2Die::InitResource()
      46              : {
      47              :     // 创建Variable,用于交换地址及token
      48            0 :     u32 transportId = 0;
      49            0 :     if (transports.size() == 0) {
      50            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die transports is empty"));
      51              :     }
      52            0 :     virRankSize = transports.size() + 1;
      53              : 
      54            0 :     for (u64 id = 0; id < transports.size(); id++) {
      55              :         // 非本地,使用远端Variable
      56            0 :         CHK_PRT_RET(transports[transportId] == nullptr,
      57              :                     HCCL_ERROR("[CcuContextAllToAllMesh1D2Die] Algorithm transport ptr is null"), );
      58            0 :         input_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_0));
      59            0 :         output_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_1));
      60            0 :         token_.push_back(CreateVariable((*transports[transportId]), CKE_IDX_2));
      61            0 :         transportId++;
      62              :     }
      63              :     // 最后一个位置放自己地址
      64            0 :     input_.push_back(CreateVariable());
      65            0 :     output_.push_back(CreateVariable());
      66            0 :     token_.push_back(CreateVariable());
      67              : 
      68            0 :     sliceSize_         = CreateVariable();
      69            0 :     inputSliceStride_  = CreateVariable();
      70            0 :     outputoffset_ = CreateVariable();
      71            0 :     outBuffBaseOff_    = CreateVariable();
      72            0 :     groupOpSize_       = CreateGroupOpSize();
      73              : 
      74            0 :     moConfig.loopCount    = 8;                           // loop展开8次、16次
      75            0 :     moConfig.msInterleave = LOCAL_COPY_MS;               // 一个loop 8个MS
      76            0 :     moConfig.memSlice     = LOCAL_COPY_MS * CCU_MS_SIZE; // 32k
      77            0 :     if (moRes.executor.size() == 0) {
      78            0 :         moRes.executor   = CreateBlockExecutor(moConfig.loopCount);
      79            0 :         moRes.maskSignal = CreateBlockMaskSignal(moConfig.loopCount);
      80            0 :         moRes.ccuBuffer  = CreateBlockCcuBuffer(moConfig.loopCount * moConfig.msInterleave);
      81              :     }
      82              : 
      83            0 :     logicRankSize = withMyRank_ ? transports.size() + 1 : transports.size();
      84            0 :     signalNum_    = (rankSize_ + bitNumPerCKE_ - 1) / bitNumPerCKE_;
      85            0 :     HCCL_INFO("[CcuContextAlltoAll2Die] CtxArg: rankId_[%u], rankSize_[%u], signalNum_[%u]", rankId_, rankSize_, signalNum_);
      86            0 :     return;
      87              : }
      88              : 
      89            0 : void CcuContextAllToAllMesh1D2Die::LoadArgs()
      90              : {
      91              :     // 从SQE load args,本rank需要的input、output地址等信息
      92              :     // inputAddr, outputAddr, tokenInfo, srcStride, srcOffset, dstOffset, groupOpSize
      93            0 :     Load(input_[virRankSize - 1]);
      94            0 :     Load(output_[virRankSize - 1]);
      95            0 :     Load(token_[virRankSize - 1]);
      96            0 :     Load(sliceSize_); // 本轮传输的分片大小
      97            0 :     Load(inputSliceStride_);
      98            0 :     Load(outputoffset_);
      99            0 :     Load(outBuffBaseOff_);
     100            0 :     Load(groupOpSize_);
     101            0 :     return;
     102              : }
     103              : 
     104            0 : void CcuContextAllToAllMesh1D2Die::PreSync()
     105              : {
     106            0 :     if (withMyRank_) {
     107            0 :         uint16_t logicId = rankId_ % logicRankSize;
     108            0 :         selfBit = 1 << logicId;
     109            0 :         allBit = ((1 << logicRankSize) - 1) & (~(1 << logicId));
     110            0 :         for (auto t : transports) {
     111              :             // (transport, param, paramID, SemID, mask)
     112            0 :             WriteVariableWithSignal(*t, output_[virRankSize - 1], OUPUT_XN_ID, CKE_IDX_1,
     113            0 :                                     selfBit); // index = 1,传递output信息
     114            0 :             WriteVariableWithSignal(*t, token_[virRankSize - 1], TOKEN_XN_ID, CKE_IDX_2,
     115            0 :                                     selfBit); // index = 2,传递token信息
     116              :         }
     117            0 :         GroupWait(*transportGroup, CKE_IDX_1, allBit); // index = 1,传递output信息
     118            0 :         GroupWait(*transportGroup, CKE_IDX_2, allBit); // index = 2,传递token信息
     119              :     } else {
     120            0 :         uint16_t selfSignalId = rankId_ / bitNumPerCKE_;
     121            0 :         uint16_t selfBit      = 1 << (rankId_ % bitNumPerCKE_);
     122            0 :         for (auto t : transports) {
     123              :             // (transport, param, paramID, SemID, mask)
     124            0 :             WriteVariableWithSignal(*t, output_[virRankSize - 1], OUPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1,
     125              :                                     selfBit); // index = 1,传递output信息
     126            0 :             WriteVariableWithSignal(*t, token_[virRankSize - 1], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2,
     127              :                                     selfBit); // index = 2,传递token信息
     128              :         }
     129            0 :         std::vector<uint16_t> waitBitVector(signalNum_, 0);
     130            0 :         for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
     131            0 :             waitBitVector[sId] = (1 << bitNumPerCKE_) - 1;
     132            0 :             if (sId == selfSignalId) {
     133            0 :                 waitBitVector[sId] = 0;
     134              :             }
     135            0 :             GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]); // index = 1,传递output信息
     136            0 :             GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]); // index = 2,传递token信息
     137              :         }
     138            0 :     }
     139            0 :     return;
     140              : }
     141              : 
     142            0 : void CcuContextAllToAllMesh1D2Die::PostSync()
     143              : {
     144            0 :     if (withMyRank_) {
     145            0 :         uint16_t logicId = rankId_ % logicRankSize;
     146            0 :         selfBit = 1 << logicId;
     147            0 :         allBit = ((1 << logicRankSize) - 1) & (~(1 << logicId));
     148            0 :         for (auto t : transports) {
     149            0 :             if (t == nullptr) {
     150            0 :                 THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::Algorithm transport ptr is null"));
     151              :             }
     152            0 :             RemotePost(*t, CKE_IDX_0, selfBit);
     153              :         }
     154            0 :         GroupWait(*transportGroup, CKE_IDX_0, allBit);
     155              :     } else {
     156            0 :         uint16_t selfSignalId = rankId_ / bitNumPerCKE_;
     157            0 :         uint16_t selfBit      = 1 << (rankId_ % bitNumPerCKE_);
     158            0 :         for (auto t : transports) {
     159            0 :             if (t == nullptr) {
     160            0 :                 THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::Algorithm transport ptr is null"));
     161              :             }
     162            0 :             RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
     163              :         }
     164            0 :         std::vector<uint16_t> waitBitVector(signalNum_, 0);
     165            0 :         for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
     166            0 :             waitBitVector[sId] = (1 << bitNumPerCKE_) - 1;
     167            0 :             if (sId == selfSignalId) {
     168            0 :                 waitBitVector[sId] = 0;
     169              :             }
     170            0 :             GroupWait(*transportGroup, CKE_IDX_0, allBit);
     171              :         }
     172            0 :     }
     173              :     
     174            0 :     return;
     175              : }
     176              : 
     177            0 : uint32_t CcuContextAllToAllMesh1D2Die::CalcDstRank(uint32_t peerId) const
     178              : {
     179            0 :     if (peerId > rankGroup_.size()) {
     180            0 :         THROW<InvalidParamsException>(
     181            0 :             StringFormat("[CcuContextAllToAllMesh1D2Die][CalcDstRank] Unexpected peerId[%u]", peerId));
     182              :     }
     183            0 :     return rankGroup_[peerId];
     184              : }
     185              : 
     186            0 : void CcuContextAllToAllMesh1D2Die::DoRepeatAllToAll()
     187              : {
     188              :     // 创建GSA, src为本地的各片HBM地址GSA列表,dst为所有对端的HBM地址GSA列表
     189            0 :     std::vector<CcuRep::Memory> src;
     190            0 :     for (uint64_t rankIdx = 0; rankIdx < logicRankSize; rankIdx++) {
     191            0 :         src.push_back(CreateMemory());
     192              :     }
     193            0 :     std::vector<CcuRep::Memory> dst;
     194            0 :     for (uint64_t rankIdx = 0; rankIdx < logicRankSize; rankIdx++) {
     195            0 :         dst.push_back(CreateMemory());
     196              :     }
     197              : 
     198              :     // 考虑stride信息
     199            0 :     for (uint64_t r = 0; r < logicRankSize; r++) {
     200            0 :         const u32 dstRank = CalcDstRank(r);
     201              : 
     202            0 :         src[r].token = token_[r];
     203            0 :         dst[r].token = token_[r];
     204              : 
     205            0 :         src[r].addr = input_[virRankSize - 1];
     206            0 :         dst[r].addr = output_[r];
     207            0 :         dst[r].addr += outputoffset_;
     208            0 :         for(uint64_t i = 0; i < dstRank; i++){
     209            0 :             src[r].addr += inputSliceStride_;
     210              :         }
     211              :     }
     212              : 
     213              :     //  all2all 数据搬运
     214            0 :     u32 transportIdx = 0;
     215            0 :     if (withMyRank_) {
     216            0 :         uint64_t allBit_ = withMyRank_ ? ((1 << logicRankSize) - 1) & (~(1 << transports.size())) : (1 << logicRankSize) - 1;
     217            0 :         CcuRep::MaskSignal locMask = CreateMaskSignal();
     218            0 :         for (uint64_t r = 0; r < logicRankSize; r++) {
     219            0 :             if (withMyRank_ && r == logicRankSize - 1) {
     220            0 :                 LocalCopyByLoopGroup(dst[r], src[r]);
     221            0 :                 continue;
     222              :             }
     223            0 :             Write(*transports[transportIdx], dst[r], src[r], sliceSize_, locMask, 1 << r);
     224            0 :             transportIdx++;
     225              :         }
     226            0 :         LocalWait(locMask, allBit_);
     227            0 :     } else {
     228            0 :         vector<CcuRep::MaskSignal> locMask;
     229            0 :         uint16_t signalNum = logicRankSize / bitNumPerCKE_;
     230            0 :         std::vector<uint16_t> waitBitVector(signalNum, 0);
     231            0 :         for (uint16_t sId = 0; sId < signalNum; sId++) {
     232            0 :             locMask.push_back(CreateMaskSignal());
     233              :         }
     234            0 :         for (uint16_t r = 0; r < logicRankSize; r++) {
     235            0 :             uint16_t rmtSignalId = r / bitNumPerCKE_;
     236            0 :             uint16_t rmtSignalBit = 1 << (r % bitNumPerCKE_);
     237            0 :             Write(*transports[transportIdx], dst[r], src[r], sliceSize_, locMask[rmtSignalId], rmtSignalBit);
     238            0 :             transportIdx++;
     239              :         }
     240            0 :         for (uint16_t sId = 0; sId < signalNum; sId++) {
     241            0 :             waitBitVector[sId] = (1 << bitNumPerCKE_) - 1;
     242            0 :             LocalWait(locMask[sId], waitBitVector[sId]);
     243              :         }
     244            0 :     }
     245            0 : }
     246              : 
     247            0 : void CcuContextAllToAllMesh1D2Die::CreateLocalCopyLoop()
     248              : {
     249            0 :     std::string loopType = "all_to_all";
     250            0 :     if (registeredLoop.find(loopType) != registeredLoop.end()) {
     251            0 :         return;
     252              :     }
     253              : 
     254            0 :     for (uint32_t index = 0; index < 2; index++) { // 需要2个Loop
     255            0 :         CcuRep::Variable  len = CreateVariable();
     256            0 :         CcuRep::Memory    src = CreateMemory();
     257            0 :         CcuRep::Memory    dst = CreateMemory();
     258            0 :         CcuRep::LoopBlock lb(this, loopType + "_localcopy_loop_" + std::to_string(index));
     259            0 :         lb(src, dst, len);
     260              : 
     261            0 :         std::vector<CcuRep::CcuBuffer> bufs;
     262            0 :         CcuRep::MaskSignal             sem = moRes.maskSignal[index];
     263            0 :         for (uint32_t i = 0; i < LOCAL_COPY_MS; i++) {
     264            0 :             bufs.push_back(moRes.ccuBuffer[i]);
     265              :         }
     266              : 
     267            0 :         LocalCopy(bufs[0], src, len, sem);
     268            0 :         LocalWait(sem);
     269            0 :         LocalCopy(dst, bufs[0], len, sem);
     270            0 :         LocalWait(sem);
     271            0 :     }
     272            0 :     registeredLoop.insert(loopType);
     273            0 :     return;
     274            0 : }
     275              : 
     276            0 : void CcuContextAllToAllMesh1D2Die::LocalCopyByLoopGroup(CcuRep::Memory dst, CcuRep::Memory src)
     277              : {
     278            0 :     CreateLocalCopyLoop();
     279              : 
     280            0 :     CCU_IF(groupOpSize_.addrOffset != 0)
     281              :     {
     282            0 :         CcuRep::Variable loopParam = CreateVariable();
     283            0 :         loopParam                  = CcuRep::GetLoopParam(0, moConfig.memSlice * moConfig.loopCount, 0);
     284            0 :         loopParam += groupOpSize_.loopParam;
     285              : 
     286            0 :         CcuRep::Variable sliceSize = CreateVariable();
     287            0 :         sliceSize                  = moConfig.memSlice;
     288            0 :         auto lc                    = Loop("all_to_all_localcopy_loop_0")(src, dst, sliceSize);
     289              : 
     290            0 :         CcuRep::Variable paraCfg   = CreateVariable();
     291            0 :         paraCfg                    = CcuRep::GetParallelParam(moConfig.loopCount - 1, 0, 1);
     292            0 :         CcuRep::Variable offsetCfg = CreateVariable();
     293            0 :         offsetCfg                  = CcuRep::GetOffsetParam(moConfig.memSlice, moConfig.msInterleave, 1);
     294            0 :         LoopGroup({lc}, {loopParam}, paraCfg, offsetCfg);
     295            0 :     }
     296              : 
     297            0 :     CCU_IF(groupOpSize_.parallelParam != 0)
     298              :     {
     299            0 :         CcuRep::Condition cond(this, groupOpSize_.parallelParam != 0);
     300              : 
     301            0 :         src.addr += groupOpSize_.addrOffset;
     302            0 :         dst.addr += groupOpSize_.addrOffset;
     303            0 :         auto lc0 = Loop("all_to_all_localcopy_loop_0")(src, dst, groupOpSize_.residual);
     304              : 
     305            0 :         src.addr += groupOpSize_.residual;
     306            0 :         dst.addr += groupOpSize_.residual;
     307            0 :         CcuRep::Variable sliceSize = CreateVariable();
     308            0 :         sliceSize                  = moConfig.memSlice;
     309            0 :         auto lc1                   = Loop("all_to_all_localcopy_loop_1")(src, dst, sliceSize);
     310              : 
     311            0 :         CcuRep::Variable loopCfg0  = CreateVariable();
     312            0 :         loopCfg0                   = CcuRep::GetLoopParam(0, 0, 1);
     313            0 :         CcuRep::Variable loopCfg1  = CreateVariable();
     314            0 :         loopCfg1                   = CcuRep::GetLoopParam(0, 0, 1);
     315            0 :         CcuRep::Variable offsetCfg = CreateVariable();
     316            0 :         offsetCfg                  = CcuRep::GetOffsetParam(moConfig.memSlice, moConfig.msInterleave, 1);
     317            0 :         LoopGroup({lc0, lc1}, {loopCfg0, loopCfg1}, groupOpSize_.parallelParam, offsetCfg);
     318            0 :     }
     319            0 : }
     320              : 
     321            0 : void CcuContextAllToAllMesh1D2Die::Algorithm()
     322              : {
     323            0 :     HCCL_INFO("[ccuAllToAllMesh1D2Die_context] AllToAllMesh1D2Die run.");
     324            0 :     InitResource();
     325              : 
     326            0 :     LoadArgs();
     327              : 
     328            0 :     PreSync();
     329              : 
     330            0 :     DoRepeatAllToAll();
     331              : 
     332            0 :     PostSync();
     333            0 :     HCCL_INFO("[ccuAllToAllMesh1D2Die_context] AllToAllMesh1D2Die end.");
     334            0 :     return;
     335              : }
     336              : 
     337            0 : std::vector<uint64_t> CcuContextAllToAllMesh1D2Die::GeneArgs(const CcuTaskArg &arg)
     338              : {
     339            0 :     const CcuTaskArgAllToAllMesh1D2Die *taskArg = dynamic_cast<const CcuTaskArgAllToAllMesh1D2Die *>(&arg);
     340            0 :     if (taskArg == nullptr) {
     341            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllToAllMesh1D2Die::taskArg ptr is null"));
     342              :     }
     343            0 :     uint64_t inputAddr         = taskArg->inputAddr_;
     344            0 :     uint64_t outputAddr        = taskArg->outputAddr_;
     345            0 :     uint64_t tokenInfo         = taskArg->token_;
     346            0 :     uint64_t sliceSize         = taskArg->sliceSize_;
     347            0 :     uint64_t inputSliceStride  = taskArg->inputSliceStride_;
     348            0 :     uint64_t outputSliceStride = taskArg->outputSliceStride_ * rankId_;
     349            0 :     uint64_t outBuffBaseOff    = taskArg->outBuffBaseOff_;
     350              : 
     351            0 :     auto goSize = CalGoSize(sliceSize);
     352            0 :     HCCL_INFO("[CcuContextAllToAllMesh1D2Die] inputAddr[%llu], outputAddr[%llu], sliceSize[%llu], "
     353              :               "inputSliceStride[%llu], outputSliceStride[%llu], outBuffBaseOff[%llu].",
     354              :               inputAddr, outputAddr, sliceSize, inputSliceStride, outputSliceStride, outBuffBaseOff);
     355              : 
     356              :     return {inputAddr,      outputAddr, tokenInfo, sliceSize, inputSliceStride, outputSliceStride,
     357            0 :             outBuffBaseOff, goSize[0],  goSize[1], goSize[2], goSize[3]};
     358            0 : }
     359              : 
     360              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1