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

Generated by: LCOV version 2.0-1