LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/prim_alg_template - temp_all_gather_concurr_mesh.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 197 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 13 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              : 
      13              : #include "temp_all_gather_concurr_mesh.h"
      14              : 
      15              : namespace Hccl {
      16            0 : TempAllGatherConcurrMesh::TempAllGatherConcurrMesh(
      17              :     const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
      18            0 :     const std::map<RankId, u32>& tempVirtRankMap)
      19            0 :     : AlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      20            0 : {}
      21              : 
      22            0 : TempAllGatherConcurrMesh::~TempAllGatherConcurrMesh() {}
      23              : 
      24            0 : HcclResult TempAllGatherConcurrMesh::CalcRes(AlgTempResReq& tempResReq)
      25              : {
      26            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      27            0 :         tempResReq.queNum += tempVTopo_[dim].size() - 1;
      28              :     }
      29              : 
      30              :     u32 myAlgRank;
      31            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      32            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
      33            0 :         for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
      34              :             // find neighbors -> virtualRank
      35            0 :             u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
      36            0 :             RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
      37            0 :             HCCL_INFO(
      38              :                 "[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], Dim [%u], NeighborRank [%d].", myRank_, dim,
      39              :                 neighborRank);
      40              : 
      41              :             // LinkNum
      42            0 :             tempResReq.links[neighborRank] = 1;
      43              :         }
      44              :     }
      45            0 :     return HcclResult::HCCL_SUCCESS;
      46              : }
      47              : 
      48              : /*
      49              : dataSize / (rankSize) --> chunkSize
      50              : dataSize / (rankSize * dimNum) --> sliceSize
      51              : 
      52              : SliceInfoVecforConcurrMesh: [1st chunk: [1st Slice, 2nd Slice], 2nd chunk: [1st Slice, 2nd Slice], ...]
      53              : */
      54              : HcclResult
      55            0 : TempAllGatherConcurrMesh::CalcSliceInfo(const AllignInfo& allignInfo, const u64 dataSize, RankSliceInfo& sliceInfoVec)
      56              : {
      57            0 :     u32 dimSize = 0;
      58            0 :     for (u32 dimIdx = 0; dimIdx < tempVTopo_.size(); dimIdx++) {
      59            0 :         if (tempVTopo_[dimIdx].size() != 1) {
      60            0 :             dimSize += 1;
      61              :         }
      62              :     }
      63            0 :     std::vector<SliceInfo> tmp(dimSize);
      64            0 :     sliceInfoVec.resize(tempRankSize_, tmp);
      65              : 
      66            0 :     if (sliceInfoVec[0].size() == 1) {
      67              :         // one-dimensional mesh
      68            0 :         CHK_RET(CalcRsAgSliceInfoMesh(myRank_, tempRankSize_, allignInfo, dataSize, sliceInfoVec));
      69              :     } else {
      70              :         // multi-dimensional mesh
      71            0 :         CHK_RET(CalcRsAgSliceInfoConcurrMesh(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
      72              :     }
      73              : 
      74            0 :     return HcclResult::HCCL_SUCCESS;
      75            0 : }
      76              : 
      77            0 : HcclResult TempAllGatherConcurrMesh::GenPrimQue(
      78              :     const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, const BuffInfo& buffInfo, const ResLinks& tempLinks,
      79              :     std::vector<PrimQuePtr>& tempPrimQues)
      80              : {
      81            0 :     opMode_ = tempFuncs.opMode;
      82            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
      83            0 :     buffInfo_ = buffInfo;
      84            0 :     HCCL_INFO(
      85              :         "[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], EnableCounterNotify [%d].", myRank_,
      86              :         enableCounterNotify_);
      87              : 
      88            0 :     queNum_ = 0;
      89            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      90            0 :         queNum_ += tempVTopo_[dim].size() - 1;
      91              :     }
      92            0 :     CHK_PRT_RET(
      93              :         queNum_ != tempPrimQues.size(),
      94              :         HCCL_ERROR("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], requiredQue Error.", myRank_),
      95              :         HcclResult::HCCL_E_INTERNAL);
      96              : 
      97              :     // Local Copy from Input to Scratch Buffer for OPBASE
      98            0 :     if ((opMode_ == OpMode::OPBASE) && tempFuncs.isForepart && !tempFuncs.forAllReduce) {
      99            0 :         CHK_RET(PreCopyOpbase(tempFuncs.usrData, tempPrimQues));
     100              :     }
     101              : 
     102              :     // Local Copy from Input to Output Buffer for OFFLOAD
     103            0 :     if ((opMode_ == OpMode::OFFLOAD) && (!tempFuncs.forAlgSeqComb)) {
     104            0 :         CHK_RET(PreCopyOffload(sliceInfoVec, tempFuncs.forAllReduce, tempPrimQues));
     105              :     }
     106              : 
     107            0 :     if (sliceInfoVec[0].size() == 1) {
     108            0 :         CHK_RET(RunOneDimMesh(sliceInfoVec, tempLinks, tempPrimQues));
     109              :     } else {
     110            0 :         CHK_RET(RunConcurrMesh(sliceInfoVec, tempLinks, tempPrimQues));
     111              :     }
     112              : 
     113              :     // LocalCopy: from scratch to output for opbase
     114            0 :     if ((opMode_ == OpMode::OPBASE) && tempFuncs.isBottom) {
     115            0 :         CHK_RET(PostCopyOpbase(tempFuncs.usrData, tempPrimQues));
     116              :     }
     117              : 
     118            0 :     return HcclResult::HCCL_SUCCESS;
     119              : }
     120              : 
     121            0 : HcclResult TempAllGatherConcurrMesh::RunOneDimMesh(
     122              :     const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<PrimQuePtr>& tempPrimQues)
     123              : {
     124              :     // semaphore sync
     125            0 :     if (queNum_ > 1) {
     126            0 :         CHK_RET(PreSyncInterQueues(tempPrimQues));
     127              :     }
     128              : 
     129              :     // locate myRank in tempVTopo -> algRank
     130              :     u32 myAlgRank;
     131            0 :     u32 validDim = (tempVTopo_[0].size() == 1) ? 1 : 0;
     132            0 :     HCCL_INFO("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], valid Dim [%u].", myRank_, validDim);
     133            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[validDim], myAlgRank));
     134              : 
     135              :     // runMesh
     136            0 :     CHK_PRT_RET(
     137              :         RunMesh(myAlgRank, tempVTopo_[validDim], sliceInfoVec, tempLinks, tempPrimQues) != HcclResult::HCCL_SUCCESS,
     138              :         HCCL_ERROR("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], unable to run the mesh algorithm.", myRank_),
     139              :         HcclResult::HCCL_E_INTERNAL);
     140              : 
     141              :     // semaphore sync
     142            0 :     if (queNum_ > 1) {
     143            0 :         CHK_RET(PostSyncInterQueues(tempPrimQues));
     144              :     }
     145              : 
     146            0 :     return HcclResult::HCCL_SUCCESS;
     147              : }
     148              : 
     149            0 : HcclResult TempAllGatherConcurrMesh::RunConcurrMesh(
     150              :     const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<PrimQuePtr>& tempPrimQues)
     151              : {
     152            0 :     std::vector<u32> myAlgRank;
     153            0 :     std::vector<std::vector<PrimQuePtr>> dimQues;
     154            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     155              :         // locate myRank in tempVTopo -> algRank
     156              :         u32 tmpAlgRank;
     157            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], tmpAlgRank));
     158            0 :         myAlgRank.push_back(tmpAlgRank);
     159              : 
     160              :         // assign queues
     161            0 :         std::vector<PrimQuePtr> tmpQue;
     162            0 :         for (u32 idx = 0; idx < tempVTopo_[dim].size() - 1; idx++) {
     163            0 :             if (dim == 0) {
     164            0 :                 tmpQue.push_back(tempPrimQues[idx]);
     165              :             } else {
     166            0 :                 tmpQue.push_back(tempPrimQues[tempVTopo_[0].size() - 1 + idx]);
     167              :             }
     168              :         }
     169            0 :         dimQues.push_back(tmpQue);
     170            0 :     }
     171              : 
     172            0 :     std::vector<PrimQuePtr> majorDimQue = {tempPrimQues[0], tempPrimQues[tempVTopo_[0].size() - 1]};
     173              : 
     174              :     // semaphore sync inter dimensions
     175            0 :     CHK_RET(PreSyncInterQueues(majorDimQue));
     176              : 
     177              :     // run concurrent Mesh Step 0
     178            0 :     u32 step = 0;
     179            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     180            0 :         CHK_RET(RunSingleDimension(step, dim, sliceInfoVec, tempLinks, dimQues[dim]));
     181              :     }
     182              : 
     183              :     // semaphore sync
     184            0 :     CHK_RET(PostSyncInterQueues(majorDimQue));
     185              : 
     186              :     // semaphore sync inter dimensions
     187            0 :     CHK_RET(PreSyncInterQueues(majorDimQue));
     188              : 
     189              :     // run concurrent Mesh Step 1
     190            0 :     step = 1;
     191            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     192            0 :         CHK_RET(RunSingleDimension(step, dim, sliceInfoVec, tempLinks, dimQues[dim]));
     193              :     }
     194              : 
     195              :     // semaphore sync
     196            0 :     CHK_RET(PostSyncInterQueues(majorDimQue));
     197              : 
     198            0 :     return HcclResult::HCCL_SUCCESS;
     199            0 : }
     200              : 
     201            0 : HcclResult TempAllGatherConcurrMesh::PreCopyOffload(
     202              :     const RankSliceInfo& sliceInfoVec, const bool forAllReduce, std::vector<PrimQuePtr>& tempPrimQues)
     203              : {
     204            0 :     u64 srcOffset = 0;
     205            0 :     if (forAllReduce) {
     206            0 :         srcOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
     207              :     }
     208              : 
     209            0 :     u64 dstOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
     210              : 
     211            0 :     u64 srcSize = 0;
     212            0 :     for (u32 dimIdx = 0; dimIdx < sliceInfoVec[0].size(); dimIdx++) {
     213            0 :         srcSize += sliceInfoVec[tempVirtRankMap_[myRank_]][dimIdx].size;
     214              :     }
     215              : 
     216            0 :     DataSlice srcSlice = DataSlice(buffInfo_.inBuffType, srcOffset + buffInfo_.inBuffBaseOff, srcSize);
     217            0 :     DataSlice dstSlice = DataSlice(buffInfo_.outBuffType, dstOffset + buffInfo_.outBuffBaseOff, srcSize);
     218            0 :     std::unique_ptr<Primitive> primLocalCopy = std::make_unique<PrimLocalCopy>(srcSlice, dstSlice);
     219            0 :     tempPrimQues[0]->Append(std::move(primLocalCopy));
     220              : 
     221            0 :     return HcclResult::HCCL_SUCCESS;
     222            0 : }
     223              : 
     224            0 : HcclResult TempAllGatherConcurrMesh::RunMesh(
     225              :     const u32 myAlgRank, const std::vector<RankId>& vTopo, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
     226              :     std::vector<PrimQuePtr>& tempPrimQues)
     227              : {
     228            0 :     for (u32 queIdx = 0; queIdx < tempPrimQues.size(); queIdx++) {
     229              :         // find neighbors -> virtualRank
     230            0 :         RankId neighborRank = vTopo[(myAlgRank + 1 + queIdx) % tempRankSize_];
     231              :         // Link
     232            0 :         LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
     233              : 
     234            0 :         u32 sendChunkIdx = tempVirtRankMap_[myRank_];
     235            0 :         u64 sendOffset = sliceInfoVec[sendChunkIdx][0].offset;
     236            0 :         u64 sendSize = sliceInfoVec[sendChunkIdx][0].size;
     237            0 :         u32 recvChunkIdx = tempVirtRankMap_[neighborRank];
     238            0 :         u64 recvOffset = sliceInfoVec[recvChunkIdx][0].offset;
     239            0 :         u64 recvSize = sliceInfoVec[recvChunkIdx][0].size;
     240              : 
     241              :         // PrimGroup
     242            0 :         std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
     243              : 
     244              :         // Send
     245            0 :         DataSlice sendLocSlice = DataSlice(buffInfo_.outBuffType, sendOffset + buffInfo_.outBuffBaseOff, sendSize);
     246            0 :         DataSlice sendRemSlice = DataSlice(buffInfo_.outBuffType, sendOffset + buffInfo_.outBuffBaseOff, sendSize);
     247              :         std::unique_ptr<Primitive> primSend
     248            0 :             = std::make_unique<PrimSend>(neighborRank, neighborLinkData, sendLocSlice, sendRemSlice, dmaMode_);
     249              : 
     250            0 :         primGroup->Append(std::move(primSend));
     251              : 
     252              :         // Recv
     253            0 :         DataSlice recvRemSlice = DataSlice(buffInfo_.outBuffType, recvOffset + buffInfo_.outBuffBaseOff, recvSize);
     254            0 :         DataSlice recvLocSlice = DataSlice(buffInfo_.outBuffType, recvOffset + buffInfo_.outBuffBaseOff, recvSize);
     255              :         std::unique_ptr<Primitive> primRecv
     256            0 :             = std::make_unique<PrimRecv>(neighborRank, neighborLinkData, recvLocSlice, recvRemSlice, dmaMode_);
     257              : 
     258            0 :         primGroup->Append(std::move(primRecv));
     259              : 
     260            0 :         tempPrimQues[queIdx]->Append(std::move(primGroup));
     261            0 :     }
     262            0 :     return HcclResult::HCCL_SUCCESS;
     263              : }
     264              : 
     265            0 : HcclResult TempAllGatherConcurrMesh::RunSingleDimension(
     266              :     const u32& step, const u32& dim, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
     267              :     std::vector<PrimQuePtr>& dimPrimQues)
     268              : {
     269            0 :     CHK_PRT_RET(
     270              :         dim > 1, HCCL_ERROR("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], invalid dim [%u].", myRank_, dim),
     271              :         HcclResult::HCCL_E_INTERNAL);
     272              : 
     273              :     // locate myRank in tempVTopo -> algRank
     274              :     u32 myAlgRank;
     275            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
     276              : 
     277            0 :     for (u32 queIdx = 0; queIdx < dimPrimQues.size(); queIdx++) {
     278              :         // semaphore sync
     279            0 :         if (dimPrimQues.size() > 1) {
     280            0 :             CHK_PRT_RET(
     281              :                 PreSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
     282              :                 HCCL_ERROR(
     283              :                     "[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], Que [%u], Semaphore "
     284              :                     "Synchronization Failed.",
     285              :                     myRank_, dimPrimQues[queIdx]->GetId()),
     286              :                 HcclResult::HCCL_E_INTERNAL);
     287              :         }
     288              : 
     289              :         // find neighbors -> virtualRank
     290            0 :         u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
     291            0 :         RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
     292              : 
     293              :         // link
     294            0 :         LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
     295              : 
     296              :         // PrimGroup
     297            0 :         std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
     298              : 
     299            0 :         std::vector<u32> sendChunkIdxs;
     300            0 :         std::vector<u32> recvChunkIdxs;
     301              : 
     302            0 :         if (step == 0) {
     303            0 :             sendChunkIdxs.push_back(tempVirtRankMap_[myRank_]);
     304            0 :             recvChunkIdxs.push_back(tempVirtRankMap_[neighborRank]);
     305              :         } else {
     306            0 :             for (u32 chunkIdx = 0; chunkIdx < tempVTopo_[1 - dim].size(); chunkIdx++) {
     307            0 :                 u32 sendChunkIdx = (dim == 0) ? (myAlgRank + chunkIdx * tempVTopo_[0].size()) :
     308            0 :                                                 (myAlgRank * tempVTopo_[0].size() + chunkIdx);
     309            0 :                 sendChunkIdxs.push_back(sendChunkIdx);
     310            0 :                 u32 recvChunkIdx = (dim == 0) ? (neighborAlgRank + chunkIdx * tempVTopo_[0].size()) :
     311            0 :                                                 (neighborAlgRank * tempVTopo_[0].size() + chunkIdx);
     312            0 :                 recvChunkIdxs.push_back(recvChunkIdx);
     313              :             }
     314              :         }
     315              : 
     316              :         // Send
     317            0 :         u32 sliceIdx = (step == 0) ? (1 - dim) : dim;
     318              :         std::unique_ptr<PrimSend> primSend
     319            0 :             = RunSend(sliceInfoVec, sendChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
     320            0 :         primGroup->Append(std::move(primSend));
     321              : 
     322              :         // Recv
     323              :         std::unique_ptr<PrimRecv> primRecv
     324            0 :             = RunRecv(sliceInfoVec, recvChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
     325            0 :         primGroup->Append(std::move(primRecv));
     326              : 
     327            0 :         dimPrimQues[queIdx]->Append(std::move(primGroup));
     328              : 
     329              :         // semaphore sync
     330            0 :         if (dimPrimQues.size() > 1) {
     331            0 :             CHK_PRT_RET(
     332              :                 PostSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
     333              :                 HCCL_ERROR(
     334              :                     "[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], Que [%u], Semaphore "
     335              :                     "Synchronization Failed.",
     336              :                     myRank_, dimPrimQues[queIdx]->GetId()),
     337              :                 HcclResult::HCCL_E_INTERNAL);
     338              :         }
     339            0 :     }
     340              : 
     341            0 :     return HcclResult::HCCL_SUCCESS;
     342              : }
     343              : 
     344            0 : std::unique_ptr<PrimSend> TempAllGatherConcurrMesh::RunSend(
     345              :     const RankSliceInfo& sliceInfoVec, const std::vector<u32>& sendChunkIdxs, const u32& sliceIdx,
     346              :     const RankId& neighborRank, const LinkData& priorLinkData)
     347              : {
     348            0 :     std::unique_ptr<PrimSend> primSend;
     349              :     u64 tmpSendOff;
     350              :     u64 tmpSendSize;
     351            0 :     for (u32 chunkIdx = 0; chunkIdx < sendChunkIdxs.size(); chunkIdx++) {
     352            0 :         if (chunkIdx == 0) {
     353              :             // first slice
     354            0 :             tmpSendOff = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
     355            0 :             tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     356            0 :         } else if (tmpSendOff + tmpSendSize == sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset) {
     357              :             // consequent slice
     358            0 :             tmpSendSize += sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     359              :         } else {
     360              :             DataSlice sendLocSlice
     361            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     362              :             DataSlice sendRemSlice
     363            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     364            0 :             if (!primSend) {
     365            0 :                 primSend.reset(new PrimSend(neighborRank, priorLinkData, sendLocSlice, sendRemSlice, dmaMode_));
     366              :             } else {
     367            0 :                 primSend->Append(sendLocSlice, sendRemSlice);
     368              :             }
     369            0 :             tmpSendOff = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
     370            0 :             tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     371              :         }
     372              : 
     373            0 :         if (chunkIdx == (sendChunkIdxs.size() - 1)) {
     374              :             DataSlice sendLocSlice
     375            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     376              :             DataSlice sendRemSlice
     377            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     378              : 
     379            0 :             if (!primSend) {
     380            0 :                 primSend.reset(new PrimSend(neighborRank, priorLinkData, sendLocSlice, sendRemSlice, dmaMode_));
     381              :             } else {
     382            0 :                 primSend->Append(sendLocSlice, sendRemSlice);
     383              :             }
     384              :         }
     385              :     }
     386              : 
     387            0 :     return primSend;
     388            0 : }
     389              : 
     390            0 : std::unique_ptr<PrimRecv> TempAllGatherConcurrMesh::RunRecv(
     391              :     const RankSliceInfo& sliceInfoVec, const std::vector<u32>& recvChunkIdxs, const u32& sliceIdx,
     392              :     const RankId& neighborRank, const LinkData& priorLinkData)
     393              : {
     394            0 :     std::unique_ptr<PrimRecv> primRecv;
     395              :     u64 tmpRecvOff;
     396              :     u64 tmpRecvSize;
     397            0 :     for (u32 chunkIdx = 0; chunkIdx < recvChunkIdxs.size(); chunkIdx++) {
     398            0 :         if (chunkIdx == 0) {
     399              :             // first slice
     400            0 :             tmpRecvOff = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
     401            0 :             tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     402            0 :         } else if (tmpRecvOff + tmpRecvSize == sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset) {
     403              :             // consequent slice
     404            0 :             tmpRecvSize += sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     405              :         } else {
     406              :             DataSlice recvRemSlice
     407            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     408              :             DataSlice recvLocSlice
     409            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     410            0 :             if (!primRecv) {
     411            0 :                 primRecv.reset(new PrimRecv(neighborRank, priorLinkData, recvLocSlice, recvRemSlice, dmaMode_));
     412              :             } else {
     413            0 :                 primRecv->Append(recvLocSlice, recvRemSlice);
     414              :             }
     415            0 :             tmpRecvOff = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
     416            0 :             tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     417              :         }
     418              : 
     419            0 :         if (chunkIdx == (recvChunkIdxs.size() - 1)) {
     420              :             DataSlice recvRemSlice
     421            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     422              :             DataSlice recvLocSlice
     423            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     424              : 
     425            0 :             if (!primRecv) {
     426            0 :                 primRecv.reset(new PrimRecv(neighborRank, priorLinkData, recvLocSlice, recvRemSlice, dmaMode_));
     427              :             } else {
     428            0 :                 primRecv->Append(recvLocSlice, recvRemSlice);
     429              :             }
     430              :         }
     431              :     }
     432              : 
     433            0 :     return primRecv;
     434            0 : }
     435              : 
     436              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1