LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/scatter - ccu_context_scatter_nhr1d_mem2mem.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 244 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_scatter_nhr1d_mem2mem.h"
      12              : namespace Hccl {
      13              : constexpr uint16_t SCRATCH_XN_ID = 1;
      14              : constexpr uint16_t TOKEN_XN_ID = 2;
      15              : constexpr uint16_t CKE_IDX_0 = 0; // 后同步
      16              : constexpr uint16_t CKE_IDX_1 = 1; // 前同步addr
      17              : constexpr uint16_t CKE_IDX_2 = 2; // 前同步token
      18              : constexpr uint16_t CKE_IDX_3 = 3; // NHR step同步信号,用于scatter后同步
      19              : constexpr uint16_t FST_AXIS_ID = 0;
      20              : constexpr uint16_t SEC_AXIS_ID = 1;
      21              : constexpr uint16_t RANK_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
      22              : 
      23            0 : CcuContextScatterNHR1DMem2Mem::CcuContextScatterNHR1DMem2Mem(
      24            0 :     const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
      25            0 :     : CcuContextAlgBase(arg, transports, group)
      26              : {
      27            0 :     const CcuCtxArgScatterNHR1D* ctxArg = dynamic_cast<const CcuCtxArgScatterNHR1D*>(&arg);
      28            0 :     rankId_ = ctxArg->rankId_;
      29            0 :     rootId_ = ctxArg->rootId_;
      30            0 :     axisId_ = ctxArg->axisId_;
      31            0 :     axisSize_ = ctxArg->axisSize_;
      32            0 :     dimSize_ = ctxArg->dimSize_[0];
      33            0 :     localAxisSignalName_ = "CcuContextScatterNHR1DMem2MemDieSync_" + std::to_string(axisId_);
      34            0 :     anotherAxisSignalName_ = "CcuContextScatterNHR1DMem2MemDieSync_" + std::to_string(1 - axisId_);
      35            0 :     stepInfoVector_ = ctxArg->stepInfoVector_;
      36            0 :     indexMap_ = ctxArg->indexMap_;
      37            0 :     localSize_ = indexMap_.size();
      38            0 :     myRankIdx_ = indexMap_.size();
      39            0 :     dataType_ = ctxArg->op_.dataType;
      40            0 :     signalNum_ = (dimSize_ + RANK_NUM_PER_CKE - 1) / RANK_NUM_PER_CKE; // 每个CKE有16个bit
      41            0 :     HCCL_INFO(
      42              :         "[CcuContextScatterNHR1DMem2Mem] CtxArg: rankId_[%u], rootId_[%u], axisId_[%u], axisSize_[%u], dimSize_[%u], "
      43              :         "localSize_[%u], "
      44              :         "signalNum_[%u], dataType[%s]",
      45              :         rankId_, rootId_, axisId_, axisSize_, dimSize_, localSize_, signalNum_, dataType_.Describe().c_str());
      46            0 : }
      47              : 
      48            0 : void CcuContextScatterNHR1DMem2Mem::LoadArgs()
      49              : {
      50            0 :     Load(input_);
      51            0 :     Load(output_);
      52            0 :     Load(token_[myRankIdx_]);
      53            0 :     Load(scratch_[myRankIdx_]);
      54            0 :     Load(die0Size_);
      55            0 :     Load(die1Size_);
      56            0 :     Load(inputSliceStride_);
      57            0 :     Load(curScratchStride_);
      58            0 :     Load(inputRepeatStride_);
      59            0 :     Load(outputRepeatStride_);
      60            0 :     Load(repeatNumVar_);
      61            0 :     Load(isOutputScratch_);
      62            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] LoadArgs run finished");
      63            0 : }
      64              : 
      65            0 : void CcuContextScatterNHR1DMem2Mem::InitResources()
      66              : {
      67            0 :     die0Size_ = CreateVariable();
      68            0 :     die1Size_ = CreateVariable();
      69            0 :     inputSliceStride_ = CreateVariable();
      70            0 :     curScratchStride_ = CreateVariable();
      71            0 :     inputRepeatStride_ = CreateVariable();
      72            0 :     outputRepeatStride_ = CreateVariable();
      73            0 :     repeatNumVar_ = CreateVariable();
      74            0 :     repeatNumVarTemp_ = CreateVariable();
      75            0 :     repeatTimeflag_ = CreateVariable();
      76            0 :     curInputOffset_ = CreateVariable();
      77            0 :     curScratchOffset_ = CreateVariable();
      78            0 :     cursliceSize_ = CreateVariable();
      79            0 :     isOutputScratch_ = CreateVariable();
      80            0 :     localSignal_ = CreateMaskSignal();
      81            0 :     if (axisSize_ > 1) {
      82            0 :         localAxisSignal_ = CreateMaskSignal();
      83            0 :         ExportMaskSignal(localAxisSignal_, localAxisSignalName_);
      84            0 :         anotherAxisSignal_ = ImportMaskSignal(anotherAxisSignalName_);
      85              :     }
      86              : 
      87            0 :     input_ = CreateVariable();
      88            0 :     for (uint32_t transportIdx = 0; transportIdx < localSize_; transportIdx++) {
      89            0 :         HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] MyRank[%u], TransportId[%u]", rankId_, transportIdx);
      90            0 :         CHK_PRT_RET(
      91              :             transports[transportIdx] == nullptr,
      92              :             HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] Algorithm transport ptr is null"), );
      93            0 :         scratch_.push_back(CreateVariable((*transports[transportIdx]), SCRATCH_XN_ID)); // 存放
      94            0 :         token_.push_back(CreateVariable((*transports[transportIdx]), TOKEN_XN_ID));
      95              :     }
      96            0 :     scratch_.push_back(CreateVariable()); // 本端放最后
      97            0 :     token_.push_back(CreateVariable());
      98            0 :     output_ = CreateVariable();
      99            0 :     srcMem_ = CreateMemory();
     100            0 :     dstMem_ = CreateMemory();
     101            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] InitResources finished");
     102              : }
     103              : 
     104            0 : void CcuContextScatterNHR1DMem2Mem::PreSync()
     105              : {
     106            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] PreSync start");
     107            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     108            0 :     uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     109            0 :     for (auto t : transports) {
     110            0 :         WriteVariableWithSignal(
     111            0 :             *t, scratch_[localSize_], SCRATCH_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1, selfBit);
     112            0 :         WriteVariableWithSignal(*t, token_[localSize_], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2, selfBit);
     113              :     }
     114            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     115            0 :     for (auto& pair : indexMap_) {
     116            0 :         uint16_t pairSignalId = pair.first / RANK_NUM_PER_CKE;
     117            0 :         uint16_t pairBit = 1 << (pair.first % RANK_NUM_PER_CKE);
     118            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     119              :     }
     120            0 :     for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
     121            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]);
     122            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]);
     123              :     }
     124            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] PreSync end");
     125            0 : }
     126              : 
     127            0 : void CcuContextScatterNHR1DMem2Mem::PostSync()
     128              : {
     129            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     130            0 :     uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     131            0 :     for (auto& t : transports) {
     132            0 :         RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
     133              :     }
     134            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     135            0 :     for (auto& pair : indexMap_) {
     136            0 :         uint16_t pairSignalId = pair.first / RANK_NUM_PER_CKE;
     137            0 :         uint16_t pairBit = 1 << (pair.first % RANK_NUM_PER_CKE);
     138            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     139              :     }
     140            0 :     for (uint32_t sId = 0; sId < waitBitVector.size(); sId++) {
     141            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_0, waitBitVector[sId]);
     142              :     }
     143            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] PostSync run finished");
     144            0 : }
     145              : 
     146            0 : void CcuContextScatterNHR1DMem2Mem::AxisSync(uint32_t signalIndex)
     147              : {
     148            0 :     const uint32_t DIE_NUM = 2;
     149            0 :     if (signalIndex > 1) {
     150            0 :         THROW<InvalidParamsException>(
     151            0 :             StringFormat("[CcuContextScatterNHR1DMem2Mem] Unexpected SignalInex[%u]", signalIndex));
     152              :     }
     153            0 :     LocalCtxPost(anotherAxisSignal_, 1 << (axisId_ + signalIndex * DIE_NUM));
     154            0 :     LocalWait(localAxisSignal_, 1 << (1 - axisId_ + signalIndex * DIE_NUM));
     155            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] AxisSync run finished");
     156            0 :     return;
     157              : }
     158              : 
     159            0 : void CcuContextScatterNHR1DMem2Mem::DoScatterNHR()
     160              : {
     161            0 :     curInputOffset_ = 0;   // input偏移
     162            0 :     curScratchOffset_ = 0; // scratch偏移
     163            0 :     for (u64 i = 0; i < dimSize_; i++) {
     164            0 :         inputOffset_.push_back(CreateVariable());
     165            0 :         inputOffset_[i] = curInputOffset_;
     166            0 :         curInputOffset_ += inputSliceStride_;
     167              :     }
     168            0 :     for (u64 i = 0; i < dimSize_; i++) {
     169            0 :         ScratchOffset_.push_back(CreateVariable());
     170            0 :         ScratchOffset_[i] = curScratchOffset_;
     171            0 :         curScratchOffset_ += curScratchStride_;
     172              :     }
     173              :     // NHR
     174            0 :     for (u64 i = 0; i < stepInfoVector_.size(); i++) {
     175            0 :         const NHRStepInfo& nhrStepInfo = stepInfoVector_[i];
     176            0 :         DoScatterNHRSingleStep(nhrStepInfo);
     177              :     }
     178              :     // scratch->output
     179            0 :     if (rankId_ == rootId_) {
     180            0 :         srcMem_.addr = input_;
     181            0 :         srcMem_.addr += inputOffset_[rankId_];
     182              :     } else {
     183            0 :         srcMem_.addr = scratch_[myRankIdx_];
     184            0 :         srcMem_.addr += ScratchOffset_[rankId_];
     185              :     }
     186            0 :     dstMem_.addr = output_;
     187            0 :     srcMem_.token = token_[myRankIdx_];
     188            0 :     dstMem_.token = token_[myRankIdx_];
     189            0 :     CcuRep::Variable repeatNumAdd = CreateVariable();
     190            0 :     repeatNumAdd = 1;
     191            0 :     repeatTimeflag_ = 0;
     192            0 :     CCU_WHILE(repeatNumVar_ != UINT64_MAX)
     193              :     {
     194            0 :         repeatNumVar_ += repeatNumAdd;
     195            0 :         CCU_IF(repeatTimeflag_ != 0)
     196              :         {
     197            0 :             if (rankId_ == rootId_) {
     198            0 :                 srcMem_.addr += inputRepeatStride_;
     199              :             } else {
     200            0 :                 srcMem_.addr += outputRepeatStride_;
     201              :             }
     202            0 :             dstMem_.addr += outputRepeatStride_;
     203            0 :         }
     204            0 :         CCU_IF(repeatTimeflag_ == 0)
     205              :         {
     206            0 :             if (axisId_ == 1) {
     207            0 :                 srcMem_.addr += die0Size_;
     208            0 :                 dstMem_.addr += die0Size_;
     209              :             }
     210            0 :         }
     211            0 :         cursliceSize_ = (axisId_ == 0) ? die0Size_ : die1Size_;
     212              :         {
     213            0 :             CCU_IF(isOutputScratch_ == 1)
     214              :             {
     215            0 :                 if (rootId_ != 0 && rankId_ == 0) {
     216            0 :                     LocalPost(localSignal_, 1 << rankId_);
     217              :                 } else {
     218            0 :                     LocalCopy(dstMem_, srcMem_, cursliceSize_, localSignal_, 1 << rankId_);
     219              :                 }
     220            0 :             }
     221            0 :             CCU_IF(isOutputScratch_ != 1) { LocalCopy(dstMem_, srcMem_, cursliceSize_, localSignal_, 1 << rankId_); }
     222            0 :             LocalWait(localSignal_, 1 << rankId_);
     223              :         }
     224            0 :         repeatTimeflag_ = 1;
     225            0 :     }
     226            0 : }
     227              : 
     228            0 : void CcuContextScatterNHR1DMem2Mem::DoScatterNHRSingleStep(const NHRStepInfo& nhrStepInfo)
     229              : {
     230            0 :     const std::vector<u32>& sendSliceIdxList = nhrStepInfo.txSliceIdxs;
     231            0 :     const std::vector<u32>& recvSliceIdxList = nhrStepInfo.rxSliceIdxs;
     232            0 :     if (recvSliceIdxList.size() != 0) {
     233            0 :         u32& fromRankIdx = indexMap_[nhrStepInfo.fromRank];
     234            0 :         CcuTransport* recvTransport = transports[fromRankIdx];
     235            0 :         uint16_t recvSignalId = nhrStepInfo.fromRank / RANK_NUM_PER_CKE;
     236            0 :         uint16_t recvBit = 1 << (nhrStepInfo.fromRank % RANK_NUM_PER_CKE);
     237            0 :         RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_3,
     238              :                    recvBit); // 后同步,等待通知写入完毕
     239              :     }
     240            0 :     if (sendSliceIdxList.size() != 0) {
     241            0 :         u32& toRankIdx = indexMap_[nhrStepInfo.toRank];
     242            0 :         u32 sendSliceIdx = 0;
     243            0 :         uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     244            0 :         uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     245            0 :         CcuTransport* sendTransport = transports[toRankIdx];
     246            0 :         for (u32 i = 0; i < sendSliceIdxList.size(); i++) {
     247            0 :             sendSliceIdx = sendSliceIdxList[i];
     248            0 :             if (i != 0) {
     249            0 :                 if (i % RANK_NUM_PER_CKE == 0) {
     250            0 :                     LocalWait(localSignal_, (1 << RANK_NUM_PER_CKE) - 1);
     251              :                 }
     252              :             }
     253            0 :             if (rankId_ == rootId_) { // root节点的源数据从input中取
     254            0 :                 srcMem_.addr = input_;
     255            0 :                 srcMem_.addr += inputOffset_[sendSliceIdx];
     256              :             } else {
     257            0 :                 srcMem_.addr = scratch_[myRankIdx_];
     258            0 :                 srcMem_.addr += ScratchOffset_[sendSliceIdx];
     259              :             }
     260            0 :             srcMem_.token = token_[myRankIdx_];
     261            0 :             dstMem_.token = token_[toRankIdx];
     262            0 :             dstMem_.addr = scratch_[toRankIdx];
     263            0 :             dstMem_.addr += ScratchOffset_[sendSliceIdx];
     264            0 :             DoSendRecvSlice(nhrStepInfo.toRank, srcMem_, dstMem_, i % RANK_NUM_PER_CKE);
     265              :         }
     266            0 :         RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_3, selfBit, true); // 后同步,通知写入完毕
     267              :     }
     268            0 :     HCCL_INFO(
     269              :         "[DoScatterNHRSingleStep] rank %u step %u, toRank=%u, fromRank=%u, nSlice=%lu", rankId_, nhrStepInfo.step,
     270              :         nhrStepInfo.toRank, nhrStepInfo.fromRank, sendSliceIdxList.size());
     271            0 : }
     272              : 
     273            0 : void CcuContextScatterNHR1DMem2Mem::DoSendRecvSlice(
     274              :     const u32& toRank, CcuRep::Memory& src, CcuRep::Memory& dst, u32 signalIndex)
     275              : {
     276            0 :     CcuTransport* sendTransport = transports[indexMap_[toRank]];
     277            0 :     CcuRep::Variable repeatNumAdd2 = CreateVariable();
     278            0 :     repeatNumAdd2 = 1;
     279            0 :     repeatTimeflag_ = 0;
     280            0 :     repeatNumVarTemp_ = repeatNumVar_;
     281            0 :     CCU_WHILE(repeatNumVarTemp_ != UINT64_MAX)
     282              :     {
     283            0 :         repeatNumVarTemp_ += repeatNumAdd2;
     284            0 :         CCU_IF(repeatTimeflag_ == 1)
     285              :         {
     286            0 :             if (rankId_ == rootId_) {
     287            0 :                 src.addr += inputRepeatStride_;
     288              :             } else {
     289            0 :                 src.addr += outputRepeatStride_;
     290              :             }
     291            0 :             dst.addr += outputRepeatStride_;
     292            0 :         }
     293            0 :         CCU_IF(repeatTimeflag_ == 0)
     294              :         {
     295            0 :             if (axisId_ == 1) {
     296            0 :                 src.addr += die0Size_;
     297            0 :                 dst.addr += die0Size_;
     298              :             }
     299            0 :         }
     300            0 :         cursliceSize_ = (axisId_ == 0) ? die0Size_ : die1Size_;
     301            0 :         Write(*sendTransport, dst, src, cursliceSize_, localSignal_, 1 << signalIndex);
     302            0 :         LocalWait(localSignal_, 1 << signalIndex);
     303            0 :         repeatTimeflag_ = 1;
     304            0 :     }
     305            0 : }
     306              : 
     307            0 : void CcuContextScatterNHR1DMem2Mem::Algorithm()
     308              : {
     309            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] ScatterNHR1D run");
     310            0 :     InitResources();
     311            0 :     LoadArgs();
     312            0 :     if (axisSize_ > 1)
     313            0 :         AxisSync(FST_AXIS_ID);
     314            0 :     PreSync();
     315            0 :     DoScatterNHR();
     316            0 :     PostSync();
     317              : 
     318            0 :     if (axisSize_ > 1)
     319            0 :         AxisSync(SEC_AXIS_ID);
     320              : 
     321            0 :     HCCL_INFO("[CcuContextScatterNHR1DMem2Mem] ScatterNHR1D end");
     322            0 :     return;
     323              : }
     324              : 
     325            0 : std::vector<uint64_t> CcuContextScatterNHR1DMem2Mem::GeneArgs(const CcuTaskArg& arg)
     326              : {
     327            0 :     const CcuTaskArgScatterNHR1D* taskArg = dynamic_cast<const CcuTaskArgScatterNHR1D*>(&arg);
     328            0 :     if (taskArg == nullptr) {
     329            0 :         THROW<NullPtrException>(StringFormat("CcuContextScatterNHR1DMem2Mem::taskArg ptr is null"));
     330              :     }
     331            0 :     uint64_t inputAddr = taskArg->inputAddr_;
     332            0 :     uint64_t outputAddr = taskArg->outputAddr_;
     333            0 :     uint64_t token = taskArg->token_;
     334            0 :     uint64_t scratchAddr = taskArg->scratchAddr_;
     335            0 :     uint64_t die0Size = taskArg->die0Size_;
     336            0 :     uint64_t die1Size = taskArg->die1Size_;
     337            0 :     uint64_t inputSliceStride = taskArg->inputSliceStride_;
     338            0 :     uint64_t curScratchStride = taskArg->sliceSize_ * taskArg->repeatNum_;
     339            0 :     uint64_t inputRepeatStride = taskArg->inputRepeatStride_;
     340            0 :     uint64_t outputRepeatStride = taskArg->outputRepeatStride_;
     341            0 :     uint64_t repeatNumVar = taskArg->repeatNumVar_;
     342            0 :     uint64_t isOutputScratch = taskArg->isOutputScratch_;
     343              : 
     344            0 :     HCCL_INFO(
     345              :         "[CcuContextScatterNHR1DMem2Mem] TaskArgs: rankId_[%llu], inputAddr[%llu], outputAddr[%llu], scratchAddr[%llu],"
     346              :         "die0Size[%llu], die1Size[%llu], inputSliceStride[%llu], curScratchStride[%llu],"
     347              :         "inputRepeatStride[%llu], outputRepeatStride[%llu],repeatNumVar[%llu]",
     348              :         rankId_, inputAddr, outputAddr, scratchAddr, die0Size, die1Size, inputSliceStride, curScratchStride,
     349              :         inputRepeatStride, outputRepeatStride, repeatNumVar);
     350              :     return {inputAddr,          outputAddr,       token,
     351              :             scratchAddr,        die0Size,         die1Size,
     352              :             inputSliceStride,   curScratchStride, inputRepeatStride,
     353            0 :             outputRepeatStride, repeatNumVar,     isOutputScratch};
     354              : }
     355              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1