LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_ccu_context/reduce - ccu_context_reduce_nhr1d_mem2mem.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 276 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 17 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_nhr1d_mem2mem.h"
      12              : 
      13              : namespace Hccl {
      14              : 
      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同步信号0,用于RS前同步,AG后同步
      19              : constexpr uint16_t CKE_IDX_4 = 4; // NHR step同步信号1,用于RS后同步
      20              : constexpr uint16_t FST_AXIS_ID = 0;
      21              : constexpr uint16_t SEC_AXIS_ID = 1;
      22              : constexpr uint16_t RANK_NUM_PER_CKE = 16; // 本rank给远端置位时应当写的CKE,16个对端一个CKE
      23              : constexpr uint16_t OUTPUT_XN_ID = 1;
      24              : constexpr uint16_t TOKEN_XN_ID = 2;
      25              : 
      26            0 : CcuContextReduceNHR1DMem2mem::CcuContextReduceNHR1DMem2mem(
      27            0 :     const CcuCtxArg& arg, const std::vector<CcuTransport*>& transports, const CcuTransportGroup& group)
      28            0 :     : CcuContextAlgBase(arg, transports, group)
      29              : {
      30            0 :     const CcuCtxArgReduceNHR1D* ctxArg = dynamic_cast<const CcuCtxArgReduceNHR1D*>(&arg);
      31            0 :     if (ctxArg == nullptr) {
      32            0 :         THROW<NullPtrException>(StringFormat("CcuContextReduceNHR1DMem2mem::ctxArg ptr is null"));
      33              :     }
      34            0 :     rankId_ = ctxArg->rankId_;
      35            0 :     rootId_ = ctxArg->rootId_;
      36            0 :     axisId_ = ctxArg->axisId_;
      37            0 :     axisSize_ = ctxArg->axisSize_;
      38            0 :     dimSize_ = ctxArg->dimSize_[0];
      39            0 :     localAxisSignalName_ = "CcuContextAllReduceNHR1DDieSync_" + std::to_string(axisId_);
      40            0 :     anotherAxisSignalName_ = "CcuContextAllReduceNHR1DDieSync_" + std::to_string(1 - axisId_);
      41            0 :     stepInfoVector_ = ctxArg->stepInfoVector_;
      42            0 :     indexMap_ = ctxArg->indexMap_;
      43            0 :     localSize_ = indexMap_.size();
      44            0 :     myRankIdx_ = indexMap_.size();
      45            0 :     dataType_ = ctxArg->op_.dataType;
      46            0 :     reduceOp_ = ctxArg->op_.reduceOp;
      47            0 :     signalNum_ = (dimSize_ + RANK_NUM_PER_CKE - 1) / RANK_NUM_PER_CKE; // 每个CKE有16个bit
      48            0 :     HCCL_INFO(
      49              :         "[CcuContextReduceNHR1DMem2mem] CtxArg: rankId_[%u], axisId_[%u], axisSize_[%u], dimSize_[%u], localSize_[%u], "
      50              :         "signalNum_[%u], dataType[%s], reduceOp[%s]",
      51              :         rankId_, axisId_, axisSize_, dimSize_, localSize_, signalNum_, dataType_.Describe().c_str(),
      52              :         reduceOp_.Describe().c_str());
      53            0 : }
      54              : 
      55            0 : void CcuContextReduceNHR1DMem2mem::LoadArgs()
      56              : {
      57            0 :     Load(input_);
      58            0 :     Load(output_[myRankIdx_]);
      59            0 :     Load(token_[myRankIdx_]);
      60            0 :     Load(isInputOutputEqual_);
      61            0 :     Load(die0Size_);
      62            0 :     Load(die1Size_);
      63            0 :     Load(die0SliceSize_);
      64            0 :     Load(die1SliceSize_);
      65            0 :     Load(die0LastSliceSize_);
      66            0 :     Load(die1LastSliceSize_);
      67            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] LoadArgs run finished");
      68            0 : }
      69              : 
      70            0 : void CcuContextReduceNHR1DMem2mem::InitResources()
      71              : {
      72            0 :     die0Size_ = CreateVariable();
      73            0 :     die1Size_ = CreateVariable();
      74            0 :     die0LastSliceSize_ = CreateVariable();
      75            0 :     die1LastSliceSize_ = CreateVariable();
      76            0 :     die0SliceSize_ = CreateVariable();
      77            0 :     die1SliceSize_ = CreateVariable();
      78            0 :     localAxisSignal_ = CreateMaskSignal();
      79            0 :     localSignal_ = CreateMaskSignal();
      80            0 :     if (axisSize_ > 1) {
      81            0 :         ExportMaskSignal(localAxisSignal_, localAxisSignalName_);
      82            0 :         anotherAxisSignal_ = ImportMaskSignal(anotherAxisSignalName_);
      83              :     }
      84              : 
      85            0 :     input_ = CreateVariable();
      86            0 :     isInputOutputEqual_ = CreateVariable();
      87            0 :     for (uint32_t transportIdx = 0; transportIdx < localSize_; transportIdx++) {
      88            0 :         HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] MyRank[%u], TransportId[%u]", rankId_, transportIdx);
      89            0 :         CHK_PRT_RET(
      90              :             transports[transportIdx] == nullptr,
      91              :             HCCL_ERROR("[CcuContextReduceNHR1DMem2mem] Algorithm transport ptr is null"), );
      92            0 :         output_.push_back(
      93            0 :             CreateVariable((*transports[transportIdx]), OUTPUT_XN_ID)); // 获取transport中id=1的Var来传递output
      94            0 :         token_.push_back(CreateVariable((*transports[transportIdx]), TOKEN_XN_ID));
      95              :     }
      96            0 :     output_.push_back(CreateVariable());
      97            0 :     token_.push_back(CreateVariable());
      98              : 
      99            0 :     srcMem_ = CreateMemory();
     100            0 :     dstMem_ = CreateMemory();
     101            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] InitResources finished");
     102              : }
     103              : 
     104            0 : void CcuContextReduceNHR1DMem2mem::PreSync()
     105              : {
     106            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] 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(*t, output_[localSize_], OUTPUT_XN_ID, selfSignalId + signalNum_ * CKE_IDX_1, selfBit);
     111            0 :         WriteVariableWithSignal(*t, token_[localSize_], TOKEN_XN_ID, selfSignalId + signalNum_ * CKE_IDX_2, selfBit);
     112              :     }
     113            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     114            0 :     for (auto& pair : indexMap_) {
     115            0 :         uint16_t pairSignalId = pair.first / RANK_NUM_PER_CKE;
     116            0 :         uint16_t pairBit = 1 << (pair.first % RANK_NUM_PER_CKE);
     117            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     118              :     }
     119            0 :     for (uint16_t sId = 0; sId < waitBitVector.size(); sId++) {
     120            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_1, waitBitVector[sId]);
     121            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_2, waitBitVector[sId]);
     122              :     }
     123            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] PreSync end");
     124            0 : }
     125              : 
     126            0 : void CcuContextReduceNHR1DMem2mem::PostSync()
     127              : {
     128            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     129            0 :     uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     130            0 :     for (auto& t : transports) {
     131            0 :         RemotePost(*t, selfSignalId + signalNum_ * CKE_IDX_0, selfBit);
     132              :     }
     133            0 :     std::vector<uint16_t> waitBitVector(signalNum_, 0);
     134            0 :     for (auto& pair : indexMap_) {
     135            0 :         uint16_t pairSignalId = pair.first / RANK_NUM_PER_CKE;
     136            0 :         uint16_t pairBit = 1 << (pair.first % RANK_NUM_PER_CKE);
     137            0 :         waitBitVector[pairSignalId] = waitBitVector[pairSignalId] | pairBit;
     138              :     }
     139            0 :     for (uint32_t sId = 0; sId < signalNum_; sId++) {
     140            0 :         GroupWait(*transportGroup, sId + signalNum_ * CKE_IDX_0, waitBitVector[sId]);
     141              :     }
     142            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] PostSync run finished");
     143            0 : }
     144              : 
     145            0 : void CcuContextReduceNHR1DMem2mem::AxisSync(uint32_t signalIndex)
     146              : {
     147            0 :     const uint32_t DIE_NUM = 2;
     148            0 :     if (signalIndex > 1) {
     149            0 :         THROW<InvalidParamsException>(
     150            0 :             StringFormat("[CcuContextReduceNHR1DMem2mem] Unexpected SignalInex[%u]", signalIndex));
     151              :     }
     152            0 :     LocalCtxPost(anotherAxisSignal_, 1 << (axisId_ + signalIndex * DIE_NUM));
     153            0 :     LocalWait(localAxisSignal_, 1 << (1 - axisId_ + signalIndex * DIE_NUM));
     154            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] AxisSync run finished");
     155            0 :     return;
     156              : }
     157              : 
     158            0 : void CcuContextReduceNHR1DMem2mem::DoReduceScatterNHR()
     159              : {
     160            0 :     const uint32_t NHR_NUM = 2;
     161            0 :     for (u64 i = 0; i < stepInfoVector_.size() / NHR_NUM; i++) {
     162            0 :         const NHRStepInfo& nhrStepInfo = stepInfoVector_[i];
     163            0 :         DoReduceScatterNHRSingleStep(nhrStepInfo);
     164              :     }
     165            0 : }
     166              : 
     167            0 : void CcuContextReduceNHR1DMem2mem::DoReduceScatterNHRSingleStep(const NHRStepInfo& nhrStepInfo)
     168              : {
     169            0 :     u32& toRankIdx = indexMap_[nhrStepInfo.toRank];
     170            0 :     u32& fromRankIdx = indexMap_[nhrStepInfo.fromRank];
     171            0 :     u32 sendSliceIdx = 0;
     172            0 :     CcuTransport* sendTransport = transports[toRankIdx];
     173            0 :     CcuTransport* recvTransport = transports[fromRankIdx];
     174            0 :     const std::vector<u32>& sendSliceIdxList = nhrStepInfo.txSliceIdxs;
     175            0 :     srcMem_.token = token_[myRankIdx_];
     176            0 :     dstMem_.token = token_[toRankIdx];
     177            0 :     uint16_t sendSignalId = nhrStepInfo.toRank / RANK_NUM_PER_CKE;
     178            0 :     uint16_t sendBit = 1 << (nhrStepInfo.toRank % RANK_NUM_PER_CKE);
     179            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     180            0 :     uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     181            0 :     if (nhrStepInfo.step != 0) {
     182              :         // 通知fromRank,可以写入
     183            0 :         RemotePost(*recvTransport, selfSignalId + signalNum_ * CKE_IDX_3, selfBit, true);
     184              : 
     185              :         // 等待toRank通知其可以写入
     186            0 :         RemoteWait(*sendTransport, sendSignalId + signalNum_ * CKE_IDX_3, sendBit);
     187              :     }
     188              : 
     189            0 :     for (u32 i = 0; i < sendSliceIdxList.size(); i++) {
     190            0 :         sendSliceIdx = sendSliceIdxList[i];
     191              : 
     192            0 :         if (i != 0) {
     193            0 :             if (i % RANK_NUM_PER_CKE == 0) {
     194            0 :                 LocalWait(localSignal_, (1 << RANK_NUM_PER_CKE) - 1);
     195              :             }
     196              :         }
     197              : 
     198            0 :         if (nhrStepInfo.step == 0) {
     199              :             // 只有第0步的源数据从input中取
     200            0 :             srcMem_.addr = input_;
     201            0 :             srcMem_.addr += sliceOffset_[sendSliceIdx];
     202              :         } else {
     203            0 :             srcMem_.addr = output_[myRankIdx_];
     204            0 :             srcMem_.addr += sliceOffset_[sendSliceIdx];
     205              :         }
     206              : 
     207            0 :         dstMem_.addr = output_[toRankIdx];
     208            0 :         dstMem_.addr += sliceOffset_[sendSliceIdx];
     209              : 
     210            0 :         DoWriteReduceSlice(nhrStepInfo.toRank, srcMem_, dstMem_, sendSliceIdx, i % RANK_NUM_PER_CKE);
     211              :     }
     212            0 :     LocalWait(localSignal_, (1 << (sendSliceIdxList.size() % RANK_NUM_PER_CKE)) - 1);
     213              :     // 通知toRank数据写入完毕
     214            0 :     RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_4, selfBit, true);
     215              :     // 等待fromRank通知数据写入完毕
     216            0 :     uint16_t recvSignalId = nhrStepInfo.fromRank / RANK_NUM_PER_CKE;
     217            0 :     uint16_t recvBit = 1 << (nhrStepInfo.fromRank % RANK_NUM_PER_CKE);
     218            0 :     RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_4, recvBit);
     219              : 
     220            0 :     HCCL_DEBUG(
     221              :         "[DoReduceScatterNHRSingleStep] rank %u step %u, toRank=%u, fromRank=%u, nSlice=%lu", rankId_, nhrStepInfo.step,
     222              :         nhrStepInfo.toRank, nhrStepInfo.fromRank, sendSliceIdxList.size());
     223            0 : }
     224              : 
     225            0 : void CcuContextReduceNHR1DMem2mem::DoWriteReduceSlice(
     226              :     const u32& toRank, CcuRep::Memory& src, CcuRep::Memory& dst, const u32& sendSliceIdx, u32 signalIndex)
     227              : {
     228            0 :     CcuTransport* sendTransport = transports[indexMap_[toRank]];
     229              :     bool islastSlice;
     230              : 
     231              :     // 添加 die1 偏移
     232            0 :     if (axisId_ == 1) {
     233            0 :         dst.addr += die0Size_;
     234            0 :         src.addr += die0Size_;
     235              :     }
     236              : 
     237              :     // allreduce切片的最后一块slice,大小可能不一致
     238            0 :     islastSlice = (sendSliceIdx + 1 == dimSize_);
     239            0 :     const CcuRep::Variable& sliceSize = axisId_ == 0 ? (islastSlice ? die0LastSliceSize_ : die0SliceSize_) :
     240              :                                                        (islastSlice ? die1LastSliceSize_ : die1SliceSize_);
     241              : 
     242            0 :     CCU_IF(sliceSize != 0)
     243              :     {
     244            0 :         WriteReduce(*sendTransport, dst, src, sliceSize, dataType_, reduceOp_, localSignal_, 1 << signalIndex);
     245            0 :     }
     246            0 :     CCU_IF(sliceSize == 0) { LocalPost(localSignal_, 1 << signalIndex); }
     247            0 : }
     248              : 
     249            0 : void CcuContextReduceNHR1DMem2mem::DoGatherNHR()
     250              : {
     251            0 :     const uint32_t NHR_NUM = 2;
     252            0 :     for (u64 i = stepInfoVector_.size() / NHR_NUM; i < stepInfoVector_.size(); i++) {
     253            0 :         const NHRStepInfo& nhrStepInfo = stepInfoVector_[i];
     254            0 :         DoGatherNHRSingleStep(nhrStepInfo);
     255              :     }
     256            0 : }
     257              : 
     258            0 : void CcuContextReduceNHR1DMem2mem::DoGatherNHRSingleStep(const NHRStepInfo& nhrStepInfo)
     259              : {
     260            0 :     u32& toRankIdx = indexMap_[nhrStepInfo.toRank];
     261            0 :     u32& fromRankIdx = indexMap_[nhrStepInfo.fromRank];
     262            0 :     CcuTransport* recvTransport = transports[fromRankIdx];
     263            0 :     CcuTransport* sendTransport = transports[toRankIdx];
     264            0 :     u32 sendSliceIdx = 0;
     265            0 :     const std::vector<u32>& sendSliceIdxList = nhrStepInfo.txSliceIdxs;
     266            0 :     srcMem_.token = token_[myRankIdx_];
     267            0 :     dstMem_.token = token_[toRankIdx];
     268              : 
     269            0 :     uint16_t selfSignalId = rankId_ / RANK_NUM_PER_CKE;
     270            0 :     uint16_t selfBit = 1 << (rankId_ % RANK_NUM_PER_CKE);
     271              : 
     272            0 :     for (u32 i = 0; i < sendSliceIdxList.size(); i++) {
     273            0 :         sendSliceIdx = sendSliceIdxList[i];
     274              : 
     275            0 :         if (i != 0) {
     276            0 :             if (i % RANK_NUM_PER_CKE == 0) {
     277            0 :                 LocalWait(localSignal_, (1 << RANK_NUM_PER_CKE) - 1);
     278              :             }
     279              :         }
     280              : 
     281            0 :         srcMem_.addr = output_[myRankIdx_];
     282            0 :         srcMem_.addr += sliceOffset_[sendSliceIdx];
     283              : 
     284            0 :         dstMem_.addr = output_[toRankIdx];
     285            0 :         dstMem_.addr += sliceOffset_[sendSliceIdx];
     286            0 :         DoSendRecvSlice(nhrStepInfo.toRank, srcMem_, dstMem_, sendSliceIdx, i % RANK_NUM_PER_CKE);
     287              :     }
     288            0 :     LocalWait(localSignal_, (1 << (sendSliceIdxList.size() % RANK_NUM_PER_CKE)) - 1);
     289              : 
     290            0 :     if (nhrStepInfo.step + 1 != stepInfoVector_.size()) { // 最后一步不需要同步
     291              :         // 通知toRank,写入完毕
     292            0 :         RemotePost(*sendTransport, selfSignalId + signalNum_ * CKE_IDX_3, selfBit, true);
     293              :         // 等待fromRank通知写入完毕
     294            0 :         uint16_t recvSignalId = nhrStepInfo.fromRank / RANK_NUM_PER_CKE;
     295            0 :         uint16_t recvBit = 1 << (nhrStepInfo.fromRank % RANK_NUM_PER_CKE);
     296            0 :         RemoteWait(*recvTransport, recvSignalId + signalNum_ * CKE_IDX_3, recvBit);
     297              :     }
     298              : 
     299            0 :     HCCL_DEBUG(
     300              :         "[DoAllGatherNHRSingleStep] rank %u step %u, toRank=%u, fromRank=%u, nSlice=%lu", rankId_, nhrStepInfo.step,
     301              :         nhrStepInfo.toRank, nhrStepInfo.fromRank, sendSliceIdxList.size());
     302            0 : }
     303              : 
     304            0 : void CcuContextReduceNHR1DMem2mem::DoSendRecvSlice(
     305              :     const u32& toRank, CcuRep::Memory& src, CcuRep::Memory& dst, const u32& sendSliceIdx, u32 signalIndex)
     306              : {
     307            0 :     CcuTransport* sendTransport = transports[indexMap_[toRank]];
     308              :     bool islastSlice;
     309              : 
     310              :     // 添加 die1 偏移
     311            0 :     if (axisId_ == 1) {
     312            0 :         src.addr += die0Size_;
     313            0 :         dst.addr += die0Size_;
     314              :     }
     315              : 
     316            0 :     islastSlice = (sendSliceIdx + 1 == dimSize_);
     317            0 :     const CcuRep::Variable& sliceSize = axisId_ == 0 ? (islastSlice ? die0LastSliceSize_ : die0SliceSize_) :
     318              :                                                        (islastSlice ? die1LastSliceSize_ : die1SliceSize_);
     319              : 
     320            0 :     CCU_IF(sliceSize != 0) { Write(*sendTransport, dst, src, sliceSize, localSignal_, 1 << signalIndex); }
     321            0 :     CCU_IF(sliceSize == 0) { LocalPost(localSignal_, 1 << signalIndex); }
     322            0 : }
     323              : 
     324            0 : void CcuContextReduceNHR1DMem2mem::LocalCopySlices()
     325              : {
     326            0 :     u32 nonTxSliceIdx = 0;
     327            0 :     CcuRep::Variable tmpSliceOffset = CreateVariable();
     328            0 :     tmpSliceOffset = 0;
     329              : 
     330            0 :     for (u64 i = 0; i < dimSize_; i++) {
     331            0 :         sliceOffset_.push_back(CreateVariable());
     332            0 :         sliceOffset_[i] = tmpSliceOffset;
     333            0 :         tmpSliceOffset += axisId_ == 0 ? die0SliceSize_ : die1SliceSize_;
     334              :     }
     335              : 
     336              :     // 当input == output时,不需要拷贝
     337            0 :     CCU_IF(isInputOutputEqual_ == 0)
     338              :     {
     339              :         // 将step0中不需要写的slice,拷贝到本rank的output中
     340            0 :         const NHRStepInfo& nhrStepInfo = stepInfoVector_[0];
     341            0 :         const std::vector<u32>& nonTxSliceIdxList = GetNonTxSliceIdxs(nhrStepInfo.txSliceIdxs);
     342            0 :         for (u32 i = 0; i < nonTxSliceIdxList.size(); i++) {
     343            0 :             nonTxSliceIdx = nonTxSliceIdxList[i];
     344              : 
     345            0 :             if (i != 0) {
     346            0 :                 if (i % RANK_NUM_PER_CKE == 0) {
     347            0 :                     LocalWait(localSignal_, (1 << RANK_NUM_PER_CKE) - 1);
     348              :                 }
     349              :             }
     350              : 
     351            0 :             srcMem_.addr = input_;
     352            0 :             srcMem_.addr += sliceOffset_[nonTxSliceIdx];
     353            0 :             srcMem_.token = token_[myRankIdx_];
     354              : 
     355            0 :             dstMem_.addr = output_[myRankIdx_];
     356            0 :             dstMem_.addr += sliceOffset_[nonTxSliceIdx];
     357            0 :             dstMem_.token = token_[myRankIdx_];
     358            0 :             DoLocalCopySlice(srcMem_, dstMem_, nonTxSliceIdx, i);
     359              :         }
     360            0 :         LocalWait(localSignal_, (1 << (nonTxSliceIdxList.size() % RANK_NUM_PER_CKE)) - 1);
     361            0 :     }
     362            0 : }
     363              : 
     364            0 : std::vector<u32> CcuContextReduceNHR1DMem2mem::GetNonTxSliceIdxs(const std::vector<u32>& txSliceIdxs) const
     365              : {
     366            0 :     std::vector<bool> isTx(dimSize_, false);
     367            0 :     for (u32 idx : txSliceIdxs) {
     368            0 :         if (idx < dimSize_) {
     369            0 :             isTx[idx] = true;
     370              :         }
     371              :     }
     372              : 
     373            0 :     std::vector<u32> nonTxSliceIdxs;
     374            0 :     for (u32 idx = 0; idx < dimSize_; ++idx) {
     375            0 :         if (!isTx[idx]) {
     376            0 :             nonTxSliceIdxs.push_back(idx);
     377              :         }
     378              :     }
     379              : 
     380            0 :     return nonTxSliceIdxs;
     381            0 : }
     382              : 
     383            0 : void CcuContextReduceNHR1DMem2mem::DoLocalCopySlice(
     384              :     CcuRep::Memory& src, CcuRep::Memory& dst, const u32& copySliceIdx, u32 signalIndex)
     385              : {
     386              :     bool islastSlice;
     387              :     // 添加 die1 偏移
     388            0 :     if (axisId_ == 1) {
     389            0 :         dst.addr += die0Size_;
     390            0 :         src.addr += die0Size_;
     391              :     }
     392              : 
     393            0 :     islastSlice = (copySliceIdx + 1 == dimSize_);
     394            0 :     const CcuRep::Variable& sliceSize = axisId_ == 0 ? (islastSlice ? die0LastSliceSize_ : die0SliceSize_) :
     395              :                                                        (islastSlice ? die1LastSliceSize_ : die1SliceSize_);
     396              : 
     397            0 :     CCU_IF(sliceSize != 0) { LocalCopy(dst, src, sliceSize, localSignal_, 1 << signalIndex); }
     398            0 :     CCU_IF(sliceSize == 0) { LocalPost(localSignal_, 1 << signalIndex); }
     399            0 : }
     400              : 
     401            0 : void CcuContextReduceNHR1DMem2mem::Algorithm()
     402              : {
     403            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] AllReduceNHR1D run");
     404              : 
     405            0 :     InitResources();
     406            0 :     LoadArgs();
     407            0 :     if (axisSize_ > 1) {
     408            0 :         AxisSync(FST_AXIS_ID);
     409              :     }
     410            0 :     LocalCopySlices();
     411            0 :     PreSync();
     412            0 :     DoReduceScatterNHR();
     413            0 :     DoGatherNHR();
     414            0 :     PostSync();
     415            0 :     if (axisSize_ > 1) {
     416            0 :         AxisSync(SEC_AXIS_ID);
     417              :     }
     418              : 
     419            0 :     HCCL_DEBUG("[CcuContextReduceNHR1DMem2mem] AllReduceNHR1D end");
     420            0 :     return;
     421              : }
     422              : 
     423            0 : std::vector<uint64_t> CcuContextReduceNHR1DMem2mem::GeneArgs(const CcuTaskArg& arg)
     424              : {
     425            0 :     const CcuTaskArgReduceNHR1D* taskArg = dynamic_cast<const CcuTaskArgReduceNHR1D*>(&arg);
     426            0 :     if (taskArg == nullptr) {
     427            0 :         THROW<NullPtrException>(StringFormat("CcuContextReduceNHR1DMem2mem::taskArg ptr is null"));
     428              :     }
     429              :     // input&output&buffer地址
     430            0 :     uint64_t inputAddr = taskArg->inputAddr_;
     431            0 :     uint64_t outputAddr = taskArg->outputAddr_;
     432            0 :     uint64_t token = taskArg->token_;
     433            0 :     uint64_t isInputOutputEqual = taskArg->isInputOutputEqual_;
     434            0 :     uint64_t die0Size = taskArg->die0Size_;
     435            0 :     uint64_t die1Size = taskArg->die1Size_;
     436            0 :     uint64_t die0SliceSize = taskArg->die0SliceSize_;
     437            0 :     uint64_t die1SliceSize = taskArg->die1SliceSize_;
     438            0 :     uint64_t die0LastSliceSize = taskArg->die0LastSliceSize_;
     439            0 :     uint64_t die1LastSliceSize = taskArg->die1LastSliceSize_;
     440              : 
     441            0 :     HCCL_INFO(
     442              :         "[CcuContextReduceNHR1DMem2mem] TaskArgs: inputAddr[%llu], outputAddr[%llu], "
     443              :         "die0Size[%llu], die1Size[%llu], die0SliceSize[%llu], die1SliceSize[%llu],"
     444              :         "die0LastSliceSize[%llu], die1LastSliceSize[%llu], isInputOutputEqual is [%lu]",
     445              :         inputAddr, outputAddr, die0Size, die1Size, die0SliceSize, die1SliceSize, die0LastSliceSize, die1LastSliceSize,
     446              :         isInputOutputEqual);
     447              : 
     448              :     return {inputAddr, outputAddr,    token,         isInputOutputEqual, die0Size,
     449            0 :             die1Size,  die0SliceSize, die1SliceSize, die0LastSliceSize,  die1LastSliceSize};
     450              : }
     451              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1