LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/reduce_scatter - ccu_context_reduce_scatter_nhr1d_mem2mem.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 249 0
Test Date: 2026-08-04 10:52:23 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_reduce_scatter_nhr1d_mem2mem.h"
      12              : 
      13              : namespace Hccl {
      14              : 
      15              : // 注意型号量变化
      16              : constexpr uint16_t INPUT_XN_ID      = 0;
      17              : constexpr uint16_t TOKEN_XN_ID      = 2;
      18              : constexpr uint16_t CKE_IDX_0        = 0;
      19              : constexpr uint16_t CKE_IDX_1        = 1;
      20              : constexpr uint16_t CKE_IDX_2        = 2;
      21              : constexpr uint16_t CKE_IDX_3        = 3;
      22              : constexpr uint16_t CKE_IDX_4        = 4;
      23              : constexpr uint16_t FST_AXIS_ID      = 0;
      24              : constexpr uint16_t SEC_AXIS_ID      = 1;
      25              : constexpr uint16_t RANK_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
      26              : constexpr uint16_t LINK_SIZE        = 2;
      27              : 
      28            0 : CcuContextReduceScatterNHR1DMem2Mem::CcuContextReduceScatterNHR1DMem2Mem(const CcuCtxArg &arg,
      29              :     const std::vector<CcuTransport *> &transports,
      30            0 :     const CcuTransportGroup &group)
      31            0 :     : CcuContextAlgBase(arg, transports, group)
      32              : {
      33            0 :     const CcuCtxArgReduceScatterNHR1D *ctxArg = dynamic_cast<const CcuCtxArgReduceScatterNHR1D *>(&arg);
      34            0 :     rankId_                               = ctxArg->rankId_;
      35            0 :     axisId_                               = ctxArg->axisId_;
      36            0 :     dimSize_                              = ctxArg->dimSize_[0];
      37            0 :     localAxisSignalName_                  = "CcuContextReduceScatterNHR1DDieSync_" + std::to_string(axisId_);
      38            0 :     anotherAxisSignalName_                = "CcuContextReduceScatterNHR1DDieSync_" + std::to_string(1 - axisId_);
      39            0 :     stepInfoVector_                       = ctxArg->stepInfoVector_;
      40            0 :     indexMap_                             = ctxArg->indexMap_;
      41            0 :     localSize_                            = indexMap_.size();
      42            0 :     myRankIdx_                            = indexMap_.size();
      43            0 :     reduceOp_                             = ctxArg->op_.reduceOp;
      44            0 :     dataType_                             = ctxArg->op_.dataType;
      45            0 :     outputDataType_                       = ctxArg->op_.outputDataType;
      46            0 :     linkNum_                              = ctxArg->linkNum_;
      47            0 :     if (outputDataType_ == DataType::INVALID) {
      48            0 :         outputDataType_ = dataType_;
      49            0 :         HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] outputDataType is [INVALID], set outputDataType to[%s]",
      50              :             outputDataType_.Describe().c_str());
      51              :     }
      52            0 :     signalNum_ = (dimSize_ + RANK_NUM_PER_CKE - 1) / RANK_NUM_PER_CKE; // 每个CKE有16个bit
      53            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] CtxArg: rankId_[%u], axisId_[%u], dimSize_[%u], localSize_[%u], "
      54              :               "dataType[%s], outputDataType[%s], reduceOp[%s], signalNum_[%u]",
      55              :               rankId_, axisId_, dimSize_, localSize_, dataType_.Describe().c_str(),
      56              :               outputDataType_.Describe().c_str(), reduceOp_.Describe().c_str(), signalNum_);
      57            0 : }
      58              : 
      59            0 : void CcuContextReduceScatterNHR1DMem2Mem::LoadArgs()
      60              : {
      61            0 :     Load(input_[myRankIdx_]);
      62            0 :     Load(output_);
      63            0 :     Load(token_[myRankIdx_]);
      64            0 :     Load(die0Size_);
      65            0 :     Load(die1Size_);
      66            0 :     Load(inputSliceStride_);
      67            0 :     Load(outputSliceStride_);
      68            0 :     Load(inputRepeatStride_);
      69            0 :     Load(outputRepeatStride_);
      70            0 :     Load(repeatNumVar_);
      71            0 :     Load(isBottom_);
      72            0 :     repeatNumVarTemp_ = repeatNumVar_;
      73            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] LoadArgs run finished");
      74            0 : }
      75              : 
      76            0 : void CcuContextReduceScatterNHR1DMem2Mem::InitResources()
      77              : {
      78            0 :     die0Size_           = CreateVariable();
      79            0 :     die1Size_           = CreateVariable();
      80            0 :     sliceSize_          = CreateVariable();
      81            0 :     inputSliceStride_   = CreateVariable();
      82            0 :     outputSliceStride_  = CreateVariable();
      83            0 :     inputRepeatStride_  = CreateVariable();
      84            0 :     outputRepeatStride_ = CreateVariable();
      85            0 :     localAxisSignal_    = CreateMaskSignal();
      86            0 :     anotherAxisSignal_  = CreateMaskSignal();
      87            0 :     localSignal_        = CreateMaskSignal();
      88            0 :     repeatNumVar_       = CreateVariable();
      89            0 :     repeatNumVarTemp_   = CreateVariable();
      90            0 :     isBottom_           = CreateVariable();
      91              : 
      92            0 :     if (linkNum_ == LINK_SIZE) {
      93            0 :         ExportMaskSignal(localAxisSignal_, localAxisSignalName_);
      94            0 :         anotherAxisSignal_ = ImportMaskSignal(anotherAxisSignalName_);
      95              :     }
      96              : 
      97            0 :     output_ = CreateVariable();
      98            0 :     for (uint32_t transportIdx = 0; transportIdx < localSize_; transportIdx++) {
      99            0 :         HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] MyRank[%u], TransportId[%u]", rankId_, transportIdx);
     100            0 :         CHK_PRT_RET(transports[transportIdx] == nullptr,
     101              :                     HCCL_ERROR("[CcuContextReduceScatterNHR1DMem2Mem] Algorithm transport ptr is null"),);
     102            0 :         input_.push_back(
     103            0 :             CreateVariable((*transports[transportIdx]), INPUT_XN_ID)); // 获取transport中id=0的Var来传递input
     104              : 
     105            0 :         token_.push_back(CreateVariable((*transports[transportIdx]), TOKEN_XN_ID));
     106              :     }
     107            0 :     input_.push_back(CreateVariable());
     108            0 :     token_.push_back(CreateVariable());
     109              : 
     110            0 :     repeatInputOffset_      = CreateVariable();
     111            0 :     repeatOutputOffset_     = CreateVariable();
     112            0 :     myrankInputSliceOffset_ = CreateVariable();
     113              : 
     114            0 :     srcMem_ = CreateMemory();
     115            0 :     dstMem_ = CreateMemory();
     116            0 :     flag_   = CreateVariable();
     117            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] InitResources finished");
     118              : }
     119              : 
     120            0 : void CcuContextReduceScatterNHR1DMem2Mem::PreSync()
     121              : {
     122            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] PreSync start");
     123              :     // 本rank用哪一个CKE
     124            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     125              :     // 本rank用CKE的哪一位
     126            0 :     uint16_t selfBit      = 1 << (rankId_ % RANK_NUM_PER_CKE);
     127            0 :     for (auto t : transports) {
     128            0 :         WriteVariableWithSignal(*t, input_[localSize_], INPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1, selfBit);
     129            0 :         WriteVariableWithSignal(*t, token_[localSize_], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2, selfBit);
     130              :     }
     131            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     132            0 :     for (auto &pair : indexMap_) {
     133            0 :         uint16_t pairSignalId       = pair.first / RANK_NUM_PER_CKE;
     134            0 :         uint16_t pairBit            = 1 << (pair.first % RANK_NUM_PER_CKE);
     135            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     136              :     }
     137            0 :     for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
     138            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]);
     139            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]);
     140              :     }
     141            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] PreSync end");
     142            0 : }
     143              : 
     144            0 : void CcuContextReduceScatterNHR1DMem2Mem::PostSync()
     145              : {
     146            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     147            0 :     uint16_t selfBit      = 1 << (rankId_ % RANK_NUM_PER_CKE);
     148            0 :     for (auto &t : transports) {
     149            0 :         RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
     150              :     }
     151            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     152            0 :     for (auto &pair : indexMap_) {
     153            0 :         uint16_t pairSignalId       = pair.first / RANK_NUM_PER_CKE;
     154            0 :         uint16_t pairBit            = 1 << (pair.first % RANK_NUM_PER_CKE);
     155            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     156              :     }
     157            0 :     for (uint32_t sId = 0; sId < waitBitVector.size(); sId++) {
     158            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_0, waitBitVector[sId]);
     159              :     }
     160            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] PostSync run finished");
     161            0 : }
     162              : 
     163            0 : void CcuContextReduceScatterNHR1DMem2Mem::AxisSync(uint32_t signalIndex)
     164              : {
     165            0 :     const uint32_t DIE_NUM = 2;
     166            0 :     if (signalIndex > 1) {
     167            0 :         THROW<InvalidParamsException>(
     168            0 :             StringFormat("[CcuContextReduceScatterNHR1DMem2Mem] Unexpected SignalInex[%u]", signalIndex));
     169              :     }
     170            0 :     LocalCtxPost(anotherAxisSignal_, 1 << (axisId_ + signalIndex * DIE_NUM));
     171            0 :     LocalWait(localAxisSignal_, 1 << (1 - axisId_ + signalIndex * DIE_NUM));
     172            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] AxisSync run finished");
     173            0 :     return;
     174              : }
     175              : 
     176            0 : void CcuContextReduceScatterNHR1DMem2Mem::DoRepeatReduceScatterNHR()
     177              : {
     178            0 :     CcuRep::Variable tmpSliceOffset   = CreateVariable();
     179            0 :     tmpSliceOffset                    = 0;
     180              :     // 用来记录每个rank要读取的rank的sliceIdx的偏移
     181              :     // 后面会用inputAddr来加上这个偏移获取sliceIdx的地址
     182            0 :     std::vector<CcuRep::Variable> inputSliceOffset;
     183            0 :     CCU_IF(isBottom_ == 1) {
     184            0 :         for (u64 i = 0; i < dimSize_; i++) {
     185            0 :             inputSliceOffset.push_back(CreateVariable());
     186            0 :             inputSliceOffset[i] = tmpSliceOffset;
     187            0 :             tmpSliceOffset += inputSliceStride_;
     188              :         }
     189            0 :     }
     190            0 :     CCU_IF(isBottom_ == 0) {
     191            0 :         for (u64 i = 0; i < dimSize_; i++) {
     192            0 :             inputSliceOffset.push_back(CreateVariable());
     193            0 :             inputSliceOffset[i] = tmpSliceOffset;
     194            0 :             tmpSliceOffset += inputRepeatStride_;
     195              :         }
     196            0 :     }
     197              : 
     198            0 :     for (auto &nhrStepInfo : stepInfoVector_) {
     199            0 :         DoRepeatReduceScatterNHRSingleStep(nhrStepInfo, inputSliceOffset);
     200              :     }
     201              :     // 因为所有的修改都是在input上进行的,所以最后需要把input上的数据搬到output上
     202            0 :     dstMem_.addr = output_;
     203            0 :     dstMem_.token = token_[myRankIdx_];
     204            0 :     srcMem_.addr = input_[myRankIdx_];
     205            0 :     srcMem_.addr += inputSliceOffset[rankId_];
     206            0 :     srcMem_.token = token_[myRankIdx_];
     207              : 
     208            0 :     CcuRep::Variable repeatNumAdd2 = CreateVariable();
     209            0 :     repeatNumAdd2  = 1;
     210            0 :     CCU_WHILE(repeatNumVar_ != UINT64_MAX) {
     211            0 :         repeatNumVar_ += repeatNumAdd2;
     212            0 :         CCU_IF(flag_ == 1) {
     213            0 :             CCU_IF(isBottom_ == 0) {
     214            0 :                 srcMem_.addr += inputSliceStride_;
     215            0 :                 dstMem_.addr += outputRepeatStride_;
     216            0 :             }
     217            0 :             CCU_IF(isBottom_ == 1) {
     218            0 :                 srcMem_.addr += inputRepeatStride_;
     219            0 :                 dstMem_.addr += outputRepeatStride_;
     220            0 :             }
     221            0 :         }
     222            0 :         CCU_IF(flag_ == 0) {
     223            0 :             if (axisId_ == 1) {
     224            0 :                 srcMem_.addr += die0Size_;
     225            0 :                 dstMem_.addr += die0Size_;
     226              :             }
     227            0 :         }
     228            0 :         CcuRep::Variable &localSliceSize = (axisId_ == 0) ? die0Size_ : die1Size_;
     229            0 :         LocalCopy(dstMem_, srcMem_, localSliceSize, localSignal_, 1);
     230            0 :         LocalWait(localSignal_, 1);
     231            0 :         flag_ = 1;
     232            0 :     }
     233            0 : }
     234              : 
     235            0 : void CcuContextReduceScatterNHR1DMem2Mem::DoRepeatReduceScatterNHRSingleStep(const NHRStepInfo &nhrStepInfo,
     236              :     const std::vector<CcuRep::Variable> &inputSliceOffset)
     237              : {
     238            0 :     u32& toRankIdx = indexMap_[nhrStepInfo.toRank];
     239            0 :     u32& fromRankIdx = indexMap_[nhrStepInfo.fromRank];
     240            0 :     CcuTransport           *sendTransport = transports[toRankIdx];
     241            0 :     CcuTransport           *recvTransport = transports[fromRankIdx];
     242            0 :     const std::vector<u32> &sendSliceIdxList  = nhrStepInfo.txSliceIdxs;
     243            0 :     dstMem_.token                         = token_[toRankIdx];
     244            0 :     srcMem_.token                         = token_[myRankIdx_];
     245              : 
     246              :     // 被写之前告诉写自己的rank自己准备好了-前同步
     247            0 :     uint16_t recvSignalIdPrev = nhrStepInfo.fromRank / RANK_NUM_PER_CKE;
     248            0 :     uint16_t recvBitPrev      = 1 << (nhrStepInfo.fromRank % RANK_NUM_PER_CKE);
     249            0 :     RemotePost(*recvTransport, recvSignalIdPrev + signalNum_ * CKE_IDX_3, recvBitPrev, true);
     250              : 
     251            0 :     uint16_t selfSignalIdPrev = rankId_ / RANK_NUM_PER_CKE;
     252            0 :     uint16_t selfBitPrev      = 1 << (rankId_ % RANK_NUM_PER_CKE);
     253            0 :     RemoteWait(*sendTransport, selfSignalIdPrev + signalNum_ * CKE_IDX_3, selfBitPrev);
     254              : 
     255            0 :     for (const u32 &sendSliceIdx : sendSliceIdxList) {
     256            0 :         dstMem_.addr = input_[toRankIdx];
     257            0 :         dstMem_.addr += inputSliceOffset[sendSliceIdx];
     258            0 :         srcMem_.addr = input_[myRankIdx_];
     259            0 :         srcMem_.addr += inputSliceOffset[sendSliceIdx];
     260            0 :         DoRepeatSendRecvSlices(nhrStepInfo.toRank, srcMem_, dstMem_);
     261              :     }
     262              : 
     263              :     // 写之后告诉对面写完了-后同步
     264            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     265            0 :     uint16_t selfBit      = 1 << (rankId_ % RANK_NUM_PER_CKE);
     266            0 :     RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_4, selfBit, true);
     267              : 
     268            0 :     uint16_t recvSignalId = nhrStepInfo.fromRank / RANK_NUM_PER_CKE;
     269            0 :     uint16_t recvBit      = 1 << (nhrStepInfo.fromRank % RANK_NUM_PER_CKE);
     270            0 :     RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_4, recvBit);
     271            0 : }
     272              : 
     273            0 : void CcuContextReduceScatterNHR1DMem2Mem::DoRepeatSendRecvSlices(const u32 &toRank, CcuRep::Memory &src,
     274              :                                                                  CcuRep::Memory &dst)
     275              : {
     276            0 :     CcuRep::Variable repeatNumAdd = CreateVariable();
     277            0 :     repeatNumAdd  = 1;
     278            0 :     flag_ = 0;
     279            0 :     CcuTransport *sendTransport = transports[indexMap_[toRank]];
     280            0 :     repeatNumVarTemp_ = repeatNumVar_;
     281            0 :     CCU_WHILE(repeatNumVarTemp_ != UINT64_MAX) {
     282            0 :         CCU_IF(repeatNumVarTemp_ != UINT64_MAX) {
     283            0 :             repeatNumVarTemp_ += repeatNumAdd;
     284            0 :         }
     285              :         
     286            0 :         CCU_IF(flag_ == 1) {
     287            0 :             CCU_IF(isBottom_ == 0) {
     288            0 :                 src.addr += inputSliceStride_;
     289            0 :                 dst.addr += inputSliceStride_;
     290            0 :             }
     291            0 :             CCU_IF(isBottom_ == 1) {
     292            0 :                 src.addr += inputRepeatStride_;
     293            0 :                 dst.addr += inputRepeatStride_;
     294            0 :             }
     295            0 :         }
     296            0 :         CCU_IF(flag_ == 0) {
     297            0 :             if (axisId_ == 1) {
     298            0 :                 src.addr += die0Size_;
     299            0 :                 dst.addr += die0Size_;
     300              :             }
     301            0 :         }
     302            0 :         sliceSize_ =  (axisId_ == 0) ? die0Size_ : die1Size_;
     303            0 :         WriteReduce(*sendTransport, dst, src, sliceSize_, dataType_,
     304            0 :                     reduceOp_, localSignal_, 1);
     305            0 :         LocalWait(localSignal_, (1 << 1) - 1);
     306            0 :         flag_ = 1;
     307            0 :     }
     308            0 :     flag_ = 0;
     309            0 : }
     310              : 
     311            0 : void CcuContextReduceScatterNHR1DMem2Mem::Algorithm()
     312              : {
     313            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] CcuContextReduceScatterNHR1DMem2Mem run.");
     314            0 :     InitResources();
     315            0 :     LoadArgs();
     316            0 :     if (linkNum_ == LINK_SIZE) {
     317            0 :         AxisSync(FST_AXIS_ID);
     318              :     }
     319            0 :     PreSync();
     320            0 :     DoRepeatReduceScatterNHR();
     321            0 :     PostSync();
     322            0 :     if (linkNum_ == LINK_SIZE) {
     323            0 :         AxisSync(SEC_AXIS_ID);
     324              :     }
     325              : 
     326            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] CcuContextReduceScatterNHR1DMem2Mem end.");
     327            0 :     return;
     328              : }
     329              : 
     330            0 : std::vector<uint64_t> CcuContextReduceScatterNHR1DMem2Mem::GeneArgs(const CcuTaskArg &arg)
     331              : {
     332            0 :     const CcuTaskArgReduceScatterNHR1D *taskArg = dynamic_cast<const CcuTaskArgReduceScatterNHR1D *>(&arg);
     333            0 :     if (taskArg == nullptr) {
     334            0 :         THROW<NullPtrException>(StringFormat("CcuContextReduceScatterNHR1DMem2Mem::taskArg ptr is null"));
     335              :     }
     336              :     // input & output & buffer地址
     337            0 :     uint64_t inputAddr          = taskArg->inputAddr_;
     338            0 :     uint64_t outputAddr         = taskArg->outputAddr_;
     339            0 :     uint64_t token              = taskArg->token_;
     340            0 :     uint64_t die0Size           = taskArg->die0Size_;
     341            0 :     uint64_t die1Size           = taskArg->die1Size_;
     342            0 :     uint64_t inputSliceStride   = taskArg->inputSliceStride_;
     343            0 :     uint64_t outputSliceStride  = taskArg->outputSliceStride_;
     344            0 :     uint64_t inputRepeatStride  = taskArg->inputRepeatStride_;
     345            0 :     uint64_t outputRepeatStride = taskArg->outputRepeatStride_;
     346            0 :     uint64_t repeatNumVar       = taskArg->repeatNum_;
     347            0 :     uint64_t isBottom           = taskArg->isBottom_;
     348              : 
     349            0 :     HCCL_INFO("[CcuContextReduceScatterNHR1DMem2Mem] TaskArgs: inputAddr[%llu], outputAddr[%llu],"
     350              :               "die0Size[%llu], die1Size[%llu],"
     351              :               "inputSliceStride[%llu], outputSliceStride[%llu], inputRepeatStride[%llu], outputRepeatStride[%llu]",
     352              :               inputAddr, outputAddr, die0Size, die1Size,
     353              :               inputSliceStride, outputSliceStride, inputRepeatStride, outputRepeatStride);
     354              : 
     355              :     return {inputAddr,          outputAddr,        token,
     356              :             die0Size,           die1Size,          inputSliceStride,
     357              :             outputSliceStride,  inputRepeatStride, outputRepeatStride,
     358            0 :             repeatNumVar,          isBottom};
     359              : }
     360              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1