LCOV - code coverage report
Current view: top level - legacy/ascend950/service/collective/alg/coll_alg_factory/alg_template/ins_alg_template - ins_temp_all_reduce_mesh_2D_two_shot.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 181 0
Test Date: 2026-08-18 17:47:01 Functions: 0.0 % 11 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              : #include "alg_data_trans_wrapper.h"
      13              : #include "ins_temp_all_reduce_mesh_2D_two_shot.h"
      14              : 
      15              : namespace Hccl {
      16            0 : InsTempAllReduceMesh2DTwoShot::InsTempAllReduceMesh2DTwoShot(
      17              :     const RankId virtualRank, const u32 tempRankSize, const std::vector<std::vector<RankId>>& tempVTopo,
      18            0 :     const std::map<RankId, u32>& tempVirtRankMap)
      19            0 :     : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      20              : {
      21            0 :     HCCL_INFO("[InsTempAllReduceMesh2DTwoShot] Init.");
      22            0 : }
      23              : 
      24            0 : InsTempAllReduceMesh2DTwoShot::~InsTempAllReduceMesh2DTwoShot() { HCCL_INFO("[InsTempAllReduceMesh2DTwoShot] exit."); }
      25              : 
      26              : /*
      27              :  * Desc: 计算资源需求
      28              :  * return: tempResReq: 资源计算结果存储,包括notify信息,links信息等
      29              :  * return: HcclResult
      30              :  */
      31            0 : HcclResult InsTempAllReduceMesh2DTwoShot::CalcRes(AlgTempResReq& tempResReq)
      32              : {
      33              :     // 1D Mesh 需要的 que Num 为 ranksize
      34            0 :     tempResReq.queNum = tempVTopo_[0].size() + tempVTopo_[1].size();
      35            0 :     tempResReq.streamNum = tempResReq.queNum;
      36            0 :     tempResReq.queNotifys = CreateQueNotifiesRequest(tempResReq.queNum, 1, 0, tempVTopo_[0].size());
      37              : 
      38            0 :     QId centerQ = 0;
      39            0 :     tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
      40            0 :     tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
      41              : 
      42              :     uint32_t myAlgRank;
      43            0 :     for (u32 dim = 0; dim < tempVTopo_.size(); dim++) {
      44            0 :         CHK_RET(GetAlgRank(myRank_, tempVTopo_[dim], myAlgRank));
      45            0 :         for (u32 queIdx = 0; queIdx < tempVTopo_[dim].size() - 1; queIdx++) {
      46            0 :             u32 neighborAlgRank = (myAlgRank + 1 + queIdx) % (tempVTopo_[dim].size());
      47            0 :             RankId neighborRank = tempVTopo_[dim][neighborAlgRank];
      48            0 :             HCCL_INFO(
      49              :                 "InsTempAllReduceMesh2DTwoShot::CalcRes Rank[%d], Dim[%u], NeighborRank[%d].", myRank_, dim,
      50              :                 neighborRank);
      51              :             // LinkNum
      52            0 :             tempResReq.links[neighborRank] = 1;
      53              :         }
      54              :     }
      55            0 :     HCCL_INFO("InsTempAllReduceMesh2DTwoShot::CalcRes done");
      56            0 :     return HcclResult::HCCL_SUCCESS;
      57              : }
      58              : 
      59              : std::vector<std::tuple<QId, QId, u32>>
      60            0 : InsTempAllReduceMesh2DTwoShot::CreateQueNotifiesRequest(u32 queueNum, u32 pairNum, QId masterIdX, QId masterIdY) const
      61              : {
      62            0 :     std::vector<std::tuple<QId, QId, u32>> notifyRequests;
      63            0 :     HCCL_DEBUG("[Create][MasterSlaveQueNotifiesRequest] queueNum[%u], pairNum[%u]", queueNum, pairNum);
      64            0 :     if (queueNum == 0 || pairNum == 0) {
      65            0 :         HCCL_INFO("[Create][MasterSlaveQueNotifiesRequest] queueNum or pairNum is zero, "
      66              :                   "return empty notifyRequests");
      67            0 :         return notifyRequests;
      68              :     };
      69              : 
      70            0 :     u32 slaveNum = queueNum - 1;
      71            0 :     HCCL_INFO("[Create][MasterSlaveQueNotifiesRequest] slavNum[%u]", slaveNum);
      72            0 :     if (slaveNum < 1 || pairNum < 1) {
      73            0 :         return notifyRequests;
      74              :     }
      75              : 
      76            0 :     notifyRequests.reserve((slaveNum + queueNum - masterIdY - 1) * pairNum);
      77              :     // masterX(master0)跟所有的stream有同步关系
      78            0 :     for (QId q = 0; q < queueNum; q++) {
      79            0 :         if (q == masterIdX) {
      80            0 :             continue;
      81              :         }
      82            0 :         for (u32 i = 0; i < pairNum; i++) {
      83            0 :             notifyRequests.emplace_back(std::make_tuple(masterIdX, q, i));
      84            0 :             notifyRequests.emplace_back(std::make_tuple(q, masterIdX, i));
      85              :         }
      86              :     }
      87              : 
      88            0 :     for (QId q = masterIdY + 1; q < queueNum; q++) {
      89            0 :         for (u32 i = 0; i < pairNum; i++) {
      90            0 :             notifyRequests.emplace_back(std::make_tuple(masterIdY, q, i));
      91            0 :             notifyRequests.emplace_back(std::make_tuple(q, masterIdY, i));
      92              :         }
      93              :     }
      94            0 :     return notifyRequests;
      95            0 : }
      96              : 
      97              : /*
      98              :  * Desc: 返回当前rank能处理的数据量和scratch buffer之间的比例关系
      99              :  * param: input: 输入数据位置
     100              :  * param: output 输出数据位置
     101              :  */
     102            0 : u32 InsTempAllReduceMesh2DTwoShot::CalcScratchMultiple(BufferType input, BufferType output) const
     103              : {
     104              :     // scratchbuffer如果能够通过ranksize规整:buffersize%2*ranksize_M*ranksize_N=0,则这里只需要返回1,最大化利用scratchbuffer
     105              : 
     106              :     // 否则返回2,使用1倍的buffer保证能缓存所有其他rank发来的数据,理论上数据被分成2*M*N块,假设有尾块,每个数据块的大小(inputcount/(2*M*N)+1)
     107              :     // 总共需要(inputCount/(2*M*N)+1)*(2*M*N)=[inputCount+2*M*N]*elembytesize,
     108              :     // 而预留的buffersize=inputcount*elembytesize,
     109              :     // 所以如果2*M*N>inputcount(预留)则缓存buffer仍然不够,但是由于最小的scratchbuffersize=1M,而2*M*N很难大于1M/2(一半数据一半缓存),
     110              :     // 所以返回2的时候要判断(2*M*N+inputcount)*elembytesize>scratchbufferSize(即预留缓存buffer+输入数据占用的buffer);2*M*N是常量,只需要增加buffersize解决
     111              :     (void)input;
     112              :     (void)output;
     113            0 :     u32 multiple = 2;
     114            0 :     return multiple;
     115              : }
     116              : 
     117            0 : HcclResult InsTempAllReduceMesh2DTwoShot::BuildSlice(
     118              :     const std::vector<RankId>& rankInfo, const u64 dataSize, const u64 chunkSize, RankSliceInfo& sliceInfoVec) const
     119              : {
     120            0 :     std::vector<SliceInfo> tmp(1);
     121            0 :     sliceInfoVec.resize(rankInfo.size(), tmp);
     122              : 
     123            0 :     u64 accumOff = 0;
     124            0 :     for (u32 rankIdx = 0; rankIdx < rankInfo.size(); rankIdx++) {
     125            0 :         u64 currChunkSize = ((dataSize - accumOff) > chunkSize) ? chunkSize : (dataSize - accumOff);
     126            0 :         SliceInfo slice = {accumOff, currChunkSize};
     127            0 :         sliceInfoVec[rankIdx][0] = slice;
     128            0 :         accumOff += currChunkSize;
     129              :     }
     130            0 :     return HcclResult::HCCL_SUCCESS;
     131            0 : }
     132              : 
     133              : /*
     134              :  * Desc: GenExtIns 算子执行入口
     135              :  * param: tempAlgParams: slice和stride信息
     136              :  * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
     137              :  * param: tempLinks: 当前rank通信链接信息
     138              :  * param: tempInsQues: 通信队列
     139              :  * return: HcclResult
     140              :  */
     141            0 : HcclResult InsTempAllReduceMesh2DTwoShot::GenExtIns(
     142              :     const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
     143              :     std::vector<InsQuePtr>& tempInsQues)
     144              : {
     145            0 :     InitInnerParams(tempFuncs, tempAlgParams, tempLinks, tempInsQues);
     146              :     // step1: reducescatter, X轴划分为M个块,每个块大小N*chunksize, Y轴划分为N个块,每个块M*chucksize
     147            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     148            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));
     149            0 :     SubStageArgs bufferInfo
     150              :         = {tempAlgParams.buffInfo.inBuffType, tempAlgParams.buffInfo.scratBuffType,
     151            0 :            tempAlgParams.buffInfo.inBuffBaseOff, 0};
     152            0 :     CHK_RET(RunReduceScatter(bufferInfo, XsliceInfoVec_, tempLinks, XtempInsQues_, tempVTopo_[0]));
     153              : 
     154            0 :     if (YDataSize_ != 0) {
     155              :         bufferInfo
     156            0 :             = {tempAlgParams.buffInfo.inBuffType, tempAlgParams.buffInfo.scratBuffType,
     157            0 :                tempAlgParams.buffInfo.inBuffBaseOff + M_ * N_ * chunkSize_, M_ * N_ * chunkSize_};
     158            0 :         CHK_RET(RunReduceScatter(bufferInfo, YsliceInfoVec_, tempLinks, YtempInsQues_, tempVTopo_[1]));
     159              :     }
     160              : 
     161              :     // step2: 换轴reducescatter
     162            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     163            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));
     164            0 :     if (XDataSizeS2_ != 0) {
     165              :         bufferInfo
     166            0 :             = {tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratBuffType, M_ * N_ * chunkSize_,
     167            0 :                M_ * N_ * chunkSize_ + M_ * chunkSize_};
     168            0 :         CHK_RET(RunReduceScatter(bufferInfo, XsliceInfoVecS2_, tempLinks, XtempInsQues_, tempVTopo_[0]));
     169              :     }
     170            0 :     if (YDataSizeS2_ != 0) {
     171            0 :         bufferInfo = {tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratBuffType, 0, N_ * chunkSize_};
     172            0 :         CHK_RET(RunReduceScatter(bufferInfo, YsliceInfoVecS2_, tempLinks, YtempInsQues_, tempVTopo_[1]));
     173              :     }
     174              : 
     175              :     // step3: allgather
     176            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     177            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));
     178            0 :     if (XDataSizeS2_ != 0) { // X轴allgather
     179              :         bufferInfo
     180            0 :             = {tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratBuffType,
     181            0 :                M_ * N_ * chunkSize_ + M_ * chunkSize_, M_ * N_ * chunkSize_};
     182            0 :         CHK_RET(RunAllgather(bufferInfo, XsliceInfoVecS2_, tempLinks, XtempInsQues_, tempVTopo_[0]));
     183              :     }
     184            0 :     if (YDataSizeS2_ != 0) { // Y轴allgather
     185            0 :         bufferInfo = {tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratBuffType, N_ * chunkSize_, 0};
     186            0 :         CHK_RET(RunAllgather(bufferInfo, YsliceInfoVecS2_, tempLinks, YtempInsQues_, tempVTopo_[1]));
     187              :     }
     188              : 
     189              :     // step4: 换轴allgather
     190            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     191            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));
     192              :     bufferInfo
     193            0 :         = {tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.outBuffType, 0,
     194            0 :            tempAlgParams.buffInfo.outBuffBaseOff};
     195            0 :     CHK_RET(RunAllgather(bufferInfo, XsliceInfoVec_, tempLinks, XtempInsQues_, tempVTopo_[0]));
     196            0 :     if (YDataSize_ != 0) {
     197              :         bufferInfo
     198            0 :             = {tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.outBuffType, M_ * N_ * chunkSize_,
     199            0 :                tempAlgParams.buffInfo.outBuffBaseOff + M_ * N_ * chunkSize_};
     200            0 :         CHK_RET(RunAllgather(bufferInfo, YsliceInfoVec_, tempLinks, YtempInsQues_, tempVTopo_[1]));
     201              :     }
     202            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     203            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));
     204            0 :     return HcclResult::HCCL_SUCCESS;
     205              : }
     206              : 
     207            0 : HcclResult InsTempAllReduceMesh2DTwoShot::InitInnerParams(
     208              :     const TempFuncs& tempFuncs, const TemplateDataParams& tempAlgParams, const ResLinks& tempLinks,
     209              :     std::vector<InsQuePtr>& tempInsQues)
     210              : {
     211              :     (void)tempLinks;
     212            0 :     HCCL_INFO("[InsTempAllReduceMesh2DTwoShot] start.");
     213            0 :     opMode_ = tempFuncs.opMode;
     214            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
     215              : 
     216            0 :     M_ = tempVTopo_[0].size();
     217            0 :     N_ = tempVTopo_[1].size();
     218            0 :     queNum_ = M_ + N_;
     219            0 :     CHK_PRT_RET(
     220              :         queNum_ != tempInsQues.size(),
     221              :         HCCL_ERROR(
     222              :             "[InsTempAllReduceMesh2DTwoShot] Rank [%d], queNum_:[%u], tempInsQues size:[%u],requiredQue Error.",
     223              :             myRank_, queNum_, tempInsQues.size()),
     224              :         HcclResult::HCCL_E_INTERNAL);
     225              : 
     226            0 :     u32 dataSizePerVolume = DataTypeSizeGet(dataType_); // 均分为2MN块
     227            0 :     u32 times = 2;
     228            0 :     chunkSize_ = RoundUp(tempAlgParams.sliceSize, (M_ * N_ * times * dataSizePerVolume)) * dataSizePerVolume;
     229            0 :     CHK_PRT_RET(
     230              :         (chunkSize_ * M_ * N_ * times) > tempAlgParams.buffInfo.scratchBuffSize,
     231              :         HCCL_ERROR(
     232              :             "[InsTempAllReduceMesh2DTwoShot]Rank [%d], Input size:[%llu], BfSize:[%llu] Insufficient buffer!", myRank_,
     233              :             tempAlgParams.sliceSize, tempAlgParams.buffInfo.scratchBuffSize),
     234              :         HcclResult::HCCL_E_INTERNAL);
     235            0 :     XtempInsQues_ = std::vector<InsQuePtr>(tempInsQues.begin(), tempInsQues.begin() + M_);
     236            0 :     YtempInsQues_ = std::vector<InsQuePtr>(tempInsQues.begin() + M_, tempInsQues.end());
     237              : 
     238            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[0], XAlgrankId_));
     239            0 :     CHK_RET(GetAlgRank(myRank_, tempVTopo_[1], YAlgrankId_));
     240              : 
     241              :     // for step1
     242            0 :     XDataSize_ = tempAlgParams.sliceSize >= M_ * N_ * chunkSize_ ? M_ * N_ * chunkSize_ : tempAlgParams.sliceSize;
     243            0 :     YDataSize_ = tempAlgParams.sliceSize - XDataSize_;
     244            0 :     BuildSlice(tempVTopo_[0], XDataSize_, N_ * chunkSize_, XsliceInfoVec_);
     245            0 :     BuildSlice(tempVTopo_[1], YDataSize_, M_ * chunkSize_, YsliceInfoVec_);
     246              : 
     247              :     // for step2
     248            0 :     YDataSizeS2_ = XsliceInfoVec_[XAlgrankId_][0].size; // 找到前一步切分时本rank负责的数据块大小
     249            0 :     XDataSizeS2_ = YsliceInfoVec_[YAlgrankId_][0].size;
     250            0 :     BuildSlice(tempVTopo_[1], YDataSizeS2_, chunkSize_, YsliceInfoVecS2_);
     251            0 :     BuildSlice(tempVTopo_[0], XDataSizeS2_, chunkSize_, XsliceInfoVecS2_);
     252            0 :     return HcclResult::HCCL_SUCCESS;
     253              : }
     254              : 
     255              : /*
     256              :  * Desc: 2D Mesh twoshot AllReduce: Scatter+reduce
     257              :  * param: sliceInfoVec: 每个rank的数据切片信息
     258              :  * param: tempLinks: 当前rank通信链接信息
     259              :  * param: tempInsQues: 通信队列
     260              :  * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
     261              :  * return: HcclResult
     262              :  */
     263            0 : HcclResult InsTempAllReduceMesh2DTwoShot::RunReduceScatter(
     264              :     SubStageArgs& subparams, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
     265              :     std::vector<InsQuePtr>& tempInsQues, const std::vector<RankId>& rankInfo) const
     266              : {
     267              :     u32 myAlgrankId;
     268            0 :     CHK_RET(GetAlgRank(myRank_, rankInfo, myAlgrankId));
     269              : 
     270            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     271              :     // scatter
     272            0 :     for (u32 rankId = 0; rankId < rankInfo.size(); rankId++) { // 写模式
     273              :         DataSlice ssrcSlice = DataSlice(
     274            0 :             subparams.inType, sliceInfoVec[rankId][0].offset + subparams.inbaseOff, sliceInfoVec[rankId][0].size);
     275              :         DataSlice sdestSlice = DataSlice(
     276            0 :             subparams.outType, myAlgrankId * sliceInfoVec[rankId][0].size + subparams.outbaesOff,
     277            0 :             sliceInfoVec[rankId][0].size);
     278            0 :         if (rankId == myAlgrankId) {
     279            0 :             if (sliceInfoVec[rankId][0].size != 0) { // 如果是本地rank,直接拷贝到scratch对应位置
     280            0 :                 CHK_PRT_RET(
     281              :                     LocalCopy(tempInsQues[rankId], ssrcSlice, sdestSlice),
     282              :                     HCCL_ERROR(
     283              :                         "[InsTempAllReduceMesh2DTwoShot][RunReduceScatter] RunAllReduce scatter LocalCopy failed"),
     284              :                     HcclResult::HCCL_E_INTERNAL);
     285              :             }
     286              :         } else {
     287            0 :             const std::vector<LinkData>& linkSendRecv = tempLinks.at(rankInfo[rankId]);
     288              :             // 发送, 未过滤size为0的情况
     289            0 :             std::vector<DataSlice> sendSrcSlices{ssrcSlice};
     290            0 :             std::vector<DataSlice> sendDestSlices{sdestSlice};
     291              :             // 接收,未过滤size为0的情况
     292              :             DataSlice rsrcSlice = DataSlice(
     293            0 :                 subparams.inType, sliceInfoVec[myAlgrankId][0].offset + subparams.inbaseOff,
     294            0 :                 sliceInfoVec[myAlgrankId][0].size);
     295              :             DataSlice rdestSlice = DataSlice(
     296            0 :                 subparams.outType, rankId * sliceInfoVec[myAlgrankId][0].size + subparams.outbaesOff,
     297            0 :                 sliceInfoVec[myAlgrankId][0].size);
     298            0 :             std::vector<DataSlice> recvSrcSlices{rsrcSlice};
     299            0 :             std::vector<DataSlice> recvDestSlices{rdestSlice};
     300            0 :             TxRxLinks sendRecvLinks(linkSendRecv[0], linkSendRecv[0]);
     301            0 :             TxRxSlicesList sendRecvSlicesList({sendSrcSlices, sendDestSlices}, {recvSrcSlices, recvDestSlices});
     302            0 :             SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     303            0 :             CHK_PRT_RET(
     304              :                 SendRecv(sendRecvInfo, tempInsQues[rankId], 0, true, DmaMode::PUT),
     305              :                 HCCL_ERROR("[InsTempAllReduceMesh2DTwoShot][RunReduceScatter] RunAllReduce scatter failed"),
     306              :                 HcclResult::HCCL_E_INTERNAL);
     307            0 :         }
     308              :     }
     309            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));        // 从流同步,等待所有并发的send和copy完成
     310            0 :     if (sliceInfoVec[myAlgrankId][0].size != 0) { // local reduce, 计算结果都放在最开始的位置
     311            0 :         DataSlice ldestSlice = DataSlice(subparams.outType, subparams.outbaesOff, sliceInfoVec[myAlgrankId][0].size);
     312            0 :         for (u32 rankId = 1; rankId < rankInfo.size(); rankId++) {
     313              :             DataSlice lsrcSlice = DataSlice(
     314            0 :                 subparams.outType, rankId * sliceInfoVec[myAlgrankId][0].size + subparams.outbaesOff,
     315            0 :                 sliceInfoVec[myAlgrankId][0].size);
     316              :             // 所有reduce操作在同一个insque中才能保序;
     317            0 :             CHK_PRT_RET(
     318              :                 LocalReduce(tempInsQues[0], lsrcSlice, ldestSlice, dataType_, redOp_),
     319              :                 HCCL_ERROR("[InsTempAllReduceMesh2DTwoShot]LocalReduce failed"), HcclResult::HCCL_E_INTERNAL);
     320              :         }
     321              :     }
     322            0 :     return HcclResult::HCCL_SUCCESS;
     323              : }
     324              : 
     325              : /*
     326              :  * Desc: 2D Mesh twoshot AllReduce: Allgather
     327              :  * param: sliceInfoVec: 每个rank的数据切片信息
     328              :  * param: tempLinks: 当前rank通信链接信息
     329              :  * param: tempInsQues: 通信队列
     330              :  * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
     331              :  * return: HcclResult
     332              :  */
     333            0 : HcclResult InsTempAllReduceMesh2DTwoShot::RunAllgather(
     334              :     SubStageArgs& subparams, const RankSliceInfo& sliceInfoVec, const ResLinks& tempLinks,
     335              :     std::vector<InsQuePtr>& tempInsQues, const std::vector<RankId>& rankInfo) const
     336              : {
     337              :     u32 myAlgrankId;
     338            0 :     CHK_RET(GetAlgRank(myRank_, rankInfo, myAlgrankId));
     339              : 
     340              :     // sync:前同步
     341            0 :     CHK_RET(PreSyncQues(tempInsQues, 0));
     342              : 
     343              :     // allgather
     344            0 :     for (u32 rankId = 0; rankId < rankInfo.size(); rankId++) {
     345            0 :         DataSlice rsrcSlice = DataSlice(subparams.inType, subparams.inbaseOff, sliceInfoVec[rankId][0].size);
     346              :         DataSlice rdestSlice = DataSlice(
     347            0 :             subparams.outType, sliceInfoVec[rankId][0].offset + subparams.outbaesOff, sliceInfoVec[rankId][0].size);
     348            0 :         if (u32(myAlgrankId) == rankId) {
     349            0 :             if (sliceInfoVec[rankId][0].size != 0) {
     350              :                 // copy本端计算的结果到user output
     351            0 :                 CHK_PRT_RET(
     352              :                     LocalCopy(tempInsQues[rankId], rsrcSlice, rdestSlice),
     353              :                     HCCL_ERROR("[InsTempAllReduceMesh2DTwoShot][RunAllgather] RunAllReduce AllGather "
     354              :                                "LocalCopy failed"),
     355              :                     HcclResult::HCCL_E_INTERNAL);
     356              :             }
     357              :         } else {
     358            0 :             const std::vector<LinkData>& linkSendRecv = tempLinks.at(rankInfo[rankId]);
     359              :             // 接收, 未过滤size为0的情况
     360            0 :             std::vector<DataSlice> recvSrcSlices{rsrcSlice};
     361            0 :             std::vector<DataSlice> recvDestSlices{rdestSlice};
     362              : 
     363              :             // 发送,未过滤size为0的情况
     364            0 :             DataSlice ssrcSlice = DataSlice(subparams.inType, subparams.inbaseOff, sliceInfoVec[myAlgrankId][0].size);
     365              :             DataSlice sdestSlice = DataSlice(
     366            0 :                 subparams.outType, sliceInfoVec[myAlgrankId][0].offset + subparams.outbaesOff,
     367            0 :                 sliceInfoVec[myAlgrankId][0].size);
     368            0 :             std::vector<DataSlice> sendSrcSlices{ssrcSlice};
     369            0 :             std::vector<DataSlice> sendDestSlices{sdestSlice};
     370              : 
     371            0 :             TxRxLinks sendRecvLinks(linkSendRecv[0], linkSendRecv[0]);
     372            0 :             TxRxSlicesList sendRecvSlicesList({sendSrcSlices, sendDestSlices}, {recvSrcSlices, recvDestSlices});
     373              : 
     374            0 :             SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     375            0 :             CHK_PRT_RET(
     376              :                 SendRecv(sendRecvInfo, tempInsQues[rankId], 0, true, DmaMode::GET),
     377              :                 HCCL_ERROR("[InsTempAllReduceMesh2DTwoShot][RunAllgather] RunAllReduce AllGather failed"),
     378              :                 HcclResult::HCCL_E_INTERNAL);
     379            0 :         }
     380              :     }
     381            0 :     CHK_RET(PostSyncQues(tempInsQues, 0));
     382            0 :     return HcclResult::HCCL_SUCCESS;
     383              : }
     384              : 
     385              : } // namespace Hccl
        

Generated by: LCOV version 2.0-1