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

Generated by: LCOV version 2.0-1