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 % 313 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 13 0

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "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(
      18              :     const RankId virtualRank, const u32 tempRankSize, 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 : u64 InsTempReduceScatterMesh2D::CalcScratchMultiple(const BufferType& inBuffType, const BufferType& outBuffType)
      31              : {
      32              :     (void)inBuffType;
      33              :     (void)outBuffType;
      34            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
      35            0 :     u64 scratchMultiple = xyMaxRankSize * (xRankSize_ + yRankSize_);
      36            0 :     return scratchMultiple;
      37              : }
      38              : 
      39            0 : HcclResult InsTempReduceScatterMesh2D::CalcResLinksMesh2D(const u32 linkNumBtwPeers, AlgTempResReq& tempResReq)
      40              : {
      41              :     u32 myAlgRank;
      42            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      43            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
      44            0 :         for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
      45            0 :             RankId neighborRank = tempVTopo_[dim][(myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size())];
      46            0 :             tempResReq.links[neighborRank] = linkNumBtwPeers;
      47              :         }
      48              :     }
      49            0 :     return HcclResult::HCCL_SUCCESS;
      50              : }
      51              : 
      52            0 : HcclResult InsTempReduceScatterMesh2D::CalcRes(AlgTempResReq& tempResReq)
      53              : {
      54              :     // Mesh 需要的 que Num 为 tempVTopo_[0].size() + tempVTopo_[1].size() - 2
      55            0 :     tempResReq.queNum = (xRankSize_ > 1 && yRankSize_ > 1) ? (xQueNum_ + yQueNum_) : 1;
      56            0 :     tempResReq.streamNum = tempResReq.queNum;
      57            0 :     tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
      58            0 :     QId centerQ = 0;
      59            0 :     tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
      60            0 :     tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
      61              :     // linkNumBtwPeers_这个在没有绕路的情况下,是设置成1
      62            0 :     CHK_PRT_RET(
      63              :         CalcResLinksMesh2D(linkNumBtwPeers_, tempResReq) != HcclResult::HCCL_SUCCESS,
      64              :         HCCL_ERROR("[CollAlgFactory] [InsTempReduceScatterMesh2D] Rank [%d], resLinks calculation error!", myRank_),
      65              :         HcclResult::HCCL_E_INTERNAL);
      66              : 
      67            0 :     return HcclResult::HCCL_SUCCESS;
      68              : }
      69              : 
      70            0 : HcclResult InsTempReduceScatterMesh2D::GenExtIns(
      71              :     const TempFuncs& tempFuncs, TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
      72              :     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(
      84              :         queNum_ != tempInsQues.size(),
      85              :         HCCL_ERROR("[CollAlgFactory] [InsTempReduceScatterMesh2D] Rank [%d], requiredQue Error.", myRank_),
      86              :         HcclResult::HCCL_E_INTERNAL);
      87            0 :     PreCopy(tempAlgParams, tempInsQues); // stream 0作为主流,负责把本卡的数据拷贝到scratchbuffer上
      88            0 :     if (queNum_ > 1) {
      89            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
      90              :     }
      91            0 :     CHK_RET(RunFirstLevel(tempLinks, tempInsQues, tempAlgParams));
      92            0 :     if (queNum_ > 1) {
      93            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
      94            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
      95              :     }
      96            0 :     CHK_RET(RunFirstReduce(tempInsQues, tempAlgParams));
      97            0 :     if (queNum_ > 1) {
      98            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
      99            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
     100              :     }
     101            0 :     RunSecondLevel(tempLinks, tempInsQues, tempAlgParams);
     102            0 :     if (queNum_ > 1) {
     103            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
     104            0 :         CHK_RET(PreSyncInterQueues(tempInsQues));
     105              :     }
     106            0 :     RunSecondReduce(tempInsQues, tempAlgParams);
     107            0 :     if (queNum_ > 1) {
     108            0 :         CHK_RET(PostSyncInterQueues(tempInsQues));
     109              :     }
     110            0 :     return HcclResult::HCCL_SUCCESS;
     111              : }
     112              : 
     113              : HcclResult
     114            0 : InsTempReduceScatterMesh2D::PreCopy(const TemplateDataParams& tempAlgParams, std::vector<InsQuePtr>& tempInsQues)
     115              : {
     116            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     117            0 :     u64 remainDataSize = tempAlgParams.sliceSize - halfDataSize_;
     118              :     // 前一半数据,将本卡数据从input拷贝到scratchbuffer
     119            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     120            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     121            0 :         for (u32 yRankId = 0; yRankId < yRankSize_; yRankId++) {
     122            0 :             u32 rankId = yRankId * xRankSize_ + xRankId_; // 同y轴平面的所有卡,
     123              :             DataSlice inputRankSlice = DataSlice(
     124              :                 tempAlgParams.buffInfo.inBuffType,
     125            0 :                 tempAlgParams.buffInfo.inBuffBaseOff + rankId * tempAlgParams.inputSliceStride
     126            0 :                     + rpt * tempAlgParams.inputRepeatStride,
     127            0 :                 halfDataSize_);
     128              :             DataSlice scratchRankSlice = DataSlice(
     129              :                 tempAlgParams.buffInfo.scratBuffType,
     130            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff
     131            0 :                     + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + xRankId_) + scratchRepeatStride,
     132            0 :                 halfDataSize_);
     133            0 :             CHK_RET(LocalCopy(tempInsQues[0], inputRankSlice, scratchRankSlice));
     134            0 :             HCCL_DEBUG(
     135              :                 "[InsTempReduceScatterMesh2D][PreCopy] myRank[%d] top inputRankSlice: %s, scratchRankSlice: %s",
     136              :                 myRank_, inputRankSlice.Describe().c_str(), scratchRankSlice.Describe().c_str());
     137              :         }
     138              :     }
     139              :     // 后一半数据,将本卡数据从input拷贝到scratchbuffer
     140            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     141            0 :         u64 scratchRepeatStride
     142            0 :             = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
     143            0 :               + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     144            0 :         for (u32 xRankId = 0; xRankId < xRankSize_; xRankId++) {
     145            0 :             u32 rankId = yRankId_ * xRankSize_ + xRankId; // 同x轴平面的所有卡,
     146              :             DataSlice inputRankSlice = DataSlice(
     147              :                 tempAlgParams.buffInfo.inBuffType,
     148            0 :                 tempAlgParams.buffInfo.inBuffBaseOff + rankId * tempAlgParams.inputSliceStride + halfDataSize_
     149            0 :                     + rpt * tempAlgParams.inputRepeatStride,
     150            0 :                 remainDataSize);
     151              :             DataSlice scratchRankSlice = DataSlice(
     152              :                 tempAlgParams.buffInfo.scratBuffType,
     153            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff
     154            0 :                     + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + yRankId_) + scratchRepeatStride,
     155            0 :                 remainDataSize);
     156            0 :             CHK_RET(LocalCopy(tempInsQues[0], inputRankSlice, scratchRankSlice));
     157            0 :             HCCL_DEBUG(
     158              :                 "[InsTempReduceScatterMesh2D][PreCopy] myRank[%d] bottom inputRankSlice: %s, scratchRankSlice: %s",
     159              :                 myRank_, inputRankSlice.Describe().c_str(), scratchRankSlice.Describe().c_str());
     160              :         }
     161              :     }
     162            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][PreCopy], copy from userIn to scratch");
     163            0 :     return HcclResult::HCCL_SUCCESS;
     164              : }
     165              : 
     166            0 : HcclResult InsTempReduceScatterMesh2D::SendRecvProcess(
     167              :     const ResLinks& tempLinks, std::vector<std::vector<DataSlice>> allSliceVec, std::vector<InsQuePtr>& tempInsQues,
     168              :     u32 remoteRank, u32 queIdx) const
     169              : {
     170            0 :     CHK_PRT_RET(
     171              :         tempInsQues.empty(), HCCL_ERROR("[InsTempReduceScatterMesh2D][SendRecvProcess] empty queue"),
     172              :         HcclResult::HCCL_E_INTERNAL);
     173            0 :     CHK_PTR_NULL(tempInsQues[0]);
     174            0 :     HCCL_DEBUG("[InsTempReduceScatterMesh2D][SendRecvProcess] SendRecvProcess start");
     175            0 :     const std::vector<LinkData>& linkRecv = tempLinks.at(remoteRank);
     176            0 :     const std::vector<LinkData>& linkSend = tempLinks.at(remoteRank);
     177              :     SendRecvInfo sendRecvInfo{
     178            0 :         {linkSend[0], linkRecv[0]}, {{allSliceVec[2], allSliceVec[3]}, {allSliceVec[0], allSliceVec[1]}}};
     179              : 
     180            0 :     CHK_PRT_THROW(
     181              :         queIdx >= tempInsQues.size(),
     182              :         HCCL_ERROR(
     183              :             "[InsTempReduceScatterMesh2D] queIdx[%u] is bigger than tempInsQues size[%zu].", queIdx,
     184              :             tempInsQues.size()),
     185              :         InvalidParamsException, "queIdx is invalid");
     186              :     // 做了DMA消减之后只支持PUT
     187            0 :     CHK_PRT_RET(
     188              :         SendRecv(sendRecvInfo, tempInsQues[queIdx], 0, true, DmaMode::PUT),
     189              :         HCCL_ERROR("[InsTempReduceScatterMesh2D] RunReduceScatter SendReduce failed"), HcclResult::HCCL_E_INTERNAL);
     190            0 :     return HcclResult::HCCL_SUCCESS;
     191            0 : }
     192              : 
     193              : // 前一半数据的先x轴 和 后一半数据的先y轴
     194            0 : HcclResult InsTempReduceScatterMesh2D::RunFirstLevel(
     195              :     const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
     196              : {
     197            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][RunFirstLevel] myRank[%d]", myRank_);
     198            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     199              :     u64 processSize;
     200            0 :     for (u32 queIdx = 0; queIdx < queNum_; queIdx++) {
     201              :         u32 remoteRank;
     202              :         u32 index;
     203            0 :         std::vector<DataSlice> rxSrcSlices;
     204            0 :         std::vector<DataSlice> rxDstSlices;
     205            0 :         std::vector<DataSlice> txSrcSlices;
     206            0 :         std::vector<DataSlice> txDstSlices;
     207            0 :         if (queIdx < xQueNum_) { // 前xRankSize-1个stream,首先拉取前一半数据
     208            0 :             index = (xRankId_ + 1 + queIdx) % (tempVTopo_[0].size());
     209            0 :             remoteRank = tempVTopo_[0][index];
     210            0 :             processSize = halfDataSize_;
     211            0 :             HCCL_DEBUG(
     212              :                 "[InsTempReduceScatterMesh2D][RunFirstLevel] queID < xQueNum myRank[%d] toRank[%u] fromRank[%u] "
     213              :                 "rpt[%u], index[%u]",
     214              :                 myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum, index);
     215            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     216            0 :                 u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     217            0 :                 for (u32 yRankId = 0; yRankId < yRankSize_; yRankId++) {
     218            0 :                     u32 readRankId = yRankId * xRankSize_ + xRankId_;
     219            0 :                     u32 writeRankId = yRankId * xRankSize_ + index;
     220              :                     // 数据从其他卡,传输到本卡,接收数据
     221            0 :                     rxSrcSlices.emplace_back(
     222            0 :                         tempAlgParams.buffInfo.inBuffType,
     223            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + readRankId * tempAlgParams.inputSliceStride
     224            0 :                             + rpt * tempAlgParams.inputRepeatStride,
     225              :                         processSize);
     226            0 :                     rxDstSlices.emplace_back(
     227            0 :                         tempAlgParams.buffInfo.scratBuffType,
     228            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff
     229            0 :                             + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + index) + scratchRepeatStride,
     230              :                         processSize);
     231            0 :                     txSrcSlices.emplace_back(
     232            0 :                         tempAlgParams.buffInfo.inBuffType,
     233            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + writeRankId * tempAlgParams.inputSliceStride
     234            0 :                             + rpt * tempAlgParams.inputRepeatStride,
     235              :                         processSize);
     236            0 :                     txDstSlices.emplace_back(
     237            0 :                         tempAlgParams.buffInfo.scratBuffType,
     238            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff
     239            0 :                             + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId + xRankId_)
     240            0 :                             + scratchRepeatStride,
     241              :                         processSize);
     242            0 :                     HCCL_DEBUG(
     243              :                         "[InsTempReduceScatterMesh2D][RunFirstLevel] queID < xQueNum myRank[%d] *****sendrecv*****, "
     244              :                         "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
     245              :                         myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     246              :                         txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     247              :                 }
     248              :             }
     249              :         } else { // 后yRankSize-1个stream,首先拉取后一半数据
     250            0 :             index = (yRankId_ + 1 + queIdx - xQueNum_) % (tempVTopo_[1].size());
     251            0 :             remoteRank = tempVTopo_[1][index];
     252            0 :             processSize = tempAlgParams.sliceSize - halfDataSize_;
     253            0 :             HCCL_DEBUG(
     254              :                 "[InsTempReduceScatterMesh2D][RunFirstLevel] queId >= xQueNum myRank[%d] toRank[%u] fromRank[%u], "
     255              :                 "rpt[%u], index[%u]",
     256              :                 myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum, index);
     257            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     258            0 :                 u64 scratchRepeatStride
     259            0 :                     = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
     260            0 :                       + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     261            0 :                 for (u32 xRankId = 0; xRankId < xRankSize_; xRankId++) {
     262            0 :                     u32 readRankId = yRankId_ * xRankSize_ + xRankId; // 同x轴平面的所有卡,
     263            0 :                     u32 writeRankId = index * xRankSize_ + xRankId;
     264            0 :                     rxSrcSlices.emplace_back(
     265            0 :                         tempAlgParams.buffInfo.inBuffType,
     266            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + readRankId * tempAlgParams.inputSliceStride
     267            0 :                             + halfDataSize_ + rpt * tempAlgParams.inputRepeatStride,
     268              :                         processSize);
     269            0 :                     rxDstSlices.emplace_back(
     270            0 :                         tempAlgParams.buffInfo.scratBuffType,
     271            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff
     272            0 :                             + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + index) + scratchRepeatStride,
     273              :                         processSize);
     274            0 :                     txSrcSlices.emplace_back(
     275            0 :                         tempAlgParams.buffInfo.inBuffType,
     276            0 :                         tempAlgParams.buffInfo.inBuffBaseOff + writeRankId * tempAlgParams.inputSliceStride
     277            0 :                             + halfDataSize_ + rpt * tempAlgParams.inputRepeatStride,
     278              :                         processSize);
     279            0 :                     txDstSlices.emplace_back(
     280            0 :                         tempAlgParams.buffInfo.scratBuffType, // tempAlgParams.buffInfo.scratBuffType,
     281            0 :                         tempAlgParams.buffInfo.scratchBuffBaseOff
     282            0 :                             + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId + yRankId_)
     283            0 :                             + scratchRepeatStride,
     284              :                         processSize);
     285            0 :                     HCCL_DEBUG(
     286              :                         "[InsTempReduceScatterMesh2D][RunFirstLevel] queId >= xQueNum myRank[%d] *****sendrecv*****, "
     287              :                         "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
     288              :                         myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     289              :                         txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     290              :                 }
     291              :             }
     292              :         }
     293            0 :         if (processSize == 0) {
     294            0 :             continue;
     295              :         }
     296            0 :         std::vector<std::vector<DataSlice>> allSliceVec = {rxSrcSlices, rxDstSlices, txSrcSlices, txDstSlices};
     297            0 :         CHK_RET(SendRecvProcess(tempLinks, allSliceVec, tempInsQues, remoteRank, queIdx));
     298            0 :     }
     299            0 :     return HcclResult::HCCL_SUCCESS;
     300            0 : }
     301              : 
     302              : HcclResult
     303            0 : InsTempReduceScatterMesh2D::RunFirstReduce(std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
     304              : {
     305            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] rpt[%u]", myRank_, tempAlgParams.repeatNum);
     306            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     307            0 :     u64 processSize = 0;
     308              :     // 这里的stream 0和stream xRankSize-1分别负责前一半数据与后一半数据的本地reduce
     309            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     310            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     311            0 :         for (u32 tmpRank = 0; tmpRank < yRankSize_; tmpRank++) { // 前一半数据做local reduce,由这部分的第一个stream做
     312            0 :             processSize = halfDataSize_;
     313            0 :             for (u32 dataIdx = 1; dataIdx < xRankSize_;
     314              :                  dataIdx++) { // 原始这个位置已经有数据了,因此从后一片数据开始累加
     315              :                 DataSlice srcDataSlice = DataSlice(
     316              :                     tempAlgParams.buffInfo.scratBuffType,
     317            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     318            0 :                         + (xyMaxRankSize * tmpRank + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     319            0 :                     processSize);
     320              :                 DataSlice dstDataSlice = DataSlice(
     321              :                     tempAlgParams.buffInfo.scratBuffType,
     322            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     323            0 :                         + (xyMaxRankSize * tmpRank) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     324            0 :                     processSize);
     325            0 :                 HCCL_DEBUG(
     326              :                     "[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] queId < xQueNum *****LocalReduce*****, "
     327              :                     "srcDataSlice: %s, dstDataSlice: %s",
     328              :                     myRank_, srcDataSlice.Describe().c_str(), dstDataSlice.Describe().c_str());
     329            0 :                 CHK_RET(LocalReduce(tempInsQues[0], srcDataSlice, dstDataSlice, dataType_, redOp_));
     330              :             }
     331              :         }
     332              :     }
     333            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     334            0 :         u64 scratchRepeatStride
     335            0 :             = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
     336            0 :               + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     337            0 :         for (u32 tmpRank = 0; tmpRank < xRankSize_; tmpRank++) { // 后一半数据做local reduce,由这部分的第一个stream做
     338            0 :             processSize = tempAlgParams.sliceSize - halfDataSize_;
     339            0 :             for (u32 dataIdx = 1; dataIdx < yRankSize_;
     340              :                  dataIdx++) { // 原始这个位置已经有数据了,因此从后一片数据开始累加
     341              :                 DataSlice srcDataSlice = DataSlice(
     342              :                     tempAlgParams.buffInfo.scratBuffType,
     343            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     344            0 :                         + (xyMaxRankSize * tmpRank + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     345            0 :                     processSize);
     346              :                 DataSlice dstDataSlice = DataSlice(
     347              :                     tempAlgParams.buffInfo.scratBuffType,
     348            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     349            0 :                         + (xyMaxRankSize * tmpRank) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     350            0 :                     processSize);
     351            0 :                 HCCL_DEBUG(
     352              :                     "[InsTempReduceScatterMesh2D][RunFirstReduce] myRank[%d] queId >= xQueNum *****LocalReduce*****, "
     353              :                     "srcDataSlice: %s, dstDataSlice: %s",
     354              :                     myRank_, srcDataSlice.Describe().c_str(), dstDataSlice.Describe().c_str());
     355            0 :                 CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcDataSlice, dstDataSlice, dataType_, redOp_));
     356              :             }
     357              :         }
     358              :     }
     359            0 :     return HcclResult::HCCL_SUCCESS;
     360              : }
     361              : 
     362              : // 后一半数据的后x轴 和 前一半数据的后y轴
     363            0 : HcclResult InsTempReduceScatterMesh2D::RunSecondLevel(
     364              :     const ResLinks& tempLinks, std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
     365              : {
     366            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     367              :     u64 processSize;
     368            0 :     for (u32 queIdx = 0; queIdx < queNum_; queIdx++) {
     369              :         u32 remoteRank;
     370              :         u32 index;
     371            0 :         std::vector<DataSlice> rxSrcSlices;
     372            0 :         std::vector<DataSlice> rxDstSlices;
     373            0 :         std::vector<DataSlice> txSrcSlices;
     374            0 :         std::vector<DataSlice> txDstSlices;
     375            0 :         if (queIdx < xQueNum_) { // 前xRankSize-1个stream,后一半数据
     376            0 :             index = (xRankId_ + 1 + queIdx) % (tempVTopo_[0].size());
     377            0 :             remoteRank = tempVTopo_[0][index];
     378            0 :             HCCL_DEBUG(
     379              :                 "[InsTempReduceScatterMesh2D][RunSecondLevel] queIdx < xQueNum myRank[%d] toRank[%u] fromRank[%u]",
     380              :                 myRank_, remoteRank, remoteRank);
     381            0 :             processSize = tempAlgParams.sliceSize - halfDataSize_;
     382              :             // 这里过来的数据,直接按照queIdx的顺序放置,不一定是按照rankId顺序排列的
     383            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     384            0 :                 u64 scratchRepeatStride
     385            0 :                     = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
     386            0 :                       + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     387              :                 DataSlice rxSrcSlice = DataSlice(
     388              :                     tempAlgParams.buffInfo.scratBuffType,
     389            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     390            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId_) + scratchRepeatStride,
     391            0 :                     processSize);
     392              :                 DataSlice rxDstSlice = DataSlice(
     393              :                     tempAlgParams.buffInfo.scratBuffType,
     394            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     395            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankId_ + queIdx + 1)
     396              :                         + scratchRepeatStride,
     397            0 :                     processSize);
     398              :                 DataSlice txSrcSlice = DataSlice(
     399              :                     tempAlgParams.buffInfo.scratBuffType,
     400            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     401            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * index) + scratchRepeatStride,
     402            0 :                     processSize);
     403              :                 DataSlice txDstSlice = DataSlice(
     404              :                     tempAlgParams.buffInfo.scratBuffType,
     405            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     406            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * index + queIdx + 1) + scratchRepeatStride,
     407            0 :                     processSize);
     408              : 
     409            0 :                 rxSrcSlices.emplace_back(rxSrcSlice);
     410            0 :                 rxDstSlices.emplace_back(rxDstSlice);
     411            0 :                 txSrcSlices.emplace_back(txSrcSlice);
     412            0 :                 txDstSlices.emplace_back(txDstSlice);
     413              : 
     414            0 :                 HCCL_DEBUG(
     415              :                     "[InsTempReduceScatterMesh2D][RunSecondLevel] queId < xQueNum myRank[%d] *****sendrecv*****, "
     416              :                     "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
     417              :                     myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     418              :                     txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     419              :             }
     420              :         } else { // 后yRankSize-1个stream,前一半数据
     421            0 :             index = (yRankId_ + 1 + queIdx - xQueNum_) % (tempVTopo_[1].size());
     422            0 :             remoteRank = tempVTopo_[1][index];
     423            0 :             HCCL_DEBUG(
     424              :                 "[InsTempReduceScatterMesh2D][RunSecondLevel] queId >= xQueNum myRank[%d] toRank[%u] fromRank[%u] "
     425              :                 "rpt[%u]",
     426              :                 myRank_, remoteRank, remoteRank, tempAlgParams.repeatNum);
     427            0 :             processSize = halfDataSize_;
     428            0 :             for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     429            0 :                 u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     430              :                 DataSlice rxSrcSlice = DataSlice(
     431              :                     tempAlgParams.buffInfo.scratBuffType,
     432            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     433            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId_) + scratchRepeatStride,
     434            0 :                     processSize);
     435              :                 DataSlice rxDstSlice = DataSlice(
     436              :                     tempAlgParams.buffInfo.scratBuffType,
     437            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     438            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankId_ + queIdx - xQueNum_ + 1)
     439              :                         + scratchRepeatStride,
     440            0 :                     processSize);
     441              :                 DataSlice txSrcSlice = DataSlice(
     442              :                     tempAlgParams.buffInfo.scratBuffType,
     443            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     444            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * index) + scratchRepeatStride,
     445            0 :                     processSize);
     446              :                 DataSlice txDstSlice = DataSlice(
     447              :                     tempAlgParams.buffInfo.scratBuffType,
     448            0 :                     tempAlgParams.buffInfo.scratchBuffBaseOff
     449            0 :                         + tempAlgParams.outputSliceStride * (xyMaxRankSize * index + queIdx - xQueNum_ + 1)
     450              :                         + scratchRepeatStride,
     451            0 :                     processSize);
     452              : 
     453            0 :                 rxSrcSlices.emplace_back(rxSrcSlice);
     454            0 :                 rxDstSlices.emplace_back(rxDstSlice);
     455            0 :                 txSrcSlices.emplace_back(txSrcSlice);
     456            0 :                 txDstSlices.emplace_back(txDstSlice);
     457              : 
     458            0 :                 HCCL_DEBUG(
     459              :                     "[InsTempReduceScatterMesh2D][RunSecondLevel] queId >= xQueNum myRank[%d] *****sendrecv*****, "
     460              :                     "rxSrcSlice: %s, rxDstSlice: %s, txSrcSlice: %s, txDstSlice: %s",
     461              :                     myRank_, rxSrcSlices.back().Describe().c_str(), rxDstSlices.back().Describe().c_str(),
     462              :                     txSrcSlices.back().Describe().c_str(), txDstSlices.back().Describe().c_str());
     463              :             }
     464              :         }
     465            0 :         if (processSize == 0) {
     466            0 :             continue;
     467              :         }
     468            0 :         std::vector<std::vector<DataSlice>> allSliceVec = {rxSrcSlices, rxDstSlices, txSrcSlices, txDstSlices};
     469            0 :         CHK_RET(SendRecvProcess(tempLinks, allSliceVec, tempInsQues, remoteRank, queIdx));
     470            0 :     }
     471            0 :     return HcclResult::HCCL_SUCCESS;
     472            0 : }
     473              : 
     474            0 : HcclResult InsTempReduceScatterMesh2D::RunSecondReduce(
     475              :     std::vector<InsQuePtr>& tempInsQues, const TemplateDataParams& tempAlgParams)
     476              : {
     477            0 :     HCCL_INFO("[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] rpt[%u]", myRank_, tempAlgParams.repeatNum);
     478            0 :     u32 xyMaxRankSize = max(xRankSize_, yRankSize_);
     479            0 :     u64 processSize = 0;
     480              :     // 这里的stream 0和stream xRankSize-1分别负责后一半数据与前一半数据的本地reduce
     481              :     // 后一半数据做local reduce,由这部分的第一个stream做
     482            0 :     processSize = tempAlgParams.sliceSize - halfDataSize_;
     483            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     484            0 :         u64 scratchRepeatStride
     485            0 :             = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * tempAlgParams.repeatNum)
     486            0 :               + tempAlgParams.outputSliceStride * (xyMaxRankSize * xRankSize_ * rpt);
     487            0 :         for (u32 dataIdx = 0; dataIdx < xRankSize_; dataIdx++) {
     488              :             DataSlice srcSecDataSlice = DataSlice(
     489              :                 tempAlgParams.buffInfo.scratBuffType,
     490            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff
     491            0 :                     + (xyMaxRankSize * xRankId_ + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     492            0 :                 processSize);
     493            0 :             u64 outOffset
     494            0 :                 = tempAlgParams.buffInfo.outBuffBaseOff + halfDataSize_ + rpt * tempAlgParams.outputRepeatStride;
     495              :             DataSlice dstSecDataSlice = DataSlice(
     496              :                 tempAlgParams.buffInfo.outBuffType, // BufferType::OUTPUT,
     497            0 :                 outOffset, processSize);
     498            0 :             HCCL_DEBUG(
     499              :                 "[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] queId < xQueNum *****LocalReduce*****, "
     500              :                 "srcDataSlice: %s, dstDataSlice: %s",
     501              :                 myRank_, srcSecDataSlice.Describe().c_str(), dstSecDataSlice.Describe().c_str());
     502            0 :             if (srcSecDataSlice != dstSecDataSlice) {
     503            0 :                 if (dataIdx == 0) {
     504            0 :                     CHK_RET(LocalCopy(tempInsQues[0], srcSecDataSlice, dstSecDataSlice));
     505              :                 } else {
     506            0 :                     CHK_RET(LocalReduce(tempInsQues[0], srcSecDataSlice, dstSecDataSlice, dataType_, redOp_));
     507              :                 }
     508              :             }
     509              :         }
     510              :     }
     511              :     // 前一半数据做local reduce,由这部分的第一个stream做
     512            0 :     processSize = halfDataSize_;
     513            0 :     for (u32 rpt = 0; rpt < tempAlgParams.repeatNum; rpt++) {
     514            0 :         u64 scratchRepeatStride = tempAlgParams.outputSliceStride * (xyMaxRankSize * yRankSize_ * rpt);
     515            0 :         bool hasInplace = false;
     516            0 :         std::vector<DataSlice> srcFirDataSlices;
     517            0 :         std::vector<DataSlice> dstFirDataSlices;
     518            0 :         for (u32 dataIdx = 0; dataIdx < yRankSize_; dataIdx++) {
     519              :             DataSlice srcFirDataSlice = DataSlice(
     520              :                 tempAlgParams.buffInfo.scratBuffType,
     521            0 :                 tempAlgParams.buffInfo.scratchBuffBaseOff
     522            0 :                     + (xyMaxRankSize * yRankId_ + dataIdx) * tempAlgParams.outputSliceStride + scratchRepeatStride,
     523            0 :                 processSize);
     524            0 :             u64 outOffset = tempAlgParams.buffInfo.outBuffBaseOff + rpt * tempAlgParams.outputRepeatStride;
     525              :             DataSlice dstFirDataSlice = DataSlice(
     526              :                 tempAlgParams.buffInfo.outBuffType, // BufferType::OUTPUT,
     527            0 :                 outOffset, processSize);
     528            0 :             HCCL_DEBUG(
     529              :                 "[InsTempReduceScatterMesh2D][RunSecondReduce] myRank[%d] queId >= xQueNum *****LocalReduce*****, "
     530              :                 "srcDataSlice: %s, dstDataSlice: %s",
     531              :                 myRank_, srcFirDataSlice.Describe().c_str(), dstFirDataSlice.Describe().c_str());
     532            0 :             if (srcFirDataSlice != dstFirDataSlice) {
     533              : #if DATASLICE_ONE
     534            0 :                 srcFirDataSlices.push_back(srcFirDataSlice);
     535            0 :                 dstFirDataSlices.push_back(dstFirDataSlice);
     536              : #else
     537              :                 if (dataIdx == 0) {
     538              :                     CHK_RET(LocalCopy(tempInsQues[xQueNum_], srcFirDataSlice, dstFirDataSlice));
     539              :                 } else {
     540              :                     CHK_RET(LocalReduce(tempInsQues[xQueNum_], srcFirDataSlice, dstFirDataSlice, dataType_, redOp_));
     541              :                 }
     542              : #endif
     543              :             } else {
     544            0 :                 hasInplace = true;
     545              :             }
     546              :         }
     547            0 :         for (u32 dataIdx = 0; dataIdx < srcFirDataSlices.size(); dataIdx++) {
     548            0 :             if (!hasInplace && dataIdx == 0) {
     549            0 :                 CHK_RET(LocalCopy(tempInsQues[xQueNum_], srcFirDataSlices[dataIdx], dstFirDataSlices[dataIdx]));
     550            0 :             } else {
     551            0 :                 CHK_RET(LocalReduce(
     552              :                     tempInsQues[xQueNum_], srcFirDataSlices[dataIdx], dstFirDataSlices[dataIdx], dataType_, redOp_));
     553              :             }
     554              :         }
     555            0 :     }
     556            0 :     return HcclResult::HCCL_SUCCESS;
     557              : }
     558              : 
     559              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1