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_1D_two_shot.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 144 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 10 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_all_reduce_mesh_1D_two_shot.h"
      12              : 
      13              : #include "log.h"
      14              : #include "alg_data_trans_wrapper.h"
      15              : 
      16              : namespace Hccl {
      17            0 : InsTempAllReduceMesh1DTwoShot::InsTempAllReduceMesh1DTwoShot(const RankId virtualRank, const u32 tempRankSize,
      18            0 :     const std::vector<std::vector<RankId>> &tempVTopo, const std::map<RankId, u32> &tempVirtRankMap)
      19            0 :     : InsAlgTemplateBase(virtualRank, tempRankSize, tempVTopo, tempVirtRankMap)
      20              : {
      21            0 : }
      22              : 
      23            0 : InsTempAllReduceMesh1DTwoShot::~InsTempAllReduceMesh1DTwoShot()
      24              : {
      25            0 : }
      26              : 
      27              : /*
      28              :  * Desc: 计算资源需求
      29              :  * return: tempResReq: 资源计算结果存储,包括notify信息,links信息等
      30              :  * return: HcclResult
      31              :  */
      32            0 : HcclResult InsTempAllReduceMesh1DTwoShot::CalcRes(AlgTempResReq &tempResReq)
      33              : {
      34            0 :         CHK_PRT_RET(CalcResLinksMesh(myRank_, tempRankSize_, tempVTopo_, linkNumBtwPeers_, tempResReq) != HcclResult::HCCL_SUCCESS,
      35              :         HCCL_ERROR("[CollAlgFactory] [InsTempAllReduceMesh1DTwoShot] Rank [%d], resLinks calculation error!", myRank_),
      36              :         HcclResult::HCCL_E_INTERNAL);
      37            0 :     auto& linkReq = tempResReq.links;
      38            0 :     u32 pathNum = 0;
      39            0 :     for (auto resReqIter = linkReq.begin(); resReqIter != linkReq.end(); resReqIter++) {
      40            0 :         auto remoteRank = resReqIter->first;
      41            0 :         if (rank2PathNumMap_.find(remoteRank) == rank2PathNumMap_.end() || rank2PathNumMap_[remoteRank] == 0) {
      42            0 :             HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot] No path to remoteRank[%d]", remoteRank);
      43            0 :             return HcclResult::HCCL_E_INTERNAL;
      44              :         }
      45            0 :         if (pathNum == 0) {
      46            0 :             pathNum = rank2PathNumMap_[remoteRank];
      47            0 :         } else if (rank2PathNumMap_[remoteRank] != pathNum) {
      48            0 :             HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot] Inconsistency pathNum to remoteRanks, Previous consistent pathNum=[%u], mismatched "
      49              :                        "remoteRank=[%d], pathNum=[%u]",
      50              :                 pathNum,
      51              :                 remoteRank,
      52              :                 rank2PathNumMap_[remoteRank]);
      53            0 :             return HcclResult::HCCL_E_INTERNAL;
      54              :         }
      55            0 :         resReqIter->second = pathNum;
      56              :     }
      57              :  
      58              :     // 1D Mesh 需要的 que Num 为 ranksize
      59            0 :     tempResReq.queNum = tempVTopo_[0].size() * pathNum;
      60            0 :     HCCL_INFO("[InsTempAllReduceMesh1DTwoShot] tempResReq.queNum = %u", tempResReq.queNum);
      61            0 :     tempResReq.streamNum = tempResReq.queNum;
      62            0 :     tempResReq.queNotifys = CreateMasterSlaveQueNotifiesRequest(tempResReq.queNum);
      63              : 
      64            0 :     QId centerQ = 0;
      65            0 :     tempResReq.localWaitGroupCntNotify.emplace_back(centerQ, 0);
      66            0 :     tempResReq.localBcastPostCntNotify.emplace_back(centerQ, 0);
      67              : 
      68            0 :     return HcclResult::HCCL_SUCCESS;
      69              : }
      70              : 
      71              : /*
      72              :  * Desc: 将数据按照rank切分为chuck 块,给后续的allreduce操作使用
      73              :  * param: dataSize: 待处理的输入数据大小
      74              :  * return: sliceInfoVec: 存储数据切分结果
      75              :  * return: HcclResult
      76              :  */
      77            0 : HcclResult InsTempAllReduceMesh1DTwoShot::CalcSlice(const u64 dataSize, const u64 baseOff, RankSliceInfo &sliceInfoVec)
      78              : {
      79            0 :     std::vector<SliceInfo> tmp(tempVTopo_.size());
      80            0 :     sliceInfoVec.resize(tempRankSize_, tmp);
      81              : 
      82            0 :     u64 unitAllignSize = DataTypeSizeGet(dataType_);
      83            0 :     u64 chunkSize = RoundUp(dataSize, (tempRankSize_ * unitAllignSize)) * unitAllignSize;
      84              : 
      85            0 :     u64 accumOff = 0;
      86            0 :     for (u32 rankIdx = 0; rankIdx < tempRankSize_; rankIdx++) {
      87            0 :         u64 currChunkSize = ((dataSize - accumOff) > chunkSize) ? chunkSize : (dataSize - accumOff);
      88            0 :         SliceInfo slice = {accumOff + baseOff, currChunkSize};
      89            0 :         sliceInfoVec[rankIdx][0]=slice;
      90            0 :         accumOff += currChunkSize;
      91              :     }
      92              : 
      93            0 :     CHK_PRT_RET((sliceInfoVec[tempRankSize_ - 1][0].offset + sliceInfoVec[tempRankSize_ - 1][0].size != baseOff + dataSize),
      94              :         HCCL_ERROR("[InsAllReduceCombExecutor] chunkSize:[%llu], Rank:[%d], SliceInfo calculation error!", chunkSize, myRank_),
      95              :         HcclResult::HCCL_E_INTERNAL);
      96            0 :     return HcclResult::HCCL_SUCCESS;
      97            0 : }
      98              : 
      99              : /*
     100              : * Desc: 返回当前rank能处理的数据量和scratch buffer之间的比例关系
     101              : * param: input: 输入数据位置
     102              : * param: output 输出数据位置
     103              : */
     104            0 :  u32 InsTempAllReduceMesh1DTwoShot::CalcScratchMultiple(BufferType input, BufferType output) const
     105              :  {
     106              :     (void)input;
     107              :     (void)output;
     108            0 :     u32 multiple = 2;
     109            0 :     return multiple;
     110              :  }
     111              : 
     112              : /*
     113              :  * Desc: GenExtIns 算子执行入口
     114              :  * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
     115              :  * param: tempAlgParams: 每个rank的数据切片信息
     116              :  * param: tempLinks: 当前rank通信链接信息
     117              :  * param: tempInsQues: 通信队列
     118              :  * return: HcclResult
     119              :  */
     120            0 : HcclResult InsTempAllReduceMesh1DTwoShot::GenExtIns(const TempFuncs &tempFuncs, const TemplateDataParams &tempAlgParams,
     121              :     const ResLinks &tempLinks, std::vector<InsQuePtr> &tempInsQues)
     122              : {
     123            0 :     HCCL_INFO("[InsTempAllReduceMesh1DTwoShot] start.");
     124            0 :     opMode_ = tempFuncs.opMode;
     125            0 :     enableCounterNotify_ = tempFuncs.enableCounterNotify;
     126              : 
     127            0 :     uint32_t linkNum = tempLinks.begin()->second.size();
     128            0 :     CHK_PRT_RET( tempInsQues.size() != tempVTopo_[0].size() * linkNum,
     129              :         HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot] RankSize [%lu], linkNum_:[%u], tempInsQues size:[%zu],requiredQue Error.",
     130              :             tempVTopo_[0].size(),
     131              :             linkNum,
     132              :             tempInsQues.size()),
     133              :         HcclResult::HCCL_E_INTERNAL);
     134              : 
     135            0 :     u64 dataSizePerVolume = DataTypeSizeGet(dataType_);
     136            0 :     CHK_PRT_RET((tempRankSize_ * dataSizePerVolume) + tempAlgParams.sliceSize > tempAlgParams.buffInfo.scratchBuffSize,
     137              :         HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot]Rank [%d], Input size:[%llu], BfSize:[%llu] Insufficient buffer!",
     138              :             myRank_,
     139              :             tempAlgParams.sliceSize,
     140              :             tempAlgParams.buffInfo.scratchBuffSize),
     141              :         HcclResult::HCCL_E_INTERNAL);
     142              : 
     143            0 :     u64 inBuffSize = tempAlgParams.sliceSize;
     144            0 :     std::vector<RankSliceInfo> sliceInfoVecForAllLinks(linkNum);
     145            0 :     std::vector<float> dataSplitRate(linkNum);
     146            0 :     CHK_RET(CalcDataSplitRateForLinks(tempLinks.begin()->second, dataSplitRate));
     147            0 :     u64 typeSize = DataTypeSizeGet(dataType_);
     148            0 :     u64 dataCnt = inBuffSize / typeSize;
     149            0 :     std::vector<u64> inBuffSizeForLinks(linkNum);
     150            0 :     std::vector<u64> baseOff(linkNum);
     151            0 :     u64 offset = 0;
     152            0 :     for (u32 linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     153            0 :         if (linkIdx != linkNum - 1) {
     154            0 :             inBuffSizeForLinks[linkIdx] =
     155            0 :                 static_cast<u64>(static_cast<float>(dataCnt) * dataSplitRate[linkIdx]) * typeSize;
     156              :         } else {
     157            0 :             inBuffSizeForLinks[linkIdx] = inBuffSize - offset;
     158              :         }
     159              :         //对每块输入都CalcSlice
     160            0 :         CHK_RET(CalcSlice(inBuffSizeForLinks[linkIdx], offset, sliceInfoVecForAllLinks[linkIdx]));
     161            0 :         offset += inBuffSizeForLinks[linkIdx];
     162              :     }
     163              : 
     164            0 :     std::vector<std::vector<InsQuePtr>> tempInsQuesVec;
     165            0 :     u32 queNumPerlink = tempInsQues.size() / linkNum;
     166            0 :     HCCL_INFO("tempInsQues.size()=%zu,queNumPerlink=%u",tempInsQues.size(),queNumPerlink);
     167            0 :     for (uint32_t linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     168            0 :         tempInsQuesVec.emplace_back(tempInsQues.begin() + linkIdx * queNumPerlink, tempInsQues.begin() + (linkIdx+1) * queNumPerlink);
     169              :     }
     170            0 :     u32 mainQueIdx = 0;
     171            0 :     CHK_RET(PreSyncQues(tempInsQues, mainQueIdx));
     172            0 :     HCCL_INFO("[InsTempAllReduce1DMeshTwoShot][PreCopy] Rank [%d].", myRank_);
     173            0 :     for (uint32_t linkIdx = 0; linkIdx < linkNum; linkIdx++) {
     174            0 :         CHK_RET(RunAllReduceScatter(sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQuesVec[linkIdx], tempAlgParams, linkIdx));
     175            0 :         CHK_RET(RunAllReduceAllgather(sliceInfoVecForAllLinks[linkIdx], tempLinks, tempInsQuesVec[linkIdx], tempAlgParams, linkIdx));
     176              :     }
     177            0 :     HCCL_INFO("[InsTempAllReduce1DMeshTwoShot][PostCopy] Rank [%d].", myRank_);
     178              :     // 流间后同步,从流通知主流
     179            0 :     CHK_RET(PostSyncQues(tempInsQues, mainQueIdx));
     180              :  
     181            0 :     return HcclResult::HCCL_SUCCESS;
     182            0 : }
     183              : 
     184              : /*
     185              :  * Desc: 1D Mesh twoshot AllReduce: Scatter+reduce
     186              :  * param: sliceInfoVec: 每个rank的数据切片信息
     187              :  * param: tempLinks: 当前rank通信链接信息
     188              :  * param: tempInsQues: 通信队列
     189              :  * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
     190              :  * return: HcclResult
     191              :  */
     192            0 : HcclResult InsTempAllReduceMesh1DTwoShot::RunAllReduceScatter(const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
     193              :     std::vector<InsQuePtr> &tempInsQues, const TemplateDataParams &tempAlgParams, u32 linkIdx)
     194              : {
     195            0 :     u64 inBuffBaseOff = tempAlgParams.buffInfo.inBuffBaseOff;
     196            0 :     u32 MyAlgRank = tempVirtRankMap_[myRank_];
     197              :     // scatter
     198            0 :     for (u32 rankId = 0; rankId < tempRankSize_; rankId++) {
     199              :         DataSlice ssrcSlice = DataSlice(
     200            0 :             tempAlgParams.buffInfo.inBuffType, sliceInfoVec[rankId][0].offset + inBuffBaseOff, sliceInfoVec[rankId][0].size);
     201              :         DataSlice sdestSlice = DataSlice(
     202            0 :             tempAlgParams.buffInfo.scratBuffType, MyAlgRank * sliceInfoVec[rankId][0].size + tempAlgParams.buffInfo.scratchBuffBaseOff, sliceInfoVec[rankId][0].size);
     203            0 :         if (rankId == u32(MyAlgRank)) {
     204            0 :             if(sliceInfoVec[rankId][0].size != 0){
     205              :             // 如果是本地rank,直接拷贝到scratch对应位置
     206            0 :             CHK_PRT_RET(LocalCopy(tempInsQues[rankId], ssrcSlice, sdestSlice),
     207              :                 HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceScatter] RunAllReduce scatter LocalCopy failed"),
     208              :                 HcclResult::HCCL_E_INTERNAL);
     209              :             }
     210              :         } else {
     211            0 :             const std::vector<LinkData> &linkSendRecv = tempLinks.at(GetRankFromMap(rankId));
     212              :             // 发送, 未过滤size为0的情况
     213            0 :             std::vector<DataSlice> sendSrcSlices{ssrcSlice};
     214            0 :             std::vector<DataSlice> sendDestSlices{sdestSlice};
     215              : 
     216              :             // 接收,未过滤size为0的情况
     217              :             DataSlice rsrcSlice = DataSlice(
     218            0 :                 tempAlgParams.buffInfo.inBuffType, sliceInfoVec[MyAlgRank][0].offset + inBuffBaseOff, sliceInfoVec[MyAlgRank][0].size);
     219              :             DataSlice rdestSlice = DataSlice(
     220            0 :                 tempAlgParams.buffInfo.scratBuffType, rankId * sliceInfoVec[MyAlgRank][0].size + tempAlgParams.buffInfo.scratchBuffBaseOff, sliceInfoVec[MyAlgRank][0].size);
     221            0 :             std::vector<DataSlice> recvSrcSlices{rsrcSlice};
     222            0 :             std::vector<DataSlice> recvDestSlices{rdestSlice};
     223              : 
     224            0 :             TxRxLinks sendRecvLinks(linkSendRecv[linkIdx], linkSendRecv[linkIdx]);
     225            0 :             TxRxSlicesList sendRecvSlicesList({sendSrcSlices, sendDestSlices}, {recvSrcSlices, recvDestSlices});
     226              : 
     227            0 :             SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     228            0 :             CHK_PRT_RET(SendRecv(sendRecvInfo, tempInsQues[rankId],0, true, DmaMode::PUT),
     229              :                 HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceScatter] RunAllReduce scatter failed"),
     230              :                 HcclResult::HCCL_E_INTERNAL);
     231            0 :         }
     232              :     }
     233              :     //从流同步,等待所有并发的send和copy完成
     234            0 :     PostSyncInterQueues(tempInsQues);
     235              : 
     236              :     // local reduce
     237            0 :     if (sliceInfoVec[MyAlgRank][0].size != 0) {
     238            0 :        DataSlice ldestSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratchBuffBaseOff, sliceInfoVec[MyAlgRank][0].size);
     239            0 :         for (u32 rankId = 1; rankId < tempRankSize_; rankId++) {
     240              :             DataSlice lsrcSlice = DataSlice(
     241            0 :                 tempAlgParams.buffInfo.scratBuffType, rankId * sliceInfoVec[MyAlgRank][0].size + tempAlgParams.buffInfo.scratchBuffBaseOff, sliceInfoVec[MyAlgRank][0].size);
     242              :             // 所有reduce操作在同一个insque中才能保序;
     243            0 :             CHK_PRT_RET(LocalReduce(tempInsQues[0], lsrcSlice, ldestSlice, dataType_, redOp_),
     244              :                 HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceScatter] RunAllReduce reduce LocalReduce failed"),
     245              :                 HcclResult::HCCL_E_INTERNAL);
     246              :         }
     247              :     }
     248              :    
     249            0 :     return HcclResult::HCCL_SUCCESS;
     250              : }
     251              : 
     252              : /*
     253              :  * Desc: 1D Mesh twoshot AllReduce: Allgather
     254              :  * param: sliceInfoVec: 每个rank的数据切片信息 
     255              :  * param: tempLinks: 当前rank通信链接信息
     256              :  * param: tempInsQues: 通信队列
     257              :  * param: tempFuncs: 辅助信息包括userIn/OutSlices, opMode等标记信息
     258              :  * return: HcclResult
     259              :  */
     260            0 : HcclResult InsTempAllReduceMesh1DTwoShot::RunAllReduceAllgather(const RankSliceInfo &sliceInfoVec, const ResLinks &tempLinks,
     261              :     std::vector<InsQuePtr> &tempInsQues, const TemplateDataParams &tempAlgParams, u32 linkIdx)
     262              : {
     263            0 :     u64 outBuffBaseOff = tempAlgParams.buffInfo.outBuffBaseOff; 
     264              :     // sync:前同步
     265            0 :     PreSyncInterQueues(tempInsQues);
     266            0 :     u32 MyAlgRank = tempVirtRankMap_[myRank_];
     267              :     // allgather
     268            0 :     for (u32 rankId = 0; rankId < tempRankSize_; rankId++) {
     269            0 :         DataSlice rsrcSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratchBuffBaseOff, sliceInfoVec[rankId][0].size);
     270            0 :         DataSlice rdestSlice = DataSlice(tempAlgParams.buffInfo.outBuffType, sliceInfoVec[rankId][0].offset + outBuffBaseOff, sliceInfoVec[rankId][0].size);
     271            0 :         if (u32(MyAlgRank) == rankId ) {
     272            0 :             if (sliceInfoVec[rankId][0].size != 0) {
     273              :                 // copy本端计算的结果到user output
     274            0 :                 CHK_PRT_RET(LocalCopy(tempInsQues[rankId], rsrcSlice, rdestSlice),
     275              :                 HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceAllgather] RunAllReduce AllGather LocalCopy failed"),
     276              :                 HcclResult::HCCL_E_INTERNAL);
     277              :             }
     278              :         } else {
     279            0 :             const std::vector<LinkData> &linkSendRecv = tempLinks.at(GetRankFromMap(rankId));
     280              :             // 接收, 未过滤size为0的情况
     281            0 :             std::vector<DataSlice> recvSrcSlices{rsrcSlice};
     282            0 :             std::vector<DataSlice> recvDestSlices{rdestSlice};
     283              : 
     284              :             // 发送,未过滤size为0的情况
     285            0 :             DataSlice ssrcSlice = DataSlice(tempAlgParams.buffInfo.scratBuffType, tempAlgParams.buffInfo.scratchBuffBaseOff, sliceInfoVec[MyAlgRank][0].size);
     286            0 :             DataSlice sdestSlice = DataSlice(tempAlgParams.buffInfo.outBuffType, sliceInfoVec[MyAlgRank][0].offset  + outBuffBaseOff, sliceInfoVec[MyAlgRank][0].size);
     287            0 :             std::vector<DataSlice> sendSrcSlices{ssrcSlice};
     288            0 :             std::vector<DataSlice> sendDestSlices{sdestSlice};
     289              : 
     290            0 :             TxRxLinks sendRecvLinks(linkSendRecv[linkIdx], linkSendRecv[linkIdx]);
     291            0 :             TxRxSlicesList sendRecvSlicesList({sendSrcSlices, sendDestSlices}, {recvSrcSlices, recvDestSlices});
     292              : 
     293            0 :             SendRecvInfo sendRecvInfo(sendRecvLinks, sendRecvSlicesList);
     294            0 :             CHK_PRT_RET(SendRecv(sendRecvInfo, tempInsQues[rankId],0, true, DmaMode::GET),
     295              :                 HCCL_ERROR("[InsTempAllReduceMesh1DTwoShot][RunAllReduceAllgather] RunAllReduce AllGather failed"),
     296              :                 HcclResult::HCCL_E_INTERNAL);
     297            0 :         }
     298              :     }
     299            0 :     return HcclResult::HCCL_SUCCESS;
     300              : }
     301              : 
     302            0 : RankId InsTempAllReduceMesh1DTwoShot::GetRankFromMap(const u32 rankIdx)
     303              : {
     304            0 :     RankId rank = -1;
     305            0 :     HCCL_INFO("[InsTempAllReduceMesh1DTwoShot] GetRankFromMap");
     306            0 :     for (auto &pair : tempVirtRankMap_) {
     307            0 :         if (pair.second == rankIdx) {
     308            0 :             rank = pair.first;
     309            0 :             break;
     310              :         }
     311              :     }
     312            0 :     return rank;
     313              : }
     314              : }  // namespace Hccl
        

Generated by: LCOV version 2.0-1