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

Generated by: LCOV version 2.0-1