LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/ins_alg_template - ins_temp_reduce_scatter_mesh_2D.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 284 0
Test Date: 2026-07-28 12:11:00 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 "ins_temp_reduce_scatter_mesh_2D.h"
      12              : 
      13              : #include "log.h"
      14              : #include "alg_data_trans_wrapper.h"
      15              : 
      16              : namespace Hccl {
      17            0 : InsTempReduceScatterMesh2D::InsTempReduceScatterMesh2D(const RankId virtualRank, const u32 tempRankSize,
      18              :                                                        const std::vector<std::vector<RankId>> &tempVTopo,
      19            0 :                                                        const std::map<RankId, u32>            &tempVirtRankMap)
      20            0 :     : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      21              : {
      22            0 :     xQueNum_ = tempVTopo_[0].size() - 1; // x轴的卡数-1
      23            0 :     yQueNum_ = tempVTopo_[1].size() - 1; // y轴的卡数-1
      24            0 :     xRankSize_ = tempVTopo_[0].size(); // x轴的卡数
      25            0 :     yRankSize_ = tempVTopo_[1].size(); // y轴的卡数
      26            0 : }
      27              : 
      28            0 : InsTempReduceScatterMesh2D::~InsTempReduceScatterMesh2D()
      29              : {
      30            0 : }
      31              : 
      32            0 : u64 InsTempReduceScatterMesh2D::CalcScratchMultiple(const BufferType &inBuffType, const BufferType &outBuffType)
      33              : {
      34              :     (void)inBuffType;
      35              :     (void)outBuffType;
      36            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
      37            0 :     u64 scratchMultiple = xyMaxRankSize * (xRankSize_ + yRankSize_);
      38            0 :     return scratchMultiple;
      39              : }
      40              : 
      41            0 : HcclResult InsTempReduceScatterMesh2D::CalcResLinksMesh2D(const u32 linkNumBtwPeers, AlgTempResReq &tempResReq)
      42              : {
      43              :     u32 myAlgRank;
      44            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      45            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
      46            0 :         for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
      47            0 :             RankId neighborRank = tempVTopo_[dim][(myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size())];
      48            0 :             tempResReq.links[neighborRank] = linkNumBtwPeers;
      49              :         }
      50              :     }
      51            0 :     return HcclResult::HCCL_SUCCESS;
      52              : }
      53              : 
      54            0 : HcclResult InsTempReduceScatterMesh2D::CalcRes(AlgTempResReq &tempResReq)
      55              : {
      56              :     // Mesh 需要的 que Num 为 tempVTopo_[0].size() + tempVTopo_[1].size() - 2
      57            0 :     tempResReq.queNum = (xRankSize_ > 1 && yRankSize_ > 1) ? (xQueNum_ + yQueNum_): 1;
      58            0 :     tempResReq.streamNum = tempResReq.queNum;
      59            0 :     tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
      60            0 :     QId centerQ = 0;
      61            0 :     tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
      62            0 :     tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
      63              :     // linkNumBtwPeers_这个在没有绕路的情况下,是设置成1
      64            0 :     CHK_PRT_RET(CalcResLinksMesh2D(linkNumBtwPeers_, tempResReq) != HcclResult::HCCL_SUCCESS,
      65              :                 HCCL_ERROR("[CollAlgFactory] [InsTempReduceScatterMesh2D] Rank [%d], resLinks calculation error!", myRank_),
      66              :                 HcclResult::HCCL_E_INTERNAL);
      67              : 
      68            0 :     return HcclResult::HCCL_SUCCESS;
      69              : }
      70              : 
      71            0 : HcclResult InsTempReduceScatterMesh2D::GenExtIns(const TempFuncs &tempFuncs, TemplateDataParams &tempAlgParams,
      72              :                                                  const ResLinks &tempLinks, std::vector<InsQuePtr> &tempInsQues)
      73              : {
      74            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[0], xRankId_)); // 得到当前卡在x轴上的编号
      75            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[1], yRankId_)); // 得到当前卡在y轴上的编号
      76            0 :     opMode_              = tempFuncs.opMode;
      77            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
      78            0 :     queNum_ = xQueNum_ + yQueNum_;
      79            0 :     u64 sliceNum = tempAlgParams.sliceSize / DataTypeSizeGet(dataType_); // 先计算得到本次迭代处理的数据量
      80            0 :     halfDataSize_ = sliceNum / PARALLEL_SIZE * DataTypeSizeGet(dataType_); // 前一半数据的size
      81            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D] Run Start");
      82              :     // 这里不支持绕路的时候,应该就用原始的tempInsQues就行
      83            0 :     CHK_PRT_RET(queNum_ != tempInsQues.size(),
      84              :                 HCCL_ERROR("[CollAlgFactory] [InsTempReduceScatterMesh2D] Rank [%d], requiredQue Error.", myRank_),
      85              :                 HcclResult::HCCL_E_INTERNAL);
      86            0 :     PreCopy(tempAlgParams, tempInsQues); // stream 0作为主流,负责把本卡的数据拷贝到scratchbuffer上
      87            0 :     if (queNum_ > 1) {
      88            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
      89              :     }
      90            0 :     CHK_RET(RunFirstLevel(tempLinks, tempInsQues, tempAlgParams));
      91            0 :     if (queNum_ > 1) {
      92            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
      93            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
      94              :     }
      95            0 :     CHK_RET(RunFirstReduce(tempInsQues, tempAlgParams));
      96            0 :     if (queNum_ > 1) {
      97            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
      98            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
      99              :     }
     100            0 :     RunSecondLevel(tempLinks, tempInsQues, tempAlgParams);
     101            0 :     if (queNum_ > 1) {
     102            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
     103            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
     104              :     }
     105            0 :     RunSecondReduce(tempInsQues, tempAlgParams);
     106            0 :     if (queNum_ > 1) {
     107            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
     108              :     }
     109            0 :     return HcclResult::HCCL_SUCCESS;
     110              : }
     111              : 
     112            0 : HcclResult InsTempReduceScatterMesh2D::PreCopy(const TemplateDataParams &tempAlgParams, std::vector<InsQuePtr> &tempInsQues)
     113              : {
     114            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     115            0 :     u64 remainDataSize = tempAlgParams.sliceSize - halfDataSize_;
     116              :     // 前一半数据,将本卡数据从input拷贝到scratchbuffer
     117            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     118            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     119            0 :         for (u32 yRankId = 0; yRankId < yRankSize_; yRankId++) {
     120            0 :             u32 rankId = yRankId * xRankSize_ + xRankId_;  // 同y轴平面的所有卡,
     121              :             DataSlice inputRankSlice = DataSlice(tempAlgParams.buffInfo.inBuffType,
     122            0 :                 tempAlgParams.buffInfo.inBuffBaseOff + rankId * tempAlgParams.inputSliceStride +
     123            0 :                     rpt * tempAlgParams.inputRepeatStride, halfDataSize_);
     124              :             DataSlice scratchRankSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     125            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff +
     126            0 :                     tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + xRankId_) + scratchRepeatStride, halfDataSize_);
     127            0 :             CHK_RET(LocalCopy(tempInsQues[0], inputRankSlice, scratchRankSlice));
     128            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][PreCopy] myRank[%d] top inputRankSlice: %s, scratchRankSlice: %s",
     129              :                 myRank_, inputRankSlice.Describe().c_str(), scratchRankSlice.Describe().c_str());
     130              :         }
     131              :     }
     132              :     // 后一半数据,将本卡数据从input拷贝到scratchbuffer
     133            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     134            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum) +
     135            0 :             tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     136            0 :         for (u32 xRankId = 0; xRankId < xRankSize_; xRankId++) {
     137            0 :             u32 rankId = yRankId_ * xRankSize_ + xRankId;  // 同x轴平面的所有卡,
     138              :             DataSlice inputRankSlice = DataSlice(tempAlgParams.buffInfo.inBuffType,
     139            0 :                 tempAlgParams.buffInfo.inBuffBaseOff + rankId * tempAlgParams.inputSliceStride + halfDataSize_ +
     140            0 :                     rpt * tempAlgParams.inputRepeatStride, remainDataSize);
     141              :             DataSlice scratchRankSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     142            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff +
     143            0 :                     tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + yRankId_) + scratchRepeatStride,
     144            0 :                     remainDataSize);
     145            0 :             CHK_RET(LocalCopy(tempInsQues[0], inputRankSlice, scratchRankSlice));
     146            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][PreCopy] myRank[%d] bottom inputRankSlice: %s, scratchRankSlice: %s",
     147              :                 myRank_, inputRankSlice.Describe().c_str(), scratchRankSlice.Describe().c_str());
     148              :         }
     149              :     }
     150            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][PreCopy], copy from userIn to scratch");
     151            0 :     return HcclResult::HCCL_SUCCESS;
     152              : }
     153              : 
     154            0 : HcclResult InsTempReduceScatterMesh2D::SendRecvProcess(const ResLinks &tempLinks, std::vector<std::vector<DataSlice>> allSliceVec,
     155              :                                                        std::vector<InsQuePtr> &tempInsQues, u32 remoteRank, u32 queIdx) const
     156              : {
     157            0 :     CHK_PRT_RET(tempInsQues.empty(),
     158              :         HCCL_ERROR("[InsTempReduceScatterMesh2D][SendRecvProcess] empty queue"), HcclResult::HCCL_E_INTERNAL);
     159            0 :     CHK_PTR_NULL(tempInsQues[0]);
     160            0 :     HCCL_DEBUG("[InsTempReduceScatterMesh2D][SendRecvProcess] SendRecvProcess start");
     161            0 :     const std::vector<LinkData> &linkRecv = tempLinks.at(remoteRank);
     162            0 :     const std::vector<LinkData> &linkSend = tempLinks.at(remoteRank);
     163            0 :     SendRecvInfo sendRecvInfo{{linkSend[0], linkRecv[0]},
     164            0 :                                 {{allSliceVec[2], allSliceVec[3]}, {allSliceVec[0], allSliceVec[1]}}};
     165              : 
     166            0 :     CHK_PRT_THROW(queIdx >= tempInsQues.size(),
     167              :                     HCCL_ERROR("[InsTempReduceScatterMesh2D] queIdx[%u] is bigger than tempInsQues size[%zu].", queIdx,
     168              :                                 tempInsQues.size()),
     169              :                     InvalidParamsException, "queIdx is invalid");                                
     170              :     // 做了DMA消减之后只支持PUT
     171            0 :     CHK_PRT_RET(SendRecv(sendRecvInfo, tempInsQues[queIdx], 0, true, DmaMode::PUT),
     172              :                 HCCL_ERROR("[InsTempReduceScatterMesh2D] RunReduceScatter SendReduce failed"),
     173              :                 HcclResult::HCCL_E_INTERNAL);
     174            0 :     return HcclResult::HCCL_SUCCESS;
     175            0 : }
     176              : 
     177              : // 前一半数据的先x轴 和 后一半数据的先y轴
     178            0 : HcclResult InsTempReduceScatterMesh2D::RunFirstLevel(const ResLinks &tempLinks, std::vector<InsQuePtr> &tempInsQues,
     179              :                                                      const TemplateDataParams &tempAlgParams)
     180              : {
     181            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][RunFirstLevel] myRank[%d]", myRank_);
     182            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     183              :     u64 processSize;
     184            0 :     for (u32 queIdx = 0; queIdx < queNum_; queIdx++) {
     185              :         u32 remoteRank;
     186              :         u32 index;
     187            0 :         std::vector<DataSlice> rxSrcSlices;
     188            0 :         std::vector<DataSlice> rxDstSlices;
     189            0 :         std::vector<DataSlice> txSrcSlices;
     190            0 :         std::vector<DataSlice> txDstSlices;
     191            0 :         if (queIdx < xQueNum_) {  // 前xRankSize-1个stream,首先拉取前一半数据
     192            0 :             index = (xRankId_ + 1 + queIdx) % (tempVTopo_[0].size());
     193            0 :             remoteRank = tempVTopo_[0][index];
     194            0 :             processSize = halfDataSize_;
     195            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunFirstLevel] queID < xQueNum myRank[%d] toRank[%u] fromRank[%u] rpt[%u], index[%u]",
     196              :                 myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum, index);
     197            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     198            0 :                 u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     199            0 :                 for (u32 yRankId = 0; yRankId < yRankSize_; yRankId++) {
     200            0 :                     u32 readRankId = yRankId * xRankSize_ + xRankId_;
     201            0 :                     u32 writeRankId = yRankId * xRankSize_ + index;
     202              :                     // 数据从其他卡,传输到本卡,接收数据
     203            0 :                     rxSrcSlices.emplace_back(tempAlgParams.buffInfo.inBuffType,
     204            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + readRankId * tempAlgParams.inputSliceStride +
     205            0 :                             rpt * tempAlgParams.inputRepeatStride, processSize);
     206            0 :                     rxDstSlices.emplace_back(tempAlgParams.buffInfo.scratBuffType,
     207            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff +
     208            0 :                             tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + index) + scratchRepeatStride, processSize);
     209            0 :                     txSrcSlices.emplace_back(tempAlgParams.buffInfo.inBuffType,
     210            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + writeRankId * tempAlgParams.inputSliceStride +
     211            0 :                             rpt * tempAlgParams.inputRepeatStride, processSize);
     212            0 :                     txDstSlices.emplace_back(tempAlgParams.buffInfo.scratBuffType,
     213            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff +
     214            0 :                             tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + xRankId_) + scratchRepeatStride, processSize);
     215            0 :                     HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunFirstLevel] queID < xQueNum myRank[%d] *****sendrecv*****, "
     216              :                         "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s", myRank_,
     217              :                         rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     218              :                         txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     219              :                 }
     220              :             }
     221              :         } else {  // 后yRankSize-1个stream,首先拉取后一半数据
     222            0 :             index = (yRankId_ + 1 + queIdx - xQueNum_) % (tempVTopo_[1].size());
     223            0 :             remoteRank = tempVTopo_[1][index];
     224            0 :             processSize = tempAlgParams.sliceSize - halfDataSize_;
     225            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunFirstLevel] queId >= xQueNum myRank[%d] toRank[%u] fromRank[%u], rpt[%u], index[%u]",
     226              :                 myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum, index);
     227            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     228            0 :                 u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum) +
     229            0 :                     tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     230            0 :                 for (u32 xRankId = 0; xRankId < xRankSize_; xRankId++) {
     231            0 :                     u32 readRankId = yRankId_ * xRankSize_ + xRankId;  // 同x轴平面的所有卡,
     232            0 :                     u32 writeRankId = index * xRankSize_ + xRankId;
     233            0 :                     rxSrcSlices.emplace_back(tempAlgParams.buffInfo.inBuffType,
     234            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + readRankId * tempAlgParams.inputSliceStride +
     235            0 :                             halfDataSize_ + rpt * tempAlgParams.inputRepeatStride, processSize);
     236            0 :                     rxDstSlices.emplace_back(tempAlgParams.buffInfo.scratBuffType,
     237            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff +
     238            0 :                             tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + index) + scratchRepeatStride,
     239              :                             processSize);
     240            0 :                     txSrcSlices.emplace_back(tempAlgParams.buffInfo.inBuffType,
     241            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + writeRankId * tempAlgParams.inputSliceStride +
     242            0 :                             halfDataSize_ + rpt * tempAlgParams.inputRepeatStride, processSize);
     243            0 :                     txDstSlices.emplace_back(tempAlgParams.buffInfo.scratBuffType,//tempAlgParams.buffInfo.scratBuffType,
     244            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff +
     245            0 :                             tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + yRankId_) + scratchRepeatStride,
     246              :                             processSize);
     247            0 :                     HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunFirstLevel] queId >= xQueNum myRank[%d] *****sendrecv*****, "
     248              :                         "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s", myRank_,
     249              :                         rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     250              :                         txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     251              :                 }
     252              :             }
     253              :         }
     254            0 :         if (processSize == 0) {
     255            0 :             continue;
     256              :         }
     257            0 :         std::vector<std::vector<DataSlice>> allSliceVec = {rxSrcSlices, rxDstSlices, txSrcSlices, txDstSlices};
     258            0 :         CHK_RET(SendRecvProcess(tempLinks, allSliceVec, tempInsQues, remoteRank, queIdx));
     259            0 :     }
     260            0 :     return HcclResult::HCCL_SUCCESS;
     261            0 : }
     262              : 
     263            0 : HcclResult InsTempReduceScatterMesh2D::RunFirstReduce(std::vector<InsQuePtr> &tempInsQues, const TemplateDataParams &tempAlgParams)
     264              : {
     265            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] rpt[%u]", myRank_, tempAlgParams.repeatNum);
     266            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     267            0 :     u64 processSize = 0;
     268              :     // 这里的stream 0和stream xRankSize-1分别负责前一半数据与后一半数据的本地reduce
     269            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     270            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     271            0 :         for (u32 tmpRank = 0; tmpRank < yRankSize_; tmpRank++) {  // 前一半数据做local reduce,由这部分的第一个stream做
     272            0 :             processSize = halfDataSize_;
     273            0 :             for (u32 dataIdx = 1; dataIdx < xRankSize_; dataIdx++) {  // 原始这个位置已经有数据了,因此从后一片数据开始累加
     274              :                 DataSlice srcDataSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     275            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     276            0 :                         (xyMaxRankSize * tmpRank + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride, processSize);
     277              :                 DataSlice dstDataSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     278            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     279            0 :                         (xyMaxRankSize * tmpRank) * tempAlgParams.outputSliceStride + scratchRepeatStride, processSize);
     280            0 :                 HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] queId < xQueNum *****LocalReduce*****, "
     281              :                     "srcDataSlice: %s, dstDataSlice: %s", myRank_, srcDataSlice.Describe().c_str(),
     282              :                     dstDataSlice.Describe().c_str());
     283            0 :                 CHK_RET(LocalReduce(tempInsQues[0], srcDataSlice, dstDataSlice, dataType_, redOp_));
     284              :             }
     285              :         }
     286              :     }
     287            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     288            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum) +
     289            0 :             tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     290            0 :         for (u32 tmpRank = 0; tmpRank < xRankSize_; tmpRank++) {  // 后一半数据做local reduce,由这部分的第一个stream做
     291            0 :             processSize = tempAlgParams.sliceSize - halfDataSize_;
     292            0 :             for (u32 dataIdx = 1; dataIdx < yRankSize_; dataIdx++) {  // 原始这个位置已经有数据了,因此从后一片数据开始累加
     293              :                 DataSlice srcDataSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     294            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     295            0 :                         (xyMaxRankSize * tmpRank + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     296            0 :                     processSize);
     297              :                 DataSlice dstDataSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     298            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     299            0 :                         (xyMaxRankSize * tmpRank) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     300            0 :                     processSize);
     301            0 :                 HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] queId >= xQueNum *****LocalReduce*****, "
     302              :                     "srcDataSlice: %s, dstDataSlice: %s", myRank_, srcDataSlice.Describe().c_str(),
     303              :                     dstDataSlice.Describe().c_str());
     304            0 :                 CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcDataSlice, dstDataSlice, dataType_, redOp_));
     305              :             }
     306              :         }
     307              :     }
     308            0 :     return HcclResult::HCCL_SUCCESS;
     309              : }
     310              : 
     311              : // 后一半数据的后x轴 和 前一半数据的后y轴
     312            0 : HcclResult InsTempReduceScatterMesh2D::RunSecondLevel(const ResLinks &tempLinks, std::vector<InsQuePtr> &tempInsQues,
     313              :                                                       const TemplateDataParams &tempAlgParams)
     314              : {
     315            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     316              :     u64 processSize;
     317            0 :     for (u32 queIdx = 0; queIdx < queNum_; queIdx++) {
     318              :         u32 remoteRank;
     319              :         u32 index;
     320            0 :         std::vector<DataSlice> rxSrcSlices;
     321            0 :         std::vector<DataSlice> rxDstSlices;
     322            0 :         std::vector<DataSlice> txSrcSlices;
     323            0 :         std::vector<DataSlice> txDstSlices;
     324            0 :         if (queIdx < xQueNum_) { // 前xRankSize-1个stream,后一半数据
     325            0 :             index = (xRankId_ + 1 + queIdx) % (tempVTopo_[0].size());
     326            0 :             remoteRank = tempVTopo_[0][index];
     327            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunSecondLevel] queIdx < xQueNum myRank[%d] toRank[%u] fromRank[%u]",
     328              :                 myRank_, remoteRank, remoteRank);
     329            0 :             processSize = tempAlgParams.sliceSize - halfDataSize_;
     330              :             // 这里过来的数据,直接按照queIdx的顺序放置,不一定是按照rankId顺序排列的
     331            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     332            0 :                 u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum) +
     333            0 :                     tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     334              :                 DataSlice rxSrcSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     335            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     336            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId_) + scratchRepeatStride, processSize);
     337              :                 DataSlice rxDstSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     338            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     339            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId_ + queIdx + 1) + scratchRepeatStride, processSize);
     340              :                 DataSlice txSrcSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     341            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     342            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * index) + scratchRepeatStride, processSize);
     343              :                 DataSlice txDstSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     344            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     345            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * index + queIdx + 1) + scratchRepeatStride, processSize);
     346              : 
     347            0 :                 rxSrcSlices.emplace_back(rxSrcSlice);
     348            0 :                 rxDstSlices.emplace_back(rxDstSlice);
     349            0 :                 txSrcSlices.emplace_back(txSrcSlice);
     350            0 :                 txDstSlices.emplace_back(txDstSlice);
     351              : 
     352            0 :                 HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunSecondLevel] queId < xQueNum myRank[%d] *****sendrecv*****, "
     353              :                     "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s", myRank_,
     354              :                     rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     355              :                     txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     356              :             }
     357              :         } else { // 后yRankSize-1个stream,前一半数据
     358            0 :             index = (yRankId_ + 1 + queIdx - xQueNum_) % (tempVTopo_[1].size());
     359            0 :             remoteRank = tempVTopo_[1][index];
     360            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunSecondLevel] queId >= xQueNum myRank[%d] toRank[%u] fromRank[%u] rpt[%u]",
     361              :                 myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum);
     362            0 :             processSize = halfDataSize_;
     363            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     364            0 :                 u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     365              :                 DataSlice rxSrcSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     366            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     367            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId_) + scratchRepeatStride, processSize);
     368              :                 DataSlice rxDstSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     369            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     370            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId_ + queIdx - xQueNum_ + 1) + scratchRepeatStride, processSize);
     371              :                 DataSlice txSrcSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     372            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     373            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * index) + scratchRepeatStride, processSize);
     374              :                 DataSlice txDstSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     375            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff +
     376            0 :                         tempAlgParams.outputSliceStride * (xyMaxRankSize * index + queIdx - xQueNum_ + 1) + scratchRepeatStride, processSize);
     377              : 
     378            0 :                 rxSrcSlices.emplace_back(rxSrcSlice);
     379            0 :                 rxDstSlices.emplace_back(rxDstSlice);
     380            0 :                 txSrcSlices.emplace_back(txSrcSlice);
     381            0 :                 txDstSlices.emplace_back(txDstSlice);
     382              : 
     383            0 :                 HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunSecondLevel] queId >= xQueNum myRank[%d] *****sendrecv*****, "
     384              :                     "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s", myRank_,
     385              :                     rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     386              :                     txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     387              :             }
     388              :         }
     389            0 :         if (processSize == 0) {
     390            0 :             continue;
     391              :         }
     392            0 :         std::vector<std::vector<DataSlice>> allSliceVec = {rxSrcSlices, rxDstSlices, txSrcSlices, txDstSlices};
     393            0 :         CHK_RET(SendRecvProcess(tempLinks, allSliceVec, tempInsQues, remoteRank, queIdx));
     394            0 :     }
     395            0 :     return HcclResult::HCCL_SUCCESS;
     396            0 : }
     397              : 
     398            0 : HcclResult InsTempReduceScatterMesh2D::RunSecondReduce(std::vector<InsQuePtr> &tempInsQues, const TemplateDataParams &tempAlgParams)
     399              : {
     400            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] rpt[%u]", myRank_, tempAlgParams.repeatNum);
     401            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     402            0 :     u64 processSize = 0;
     403              :     // 这里的stream 0和stream xRankSize-1分别负责后一半数据与前一半数据的本地reduce
     404              :     // 后一半数据做local reduce,由这部分的第一个stream做
     405            0 :     processSize = tempAlgParams.sliceSize - halfDataSize_;
     406            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     407            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum) +
     408            0 :             tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     409            0 :         for (u32 dataIdx = 0; dataIdx < xRankSize_; dataIdx++) {
     410              :             DataSlice srcSecDataSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     411            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff +
     412            0 :                     (xyMaxRankSize * xRankId_ + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride, processSize);
     413            0 :             u64 outOffset = tempAlgParams.buffInfo.outBuffBaseOff + halfDataSize_ + rpt * tempAlgParams.outputRepeatStride;
     414              :             DataSlice dstSecDataSlice = DataSlice(tempAlgParams.buffInfo.outBuffType,   // BufferType::OUTPUT,
     415            0 :                 outOffset, processSize);
     416            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] queId < xQueNum *****LocalReduce*****, "
     417              :                 "srcDataSlice: %s, dstDataSlice: %s", myRank_, srcSecDataSlice.Describe().c_str(),
     418              :                 dstSecDataSlice.Describe().c_str());
     419            0 :             if (srcSecDataSlice != dstSecDataSlice) {
     420            0 :                 if (dataIdx == 0) {
     421            0 :                     CHK_RET(LocalCopy(tempInsQues[0], srcSecDataSlice, dstSecDataSlice));
     422              :                 } else {
     423            0 :                     CHK_RET(LocalReduce(tempInsQues[0], srcSecDataSlice, dstSecDataSlice, dataType_, redOp_));
     424              :                 }
     425              :             }
     426              :         }
     427              :     }
     428              :     // 前一半数据做local reduce,由这部分的第一个stream做
     429            0 :     processSize = halfDataSize_;
     430            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     431            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     432            0 :         bool hasInplace = false;
     433            0 :         std::vector<DataSlice> srcFirDataSlices;
     434            0 :         std::vector<DataSlice> dstFirDataSlices;
     435            0 :         for (u32 dataIdx = 0; dataIdx < yRankSize_; dataIdx++) {
     436              :             DataSlice srcFirDataSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType,
     437            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff +
     438            0 :                     (xyMaxRankSize * yRankId_ + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride, processSize);
     439            0 :             u64 outOffset = tempAlgParams.buffInfo.outBuffBaseOff + rpt * tempAlgParams.outputRepeatStride;
     440              :             DataSlice dstFirDataSlice = DataSlice(tempAlgParams.buffInfo.outBuffType,   // BufferType::OUTPUT,
     441            0 :                 outOffset, processSize);
     442            0 :             HCCL_DEBUG("[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] queId >= xQueNum *****LocalReduce*****, "
     443              :                 "srcDataSlice: %s, dstDataSlice: %s", myRank_, srcFirDataSlice.Describe().c_str(),
     444              :                 dstFirDataSlice.Describe().c_str());
     445            0 :             if (srcFirDataSlice != dstFirDataSlice) {
     446              : #if DATASLICE_ONE
     447            0 :                 srcFirDataSlices.push_back(srcFirDataSlice);
     448            0 :                 dstFirDataSlices.push_back(dstFirDataSlice);
     449              : #else
     450              :                 if (dataIdx == 0) {
     451              :                     CHK_RET(LocalCopy(tempInsQues[xQueNum_], srcFirDataSlice, dstFirDataSlice));
     452              :                 } else {
     453              :                     CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcFirDataSlice, dstFirDataSlice, dataType_, redOp_));
     454              :                 }
     455              : #endif
     456              :             } else {
     457            0 :                 hasInplace = true;
     458              :             }
     459              :         }
     460            0 :         for (u32 dataIdx = 0; dataIdx < srcFirDataSlices.size(); dataIdx++) {
     461            0 :             if (!hasInplace && dataIdx == 0) {
     462            0 :                 CHK_RET(LocalCopy(tempInsQues[xQueNum_], srcFirDataSlices[dataIdx], dstFirDataSlices[dataIdx]));
     463            0 :             } else {
     464            0 :                 CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcFirDataSlices[dataIdx], dstFirDataSlices[dataIdx], dataType_, redOp_));
     465              :             }
     466              :         }
     467            0 :     }
     468            0 :     return HcclResult::HCCL_SUCCESS;
     469              : }
     470              : 
     471              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1