LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/ins_alg_template - ins_temp_all_reduce_nhr.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 235 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 15 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 "log.h"
      12              : #include "alg_data_trans_wrapper.h"
      13              : #include "ins_temp_all_reduce_nhr.h"
      14              : 
      15              : namespace Hccl {
      16            0 : InsTempAllReduceNHR::InsTempAllReduceNHR(
      17              :     const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
      18            0 :     const std::map<RankId, u32>& tempVirtRankMap)
      19            0 :     : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      20            0 : {}
      21              : 
      22            0 : InsTempAllReduceNHR::~InsTempAllReduceNHR() {}
      23              : 
      24            0 : HcclResult InsTempAllReduceNHR::CalcRes(AlgTempResReq& tempResReq)
      25              : {
      26              :     // NHR 需要的 que Num 为 1
      27            0 :     CHK_PRT_RET(
      28              :         CalcResLinksNHR(myRank_, tempRankSize_, tempVTopo_, tempResReq) != HcclResult::HCCL_SUCCESS,
      29              :         HCCL_ERROR("[CollAlgFactory] [InsTempAllReduceNHR] Rank [%d], resLinks calculation error!", myRank_),
      30              :         HcclResult::HCCL_E_INTERNAL);
      31            0 :     auto& linkReq = tempResReq.links;
      32            0 :     u32 pathNum = 0;
      33            0 :     for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
      34            0 :         auto remoteRank = resReqIter->first;
      35            0 :         if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_[remoteRank] == 0) {
      36            0 :             HCCL_ERROR("[InsTempAllReduceNHR] No path to remoteRank[%d]", remoteRank);
      37            0 :             return HcclResult::HCCL_E_INTERNAL;
      38              :         }
      39            0 :         if (pathNum == 0) {
      40            0 :             pathNum = rank2PathNumMap_[remoteRank];
      41            0 :         } else if (rank2PathNumMap_[remoteRank] != pathNum) {
      42            0 :             HCCL_ERROR(
      43              :                 "[InsTempAllReduceNHR] Inconsistency pathNum to remoteRanks, Previous consistent pathNum=[%u], "
      44              :                 "mismatched "
      45              :                 "remoteRank=[%d], pathNum=[%u]",
      46              :                 pathNum, remoteRank, rank2PathNumMap_[remoteRank]);
      47            0 :             return HcclResult::HCCL_E_INTERNAL;
      48              :         }
      49            0 :         resReqIter->second = pathNum;
      50              :     }
      51            0 :     tempResReq.queNum = 1 * pathNum;
      52            0 :     HCCL_INFO("[InsTempAllReduceNHR] tempResReq.queNum = %u", tempResReq.queNum);
      53            0 :     tempResReq.streamNum = tempResReq.queNum;
      54            0 :     tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
      55              : 
      56            0 :     return HcclResult::HCCL_SUCCESS;
      57              : }
      58              : 
      59              : /*
      60              :  * Desc: 将数据按照rank切分为chuck 块,给后续的allreduce操作使用
      61              :  * param: dataSize: 待处理的输入数据大小
      62              :  * return: sliceInfoVec: 存储数据切分结果
      63              :  * return: HcclResult
      64              :  */
      65            0 : HcclResult InsTempAllReduceNHR::CalcSlice(const u64 dataSize, const u64 baseOff, RankSliceInfo& sliceInfoVec)
      66              : {
      67            0 :     std::vector<SliceInfo> tmp(tempVTopo_.size());
      68            0 :     sliceInfoVec.resize(tempRankSize_, tmp);
      69              : 
      70            0 :     u64 unitAllignSize = DataTypeSizeGet(dataType_);
      71            0 :     u64 chunkSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
      72              : 
      73            0 :     u64 accumOff = 0;
      74            0 :     for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
      75            0 :         u64 currChunkSize = ((dataSize - accumOff) > chunkSize) ? chunkSize : (dataSize - accumOff);
      76            0 :         SliceInfo slice = {accumOff + baseOff, currChunkSize};
      77            0 :         sliceInfoVec[rankIdx][0] = slice;
      78            0 :         accumOff += currChunkSize;
      79              :     }
      80              : 
      81            0 :     CHK_PRT_RET(
      82              :         (sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != baseOff + dataSize),
      83              :         HCCL_ERROR(
      84              :             "[InsTempAllReduceNHR] chunkSize:[%llu], Rank:[%d], SliceInfo calculation error!", chunkSize, myRank_),
      85              :         HcclResult::HCCL_E_INTERNAL);
      86            0 :     return HcclResult::HCCL_SUCCESS;
      87            0 : }
      88              : 
      89              : /*
      90              :  * Desc: 返回当前rank能处理的数据量和scratch buffer之间的比例关系
      91              :  * param: input: 输入数据位置
      92              :  * param: output 输出数据位置
      93              :  */
      94            0 : u32 InsTempAllReduceNHR::CalcScratchMultiple(BufferType input, BufferType output)
      95              : {
      96              :     (void)input;
      97              :     (void)output;
      98              :     // 单算子模式,cclBuffer和usrIn一样大,图模式,不需要cclBuffer
      99            0 :     u32 multiple = 0;
     100            0 :     if (op_.opMode == OpMode::OPBASE) {
     101            0 :         multiple = 1;
     102              :     }
     103              : 
     104            0 :     return multiple;
     105              : }
     106              : 
     107            0 : HcclResult InsTempAllReduceNHR::GenExtIns(
     108              :     const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
     109              :     std::vector<InsQuePtr>& tempInsQues)
     110              : {
     111            0 :     HCCL_INFO("[InsTempAllReduceNHR][GenExtIns] AllReduceNHR begin: rank[%d] start", myRank_);
     112            0 :     if (IsPcieLink(tempLinks)) {
     113            0 :         dmaMode_ = DmaMode::GET;
     114              :     }
     115              : 
     116            0 :     opMode_ = tempFuncs.opMode;
     117            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
     118              : 
     119            0 :     uint32_t linkNum = tempLinks.begin()->second.size();
     120              :     // 流的数量不能少于linkNum
     121            0 :     CHK_PRT_RET(
     122              :         linkNum > tempInsQues.size(),
     123              :         HCCL_ERROR("[CollAlgFactory] [InsTempAllReduceNHR] Rank [%d], requiredQue Error.", myRank_),
     124              :         HcclResult::HCCL_E_INTERNAL);
     125              : 
     126            0 :     std::vector<float> dataSplitRate(linkNum);
     127            0 :     CHK_RET(CalcDataSplitRateForLinks(tempLinks.begin()->second, dataSplitRate));
     128              : 
     129            0 :     RankSliceInfo sliceInfoVec;
     130            0 :     CHK_RET(CalcSlice(tempAlgParams.sliceSize, 0, sliceInfoVec));
     131              : 
     132              :     // 将一个RankSliceInfo,拆分成linkNum 个RankSliceInfo
     133            0 :     u64 typeSize = DataTypeSizeGet(dataType_);
     134            0 :     std::vector<RankSliceInfo> sliceInfoVecForAllLinks(linkNum);
     135            0 :     for (auto sliceInfoPerRank : sliceInfoVec) {
     136            0 :         std::vector<std::vector<SliceInfo>> sliceInfoPerRankForAllLinks(linkNum);
     137            0 :         for (auto sliceInfo : sliceInfoPerRank) {
     138            0 :             u64 size = sliceInfo.size;
     139            0 :             u64 offset = sliceInfo.offset;
     140            0 :             u64 AccSize = 0;
     141            0 :             u64 dataCnt = size / typeSize;
     142            0 :             vector<SliceInfo> sliceInfoForAllLinks(linkNum);
     143            0 :             for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     144            0 :                 if (linkIdx != linkNum - 1) {
     145            0 :                     sliceInfoForAllLinks[linkIdx].size
     146            0 :                         = static_cast<u64>(static_cast<float>(dataCnt) * dataSplitRate[linkIdx]) * typeSize;
     147              :                 } else {
     148            0 :                     sliceInfoForAllLinks[linkIdx].size = size - AccSize;
     149              :                 }
     150            0 :                 sliceInfoForAllLinks[linkIdx].offset = offset + AccSize;
     151            0 :                 AccSize += sliceInfoForAllLinks[linkIdx].size;
     152              :             }
     153            0 :             for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     154            0 :                 sliceInfoPerRankForAllLinks[linkIdx].emplace_back(sliceInfoForAllLinks[linkIdx]);
     155              :             }
     156            0 :         }
     157            0 :         for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     158            0 :             sliceInfoVecForAllLinks[linkIdx].emplace_back(sliceInfoPerRankForAllLinks[linkIdx]);
     159              :         }
     160            0 :     }
     161              : 
     162              :     // 预拷贝
     163            0 :     CHK_RET(PreCopy(tempAlgParams, tempInsQues));
     164              : 
     165            0 :     u32 mainQueIdx = 0;
     166              :     // 流间前同步,主流通知从流,只有一个流则不做任何事
     167            0 :     CHK_RET(PreSyncQues(tempInsQues, mainQueIdx));
     168              : 
     169              :     // 主从流执行nhr
     170              :     // 待修改数据切分方式
     171            0 :     for (uint32_t linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     172            0 :         CHK_RET(RunReduceScatter(sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQues, linkIdx));
     173            0 :         CHK_RET(PrepareDataForAllGather(sliceInfoVecForAllLinks[linkIdx], tempInsQues, linkIdx));
     174            0 :         CHK_RET(RunAllGather(sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQues, linkIdx));
     175              :     }
     176              :     // 流间后同步,从流通知主流
     177            0 :     CHK_RET(PostSyncQues(tempInsQues, mainQueIdx));
     178              :     // 结果拷贝
     179            0 :     CHK_RET(PostCopy(tempAlgParams, tempInsQues));
     180            0 :     HCCL_INFO("[InsTempAllReduceNHR][GenExtIns] AllReduceNHR finished: rank[%d] end", myRank_);
     181            0 :     return HcclResult::HCCL_SUCCESS;
     182            0 : }
     183              : 
     184            0 : HcclResult InsTempAllReduceNHR::PreCopy(const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues)
     185              : {
     186              :     // 单算子模式,需要先将数据拷贝到cclBuffer
     187            0 :     if (opMode_ == OpMode::OPBASE) {
     188            0 :         nhrInBuffType_ = BufferType::SCRATCH;
     189            0 :         nhrInBuffBaseOff_ = tempAlgParams.buffInfo.inBuffBaseOff;
     190              : 
     191            0 :         if (tempAlgParams.buffInfo.inBuffType != BufferType::SCRATCH) {
     192            0 :             HCCL_INFO("[InsTempAllReduceNHR][PreCopy] Opbase copy from userIn to scratchBuffer");
     193              :             DataSlice usrInSlices = DataSlice(
     194            0 :                 tempAlgParams.buffInfo.inBuffType, tempAlgParams.buffInfo.inBuffBaseOff, tempAlgParams.sliceSize);
     195              :             DataSlice scratchSlices
     196            0 :                 = DataSlice(BufferType::SCRATCH, tempAlgParams.buffInfo.scratchBuffBaseOff, tempAlgParams.sliceSize);
     197            0 :             CHK_RET(LocalCopy(tempInsQues[0], usrInSlices, scratchSlices));
     198              : 
     199            0 :             nhrInBuffBaseOff_ = tempAlgParams.buffInfo.scratchBuffBaseOff;
     200              :         } else {
     201            0 :             HCCL_INFO("[InsTempAllReduceNHR][PreCopy] skip precopy");
     202              :         }
     203              :     } else {
     204            0 :         HCCL_INFO("[InsTempAllReduceNHR][PreCopy] offload skip precopy");
     205            0 :         nhrInBuffType_ = tempAlgParams.buffInfo.inBuffType;
     206            0 :         nhrInBuffBaseOff_ = tempAlgParams.buffInfo.inBuffBaseOff;
     207              :     }
     208              : 
     209            0 :     nhrOutBuffType_ = tempAlgParams.buffInfo.outBuffType;
     210            0 :     nhrOutBuffBaseOff_ = tempAlgParams.buffInfo.outBuffBaseOff;
     211              : 
     212            0 :     return HcclResult::HCCL_SUCCESS;
     213              : }
     214              : 
     215              : // 将reduceScatter之后的数据先放到usrOut
     216            0 : HcclResult InsTempAllReduceNHR::PrepareDataForAllGather(
     217              :     const RankSliceInfo& sliceInfoVec, std::vector<InsQuePtr>& tempInsQues, u32 linkIdx)
     218              : {
     219              :     // 如果是单算子模式,在原来的位置要先做完allGather,然后postCopy把数据放到usrOut
     220              :     // 如果是图模式,直接把数据放到usrOUt,然后在usrOut上做allGather
     221            0 :     HCCL_INFO("[InsTempAllReduceNHR][PrepareDataForAllGather] prepare data for allGather");
     222              : 
     223            0 :     if (opMode_ == OpMode::OFFLOAD) {
     224            0 :         u64 size = sliceInfoVec[tempVirtRankMap_[myRank_]][0].size;
     225            0 :         u64 srcOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
     226            0 :         u64 dstOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
     227            0 :         DataSlice srcSlice = DataSlice(nhrInBuffType_, nhrInBuffBaseOff_ + srcOffset, size);
     228            0 :         DataSlice dstSlice = DataSlice(nhrOutBuffType_, nhrOutBuffBaseOff_ + dstOffset, size);
     229            0 :         CHK_RET(LocalCopy(tempInsQues[linkIdx], srcSlice, dstSlice));
     230              : 
     231            0 :         nhrInBuffType_ = nhrOutBuffType_;
     232            0 :         nhrInBuffBaseOff_ = nhrOutBuffBaseOff_;
     233              :     }
     234              : 
     235            0 :     return HcclResult::HCCL_SUCCESS;
     236              : }
     237              : 
     238            0 : HcclResult InsTempAllReduceNHR::PostCopy(const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues)
     239              : {
     240              :     // 单算子模式,需要将数据拷贝到usrOut
     241            0 :     if (opMode_ == OpMode::OPBASE) {
     242            0 :         HCCL_INFO("[InsTempAllReduceNHR][PostCopy] Opbase copy from scratchBuffer to userOut");
     243            0 :         DataSlice scratchSlices = DataSlice(nhrInBuffType_, nhrInBuffBaseOff_, tempAlgParams.sliceSize);
     244            0 :         DataSlice usrOutSlices = DataSlice(nhrOutBuffType_, nhrOutBuffBaseOff_, tempAlgParams.sliceSize);
     245            0 :         CHK_RET(LocalCopy(tempInsQues[0], scratchSlices, usrOutSlices));
     246              :     } else {
     247            0 :         HCCL_INFO("[InsTempAllReduceNHR][PostCopy] offload skip postcopy");
     248              :     }
     249              : 
     250            0 :     return HcclResult::HCCL_SUCCESS;
     251              : }
     252              : 
     253            0 : HcclResult InsTempAllReduceNHR::RunReduceScatter(
     254              :     const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, u32 linkIdx)
     255              : {
     256            0 :     std::vector<AicpuNHRStepInfo> stepInfoList;
     257            0 :     GetStepInfoList(stepInfoList);
     258            0 :     for (auto& stepInfo : stepInfoList) {
     259            0 :         HCCL_DEBUG(
     260              :             "[InsTempAllReduceNHR][RunReduceScatter] step[%u], myRank[%u], toRank[%u], fromRank[%u], nSlices[%u].",
     261              :             stepInfo.step, stepInfo.myRank, stepInfo.toRank, stepInfo.fromRank, stepInfo.nSlices);
     262              : 
     263            0 :         const std::vector<LinkData>& linkRecv = tempLinks.at(GetRankFromMap(stepInfo.fromRank));
     264            0 :         const std::vector<LinkData>& linkSend = tempLinks.at(GetRankFromMap(stepInfo.toRank));
     265            0 :         std::vector<DataSlice> txSlices;
     266            0 :         std::vector<DataSlice> rxSlices;
     267              : 
     268              :         // 在 nhrInBuffType_ 上进行 ReduceScatter 操作
     269            0 :         for (u32 i = 0; i < stepInfo.nSlices; i++) {
     270            0 :             u64 txOffset = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
     271            0 :             u64 txSize = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].size;
     272            0 :             u64 rxOffset = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
     273            0 :             u64 rxSize = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].size;
     274            0 :             DataSlice txSlice = DataSlice(nhrInBuffType_, txOffset, txSize);
     275            0 :             DataSlice rxSlice = DataSlice(nhrInBuffType_, rxOffset, rxSize);
     276            0 :             txSlices.push_back(txSlice);
     277            0 :             rxSlices.push_back(rxSlice);
     278              :         }
     279              :         SendRecvReduceInfo sendRecvReduceInfo{
     280            0 :             {linkSend[linkIdx], linkRecv[linkIdx]}, {{txSlices, txSlices}, {rxSlices, rxSlices}}, dataType_, redOp_};
     281            0 :         CHK_PRT_RET(
     282              :             SendRecvReduce(sendRecvReduceInfo, tempInsQues[linkIdx], 0, true, dmaMode_) != HcclResult::HCCL_SUCCESS,
     283              :             HCCL_ERROR("[InsTempAllReduceNHR] RunReduceScatter SendRecvReduce failed"), HcclResult::HCCL_E_INTERNAL);
     284            0 :     }
     285            0 :     return HcclResult::HCCL_SUCCESS;
     286            0 : }
     287              : 
     288            0 : HcclResult InsTempAllReduceNHR::RunAllGather(
     289              :     const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, u32 linkIdx)
     290              : {
     291            0 :     u32 nSteps = GetNHRStepNum(tempRankSize_);
     292            0 :     for (u32 step = 0; step < nSteps; step++) {
     293            0 :         AicpuNHRStepInfo stepInfo;
     294            0 :         CHK_RET(GetStepInfo(step, nSteps, stepInfo));
     295              : 
     296            0 :         const std::vector<LinkData>& linkRecv = tempLinks.at(GetRankFromMap(stepInfo.fromRank));
     297            0 :         const std::vector<LinkData>& linkSend = tempLinks.at(GetRankFromMap(stepInfo.toRank));
     298              : 
     299            0 :         std::vector<DataSlice> txSlices;
     300            0 :         std::vector<DataSlice> rxSlices;
     301              : 
     302            0 :         HCCL_DEBUG(
     303              :             "[InsTempAllReduceNHR] rank[%d] rankSize[%u] recvFrom[%u] sendTo[%u] step[%u] nSteps[%u] nSlices[%u]",
     304              :             myRank_, tempRankSize_, stepInfo.fromRank, stepInfo.toRank, step, nSteps, stepInfo.nSlices);
     305              : 
     306            0 :         for (u32 i = 0; i < stepInfo.nSlices; i++) {
     307            0 :             u64 txOffset = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
     308            0 :             u64 txSize = sliceInfoVec[stepInfo.txSliceIdxs[i]][0].size;
     309            0 :             u64 rxOffset = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].offset + nhrInBuffBaseOff_;
     310            0 :             u64 rxSize = sliceInfoVec[stepInfo.rxSliceIdxs[i]][0].size;
     311            0 :             DataSlice txSlice = DataSlice(nhrInBuffType_, txOffset, txSize);
     312            0 :             DataSlice rxSlice = DataSlice(nhrInBuffType_, rxOffset, rxSize);
     313            0 :             txSlices.push_back(txSlice);
     314            0 :             rxSlices.push_back(rxSlice);
     315              :         }
     316              : 
     317            0 :         TxRxLinks sendRecvLinks(linkSend[linkIdx], linkRecv[linkIdx]);
     318            0 :         TxRxSlicesList sendRecvSlicesList({txSlices, txSlices}, {rxSlices, rxSlices});
     319              : 
     320            0 :         SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     321            0 :         CHK_PRT_RET(
     322              :             SendRecv(sendRecvInfo, tempInsQues[linkIdx], 0, true, dmaMode_) != HcclResult::HCCL_SUCCESS,
     323              :             HCCL_ERROR("[InsTempAllReduceNHR] RunAllGather send/recv failed"), HcclResult::HCCL_E_INTERNAL);
     324            0 :     }
     325            0 :     return HcclResult::HCCL_SUCCESS;
     326              : }
     327              : 
     328            0 : HcclResult InsTempAllReduceNHR::GetStepInfo(u32 step, u32 nSteps, AicpuNHRStepInfo& stepInfo)
     329              : {
     330            0 :     u32 rankIdx = tempVirtRankMap_[myRank_];
     331            0 :     stepInfo.txSliceIdxs.clear();
     332            0 :     stepInfo.rxSliceIdxs.clear();
     333            0 :     stepInfo.step = step;
     334            0 :     stepInfo.myRank = rankIdx;
     335              : 
     336              :     // 计算通信对象
     337            0 :     u32 deltaRank = 1 << (nSteps - 1 - step);
     338            0 :     u32 recvFrom = (rankIdx + tempRankSize_ - deltaRank) % tempRankSize_;
     339            0 :     u32 sendTo = (rankIdx + deltaRank) % tempRankSize_;
     340              : 
     341              :     // 数据份数和数据编号增量
     342            0 :     u32 nSlices = (tempRankSize_ - 1 + (1 << (nSteps - 1 - step))) / (1 << (nSteps - step));
     343            0 :     u32 deltaSliceIndex = 1 << (nSteps - step);
     344            0 :     u32 txSliceIdx = rankIdx;
     345            0 :     u32 rxSliceIdx = (rankIdx - (1 << (nSteps - 1 - step)) + tempRankSize_) % tempRankSize_;
     346              : 
     347            0 :     stepInfo.nSlices = nSlices;
     348            0 :     stepInfo.toRank = sendTo;
     349            0 :     stepInfo.fromRank = recvFrom;
     350              : 
     351            0 :     for (u32 i = 0; i < nSlices; i++) {
     352            0 :         stepInfo.txSliceIdxs.push_back(txSliceIdx);
     353            0 :         stepInfo.rxSliceIdxs.push_back(rxSliceIdx);
     354              : 
     355            0 :         HCCL_DEBUG("[InsTempAllReduceNHR][GetStepInfo] i[%u] txSliceIdx[%u] rxSliceIdx[%u]", i, txSliceIdx, rxSliceIdx);
     356              : 
     357            0 :         txSliceIdx = (txSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
     358            0 :         rxSliceIdx = (rxSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
     359              :     }
     360            0 :     return HcclResult::HCCL_SUCCESS;
     361              : }
     362              : 
     363              : //  计算每轮收发的对端以及slice编号
     364            0 : HcclResult InsTempAllReduceNHR::GetStepInfoList(std::vector<AicpuNHRStepInfo>& stepInfoList)
     365              : {
     366              :     // 将本 rank 号转换成算法使用的索引号
     367            0 :     u32 rankIdx = tempVirtRankMap_[myRank_];
     368            0 :     stepInfoList.clear();
     369              : 
     370            0 :     u32 nSteps = GetNHRStepNum(tempRankSize_);
     371            0 :     stepInfoList.resize(nSteps);
     372            0 :     for (u32 step = 0; step < nSteps; step++) {
     373              :         // 计算通信对象
     374            0 :         u32 deltaRank = 1 << step;
     375            0 :         u32 sendTo = (rankIdx + tempRankSize_ - deltaRank) % tempRankSize_;
     376            0 :         u32 recvFrom = (rankIdx + deltaRank) % tempRankSize_;
     377              : 
     378              :         // 数据份数和数据编号增量
     379            0 :         u32 nSlices = (tempRankSize_ - 1 + (1 << step)) / (1 << (step + 1));
     380            0 :         u32 deltaSliceIndex = 1 << (step + 1);
     381            0 :         u32 txSliceIdx = sendTo;
     382            0 :         u32 rxSliceIdx = rankIdx;
     383              : 
     384            0 :         AicpuNHRStepInfo& currStepInfo = stepInfoList[step];
     385            0 :         currStepInfo.step = step;
     386            0 :         currStepInfo.myRank = rankIdx;
     387            0 :         currStepInfo.nSlices = nSlices;
     388            0 :         currStepInfo.toRank = sendTo;
     389            0 :         currStepInfo.fromRank = recvFrom;
     390              : 
     391              :         // 计算本rank在每轮收/发中的slice编号
     392            0 :         currStepInfo.txSliceIdxs.reserve(nSlices);
     393            0 :         currStepInfo.rxSliceIdxs.reserve(nSlices);
     394            0 :         for (u32 i = 0; i < nSlices; i++) {
     395            0 :             currStepInfo.txSliceIdxs.push_back(txSliceIdx);
     396            0 :             currStepInfo.rxSliceIdxs.push_back(rxSliceIdx);
     397            0 :             HCCL_DEBUG(
     398              :                 "[InsTempAllReduceNHR][GetStepInfoList] i[%u] txSliceIdx[%u] rxSliceIdx[%u]", i, txSliceIdx,
     399              :                 rxSliceIdx);
     400            0 :             txSliceIdx = (txSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
     401            0 :             rxSliceIdx = (rxSliceIdx + tempRankSize_ - deltaSliceIndex) % tempRankSize_;
     402              :         }
     403              :     }
     404            0 :     return HcclResult::HCCL_SUCCESS;
     405              : }
     406              : 
     407            0 : RankId InsTempAllReduceNHR::GetRankFromMap(const u32 rankIdx)
     408              : {
     409            0 :     RankId rank = -1;
     410            0 :     HCCL_INFO("[InsTempAllReduceNHR] GetRankFromMap");
     411            0 :     for (auto& pair : tempVirtRankMap_) {
     412            0 :         if (pair.second == rankIdx) {
     413            0 :             rank = pair.first;
     414            0 :             break;
     415              :         }
     416              :     }
     417            0 :     return rank;
     418              : }
     419              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1