LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/prim_alg_template - temp_reduce_scatter_concurr_mesh.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 246 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 14 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_reduce_scatter_concurr_mesh.h"
      14              : 
      15              : namespace Hccl {
      16            0 : TempReduceScatterConcurrMesh::TempReduceScatterConcurrMesh(
      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 : TempReduceScatterConcurrMesh::~TempReduceScatterConcurrMesh() {}
      23              : 
      24            0 : HcclResult TempReduceScatterConcurrMesh::CalcRes(
      25              :     const bool forAllReduce, AlgTempResReq& tempResReq, u32& requiredScratchMultiplier)
      26              : {
      27              :     (void)forAllReduce;
      28            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      29            0 :         tempResReq.queNum += tempVTopo_[dim].size() - 1;
      30              :     }
      31            0 :     requiredScratchMultiplier = tempRankSize_;
      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(
      41              :                 "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], Dim [%u], NeighborRank [%d].", myRank_, dim,
      42              :                 neighborRank);
      43              : 
      44              :             // LinkNum
      45            0 :             tempResReq.links[neighborRank] = 1;
      46              :         }
      47              :     }
      48              : 
      49            0 :     return HcclResult::HCCL_SUCCESS;
      50              : }
      51              : 
      52              : /*
      53              : dataSize / (rankSize) --> chunkSize
      54              : dataSize / (rankSize * dimNum) --> sliceSize
      55              : 
      56              : SliceInfoVecforConcurrMesh: [1st chunk: [1st Slice, 2nd Slice], 2nd chunk: [1st Slice, 2nd Slice], ...]
      57              : */
      58            0 : HcclResult TempReduceScatterConcurrMesh::CalcSliceInfo(
      59              :     const AllignInfo& allignInfo, const bool forAllReduce, const u64 dataSize, RankSliceInfo& sliceInfoVec)
      60              : {
      61            0 :     u32 dimSize = 0;
      62            0 :     for (u32 dimIdx = 0; dimIdx < tempVTopo_.size(); dimIdx++) {
      63            0 :         if (tempVTopo_[dimIdx].size() != 1) {
      64            0 :             dimSize += 1;
      65              :         }
      66              :     }
      67            0 :     std::vector<SliceInfo> tmp(dimSize);
      68            0 :     sliceInfoVec.resize(tempRankSize_, tmp);
      69              : 
      70            0 :     if (forAllReduce) {
      71              :         // for allreduce, dataSize = total dataSize
      72            0 :         CHK_RET(CalcSliceInfoAllReduce(allignInfo, dataSize, sliceInfoVec));
      73              :     } else {
      74              :         // for reduce scatter, dataSize = chunkSize
      75            0 :         if (sliceInfoVec[0].size() == 1) {
      76              :             // one-dimensional mesh
      77            0 :             CHK_RET(CalcRsAgSliceInfoMesh(myRank_, tempRankSize_, allignInfo, dataSize, sliceInfoVec));
      78              :         } else {
      79              :             // multi-dimensional mesh
      80            0 :             CHK_RET(CalcRsAgSliceInfoConcurrMesh(myRank_, tempVTopo_, allignInfo, dataSize, sliceInfoVec));
      81              :         }
      82              :     }
      83              : 
      84            0 :     return HcclResult::HCCL_SUCCESS;
      85            0 : }
      86              : 
      87            0 : HcclResult TempReduceScatterConcurrMesh::CalcSliceInfoAllReduce(
      88              :     const AllignInfo& allignInfo, const u64 dataSize, RankSliceInfo& sliceInfoVec)
      89              : {
      90              :     u64 unitAllignSize;
      91            0 :     CHK_RET(GetUnitAllignSize(allignInfo, unitAllignSize));
      92              : 
      93            0 :     u64 rankDataSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
      94              : 
      95            0 :     if (sliceInfoVec[0].size() == 1) {
      96              :         // one dimensional mesh
      97            0 :         u64 resDataSize = dataSize;
      98            0 :         for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
      99            0 :             u64 currChunkSize = (resDataSize > rankDataSize) ? rankDataSize : resDataSize;
     100            0 :             SliceInfo slice = {dataSize - resDataSize, currChunkSize};
     101            0 :             sliceInfoVec[rankIdx][0] = slice;
     102            0 :             resDataSize -= currChunkSize;
     103              :         }
     104              : 
     105            0 :         CHK_PRT_RET(
     106              :             (sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != dataSize),
     107              :             HCCL_ERROR(
     108              :                 "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], SliceInfo calculation error for "
     109              :                 "AllReduce ConcurrMesh!",
     110              :                 myRank_),
     111              :             HcclResult::HCCL_E_INTERNAL);
     112              :     } else {
     113            0 :         u32 dimSize0 = tempVTopo_[0].size();
     114            0 :         u32 dimSize1 = tempVTopo_[1].size();
     115              : 
     116            0 :         u64 resDataSize = dataSize;
     117            0 :         for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
     118            0 :             u64 currChunkSize = (resDataSize > rankDataSize) ? rankDataSize : resDataSize;
     119            0 :             u64 sliceSize0 = min(
     120              :                 currChunkSize,
     121            0 :                 RoundUp(currChunkSize, (dimSize0 + dimSize1) * unitAllignSize) * dimSize0 * unitAllignSize);
     122            0 :             SliceInfo slice0 = {dataSize - resDataSize, sliceSize0};
     123            0 :             sliceInfoVec[rankIdx][0] = slice0;
     124            0 :             resDataSize -= sliceSize0;
     125              : 
     126            0 :             u64 sliceSize1 = currChunkSize - sliceSize0;
     127            0 :             SliceInfo slice1 = {dataSize - resDataSize, sliceSize1};
     128            0 :             sliceInfoVec[rankIdx][1] = slice1;
     129            0 :             resDataSize -= sliceSize1;
     130              :         }
     131              : 
     132            0 :         CHK_PRT_RET(
     133              :             (sliceInfoVec[tempRankSize_ - 1][1].offset + sliceInfoVec[tempRankSize_ - 1][1].size != dataSize),
     134              :             HCCL_ERROR(
     135              :                 "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], SliceInfo calculation error "
     136              :                 "for AllReduce ConcurrMesh!",
     137              :                 myRank_),
     138              :             HcclResult::HCCL_E_INTERNAL);
     139              :     }
     140              : 
     141            0 :     return HcclResult::HCCL_SUCCESS;
     142              : }
     143              : 
     144            0 : HcclResult TempReduceScatterConcurrMesh::GenPrimQue(
     145              :     const TempFuncs& tempFuncs, const RankSliceInfo& sliceInfoVec, const BuffInfo& buffInfo, const ResLinks& tempLinks,
     146              :     std::vector<PrimQuePtr>& tempPrimQues)
     147              : {
     148            0 :     opMode_ = tempFuncs.opMode;
     149            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
     150            0 :     buffInfo_ = buffInfo;
     151            0 :     HCCL_INFO(
     152              :         "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], EnableCounterNotify [%d].", myRank_,
     153              :         enableCounterNotify_);
     154              : 
     155            0 :     queNum_ = 0;
     156            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     157            0 :         queNum_ += tempVTopo_[dim].size() - 1;
     158              :     }
     159            0 :     CHK_PRT_RET(
     160              :         queNum_ != tempPrimQues.size(),
     161              :         HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], requiredQue Error.", myRank_),
     162              :         HcclResult::HCCL_E_INTERNAL);
     163              : 
     164              :     // LocalCopy: from input to scratch In Buffer for OPBASE
     165            0 :     if ((opMode_ == OpMode::OPBASE) && tempFuncs.isForepart) {
     166            0 :         CHK_RET(PreCopyOpbase(tempFuncs.usrData, tempPrimQues));
     167              :     }
     168              : 
     169            0 :     if (sliceInfoVec[0].size() == 1) {
     170            0 :         CHK_RET(RunOneDimMesh(sliceInfoVec, tempLinks, tempPrimQues));
     171              :     } else {
     172            0 :         CHK_RET(RunConcurrMesh(sliceInfoVec, tempLinks, tempPrimQues));
     173              :     }
     174              : 
     175              :     // LocalCopy for standalone reducescatter in Offload Mode
     176            0 :     if ((opMode_ == OpMode::OFFLOAD) && !tempFuncs.forAllReduce && !tempFuncs.forAlgSeqComb) {
     177            0 :         CHK_RET(PostCopyOffload(sliceInfoVec, tempPrimQues));
     178              :     }
     179              : 
     180              :     // LocalCopy from scratch to output for Opbase
     181            0 :     if ((opMode_ == OpMode::OPBASE) && tempFuncs.isBottom && !tempFuncs.forAllReduce) {
     182            0 :         CHK_RET(PostCopyOpbase(tempFuncs.usrData, tempPrimQues));
     183              :     }
     184              : 
     185            0 :     return HcclResult::HCCL_SUCCESS;
     186              : }
     187              : 
     188            0 : HcclResult TempReduceScatterConcurrMesh::RunOneDimMesh(
     189              :     const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<PrimQuePtr>& tempPrimQues)
     190              : {
     191              :     // semaphore sync
     192            0 :     if (queNum_ > 1) {
     193            0 :         CHK_RET(PreSyncInterQueues(tempPrimQues));
     194              :     }
     195              : 
     196              :     // locate myRank in tempVTopo -> algRank
     197              :     u32 myAlgRank;
     198            0 :     u32 validDim = (tempVTopo_[0].size() == 1) ? 1 : 0;
     199            0 :     HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], valid Dim [%u].", myRank_, validDim);
     200            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[validDim], myAlgRank));
     201              : 
     202              :     // runMesh
     203            0 :     CHK_PRT_RET(
     204              :         RunMesh(myAlgRank, tempVTopo_[validDim], sliceInfoVec, tempLinks, tempPrimQues) != HcclResult::HCCL_SUCCESS,
     205              :         HCCL_ERROR(
     206              :             "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], unable to run the mesh algorithm.", myRank_),
     207              :         HcclResult::HCCL_E_INTERNAL);
     208              : 
     209              :     // semaphore sync
     210            0 :     if (queNum_ > 1) {
     211            0 :         CHK_RET(PostSyncInterQueues(tempPrimQues));
     212              :     }
     213              : 
     214            0 :     return HcclResult::HCCL_SUCCESS;
     215              : }
     216              : 
     217            0 : HcclResult TempReduceScatterConcurrMesh::RunMesh(
     218              :     const u32 myAlgRank, const std::vector<RankId>& vTopo, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
     219              :     std::vector<PrimQuePtr>& tempPrimQues)
     220              : {
     221            0 :     for (u32 queIdx = 0; queIdx < tempPrimQues.size(); queIdx++) {
     222              :         // find neighbors -> virtualRank
     223            0 :         RankId neighborRank = vTopo[(myAlgRank + 1 + queIdx) % tempRankSize_];
     224              :         // Link
     225            0 :         LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
     226              : 
     227            0 :         u32 sendChunkIdx = tempVirtRankMap_[neighborRank];
     228            0 :         u64 sendOffset = sliceInfoVec[sendChunkIdx][0].offset;
     229            0 :         u64 sendSize = sliceInfoVec[sendChunkIdx][0].size;
     230            0 :         u32 recvChunkIdx = tempVirtRankMap_[myRank_];
     231            0 :         u64 recvOffset = sliceInfoVec[recvChunkIdx][0].offset;
     232            0 :         u64 recvSize = sliceInfoVec[recvChunkIdx][0].size;
     233              : 
     234              :         // PrimGroup
     235            0 :         std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
     236              : 
     237              :         // SendReduce
     238            0 :         DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, sendOffset + buffInfo_.inBuffBaseOff, sendSize);
     239              :         DataSlice sendRemSrcSlice
     240            0 :             = DataSlice(buffInfo_.scratBuffType, sendOffset + buffInfo_.scratchBuffBaseOff, sendSize);
     241            0 :         DataSlice sendRemDstSlice = DataSlice(buffInfo_.inBuffType, sendOffset + buffInfo_.inBuffBaseOff, sendSize);
     242            0 :         std::unique_ptr<Primitive> primSendReduce = std::make_unique<PrimSendReduce>(
     243            0 :             neighborRank, neighborLinkData, sendLocSlice, sendRemSrcSlice, sendRemDstSlice, dataType_, redOp_,
     244            0 :             dmaMode_);
     245              : 
     246            0 :         primGroup->Append(std::move(primSendReduce));
     247              : 
     248              :         // RecvReduce
     249            0 :         DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, recvOffset + buffInfo_.inBuffBaseOff, recvSize);
     250              :         DataSlice recvLocSrcSlice
     251            0 :             = DataSlice(buffInfo_.scratBuffType, recvOffset + buffInfo_.scratchBuffBaseOff, recvSize);
     252            0 :         DataSlice recvLocDstSlice = DataSlice(buffInfo_.inBuffType, recvOffset + buffInfo_.inBuffBaseOff, recvSize);
     253            0 :         std::unique_ptr<Primitive> primRecvReduce = std::make_unique<PrimRecvReduce>(
     254            0 :             neighborRank, neighborLinkData, recvRemSlice, recvLocSrcSlice, recvLocDstSlice, dataType_, redOp_,
     255            0 :             dmaMode_);
     256              : 
     257            0 :         primGroup->Append(std::move(primRecvReduce));
     258              : 
     259            0 :         tempPrimQues[queIdx]->Append(std::move(primGroup));
     260            0 :     }
     261            0 :     return HcclResult::HCCL_SUCCESS;
     262              : }
     263              : 
     264            0 : HcclResult TempReduceScatterConcurrMesh::RunConcurrMesh(
     265              :     const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks, std::vector<PrimQuePtr>& tempPrimQues)
     266              : {
     267            0 :     std::vector<std::vector<PrimQuePtr>> dimQues;
     268            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     269              :         // assign queues
     270            0 :         std::vector<PrimQuePtr> tmpQue;
     271            0 :         for (u32 idx = 0; idx < tempVTopo_[dim].size() - 1; idx++) {
     272            0 :             if (dim == 0) {
     273            0 :                 tmpQue.push_back(tempPrimQues[idx]);
     274              :             } else {
     275            0 :                 tmpQue.push_back(tempPrimQues[tempVTopo_[0].size() - 1 + idx]);
     276              :             }
     277              :         }
     278            0 :         dimQues.push_back(tmpQue);
     279            0 :     }
     280              : 
     281            0 :     std::vector<PrimQuePtr> majorDimQue = {tempPrimQues[0], tempPrimQues[tempVTopo_[0].size() - 1]};
     282              : 
     283              :     // semaphore sync inter dimensions
     284            0 :     CHK_RET(PreSyncInterQueues(majorDimQue));
     285              : 
     286              :     // run concurrent Mesh Step 0
     287            0 :     u32 step = 0;
     288            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     289            0 :         CHK_RET(RunSingleDimension(step, dim, sliceInfoVec, tempLinks, dimQues[dim]));
     290              :     }
     291              : 
     292              :     // semaphore sync
     293            0 :     CHK_RET(PostSyncInterQueues(majorDimQue));
     294              : 
     295              :     // semaphore sync inter dimensions
     296            0 :     CHK_RET(PreSyncInterQueues(majorDimQue));
     297              : 
     298              :     // run concurrent Mesh Step 1
     299            0 :     step = 1;
     300            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
     301            0 :         CHK_RET(RunSingleDimension(step, dim, sliceInfoVec, tempLinks, dimQues[dim]));
     302              :     }
     303              : 
     304              :     // semaphore sync
     305            0 :     CHK_RET(PostSyncInterQueues(majorDimQue));
     306              : 
     307            0 :     return HcclResult::HCCL_SUCCESS;
     308            0 : }
     309              : 
     310            0 : HcclResult TempReduceScatterConcurrMesh::RunSingleDimension(
     311              :     const u32& step, const u32& dim, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
     312              :     std::vector<PrimQuePtr>& dimPrimQues)
     313              : {
     314            0 :     CHK_PRT_RET(
     315              :         dim > 1,
     316              :         HCCL_ERROR("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], invalid dim [%u].", myRank_, dim),
     317              :         HcclResult::HCCL_E_INTERNAL);
     318              : 
     319              :     // locate myRank in tempVTopo -> algRank
     320              :     u32 myAlgRank;
     321            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
     322              : 
     323            0 :     for (u32 queIdx = 0; queIdx < dimPrimQues.size(); queIdx++) {
     324              :         // semaphore sync
     325            0 :         if (dimPrimQues.size() > 1) {
     326            0 :             CHK_PRT_RET(
     327              :                 PreSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
     328              :                 HCCL_ERROR(
     329              :                     "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], Que [%u], Semaphore "
     330              :                     "Synchronization Failed.",
     331              :                     myRank_, dimPrimQues[queIdx]->GetId()),
     332              :                 HcclResult::HCCL_E_INTERNAL);
     333              :         }
     334              : 
     335              :         // find neighbors -> virtualRank
     336            0 :         u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
     337            0 :         RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
     338              : 
     339              :         // link
     340            0 :         LinkData neighborLinkData = tempLinks.at(neighborRank)[0];
     341            0 :         HCCL_INFO(
     342              :             "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], [%u]-th Que, queId [%u], neighborRank [%d].",
     343              :             myRank_, queIdx, dimPrimQues[queIdx]->GetId(), neighborRank);
     344              : 
     345              :         // PrimGroup
     346            0 :         std::unique_ptr<PrimGroup> primGroup = std::make_unique<PrimGroup>();
     347              : 
     348            0 :         std::vector<u32> sendChunkIdxs;
     349            0 :         std::vector<u32> recvChunkIdxs;
     350              : 
     351            0 :         if (step == 0) {
     352            0 :             for (u32 chunkIdx = 0; chunkIdx < tempVTopo_[1 - dim].size(); chunkIdx++) {
     353            0 :                 u32 sendChunkIdx = (dim == 0) ? (neighborAlgRank + chunkIdx * tempVTopo_[0].size()) :
     354            0 :                                                 (neighborAlgRank * tempVTopo_[0].size() + chunkIdx);
     355            0 :                 sendChunkIdxs.push_back(sendChunkIdx);
     356            0 :                 u32 recvChunkIdx = (dim == 0) ? (myAlgRank + chunkIdx * tempVTopo_[0].size()) :
     357            0 :                                                 (myAlgRank * tempVTopo_[0].size() + chunkIdx);
     358            0 :                 recvChunkIdxs.push_back(recvChunkIdx);
     359              :             }
     360              :         } else {
     361            0 :             sendChunkIdxs.push_back(tempVirtRankMap_[neighborRank]);
     362            0 :             recvChunkIdxs.push_back(tempVirtRankMap_[myRank_]);
     363              :         }
     364              : 
     365              :         // SendReduce
     366            0 :         u32 sliceIdx = (step == 0) ? dim : (1 - dim);
     367              :         std::unique_ptr<PrimSendReduce> primSendReduce
     368            0 :             = RunSendReduce(sliceInfoVec, sendChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
     369            0 :         primGroup->Append(std::move(primSendReduce));
     370              : 
     371              :         // RecvReduce
     372              :         std::unique_ptr<PrimRecvReduce> primRecvReduce
     373            0 :             = RunRecvReduce(sliceInfoVec, recvChunkIdxs, sliceIdx, neighborRank, neighborLinkData);
     374            0 :         primGroup->Append(std::move(primRecvReduce));
     375              : 
     376            0 :         dimPrimQues[queIdx]->Append(std::move(primGroup));
     377              : 
     378              :         // semaphore sync
     379            0 :         if (dimPrimQues.size() > 1) {
     380            0 :             CHK_PRT_RET(
     381              :                 PostSync(queIdx, dimPrimQues) != HcclResult::HCCL_SUCCESS,
     382              :                 HCCL_ERROR(
     383              :                     "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], Que [%u], Semaphore "
     384              :                     "Synchronization Failed.",
     385              :                     myRank_, dimPrimQues[queIdx]->GetId()),
     386              :                 HcclResult::HCCL_E_INTERNAL);
     387              :         }
     388            0 :     }
     389              : 
     390            0 :     return HcclResult::HCCL_SUCCESS;
     391              : }
     392              : 
     393            0 : std::unique_ptr<PrimSendReduce> TempReduceScatterConcurrMesh::RunSendReduce(
     394              :     const RankSliceInfo& sliceInfoVec, const std::vector<u32>& sendChunkIdxs, const u32& sliceIdx,
     395              :     const RankId& neighborRank, const LinkData& priorLinkData)
     396              : {
     397            0 :     std::unique_ptr<PrimSendReduce> primSendReduce;
     398              :     u64 tmpSendOff;
     399              :     u64 tmpSendSize;
     400            0 :     for (u32 chunkIdx = 0; chunkIdx < sendChunkIdxs.size(); chunkIdx++) {
     401            0 :         if (chunkIdx == 0) {
     402              :             // first slice
     403            0 :             tmpSendOff = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
     404            0 :             tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     405            0 :         } else if (tmpSendOff + tmpSendSize == sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset) {
     406              :             // consequent slice
     407            0 :             tmpSendSize += sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     408              :         } else {
     409            0 :             DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
     410              :             DataSlice sendRemSrcSlice
     411            0 :                 = DataSlice(buffInfo_.scratBuffType, tmpSendOff + buffInfo_.scratchBuffBaseOff, tmpSendSize);
     412              :             DataSlice sendRemDstSlice
     413            0 :                 = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
     414            0 :             if (!primSendReduce) {
     415            0 :                 primSendReduce.reset(new PrimSendReduce(
     416            0 :                     neighborRank, priorLinkData, sendLocSlice, sendRemSrcSlice, sendRemDstSlice, dataType_, redOp_,
     417            0 :                     dmaMode_));
     418              :             } else {
     419            0 :                 primSendReduce->Append(sendLocSlice, sendRemSrcSlice, sendRemDstSlice);
     420              :             }
     421            0 :             tmpSendOff = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].offset;
     422            0 :             tmpSendSize = sliceInfoVec[sendChunkIdxs[chunkIdx]][sliceIdx].size;
     423              :         }
     424              : 
     425            0 :         if (chunkIdx == (sendChunkIdxs.size() - 1)) {
     426            0 :             HCCL_INFO("[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk.", myRank_);
     427            0 :             DataSlice sendLocSlice = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
     428              :             DataSlice sendRemSrcSlice
     429            0 :                 = DataSlice(buffInfo_.scratBuffType, tmpSendOff + buffInfo_.scratchBuffBaseOff, tmpSendSize);
     430              :             DataSlice sendRemDstSlice
     431            0 :                 = DataSlice(buffInfo_.inBuffType, tmpSendOff + buffInfo_.inBuffBaseOff, tmpSendSize);
     432            0 :             if (!primSendReduce) {
     433            0 :                 HCCL_INFO(
     434              :                     "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is a non-consecutive chunk.",
     435              :                     myRank_);
     436            0 :                 primSendReduce.reset(new PrimSendReduce(
     437            0 :                     neighborRank, priorLinkData, sendLocSlice, sendRemSrcSlice, sendRemDstSlice, dataType_, redOp_,
     438            0 :                     dmaMode_));
     439              :             } else {
     440            0 :                 HCCL_INFO(
     441              :                     "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is consecutive.", myRank_);
     442            0 :                 primSendReduce->Append(sendLocSlice, sendRemSrcSlice, sendRemDstSlice);
     443              :             }
     444              :         }
     445              :     }
     446              : 
     447            0 :     return primSendReduce;
     448            0 : }
     449              : 
     450            0 : std::unique_ptr<PrimRecvReduce> TempReduceScatterConcurrMesh::RunRecvReduce(
     451              :     const RankSliceInfo& sliceInfoVec, const std::vector<u32>& recvChunkIdxs, const u32& sliceIdx,
     452              :     const RankId& neighborRank, const LinkData& priorLinkData)
     453              : {
     454            0 :     std::unique_ptr<PrimRecvReduce> primRecvReduce;
     455              :     u64 tmpRecvOff;
     456              :     u64 tmpRecvSize;
     457            0 :     for (u32 chunkIdx = 0; chunkIdx < recvChunkIdxs.size(); chunkIdx++) {
     458            0 :         if (chunkIdx == 0) {
     459              :             // first slice
     460            0 :             tmpRecvOff = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
     461            0 :             tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     462            0 :         } else if (tmpRecvOff + tmpRecvSize == sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset) {
     463              :             // consequent slice
     464            0 :             tmpRecvSize += sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     465              :         } else {
     466            0 :             DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
     467              :             DataSlice recvLocSrcSlice
     468            0 :                 = DataSlice(buffInfo_.scratBuffType, tmpRecvOff + buffInfo_.scratchBuffBaseOff, tmpRecvSize);
     469              :             DataSlice recvLocDstSlice
     470            0 :                 = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
     471              : 
     472            0 :             if (!primRecvReduce) {
     473            0 :                 HCCL_INFO(
     474              :                     "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is a non-consecutive chunk.",
     475              :                     myRank_);
     476            0 :                 primRecvReduce.reset(new PrimRecvReduce(
     477            0 :                     neighborRank, priorLinkData, recvRemSlice, recvLocSrcSlice, recvLocDstSlice, dataType_, redOp_,
     478            0 :                     dmaMode_));
     479              :             } else {
     480            0 :                 HCCL_INFO(
     481              :                     "[CollAlgFactory] [TempReduceScatterConcurrMesh] Rank [%d], last chunk is consecutive.", myRank_);
     482            0 :                 primRecvReduce->Append(recvRemSlice, recvLocSrcSlice, recvLocDstSlice);
     483              :             }
     484            0 :             tmpRecvOff = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].offset;
     485            0 :             tmpRecvSize = sliceInfoVec[recvChunkIdxs[chunkIdx]][sliceIdx].size;
     486              :         }
     487              : 
     488            0 :         if (chunkIdx == (recvChunkIdxs.size() - 1)) {
     489            0 :             DataSlice recvRemSlice = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
     490              :             DataSlice recvLocSrcSlice
     491            0 :                 = DataSlice(buffInfo_.scratBuffType, tmpRecvOff + buffInfo_.scratchBuffBaseOff, tmpRecvSize);
     492              :             DataSlice recvLocDstSlice
     493            0 :                 = DataSlice(buffInfo_.inBuffType, tmpRecvOff + buffInfo_.inBuffBaseOff, tmpRecvSize);
     494              : 
     495            0 :             if (!primRecvReduce) {
     496            0 :                 primRecvReduce.reset(new PrimRecvReduce(
     497            0 :                     neighborRank, priorLinkData, recvRemSlice, recvLocSrcSlice, recvLocDstSlice, dataType_, redOp_,
     498            0 :                     dmaMode_));
     499              :             } else {
     500            0 :                 primRecvReduce->Append(recvRemSlice, recvLocSrcSlice, recvLocDstSlice);
     501              :             }
     502              :         }
     503              :     }
     504            0 :     return primRecvReduce;
     505            0 : }
     506              : 
     507              : HcclResult
     508            0 : TempReduceScatterConcurrMesh::PostCopyOffload(const RankSliceInfo& sliceInfoVec, std::vector<PrimQuePtr>& tempPrimQues)
     509              : {
     510            0 :     u64 srcOffset = sliceInfoVec[tempVirtRankMap_[myRank_]][0].offset;
     511            0 :     u64 srcSize = 0;
     512            0 :     for (u32 dimIdx = 0; dimIdx < sliceInfoVec[0].size(); dimIdx++) {
     513            0 :         srcSize += sliceInfoVec[tempVirtRankMap_[myRank_]][dimIdx].size;
     514              :     }
     515            0 :     u64 dstOffset = 0;
     516            0 :     DataSlice srcSlice = DataSlice(buffInfo_.inBuffType, srcOffset + buffInfo_.inBuffBaseOff, srcSize);
     517            0 :     DataSlice dstSlice = DataSlice(buffInfo_.outBuffType, dstOffset + buffInfo_.outBuffBaseOff, srcSize);
     518            0 :     std::unique_ptr<Primitive> primLocalCopy = std::make_unique<PrimLocalCopy>(srcSlice, dstSlice);
     519            0 :     tempPrimQues[0]->Append(std::move(primLocalCopy));
     520              : 
     521            0 :     return HcclResult::HCCL_SUCCESS;
     522            0 : }
     523              : 
     524              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1