LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/all_gather - ccu_context_all_gather_nhr1d_mem2mem.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 225 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 11 0

            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 "ccu_context_all_gather_nhr1d_mem2mem.h"
      12              : 
      13              : namespace Hccl {
      14              : 
      15              : constexpr uint16_t OUTPUT_XN_ID = 1;
      16              : constexpr uint16_t TOKEN_XN_ID = 2;
      17              : constexpr uint16_t FST_AXIS_ID = 0;
      18              : constexpr uint16_t SEC_AXIS_ID = 1;
      19              : constexpr uint16_t CKE_IDX_0 = 0;
      20              : constexpr uint16_t CKE_IDX_1 = 1;
      21              : constexpr uint16_t CKE_IDX_2 = 2;
      22              : constexpr uint16_t CKE_IDX_3 = 3;
      23              : constexpr uint16_t CKE_IDX_4 = 4;
      24              : constexpr uint16_t BIT_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
      25              : 
      26            0 : CcuContextAllGatherNHR1D::CcuContextAllGatherNHR1D(
      27            0 :     const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
      28            0 :     : CcuContextAlgBase(arg, transports, group)
      29              : {
      30            0 :     const CcuCtxArgAllGatherNHR1D* ctxArg = dynamic_cast<const CcuCtxArgAllGatherNHR1D*>(&arg);
      31            0 :     if (ctxArg == nullptr) {
      32            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllGatherNHR1D::ctxArg ptr is null"));
      33              :     }
      34            0 :     rankId_ = ctxArg->rankId_;
      35            0 :     axisId_ = ctxArg->axisId_;
      36            0 :     axisSize_ = ctxArg->axisSize_;
      37            0 :     dimSize_ = ctxArg->dimSize_[0];
      38            0 :     localAxisSignalName_ = "CcuContextAllGatherNHR1DDieSync_" + std::to_string(axisId_);
      39            0 :     anotherAxisSignalName_ = "CcuContextAllGatherNHR1DDieSync_" + std::to_string(1 - axisId_);
      40            0 :     stepInfoVector_ = ctxArg->stepInfoVector_;
      41            0 :     indexMap_ = ctxArg->indexMap_;
      42            0 :     localSize_ = indexMap_.size();
      43            0 :     myRankIdx_ = indexMap_.size();
      44            0 :     signalNum_ = (dimSize_ + BIT_NUM_PER_CKE - 1) / BIT_NUM_PER_CKE; // 每个CKE有16个bit
      45            0 :     HCCL_INFO(
      46              :         "[CcuContextAllGatherNHR1D] CtxArg: rankId_[%u], axisId_[%u], axisSize_[%u], dimSize_[%u], localSize_[%u], "
      47              :         "signalNum_[%u]",
      48              :         rankId_, axisId_, axisSize_, dimSize_, localSize_, signalNum_);
      49            0 : }
      50              : 
      51            0 : void CcuContextAllGatherNHR1D::LoadArgs()
      52              : {
      53            0 :     Load(input_);
      54            0 :     Load(output_[myRankIdx_]);
      55            0 :     Load(token_[myRankIdx_]);
      56            0 :     Load(die0Size_);
      57            0 :     Load(die1Size_);
      58            0 :     Load(repeatNum_);
      59            0 :     Load(inputSliceStride_);
      60            0 :     Load(outputSliceStride_);
      61            0 :     Load(inputRepeatStride_);
      62            0 :     Load(outputRepeatStride_);
      63            0 :     Load(isInputOutputEqual_);
      64              : 
      65            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] LoadArgs run finished");
      66            0 : }
      67              : 
      68            0 : void CcuContextAllGatherNHR1D::InitResources()
      69              : {
      70            0 :     die0Size_ = CreateVariable();
      71            0 :     die1Size_ = CreateVariable();
      72            0 :     inputSliceStride_ = CreateVariable();
      73            0 :     outputSliceStride_ = CreateVariable();
      74            0 :     inputRepeatStride_ = CreateVariable();
      75            0 :     outputRepeatStride_ = CreateVariable();
      76            0 :     repeatNum_ = CreateVariable();
      77            0 :     tmpCopyRepeatNum_ = CreateVariable();
      78            0 :     repeatTimeflag_ = CreateVariable();
      79            0 :     isInputOutputEqual_ = CreateVariable();
      80            0 :     myrankInputSliceOffset_ = CreateVariable();
      81            0 :     tmpSliceOffset_ = CreateVariable();
      82            0 :     for (u64 i = 0; i < dimSize_; i++) {
      83            0 :         outputSliceOffset_.push_back(CreateVariable());
      84              :     }
      85            0 :     constVar1_ = CreateVariable();
      86            0 :     constVar1_ = 1;
      87              : 
      88            0 :     localSignal_ = CreateMaskSignal();
      89            0 :     localAxisSignal_ = CreateMaskSignal();
      90              : 
      91            0 :     if (axisSize_ > 1) {
      92            0 :         ExportMaskSignal(localAxisSignal_, localAxisSignalName_);
      93            0 :         anotherAxisSignal_ = ImportMaskSignal(anotherAxisSignalName_);
      94              :     }
      95              : 
      96            0 :     input_ = CreateVariable();
      97            0 :     for (uint32_t transportIdx = 0; transportIdx < localSize_; transportIdx++) {
      98            0 :         HCCL_DEBUG("[CcuContextAllGatherNHR1D] MyRank[%u], TransportId[%u]", rankId_, transportIdx);
      99            0 :         CHK_PRT_RET(
     100              :             transports[transportIdx] == nullptr,
     101              :             HCCL_ERROR("[CcuContextAllGatherNHR1D] Algorithm transport ptr is null"), );
     102            0 :         output_.push_back(
     103            0 :             CreateVariable((*transports[transportIdx]), OUTPUT_XN_ID)); // 获取transport中id=1的Var来传递output
     104            0 :         token_.push_back(CreateVariable((*transports[transportIdx]), TOKEN_XN_ID));
     105              :     }
     106            0 :     output_.push_back(CreateVariable());
     107            0 :     token_.push_back(CreateVariable());
     108              : 
     109            0 :     srcMem_ = CreateMemory();
     110            0 :     dstMem_ = CreateMemory();
     111            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] InitResources finished");
     112              : }
     113              : 
     114            0 : void CcuContextAllGatherNHR1D::PreSync()
     115              : {
     116            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] PreSync start");
     117            0 :     uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
     118            0 :     uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
     119            0 :     for (auto t : transports) {
     120            0 :         WriteVariableWithSignal(*t, output_[localSize_], OUTPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1, selfBit);
     121            0 :         WriteVariableWithSignal(*t, token_[localSize_], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2, selfBit);
     122              :     }
     123            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     124            0 :     for (auto& pair : indexMap_) {
     125            0 :         uint16_t pairSignalId = pair.first / BIT_NUM_PER_CKE;
     126            0 :         uint16_t pairBit = 1 << (pair.first % BIT_NUM_PER_CKE);
     127            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     128              :     }
     129            0 :     for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
     130            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]);
     131            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]);
     132              :     }
     133            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] PreSync end");
     134            0 : }
     135              : 
     136            0 : void CcuContextAllGatherNHR1D::PostSync()
     137              : {
     138            0 :     uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
     139            0 :     uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
     140            0 :     for (auto& t : transports) {
     141            0 :         RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
     142              :     }
     143            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     144            0 :     for (auto& pair : indexMap_) {
     145            0 :         uint16_t pairSignalId = pair.first / BIT_NUM_PER_CKE;
     146            0 :         uint16_t pairBit = 1 << (pair.first % BIT_NUM_PER_CKE);
     147            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     148              :     }
     149            0 :     for (uint32_t sId = 0; sId < signalNum_; sId++) {
     150            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_0, waitBitVector[sId]);
     151              :     }
     152            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] PostSync run finished");
     153            0 : }
     154              : 
     155            0 : void CcuContextAllGatherNHR1D::AxisSync(uint32_t signalIndex)
     156              : {
     157            0 :     const uint32_t DIE_NUM = 2;
     158            0 :     if (signalIndex > 1) {
     159            0 :         THROW<InvalidParamsException>(
     160            0 :             StringFormat("[CcuContextAllGatherNHR1D] Unexpected SignalInex[%u]", signalIndex));
     161              :     }
     162            0 :     LocalCtxPost(anotherAxisSignal_, 1 << (axisId_ + signalIndex * DIE_NUM));
     163            0 :     LocalWait(localAxisSignal_, 1 << (1 - axisId_ + signalIndex * DIE_NUM));
     164            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] AxisSync run finished");
     165            0 :     return;
     166              : }
     167              : 
     168            0 : void CcuContextAllGatherNHR1D::DoRepeatAllGatherNHR()
     169              : {
     170            0 :     tmpSliceOffset_ = 0;
     171            0 :     myrankInputSliceOffset_ = 0;
     172            0 :     for (u64 i = 0; i < rankId_; i++) {
     173            0 :         myrankInputSliceOffset_ += inputSliceStride_;
     174              :     }
     175            0 :     for (u64 i = 0; i < dimSize_; i++) {
     176            0 :         outputSliceOffset_[i] = tmpSliceOffset_;
     177            0 :         tmpSliceOffset_ += outputSliceStride_;
     178              :     }
     179            0 :     srcMem_.addr = input_;
     180            0 :     srcMem_.addr += myrankInputSliceOffset_;
     181            0 :     dstMem_.addr = output_[myRankIdx_];
     182            0 :     dstMem_.addr += outputSliceOffset_[rankId_];
     183            0 :     srcMem_.token = token_[myRankIdx_];
     184            0 :     dstMem_.token = token_[myRankIdx_];
     185            0 :     tmpCopyRepeatNum_ = repeatNum_;
     186            0 :     repeatTimeflag_ = 0;
     187            0 :     CCU_WHILE(tmpCopyRepeatNum_ != UINT64_MAX)
     188              :     {
     189            0 :         tmpCopyRepeatNum_ += constVar1_;
     190            0 :         CCU_IF(repeatTimeflag_ != 0)
     191              :         {
     192            0 :             srcMem_.addr += inputRepeatStride_;
     193            0 :             dstMem_.addr += outputRepeatStride_;
     194            0 :         }
     195            0 :         CCU_IF(repeatTimeflag_ == 0)
     196              :         {
     197            0 :             if (axisId_ == 1) {
     198            0 :                 srcMem_.addr += die0Size_;
     199            0 :                 dstMem_.addr += die0Size_;
     200              :             }
     201            0 :         }
     202            0 :         CCU_IF(isInputOutputEqual_ == 0)
     203              :         {
     204            0 :             LocalCopy(dstMem_, srcMem_, axisId_ == 0 ? die0Size_ : die1Size_, localSignal_, 1 << rankId_);
     205            0 :         }
     206            0 :         CCU_IF(isInputOutputEqual_ != 0) { LocalPost(localSignal_, 1 << rankId_); }
     207            0 :         LocalWait(localSignal_, 1 << rankId_);
     208            0 :         repeatTimeflag_ = 1;
     209            0 :     }
     210              : 
     211            0 :     for (auto& nhrStepInfo : stepInfoVector_) {
     212            0 :         DoRepeatAllGatherNHRSingleStep(nhrStepInfo);
     213              :     }
     214            0 : }
     215              : 
     216            0 : void CcuContextAllGatherNHR1D::DoRepeatAllGatherNHRSingleStep(const NHRStepInfo& nhrStepInfo)
     217              : {
     218            0 :     u32& toRankIdx = indexMap_[nhrStepInfo.toRank];
     219            0 :     u32& fromRankIdx = indexMap_[nhrStepInfo.fromRank];
     220            0 :     u32 sendSliceIdx = 0;
     221            0 :     CcuTransport* sendTransport = transports[toRankIdx];
     222            0 :     CcuTransport* recvTransport = transports[fromRankIdx];
     223            0 :     const std::vector<u32>& sendSliceIdxList = nhrStepInfo.txSliceIdxs;
     224            0 :     srcMem_.token = token_[myRankIdx_];
     225            0 :     dstMem_.token = token_[toRankIdx];
     226            0 :     for (u32 i = 0; i < sendSliceIdxList.size(); i++) { ////这里写的可能有问题
     227            0 :         sendSliceIdx = sendSliceIdxList[i];
     228            0 :         if (i != 0) {
     229            0 :             if (i % BIT_NUM_PER_CKE == 0) {
     230            0 :                 LocalWait(localSignal_, (1 << BIT_NUM_PER_CKE) - 1);
     231              :             }
     232              :         }
     233            0 :         if (nhrStepInfo.step == 0) {
     234            0 :             srcMem_.addr = input_;
     235            0 :             srcMem_.addr += myrankInputSliceOffset_;
     236              :         } else {
     237            0 :             srcMem_.addr = output_[myRankIdx_];
     238            0 :             srcMem_.addr += outputSliceOffset_[sendSliceIdx];
     239              :         }
     240            0 :         dstMem_.addr = output_[toRankIdx];
     241            0 :         dstMem_.addr += outputSliceOffset_[sendSliceIdx];
     242            0 :         DoRepeatSendRecvSlices(nhrStepInfo.toRank, srcMem_, dstMem_, i % BIT_NUM_PER_CKE);
     243              :     }
     244              : 
     245            0 :     if (nhrStepInfo.step + 1 != stepInfoVector_.size()) {
     246            0 :         uint16_t selfSignalId = rankId_ / BIT_NUM_PER_CKE;
     247            0 :         uint16_t selfBit = 1 << (rankId_ % BIT_NUM_PER_CKE);
     248            0 :         RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_3, selfBit, true);
     249            0 :         uint16_t recvSignalId = nhrStepInfo.fromRank / BIT_NUM_PER_CKE;
     250            0 :         uint16_t recvBit = 1 << (nhrStepInfo.fromRank % BIT_NUM_PER_CKE);
     251            0 :         RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_3, recvBit);
     252              :     }
     253            0 : }
     254              : 
     255            0 : void CcuContextAllGatherNHR1D::DoRepeatSendRecvSlices(
     256              :     const u32& toRank, CcuRep::Memory& src, CcuRep::Memory& dst, u32 signalIndex)
     257              : {
     258            0 :     CcuTransport* sendTransport = transports[indexMap_[toRank]];
     259            0 :     const CcuRep::Variable& sliceSize = axisId_ == 0 ? die0Size_ : die1Size_;
     260            0 :     repeatTimeflag_ = 0;
     261            0 :     tmpCopyRepeatNum_ = repeatNum_;
     262            0 :     CCU_WHILE(tmpCopyRepeatNum_ != UINT64_MAX)
     263              :     {
     264            0 :         tmpCopyRepeatNum_ += constVar1_;
     265            0 :         CCU_IF(repeatTimeflag_ == 1)
     266              :         {
     267            0 :             src.addr += inputRepeatStride_;
     268            0 :             dst.addr += outputRepeatStride_;
     269            0 :         }
     270            0 :         CCU_IF(repeatTimeflag_ == 0)
     271              :         {
     272            0 :             if (axisId_ == 1) {
     273            0 :                 src.addr += die0Size_;
     274            0 :                 dst.addr += die0Size_;
     275              :             }
     276            0 :         }
     277            0 :         Write(*sendTransport, dst, src, sliceSize, localSignal_, 1 << signalIndex);
     278            0 :         LocalWait(localSignal_, 1 << signalIndex);
     279            0 :         repeatTimeflag_ = 1;
     280            0 :     }
     281            0 : }
     282              : 
     283            0 : void CcuContextAllGatherNHR1D::Algorithm()
     284              : {
     285            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] AllgatherNHR1D run");
     286            0 :     InitResources();
     287            0 :     LoadArgs();
     288            0 :     if (axisSize_ > 1) {
     289            0 :         AxisSync(FST_AXIS_ID);
     290              :     }
     291            0 :     PreSync();
     292            0 :     DoRepeatAllGatherNHR();
     293            0 :     PostSync();
     294            0 :     if (axisSize_ > 1) {
     295            0 :         AxisSync(SEC_AXIS_ID);
     296              :     }
     297            0 :     HCCL_DEBUG("[CcuContextAllGatherNHR1D] AllgatherNHR1D end");
     298            0 :     return;
     299              : }
     300              : 
     301            0 : std::vector<uint64_t> CcuContextAllGatherNHR1D::GeneArgs(const CcuTaskArg& arg)
     302              : {
     303            0 :     const CcuTaskArgAllGatherNHR1D* taskArg = dynamic_cast<const CcuTaskArgAllGatherNHR1D*>(&arg);
     304            0 :     if (taskArg == nullptr) {
     305            0 :         THROW<NullPtrException>(StringFormat("CcuContextAllGatherNHR1D::taskArg ptr is null"));
     306              :     }
     307              :     // input&output&buffer地址
     308            0 :     uint64_t inputAddr = taskArg->inputAddr_;
     309            0 :     uint64_t outputAddr = taskArg->outputAddr_;
     310            0 :     uint64_t token = taskArg->token_;
     311            0 :     uint64_t die0Size = taskArg->die0Size_;
     312            0 :     uint64_t die1Size = taskArg->die1Size_;
     313            0 :     uint64_t repeatNum = UINT64_MAX - taskArg->repeatNum_;
     314            0 :     uint64_t inputSliceStride = taskArg->inputSliceStride_;
     315            0 :     uint64_t outputSliceStride = taskArg->outputSliceStride_;
     316            0 :     uint64_t inputRepeatStride = taskArg->inputRepeatStride_;
     317            0 :     uint64_t outputRepeatStride = taskArg->outputRepeatStride_;
     318            0 :     uint64_t isInputOutputEqual = taskArg->isInputOutputEqual_;
     319              : 
     320            0 :     HCCL_INFO(
     321              :         "[CcuContextAllGatherNHR1D] TaskArgs: inputAddr[%llu], outputAddr[%llu], "
     322              :         "die0Size[%llu], die1Size[%llu], repeatNum[%llu]"
     323              :         "inputSliceStride[%llu], outputSliceStride[%llu], inputRepeatStride[%llu], outputRepeatStride[%llu]",
     324              :         inputAddr, outputAddr, die0Size, die1Size, repeatNum, inputSliceStride, outputSliceStride, inputRepeatStride,
     325              :         outputRepeatStride);
     326              : 
     327              :     return {inputAddr,          outputAddr,        token,
     328              :             die0Size,           die1Size,          repeatNum,
     329              :             inputSliceStride,   outputSliceStride, inputRepeatStride,
     330            0 :             outputRepeatStride, isInputOutputEqual};
     331              : }
     332              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1