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

Generated by: LCOV version 2.0-1