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 % 198 0
Test Date: 2026-08-04 10:52:23 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(const RankId virtualRank, const u32 tempRankSize,
      17              :                                                    const std::vector<std::vector<RankId>> &tempVTopo,
      18            0 :                                                    const std::map<RankId, u32>            &tempVirtRankMap)
      19            0 :     : AlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      20              : {
      21            0 : }
      22              : 
      23            0 : TempAllGatherConcurrMesh::~TempAllGatherConcurrMesh()
      24              : {
      25            0 : }
      26              : 
      27            0 : HcclResult TempAllGatherConcurrMesh::CalcRes(AlgTempResReq &tempResReq)
      28              : {
      29            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      30            0 :         tempResReq.queNum += tempVTopo_[dim].size() - 1;
      31              :     }
      32              : 
      33              :     u32 myAlgRank;
      34            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      35            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
      36            0 :         for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
      37              :             // find neighbors -> virtualRank
      38            0 :             u32    neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
      39            0 :             RankId neighborRank    = tempVTopo_[dim][neighborAlgRank];
      40            0 :             HCCL_INFO("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], Dim [%u], NeighborRank [%d].", myRank_,
      41              :                        dim, neighborRank);
      42              : 
      43              :             // LinkNum
      44            0 :             tempResReq.links[neighborRank] = 1;
      45              :         }
      46              :     }
      47            0 :     return HcclResult::HCCL_SUCCESS;
      48              : }
      49              : 
      50              : /*
      51              : dataSize / (rankSize) --> chunkSize
      52              : dataSize / (rankSize * dimNum) --> sliceSize
      53              : 
      54              : SliceInfoVecforConcurrMesh: [1st chunk: [1st Slice, 2nd Slice], 2nd chunk: [1st Slice, 2nd Slice], ...]
      55              : */
      56            0 : HcclResult TempAllGatherConcurrMesh::CalcSliceInfo(const AllignInfo &allignInfo, const u64 dataSize,
      57              :                                                    RankSliceInfo &sliceInfoVec)
      58              : {
      59            0 :     u32 dimSize = 0;
      60            0 :     for (u32 dimIdx = 0; dimIdx < tempVTopo_.size(); dimIdx++) {
      61            0 :         if (tempVTopo_[dimIdx].size() != 1) {
      62            0 :             dimSize += 1;
      63              :         }
      64              :     }
      65            0 :     std::vector<SliceInfo> tmp(dimSize);
      66            0 :     sliceInfoVec.resize(tempRankSize_, tmp);
      67              : 
      68            0 :     if (sliceInfoVec[0].size() == 1) {
      69              :         // one-dimensional mesh
      70            0 :         CHK_RET(CalcRsAgSliceInfoMesh(myRank_, tempRankSize_, allignInfo, dataSize, sliceInfoVec));
      71              :     } else {
      72              :         // multi-dimensional mesh
      73            0 :         CHK_RET(CalcRsAgSliceInfoConcurrMesh(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
      74              :     }
      75              : 
      76            0 :     return HcclResult::HCCL_SUCCESS;
      77            0 : }
      78              : 
      79            0 : HcclResult TempAllGatherConcurrMesh::GenPrimQue(const TempFuncs &tempFuncs, const RankSliceInfo &sliceInfoVec,
      80              :                                                 const BuffInfo &buffInfo, const ResLinks &tempLinks,
      81              :                                                 std::vector<PrimQuePtr> &tempPrimQues)
      82              : {
      83            0 :     opMode_              = tempFuncs.opMode;
      84            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
      85            0 :     buffInfo_            = buffInfo;
      86            0 :     HCCL_INFO("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], EnableCounterNotify [%d].", myRank_,
      87              :                enableCounterNotify_);
      88              : 
      89            0 :     queNum_ = 0;
      90            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      91            0 :         queNum_ += tempVTopo_[dim].size() - 1;
      92              :     }
      93            0 :     CHK_PRT_RET(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(const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
     122              :                                                    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(const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
     150              :                                                     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(const RankSliceInfo &sliceInfoVec, const bool forAllReduce,
     202              :                                                     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(const u32 myAlgRank, const std::vector<RankId> &vTopo,
     225              :                                              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(const u32 &step, const u32 &dim,
     266              :                                                         const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
     267              :                                                         std::vector<PrimQuePtr> &dimPrimQues)
     268              : {
     269            0 :     CHK_PRT_RET(dim > 1,
     270              :                 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(PreSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
     281              :                         HCCL_ERROR("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], Que [%u], Semaphore "
     282              :                                    "Synchronization Failed.",
     283              :                                    myRank_, dimPrimQues[queIdx]->GetId()),
     284              :                         HcclResult::HCCL_E_INTERNAL);
     285              :         }
     286              : 
     287              :         // find neighbors -> virtualRank
     288            0 :         u32    neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
     289            0 :         RankId neighborRank    = tempVTopo_[dim][neighborAlgRank];
     290              : 
     291              :         // link
     292            0 :         LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
     293              : 
     294              :         // PrimGroup
     295            0 :         std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
     296              : 
     297            0 :         std::vector<u32> sendChunkIdxs;
     298            0 :         std::vector<u32> recvChunkIdxs;
     299              : 
     300            0 :         if (step == 0) {
     301            0 :             sendChunkIdxs.push_back(tempVirtRankMap_[myRank_]);
     302            0 :             recvChunkIdxs.push_back(tempVirtRankMap_[neighborRank]);
     303              :         } else {
     304            0 :             for (u32 chunkIdx = 0; chunkIdx < tempVTopo_[1 - dim].size(); chunkIdx++) {
     305            0 :                 u32 sendChunkIdx = (dim == 0) ? (myAlgRank + chunkIdx * tempVTopo_[0].size())
     306            0 :                                               : (myAlgRank * tempVTopo_[0].size() + chunkIdx);
     307            0 :                 sendChunkIdxs.push_back(sendChunkIdx);
     308            0 :                 u32 recvChunkIdx = (dim == 0) ? (neighborAlgRank + chunkIdx * tempVTopo_[0].size())
     309            0 :                                               : (neighborAlgRank * tempVTopo_[0].size() + chunkIdx);
     310            0 :                 recvChunkIdxs.push_back(recvChunkIdx);
     311              :             }
     312              :         }
     313              : 
     314              :         // Send
     315            0 :         u32                       sliceIdx = (step == 0) ? (1 - dim) : dim;
     316              :         std::unique_ptr<PrimSend> primSend
     317            0 :             = RunSend(sliceInfoVec, sendChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
     318            0 :         primGroup->Append(std::move(primSend));
     319              : 
     320              :         // Recv
     321              :         std::unique_ptr<PrimRecv> primRecv
     322            0 :             = RunRecv(sliceInfoVec, recvChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
     323            0 :         primGroup->Append(std::move(primRecv));
     324              : 
     325            0 :         dimPrimQues[queIdx]->Append(std::move(primGroup));
     326              : 
     327              :         // semaphore sync
     328            0 :         if (dimPrimQues.size() > 1) {
     329            0 :             CHK_PRT_RET(PostSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
     330              :                         HCCL_ERROR("[CollAlgFactory] [TempAllGatherConcurrMesh] Rank [%d], Que [%u], Semaphore "
     331              :                                    "Synchronization Failed.",
     332              :                                    myRank_, dimPrimQues[queIdx]->GetId()),
     333              :                         HcclResult::HCCL_E_INTERNAL);
     334              :         }
     335            0 :     }
     336              : 
     337            0 :     return HcclResult::HCCL_SUCCESS;
     338              : }
     339              : 
     340            0 : std::unique_ptr<PrimSend> TempAllGatherConcurrMesh::RunSend(const RankSliceInfo    &sliceInfoVec,
     341              :                                                             const std::vector<u32> &sendChunkIdxs, const u32 &sliceIdx,
     342              :                                                             const RankId &neighborRank, const LinkData &priorLinkData)
     343              : {
     344            0 :     std::unique_ptr<PrimSend> primSend;
     345              :     u64                       tmpSendOff;
     346              :     u64                       tmpSendSize;
     347            0 :     for (u32 chunkIdx = 0; chunkIdx < sendChunkIdxs.size(); chunkIdx++) {
     348            0 :         if (chunkIdx == 0) {
     349              :             // first slice
     350            0 :             tmpSendOff  = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
     351            0 :             tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     352            0 :         } else if (tmpSendOff + tmpSendSize == sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset) {
     353              :             // consequent slice
     354            0 :             tmpSendSize += sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     355              :         } else {
     356              :             DataSlice sendLocSlice
     357            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     358              :             DataSlice sendRemSlice
     359            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     360            0 :             if (!primSend) {
     361            0 :                 primSend.reset(new PrimSend(neighborRank, priorLinkData, sendLocSlice, sendRemSlice, dmaMode_));
     362              :             } else {
     363            0 :                 primSend->Append(sendLocSlice, sendRemSlice);
     364              :             }
     365            0 :             tmpSendOff  = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
     366            0 :             tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     367              :         }
     368              : 
     369            0 :         if (chunkIdx == (sendChunkIdxs.size() - 1)) {
     370              :             DataSlice sendLocSlice
     371            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     372              :             DataSlice sendRemSlice
     373            0 :                 = DataSlice(buffInfo_.outBuffType, tmpSendOff + buffInfo_.outBuffBaseOff, tmpSendSize);
     374              : 
     375            0 :             if (!primSend) {
     376            0 :                 primSend.reset(new PrimSend(neighborRank, priorLinkData, sendLocSlice, sendRemSlice, dmaMode_));
     377              :             } else {
     378            0 :                 primSend->Append(sendLocSlice, sendRemSlice);
     379              :             }
     380              :         }
     381              :     }
     382              : 
     383            0 :     return primSend;
     384            0 : }
     385              : 
     386            0 : std::unique_ptr<PrimRecv> TempAllGatherConcurrMesh::RunRecv(const RankSliceInfo    &sliceInfoVec,
     387              :                                                             const std::vector<u32> &recvChunkIdxs, const u32 &sliceIdx,
     388              :                                                             const RankId &neighborRank, const LinkData &priorLinkData)
     389              : {
     390            0 :     std::unique_ptr<PrimRecv> primRecv;
     391              :     u64                       tmpRecvOff;
     392              :     u64                       tmpRecvSize;
     393            0 :     for (u32 chunkIdx = 0; chunkIdx < recvChunkIdxs.size(); chunkIdx++) {
     394            0 :         if (chunkIdx == 0) {
     395              :             // first slice
     396            0 :             tmpRecvOff  = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
     397            0 :             tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     398            0 :         } else if (tmpRecvOff + tmpRecvSize == sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset) {
     399              :             // consequent slice
     400            0 :             tmpRecvSize += sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     401              :         } else {
     402              :             DataSlice recvRemSlice
     403            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     404              :             DataSlice recvLocSlice
     405            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     406            0 :             if (!primRecv) {
     407            0 :                 primRecv.reset(new PrimRecv(neighborRank, priorLinkData, recvLocSlice, recvRemSlice, dmaMode_));
     408              :             } else {
     409            0 :                 primRecv->Append(recvLocSlice, recvRemSlice);
     410              :             }
     411            0 :             tmpRecvOff  = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
     412            0 :             tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     413              :         }
     414              : 
     415            0 :         if (chunkIdx == (recvChunkIdxs.size() - 1)) {
     416              :             DataSlice recvRemSlice
     417            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     418              :             DataSlice recvLocSlice
     419            0 :                 = DataSlice(buffInfo_.outBuffType, tmpRecvOff + buffInfo_.outBuffBaseOff, tmpRecvSize);
     420              : 
     421            0 :             if (!primRecv) {
     422            0 :                 primRecv.reset(new PrimRecv(neighborRank, priorLinkData, recvLocSlice, recvRemSlice, dmaMode_));
     423              :             } else {
     424            0 :                 primRecv->Append(recvLocSlice, recvRemSlice);
     425              :             }
     426              :         }
     427              :     }
     428              : 
     429            0 :     return primRecv;
     430            0 : }
     431              : 
     432              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1