LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_reduce_scatter - coll_reduce_scatter_ring_for_910_93_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 48.7 % 513 250
Test Date: 2026-08-04 10:52:23 Functions: 62.5 % 32 20

            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 "coll_reduce_scatter_ring_for_910_93_executor.h"
      12              : #include <numeric>
      13              : #include "alg_template_register.h"
      14              : 
      15              : namespace hccl {
      16              : 
      17            3 : CollReduceScatterRingFor91093Executor::CollReduceScatterRingFor91093Executor(const HcclDispatcher dispatcher,
      18            3 :     std::unique_ptr<TopoMatcher> &topoMatcher)
      19            3 :     : CollReduceScatterExecutor(dispatcher, topoMatcher)
      20              : {
      21            3 :     DMAReduceFlag_ = (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE);
      22            3 :     desc_.deterministic = 1;
      23            3 :     desc_.level1SupportedAlgos = {
      24              :         AlgTypeLevel1::ALG_LEVEL1_NHR,
      25              :         AlgTypeLevel1::ALG_LEVEL1_NB,
      26              :         AlgTypeLevel1::ALG_LEVEL1_RING,
      27              :         AlgTypeLevel1::ALG_LEVEL1_AHC,
      28              :         AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE
      29            3 :     };
      30            3 :     desc_.level2SupportedAlgos = {
      31              :         AlgTypeLevel2::ALG_LEVEL2_NHR,
      32              :         AlgTypeLevel2::ALG_LEVEL2_NB,
      33              :         AlgTypeLevel2::ALG_LEVEL2_RING
      34            3 :     };
      35            3 : }
      36              : 
      37           16 : bool CollReduceScatterRingFor91093Executor::IsUnifiedMarch(const OpParam &param) const
      38              : {
      39           16 :     return IsSupportUnifiedMarch(param, topoType_, topoAttr_.serverNum, topoAttr_.superPodNum);
      40              : }
      41              : 
      42            6 : u64 CollReduceScatterRingFor91093Executor::CalcTotalCount(const OpParam &param) const
      43              : {
      44            6 :     if (isReduceScatterV_) {
      45            0 :         const auto *counts = static_cast<const u64 *>(param.VDataDes.counts);
      46            0 :         return std::accumulate(counts, counts + topoAttr_.userRankSize, 0ULL);
      47              :     }
      48            6 :     return param.DataDes.count * topoAttr_.userRankSize;
      49              : }
      50              : 
      51            6 : void CollReduceScatterRingFor91093Executor::ParseParam(const OpParam& param)
      52              : {
      53            6 :     tag_ = param.tag;
      54              : 
      55            6 :     const HcclDataType dataType = param.GetDataType();
      56              :     // 是否需要scratch memory
      57           16 :     if ((workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) &&
      58            6 :         isSupportSDMAReduce_ && IsSupportRDMAReduce(dataType, param.reduceType)) {
      59            4 :         scratchMemFlag_ = false;
      60              :     } else {
      61            2 :         scratchMemFlag_ = true;
      62              :     }
      63              : 
      64            6 :     HCCL_DEBUG("[CollReduceScatterRingFor91093Executor][ParseParam] tag[%s] isSupportSDMAReduce_[%u] "
      65              :         "scratchMemFlag_[%u] workflowMode_[%u]", tag_.c_str(), isSupportSDMAReduce_, scratchMemFlag_, workflowMode_);
      66              : 
      67              :     // 记录图模式总数据量
      68            6 :     totalSize_ = CalcTotalCount(param) * SIZE_TABLE[dataType];
      69            6 :     aicpuUnfoldMode_ = param.aicpuUnfoldMode;
      70            6 :     isZeroCopy_ = param.isZeroCopy;
      71            6 : }
      72              : 
      73            3 : HcclResult CollReduceScatterRingFor91093Executor::CalcScratchMemSize(u64& scratchMemSize)
      74              : {
      75            3 :     if (scratchMemFlag_) {
      76            1 :         if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      77            0 :             scratchMemSize = inCCLbufferSize_;
      78              :         } else {
      79            1 :             scratchMemSize = totalSize_;
      80              :         }
      81              :     } else {
      82            2 :         scratchMemSize = 0U;
      83              :     }
      84            3 :     HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcScratchMemSize] tag[%s] scratchMemSize[%llu] "
      85              :         "scratchMemFlag_[%u] workflowMode_[%u]", tag_.c_str(), scratchMemSize, scratchMemFlag_, workflowMode_);
      86            3 :     return HCCL_SUCCESS;
      87              : }
      88              : 
      89            3 : HcclResult CollReduceScatterRingFor91093Executor::CalcStreamNum(u32& streamNum)
      90              : {
      91            3 :     u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
      92              :         LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
      93            3 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      94            2 :         totalStreamNum *= STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
      95              :     }
      96            3 :     streamNum = totalStreamNum - 1;
      97            3 :     HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcStreamNum] tag[%s] streamNum[%u]",
      98              :         tag_.c_str(), streamNum);
      99            3 :     return HCCL_SUCCESS;
     100              : }
     101              : 
     102            3 : HcclResult CollReduceScatterRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
     103              : {
     104            3 :     TransportMemType inputType = TransportMemType::RESERVED;
     105            3 :     TransportMemType outputType = TransportMemType::RESERVED;
     106            3 :     CHK_RET(CalcTransportMemType(inputType, outputType));
     107            3 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
     108            3 :     CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
     109            3 :     CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
     110            3 :     return HCCL_SUCCESS;
     111              : }
     112              : 
     113            3 : HcclResult CollReduceScatterRingFor91093Executor::CalcTransportMemType(TransportMemType &inputType,
     114              :     TransportMemType &outputType)
     115              : {
     116            3 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
     117            2 :         inputType = TransportMemType::CCL_INPUT;
     118            2 :         if (scratchMemFlag_) {
     119            0 :             outputType = TransportMemType::SCRATCH;
     120              :         } else {
     121            2 :             outputType = TransportMemType::CCL_OUTPUT;
     122              :         }
     123              :     } else {
     124            1 :         inputType = TransportMemType::PARAM_INPUT;
     125            1 :         if (scratchMemFlag_) {
     126            1 :             outputType = TransportMemType::SCRATCH;
     127              :         } else {
     128            0 :             outputType = TransportMemType::PARAM_OUTPUT;
     129              :         }
     130              :     }
     131            3 :     HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]",
     132              :         tag_.c_str(), inputType, outputType);
     133            3 :     return HCCL_SUCCESS;
     134              : }
     135              : 
     136            3 : HcclResult CollReduceScatterRingFor91093Executor::CalcLevel0CommInfo(TransportMemType inputType,
     137              :     TransportMemType outputType,
     138              :     std::vector<LevelNSubCommTransport>& opTransport)
     139              : {
     140            3 :     CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
     141            3 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
     142            3 :     return HCCL_SUCCESS;
     143            3 : }
     144              : 
     145            3 : HcclResult CollReduceScatterRingFor91093Executor::CalcLevel2CommInfo(TransportMemType inputType,
     146              :     TransportMemType outputType,
     147              :     std::vector<LevelNSubCommTransport>& opTransport)
     148              : {
     149            3 :     if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
     150            3 :         algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE) {
     151            0 :         HCCL_INFO("[CollReduceScatterRingFor91093Executor][CalcLevel2CommInfo] select AHC bypass level2 comm calculate");
     152            0 :         return HCCL_SUCCESS;
     153              :     }
     154              : 
     155            3 :     CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX);
     156            3 :     if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     157            0 :         commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
     158            0 :         HCCL_INFO("[%s]Calc NHRCommInfo", __func__);
     159            3 :     } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     160            0 :         commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
     161            0 :         HCCL_INFO("[%s]Calc NBCommInfo", __func__);
     162              :     } else {
     163            3 :         commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
     164            3 :         HCCL_INFO("[%s]Calc RingCommInfo", __func__);
     165              :     }
     166            3 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
     167            3 :     return HCCL_SUCCESS;
     168            3 : }
     169              : 
     170            2 : u64 CollReduceScatterRingFor91093Executor::CalcLoopMaxCount(const u32 unitSize)
     171              : {
     172              :     // 中转内存单次最多能够接受的output count,放开ranksize限制
     173            2 :     u64 maxCountPerLoop = inCCLbufferSize_ / topoAttr_.userRankSize / HCCL_MIN_SLICE_ALIGN
     174            2 :         * HCCL_MIN_SLICE_ALIGN / unitSize;
     175            2 :     return maxCountPerLoop;
     176              : }
     177              : 
     178           16 : bool CollReduceScatterRingFor91093Executor::IsHugeData(const u64 curSize, OpParam *param)
     179              : {
     180              :     u32 level2RankSize;
     181           16 :     if ((algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
     182           16 :         algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE)) {
     183              :         //AHC非对称场景下没有L2
     184            0 :         level2RankSize =1;
     185              :     } else {
     186              :         // 多QP哈希散列开启且RDMA通信下,强制刷新子图
     187              :         // 这里如果CheckCommSize返回ERROR,相当于HugeData true,防止GetSubCommInfo越界
     188           16 :         CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
     189           16 :         SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     190           16 :         level2RankSize = level2CommInfo.localRankSize;
     191           16 :     }
     192              : 
     193           16 :     const u64 TBE_REDUCE_MAX_COUNT = INT32_MAX;
     194              : 
     195           16 :     u64 curCount = curSize / SIZE_TABLE[param->DataDes.dataType];
     196           16 :     bool issupportRDMAInlineReduce = IsSupportRDMAReduce(param->DataDes.dataType, param->reduceType);
     197           16 :     bool hugeData =
     198           16 :         (curSize * level2RankSize / HCCL_INTERNODE_MAX_DATA_RATE > RDMA_SEND_MAX_SIZE) ||
     199           16 :         (curSize > SDMA_SEND_MAX_SIZE) ||
     200           48 :         ((!isSupportSDMAReduce_) && (curCount > TBE_REDUCE_MAX_COUNT)) ||
     201           16 :         ((!issupportRDMAInlineReduce) && (curCount * level2RankSize / HCCL_INTERNODE_MAX_DATA_RATE > TBE_REDUCE_MAX_COUNT));
     202           16 :     return hugeData;
     203              : }
     204              : 
     205            0 : HcclResult CollReduceScatterRingFor91093Executor::RunIntraSeverReduceScatter(
     206              :     const std::string &tag, DeviceMem &inputMem, DeviceMem &outputMem,
     207              :     const u64 count, const HcclDataType &dataType, const HcclReduceOp &reductionOp,
     208              :     const std::vector<std::vector<Slice>> &multRingsSliceZero, const Stream &stream,
     209              :     s32 profStage, const u64 baseOffset, const HcomCollOpInfo *opInfo,
     210              :     const std::vector<std::vector<Slice>> &multRingsUserMemSlice, const bool disableDMAReduce)
     211              : {
     212            0 :     CHK_RET(MultiRingReduceScatter(tag, inputMem, outputMem, count, dataType, reductionOp,
     213              :         multRingsSliceZero, stream, profStage, baseOffset, opInfo, multRingsUserMemSlice, logicalLevel0plane_));
     214            0 :     return HCCL_SUCCESS;
     215              : }
     216              : 
     217           33 : void CollReduceScatterRingFor91093Executor::FillMultiRingSlice(const ExecMem &execMem,
     218              :     const std::vector<std::vector<Slice>> &multiStreamSlice, u32 sliceNum, u32 level1RankSize, u32 level2RankSize,
     219              :     const u32 ringIndex, std::vector<Slice> &dataSlice)
     220              : {
     221           99 :     for (u32 level0Idx = 0; level0Idx < sliceNum; level0Idx++) {
     222           66 :         Slice sliceTemp;
     223          132 :         for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
     224          198 :             for (u32 level1Idx = 0; level1Idx < level1RankSize; level1Idx++) {
     225          132 :                 sliceTemp.size = multiStreamSlice[ringIndex][level0Idx].size;
     226          132 :                 sliceTemp.offset = multiStreamSlice[ringIndex][level0Idx].offset +
     227          264 :                     level1Idx * sliceNum * execMem.outputMem.size() +
     228          132 :                     level2Idx * sliceNum * level1RankSize * execMem.outputMem.size();
     229          132 :                 dataSlice.push_back(sliceTemp);
     230          132 :                 HCCL_DEBUG("rank[%u] sliceTemp.size[%zu], sliceTemp.offset[%llu]", topoAttr_.userRank,
     231              :                     sliceTemp.size, sliceTemp.offset);
     232              :             }
     233              :         }
     234              :     }
     235           33 : }
     236              : 
     237           17 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel0DataSegsSlice(const ExecMem &execMem,
     238              :     std::vector<std::vector<Slice>> &multiStreamSlice, const OpParam &param, u32 ringNum, u32 sliceNum,
     239              :     u32 level1RankSize, u32 level2RankSize, HcclDataType dataType, std::vector<std::vector<Slice>> &level0DataSegsSlice)
     240              : {
     241           17 :     if (isReduceScatterV_) {
     242            0 :         return CalLevel0DataSegsSliceV(execMem, multiStreamSlice, param, ringNum, sliceNum, level1RankSize,
     243            0 :             level2RankSize, dataType, level0DataSegsSlice);
     244              :     }
     245           17 :     bool isInlineReduce = IsSupportSDMAReduce(execMem.inputMem.ptr(), execMem.scratchMem.ptr(), dataType,
     246           17 :         param.reduceType);
     247           17 :     bool useInlineReduce = isInlineReduce && algoAttr_.inlineReduceSwitchOn;
     248           17 :     std::vector<Slice> dataSegsSlice;   // 数据分成ranksize份,每份的起始偏移和大小
     249           17 :     multiStreamSlice = ReduceScatterRingSlicePrepare(ringNum, sliceNum, useInlineReduce, execMem.outputMem,
     250           17 :         dataSegsSlice, param.tag);
     251              : 
     252           50 :     for (u32 ringIndex = 0; ringIndex < multiStreamSlice.size(); ringIndex++) {
     253           33 :         std::vector<Slice> dataSlice;
     254           33 :         FillMultiRingSlice(execMem, multiStreamSlice, sliceNum, level1RankSize, level2RankSize, ringIndex, dataSlice);
     255           33 :         level0DataSegsSlice.push_back(dataSlice);
     256           33 :     }
     257           17 :     return HCCL_SUCCESS;
     258           17 : }
     259              : 
     260           17 : HcclResult CollReduceScatterRingFor91093Executor::CalUserMemDataSegsSlice(const ExecMem &execMem,
     261              :     const std::vector<std::vector<Slice>> &level0DataSegsSlice, const std::vector<std::vector<Slice>> &multiStreamSlice,
     262              :     const OpParam &param, u32 ringNum, u32 sliceNum, u32 level1RankSize, u32 level2RankSize, HcclDataType dataType,
     263              :     u32 perDataSize, HcomCollOpInfo *opInfoPtr, bool disableDMAReduce,
     264              :     std::vector<std::vector<Slice>> &multRingsUserMemSlice)
     265              : {
     266           17 :     if (isReduceScatterV_) {
     267            0 :         return CalUserMemDataSegsSliceV(execMem, param, ringNum, sliceNum, level1RankSize, level2RankSize, dataType,
     268            0 :             multRingsUserMemSlice);
     269              :     }
     270           17 :     CHK_PRT_RET(0 < param.DataDes.strideCount && param.DataDes.strideCount < param.DataDes.count,
     271              :         HCCL_ERROR("[CollReduceScatterRingFor91093Executor][KernelRun]strideCount[%llu] is smaller than opCount[%llu]",
     272              :         param.DataDes.strideCount, param.DataDes.count),
     273              :         HCCL_E_PARA);
     274           17 :     HCCL_DEBUG("[CollReduceScatterRingFor91093Executor][KernelRun]strideCount[%llu], opCount[%llu]",
     275              :         param.DataDes.strideCount, param.DataDes.count);
     276              : 
     277           17 :     u32 level0RankSize = logicalLevel0CommInfo_.localRankSize;
     278           17 :     bool ARSFlag = topoMatcher_->GetARSFlag();
     279           17 :     bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
     280              : 
     281           17 :     if (opInfoPtr == nullptr &&
     282            1 :         (!((topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING || ARSDoubleRing) &&
     283            0 :         (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB || disableDMAReduce)))) {
     284            1 :         multRingsUserMemSlice = level0DataSegsSlice;
     285              :         // 图模式,根据strideCount更新slice的offset
     286            1 :         if (param.DataDes.strideCount != 0) {
     287            0 :             CHK_RET(UpdateOffsetBasedOnStrideCount(param, multRingsUserMemSlice));
     288              :         }
     289            1 :     } else {
     290           48 :         for (u32 ringIndex = 0; ringIndex < level0DataSegsSlice.size(); ringIndex++) {
     291           32 :             std::vector<Slice> level1UserMemSlice;
     292          160 :             for (auto &cclSlice : level0DataSegsSlice[ringIndex]) {
     293          128 :                 Slice tmpSlice;
     294          128 :                 u64 count = (param.DataDes.strideCount == 0) ? param.DataDes.count : param.DataDes.strideCount;
     295          128 :                 tmpSlice.size = cclSlice.size;
     296          128 :                 CHK_PRT_RET(execMem.outputMem.size() == 0,
     297              :                     HCCL_ERROR("[CollReduceScatterRingFor91093Executor][KernelRun]cclout memsize[%llu] is zero",
     298              :                     execMem.outputMem.size()), HCCL_E_PARA);
     299          128 :                 tmpSlice.offset = (cclSlice.offset / execMem.outputMem.size()) * count * perDataSize +
     300          128 :                     multiStreamSlice[ringIndex][0].offset;
     301          128 :                 level1UserMemSlice.push_back(tmpSlice);
     302          128 :                 HCCL_DEBUG("rank[%u], ringIndex[%u], tmpSlice.offset=[%llu], size=[%llu]",
     303              :                     topoAttr_.userRank, ringIndex, tmpSlice.offset, tmpSlice.size);
     304              :             }
     305           32 :             multRingsUserMemSlice.push_back(level1UserMemSlice);
     306           32 :         }
     307              :     }
     308           17 :     return HCCL_SUCCESS;
     309              : }
     310              : 
     311           16 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel1DataSegsSlice(const ExecMem &execMem, const OpParam &param,
     312              :     CommPlane commPlaneLevel, const u32 &commIndex, u32 sliceNum, u32 level1RankSize, u32 level2RankSize,
     313              :     u32 perDataSize, std::vector<Slice> &level1DataSegsSlice)
     314              : {
     315           16 :     if (isReduceScatterV_) {
     316            0 :         return CalLevel1DataSegsSliceV(param, commPlaneLevel, commIndex, sliceNum, level1RankSize, level2RankSize,
     317            0 :             perDataSize, level1DataSegsSlice);
     318              :     }
     319           48 :     for (u32 i = 0; i < level1RankSize; i++) {
     320           32 :         Slice sliceTemp;
     321              :         u32 level1UserRank;
     322           32 :         CHK_RET(GetUserRankByRank(commPlaneLevel, commIndex, i, level1UserRank));
     323           32 :         if (level2RankSize <= 1) {
     324           32 :             sliceTemp.size = execMem.outputMem.size();
     325           32 :             sliceTemp.offset = level1UserRank * execMem.outputMem.size();
     326           32 :             level1DataSegsSlice.push_back(sliceTemp);
     327           32 :             HCCL_DEBUG("rank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]", topoAttr_.userRank, i,
     328              :                 sliceTemp.offset, sliceTemp.size);
     329              :         } else {
     330            0 :             for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
     331            0 :                 sliceTemp.size = execMem.outputMem.size();
     332            0 :                 sliceTemp.offset = (level1UserRank % (level1RankSize * sliceNum)) * execMem.outputMem.size() +
     333            0 :                         level2Idx * sliceNum * level1RankSize * execMem.outputMem.size();
     334            0 :                 level1DataSegsSlice.push_back(sliceTemp);
     335            0 :                 HCCL_DEBUG("rank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]", topoAttr_.userRank, i,
     336              :                     sliceTemp.offset, sliceTemp.size);
     337              :             }
     338              :         }
     339              :     }
     340           16 :     return HCCL_SUCCESS;
     341              : }
     342              : 
     343           17 : HcclResult CollReduceScatterRingFor91093Executor::GetLevelCommInfo()
     344              : {
     345           17 :     logicalLevel0plane_ = COMM_LEVEL0;
     346           17 :     CHK_RET(CheckCommSize(logicalLevel0plane_, COMM_INDEX_0 + 1));
     347           17 :     logicalLevel0CommInfo_ = GetSubCommInfo(logicalLevel0plane_, COMM_INDEX_0);
     348           17 :     u32 commIndex = logicalLevel0CommInfo_.localRank;
     349           34 :     bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
     350           17 :         algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     351           17 :     logicalLevel1plane_ = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
     352           17 :     CHK_RET(CheckCommSize(logicalLevel1plane_, commIndex + 1));
     353           17 :     logicalLevel1CommInfo_ = GetSubCommInfo(logicalLevel1plane_, commIndex);
     354           17 :     return HCCL_SUCCESS;
     355              : }
     356              :  
     357            0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel2DataSegsSlice(const ExecMem &execMem, const OpParam &param,
     358              :     u32 level2RankSize, u32 perDataSize, std::vector<Slice> &level2DataSegsSlice)
     359              : {
     360            0 :     if (isReduceScatterV_) {
     361            0 :         return CalLevel2DataSegsSliceV(param, level2RankSize, perDataSize, level2DataSegsSlice);
     362              :     }
     363            0 :     Slice sliceTemp;
     364            0 :     for (u32 i = 0; i < level2RankSize; i++) {
     365            0 :         sliceTemp.size = execMem.outputMem.size();
     366              :         u32 level2UserRank;
     367            0 :         CHK_RET(GetUserRankByRank(COMM_LEVEL2, COMM_INDEX_0, i, level2UserRank));
     368            0 :         sliceTemp.offset = level2UserRank * execMem.outputMem.size();
     369            0 :         level2DataSegsSlice.push_back(sliceTemp);
     370            0 :         HCCL_DEBUG("rank[%u], level2DataSegsSlice[%u].offset=%llu, size=[%llu], level2RankSize[%u]",
     371              :             topoAttr_.userRank, i, sliceTemp.offset, sliceTemp.size, level2RankSize);
     372              :     }
     373            0 :     return HCCL_SUCCESS;
     374              : }
     375              : 
     376            0 : void CollReduceScatterRingFor91093Executor::PrepareLevel0Slices(const OpParam &param, u32 sliceNum, u32 level1RankSize,
     377              :     u32 level1Index, u32 level2Index, u32 perDataSize, std::vector<Slice> &cclSegSlices)
     378              : {
     379            0 :     const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
     380              :     // 根据counts和displace计算每个rank的数据范围
     381              :     // cclSlices里的offset是cclBuffer范围内的偏移,就地计算得出,不考虑displs
     382            0 :     const u32 level1Rank = level2Index * level1RankSize * sliceNum + level1Index * sliceNum;
     383            0 :     u64 displace = std::accumulate(counts, counts + level1Rank, 0ULL);
     384            0 :     for (auto rank = 0U; rank < sliceNum; ++rank) {
     385            0 :         const u32 idx = level1Rank + rank;
     386            0 :         Slice slice;
     387            0 :         slice.size = counts[idx] * perDataSize;
     388            0 :         slice.offset = displace * perDataSize;
     389            0 :         cclSegSlices.emplace_back(slice);
     390            0 :         displace += counts[idx];
     391              :     }
     392            0 : }
     393              : 
     394            0 : void CollReduceScatterRingFor91093Executor::PrepareLevel0UserSlices(const OpParam &param, u32 sliceNum,
     395              :     u32 level1RankSize, u32 level1Index, u32 level2Index, u32 perDataSize, std::vector<Slice> &userSegSlices)
     396              : {
     397            0 :     const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
     398            0 :     const auto *displsPtr = static_cast<const u64*>(param.VDataDes.displs);
     399            0 :     const u32 level1Rank = level2Index * level1RankSize * sliceNum + level1Index * sliceNum;
     400              :     // 根据counts和displace计算每个rank的数据范围
     401              :     // userSlices里的offset是user input的偏移,使用传入的displs算得
     402            0 :     for (auto rank = 0U; rank < sliceNum; ++rank) {
     403            0 :         const u32 idx = level1Rank + rank;
     404            0 :         Slice slice;
     405            0 :         slice.size = counts[idx] * perDataSize;
     406            0 :         slice.offset = displsPtr[idx] * perDataSize;
     407            0 :         userSegSlices.emplace_back(std::move(slice));
     408              :     }
     409            0 : }
     410              : 
     411            0 : bool CollReduceScatterRingFor91093Executor::IsCceReduceAligned(const std::vector<Slice> &dataSlices) const
     412              : {
     413            0 :     for (const auto &slice : dataSlices) {
     414            0 :         if (slice.size % CCE_REDUCE_ALIGN_SIZE != 0) {
     415            0 :             return false;
     416              :         }
     417              :     }
     418            0 :     return true;
     419              : }
     420              : 
     421            0 : HcclResult CollReduceScatterRingFor91093Executor::FillMultiRingSliceV(const ExecMem &execMem, const OpParam &param,
     422              :     u32 ringNum, u32 sliceNum, u32 level1RankSize, u32 level2RankSize, HcclDataType dataType,
     423              :     std::vector<std::vector<Slice>> &level0DataSegsSlice, std::vector<std::vector<std::vector<Slice>>> &serverSlices,
     424              :     const Level0SlicesCalculator &calcLevel0Slices)
     425              : {
     426            0 :     bool isInlineReduce = IsSupportSDMAReduce(execMem.inputMem.ptr(), execMem.scratchMem.ptr(), dataType,
     427            0 :         param.reduceType);
     428            0 :     bool useInlineReduce = isInlineReduce && algoAttr_.inlineReduceSwitchOn;
     429            0 :     u32 perDataSize = 0;
     430            0 :     CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
     431            0 :     for (u32 i = 0; i < level2RankSize; i++) {
     432            0 :         for (u32 j = 0; j < level1RankSize; j++) {
     433            0 :             std::vector<Slice> dataSegsSlice;   // 数据分成rank size份,每份的起始偏移和大小
     434            0 :             calcLevel0Slices(param, sliceNum, level1RankSize, j, i, perDataSize, dataSegsSlice);
     435              : 
     436            0 :             std::vector<std::vector<Slice>> multiStreamSlices;
     437              :             // 再将每个 slice 划分为 ringNum 份
     438            0 :             if (ringNum == LEVEL0_PLANE_NUM_IN_8PRING) {
     439            0 :                 if (useInlineReduce) {
     440            0 :                     multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag);
     441            0 :                 } else if (IsCceReduceAligned(dataSegsSlice)) {
     442            0 :                     multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag);
     443              :                 } else {
     444            0 :                     multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, true);
     445              :                 }
     446            0 :             } else if (ringNum == LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE) {
     447              :                 // 双环场景,需要传入正确的 niclist (不涉及网口裁剪)
     448            0 :                 if (useInlineReduce) {
     449            0 :                     multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
     450            0 :                 } else if (IsCceReduceAligned(dataSegsSlice)) {
     451            0 :                     multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
     452              :                 } else {
     453            0 :                     multiStreamSlices = PrepareMultiRingSlice(dataSegsSlice, param.tag, true, topoAttr_.nicList);
     454              :                 }
     455              :             } else {
     456            0 :                 multiStreamSlices.push_back(dataSegsSlice);
     457              :             }
     458            0 :             serverSlices.push_back(multiStreamSlices);
     459            0 :         }
     460              :     }
     461            0 :     level0DataSegsSlice.resize(ringNum);
     462            0 :     for (u32 level0Idx = 0; level0Idx < sliceNum; level0Idx++) {
     463            0 :         for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
     464            0 :             for (u32 level1Idx = 0; level1Idx < level1RankSize; level1Idx++) {
     465            0 :                 u32 serverIdx = level2Idx * level1RankSize + level1Idx;
     466            0 :                 const auto &multiStreamSlices = serverSlices[serverIdx];
     467            0 :                 for (u32 ringIndex = 0; ringIndex < multiStreamSlices.size(); ringIndex++) {
     468            0 :                     const auto &slice = multiStreamSlices[ringIndex][level0Idx];
     469            0 :                     level0DataSegsSlice[ringIndex].push_back(slice);
     470            0 :                     HCCL_DEBUG("[RSV]rank[%u], level0[%u]level2[%u]level1[%u], ringIndex[%u] slice.offset=[%llu], "
     471              :                         "size=[%llu]", topoAttr_.userRank, level0Idx, level2Idx, level1Idx, ringIndex, slice.offset,
     472              :                         slice.size);
     473              :                 }
     474              :             }
     475              :         }
     476              :     }
     477            0 :     return HCCL_SUCCESS;
     478              : }
     479              : 
     480            0 : HcclResult CollReduceScatterRingFor91093Executor::CalUserMemDataSegsSliceV(const ExecMem &execMem,
     481              :     const OpParam &param, u32 ringNum, u32 sliceNum, u32 level1RankSize, u32 level2RankSize, HcclDataType dataType,
     482              :     std::vector<std::vector<Slice>> &multRingsUserMemSlice)
     483              : {
     484            0 :     std::vector<std::vector<std::vector<Slice>>> serverSlices;
     485            0 :     CHK_RET(FillMultiRingSliceV(execMem, param, ringNum, sliceNum, level1RankSize, level2RankSize, dataType,
     486              :         multRingsUserMemSlice, serverSlices, PrepareLevel0UserSlices));
     487            0 :     return HCCL_SUCCESS;
     488            0 : }
     489              : 
     490            0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel0DataSegsSliceV(const ExecMem &execMem,
     491              :     std::vector<std::vector<Slice>> &multiStreamSlice, const OpParam &param, u32 ringNum, u32 sliceNum,
     492              :     u32 level1RankSize, u32 level2RankSize, HcclDataType dataType, std::vector<std::vector<Slice>> &level0DataSegsSlice)
     493              : {
     494            0 :     std::vector<std::vector<std::vector<Slice>>> serverSlices;
     495            0 :     CHK_RET(FillMultiRingSliceV(execMem, param, ringNum, sliceNum, level1RankSize, level2RankSize, dataType,
     496              :         level0DataSegsSlice, serverSlices, PrepareLevel0Slices));
     497            0 :     multiStreamSlice = serverSlices[0];
     498            0 :     return HCCL_SUCCESS;
     499            0 : }
     500              : 
     501            0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel1DataSegsSliceV(const OpParam &param,
     502              :     CommPlane commPlaneLevel, const u32 &commIndex, u32 sliceNum, u32 level1RankSize, u32 level2RankSize,
     503              :     u32 perDataSize, std::vector<Slice> &level1DataSegsSlice)
     504              : {
     505            0 :     const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
     506            0 :     for (u32 i = 0; i < level1RankSize; i++) {
     507            0 :         Slice sliceTemp;
     508              :         u32 level1UserRank;
     509            0 :         CHK_RET(GetUserRankByRank(commPlaneLevel, commIndex, i, level1UserRank));
     510            0 :         if (level2RankSize <= 1) {
     511            0 :             sliceTemp.size = counts[level1UserRank] * perDataSize;
     512            0 :             sliceTemp.offset = std::accumulate(counts, counts + level1UserRank, 0ULL) * perDataSize;
     513            0 :             level1DataSegsSlice.push_back(sliceTemp);
     514            0 :             HCCL_DEBUG("[RSV]rank[%u], level1UserRank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]",
     515              :                 topoAttr_.userRank, level1UserRank, i, sliceTemp.offset, sliceTemp.size);
     516              :         } else {
     517            0 :             for (u32 level2Idx = 0; level2Idx < level2RankSize; level2Idx++) {
     518            0 :                 const u32 ranksPerServer = level1RankSize * sliceNum;
     519            0 :                 const u32 level2UserRank = level2Idx * ranksPerServer + level1UserRank % ranksPerServer;
     520            0 :                 sliceTemp.size = counts[level2UserRank] * perDataSize;
     521            0 :                 sliceTemp.offset = std::accumulate(counts, counts + level2UserRank, 0ULL) * perDataSize;
     522            0 :                 level1DataSegsSlice.push_back(sliceTemp);
     523            0 :                 HCCL_DEBUG("[RSV]rank[%u], level2UserRank[%u], level1DataSegsSlice[%u].offset=%llu, size=[%llu]",
     524              :                     topoAttr_.userRank, level2UserRank, i, sliceTemp.offset, sliceTemp.size);
     525              :             }
     526              :         }
     527              :     }
     528            0 :     return HCCL_SUCCESS;
     529              : }
     530              : 
     531            0 : HcclResult CollReduceScatterRingFor91093Executor::CalLevel2DataSegsSliceV(const OpParam &param, u32 level2RankSize,
     532              :     u32 perDataSize, std::vector<Slice> &level2DataSegsSlice)
     533              : {
     534            0 :     const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
     535            0 :     Slice sliceTemp;
     536            0 :     for (u32 i = 0; i < level2RankSize; i++) {
     537              :         u32 level2UserRank;
     538            0 :         CHK_RET(GetUserRankByRank(COMM_LEVEL2, COMM_INDEX_0, i, level2UserRank));
     539            0 :         sliceTemp.size = counts[level2UserRank] * perDataSize;
     540            0 :         sliceTemp.offset = std::accumulate(counts, counts + level2UserRank, 0ULL) * perDataSize;
     541            0 :         level2DataSegsSlice.push_back(sliceTemp);
     542            0 :         HCCL_DEBUG("[RSV]rank[%u], level2UserRank[%u], level2DataSegsSlice[%u].offset=%llu, size=[%llu]",
     543              :             topoAttr_.userRank, level2UserRank, i, sliceTemp.offset, sliceTemp.size);
     544              :     }
     545            0 :     return HCCL_SUCCESS;
     546              : }
     547              : 
     548           17 : HcomCollOpInfo CollReduceScatterRingFor91093Executor::GetHcomCollOpInfo(const OpParam &param,
     549              :     const ExecMem &execMem) const
     550              : {
     551           17 :     const u64 count = param.GetDataCount(topoAttr_.userRank);
     552           17 :     const HcclDataType dataType = param.GetDataType();
     553           17 :     const u64 strideCount = param.GetStrideCount();
     554           17 :     HcomCollOpInfo opInfo = {"", execMem.inputPtr, execMem.outputPtr, count, dataType, param.root, param.reduceType,
     555           17 :         strideCount};
     556           17 :     HCCL_DEBUG("[CollReduceScatterRingFor91093Executor][KernelRun] execMem.inputPtr[%p], execMem.outputPtr[%p], "
     557              :         "execMem.inputMem[%p], execMem.outputMem[%p], strideCount[%llu]", execMem.inputPtr, execMem.outputPtr,
     558              :         execMem.inputMem.ptr(), execMem.outputMem.ptr(), strideCount);
     559           17 :     return opInfo;
     560              : }
     561              : 
     562           16 : u64 CollReduceScatterRingFor91093Executor::CalcSrcMemOffset(const ExecMem &execMem, const OpParam &param,
     563              :     u32 perDataSize) const
     564              : {
     565           16 :     if (isReduceScatterV_) {
     566            0 :         const auto *counts = static_cast<u64 *>(param.VDataDes.counts);
     567            0 :         return std::accumulate(counts, counts + topoAttr_.userRank, 0ULL) * perDataSize;
     568              :     }
     569           16 :     return topoAttr_.userRank * execMem.outputMem.size();
     570              : }
     571              : 
     572           17 : HcclResult CollReduceScatterRingFor91093Executor::KernelRun(const OpParam &param, ExecMem &execMem)
     573              : {
     574           17 :     HCCL_CONFIG_INFO(HCCL_ALG, "[%s] executor starts, rsv[%u]", __func__, isReduceScatterV_);
     575           17 :     CHK_RET(GetLevelCommInfo()); // 获取通信域
     576           17 :     u32 perDataSize = 0;
     577           17 :     const HcclDataType dataType = param.GetDataType();
     578           17 :     CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
     579              : 
     580              :     u32 ringNum;
     581           17 :     u32 level0RankSize = logicalLevel0CommInfo_.localRankSize;
     582           17 :     bool ARSFlag = topoMatcher_->GetARSFlag();
     583           17 :     bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
     584           17 :     u32 sliceNum = logicalLevel0CommInfo_.localRankSize;
     585           17 :     u32 commIndex = logicalLevel0CommInfo_.localRank;
     586              :  
     587           17 :     if ((topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING && !IsUnifiedMarch(param) && !ARSFlag) || ARSDoubleRing) {
     588           16 :         ringNum = LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE;
     589              :     } else {
     590            1 :         ringNum = LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
     591              :     }
     592              : 
     593           34 :     bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
     594           17 :         algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     595              : 
     596           17 :     SubCommInfo level2CommInfo;
     597           17 :     if (isSelectAHC) {
     598            0 :         level2CommInfo = logicalLevel1CommInfo_;
     599            0 :         level2CommInfo.localRankSize = 1;   // AHC bypass level2
     600              :     } else {
     601           17 :         CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
     602           17 :         level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     603              :     }
     604           17 :     const u32 level2RankSize = level2CommInfo.localRankSize;
     605           17 :     const u32 level1RankSize = logicalLevel1CommInfo_.localRankSize;
     606              : 
     607              :     // 节点内reduce scatter
     608           17 :     CHK_RET(ActiveSlaveStreams(param.stream));
     609              : 
     610              :     // 计算slice
     611           17 :     std::vector<std::vector<Slice>> multiStreamSlice; // 每个stream使用的数据基于用户buffer的偏移
     612           17 :     std::vector<std::vector<Slice>> level0DataSegsSlice;
     613           17 :     CalLevel0DataSegsSlice(execMem, multiStreamSlice, param, ringNum, sliceNum, level1RankSize, level2RankSize,
     614              :         dataType, level0DataSegsSlice);
     615              : 
     616           17 :     HcomCollOpInfo opInfo = GetHcomCollOpInfo(param, execMem);
     617           17 :     HcomCollOpInfo *opInfoPtr = nullptr;
     618           17 :     if (DMAReduceFlag_) {
     619           16 :         opInfoPtr = &opInfo;
     620              :     }
     621              : 
     622           17 :     bool disableDMAReduce = algOpContext_.opRetryHandler.retryEnable &&
     623            0 :         (algOpContext_.opRetryHandler.inPlaceSupportRetryStatus == InplaceSupportRetryStatus::RETRY_1_ALLOW_NO_DMA_REDUCE_CASE1 ||
     624            0 :         algOpContext_.opRetryHandler.inPlaceSupportRetryStatus == InplaceSupportRetryStatus::RETRY_1_ALLOW_NO_DMA_REDUCE_CASE2);
     625           17 :     std::vector<std::vector<Slice>> multRingsUserMemSlice;
     626           17 :     CalUserMemDataSegsSlice(execMem, level0DataSegsSlice, multiStreamSlice, param, ringNum, sliceNum, level1RankSize,
     627              :         level2RankSize, dataType, perDataSize, opInfoPtr, disableDMAReduce, multRingsUserMemSlice);
     628              : 
     629              :     // 区分消减拷贝场景
     630           17 :     if ((topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING || ARSDoubleRing) &&
     631           16 :         (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OPS_KERNEL_INFO_LIB)) {
     632              :         // 图模式opinfo不为空
     633            0 :         HcomCollOpInfo graphModeOpInfo = {"", execMem.inputMem.ptr(), nullptr, param.GetDataCount(topoAttr_.userRank),
     634            0 :             dataType, param.root, param.reduceType, param.GetStrideCount()};
     635            0 :         CHK_RET(RunIntraSeverReduceScatter(param.tag, execMem.inputMem, execMem.scratchMem, execMem.count, dataType,
     636              :             param.reduceType, level0DataSegsSlice, param.stream, PROF_STAGE_1, 0, &graphModeOpInfo,
     637              :             multRingsUserMemSlice, disableDMAReduce));
     638           17 :     } else if (opInfoPtr != nullptr && (level1RankSize > 1 || level2RankSize > 1)) {
     639           16 :         HcomCollOpInfo opInfoByReduceScatterDMAreduce = *opInfoPtr;
     640           16 :         opInfoByReduceScatterDMAreduce.outputAddr = nullptr;
     641           16 :         CHK_RET(RunIntraSeverReduceScatter(param.tag, execMem.inputMem, execMem.scratchMem, execMem.count,
     642              :             dataType, param.reduceType, level0DataSegsSlice, param.stream, PROF_STAGE_1, 0,
     643              :             &opInfoByReduceScatterDMAreduce, multRingsUserMemSlice, disableDMAReduce));
     644           16 :     } else {
     645            1 :         CHK_RET(RunIntraSeverReduceScatter(param.tag, execMem.inputMem, execMem.scratchMem, execMem.count,
     646              :             dataType, param.reduceType, level0DataSegsSlice, param.stream, PROF_STAGE_1, 0, opInfoPtr,
     647              :             multRingsUserMemSlice, disableDMAReduce));
     648              :     }
     649              :     // 对于单server图模式的最后一步需要把数据从ccl input拷贝到ccl output上
     650           16 :     if (level1RankSize == 1 && level2RankSize == 1 && opInfoPtr == nullptr) {
     651            0 :         const u64 offset = CalcSrcMemOffset(execMem, param, perDataSize);
     652            0 :         DeviceMem srcMem = execMem.inputMem.range(offset, execMem.outputMem.size());
     653            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
     654            0 :     }
     655              : 
     656           16 :     if  (level1RankSize > 1) {
     657              :         // 节点间做reduce scatter(ring/NHR/NB)
     658           16 :         u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, dataType, param.reduceType);
     659           16 :         std::unique_ptr<AlgTemplateBase> level1TempAlg;
     660              : 
     661              :         // 计算slice
     662           16 :         std::vector<Slice> level1DataSegsSlice;
     663           16 :         CHK_RET(CalLevel1DataSegsSlice(execMem, param, logicalLevel1plane_, commIndex, sliceNum, level1RankSize,
     664              :             level2RankSize, perDataSize, level1DataSegsSlice));
     665              : 
     666           16 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     667           32 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     668           16 :                 TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     669           16 :             CHK_SMART_PTR_NULL(level1TempAlg);
     670           16 :             CHK_RET(level1TempAlg->Prepare(reduceAttr));
     671           16 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL1", __func__);
     672            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     673            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     674            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
     675            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     676            0 :             CHK_RET(level1TempAlg->Prepare(reduceAttr));
     677            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL1", __func__);
     678            0 :         } else if (isSelectAHC) {
     679              :             // 获取通信域分组信息
     680            0 :             std::vector<std::vector<std::vector<u32>>> globalSubGroups;
     681            0 :             std::map<AHCConcOpType, TemplateType> ahcAlgOption;
     682            0 :             CHK_RET(topoMatcher_->GetGlobalSubGroups(logicalLevel1plane_, globalSubGroups));
     683            0 :             topoMatcher_->GetAHCAlgOption(ahcAlgOption);
     684            0 :             if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
     685            0 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC, dispatcher_);
     686            0 :                 HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_AHC in COMM_LEVEL1", __func__);
     687              :             } else {
     688            0 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC_BROKE, dispatcher_);
     689            0 :                 HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_AHC_BROKE in COMM_LEVEL1", __func__);
     690              :             }
     691            0 :             HCCL_DEBUG("[CollReduceScatterRingFor91093Executor]runAsync for COMM_LEVEL1 ends");
     692            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     693            0 :             CHK_RET(level1TempAlg->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
     694            0 :             CHK_RET(level1TempAlg->Prepare(reduceAttr));
     695            0 :         } else {
     696            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     697            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
     698            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     699            0 :             CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
     700            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL1", __func__);
     701              :         }
     702              : 
     703           48 :         CHK_RET(level1TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count,
     704              :             dataType, param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, level1DataSegsSlice));
     705           16 :         CHK_RET(level1TempAlg->RegisterProfiler(
     706              :             (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + logicalLevel1CommInfo_.localRank,
     707              :             PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET, param.stream));
     708           16 :         CHK_RET(RunTemplate(level1TempAlg, logicalLevel1CommInfo_));
     709           16 :     }
     710              : 
     711           16 :     if (level2RankSize > 1) {
     712              :         /* ****************** 超节点间 reducescatter *******************************/
     713            0 :         u64 reduceAttr = GetReduceAttr(execMem.inputMem, execMem.scratchMem, dataType, param.reduceType);
     714              : 
     715              :         // 计算slice
     716            0 :         std::vector<Slice> level2DataSegsSlice;
     717            0 :         CHK_RET(CalLevel2DataSegsSlice(execMem, param, level2RankSize, perDataSize, level2DataSegsSlice));
     718              : 
     719            0 :         std::unique_ptr<AlgTemplateBase> level2TempAlg;
     720            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     721            0 :             level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     722            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
     723            0 :             CHK_SMART_PTR_NULL(level2TempAlg);
     724            0 :             CHK_RET(level2TempAlg->Prepare(reduceAttr));
     725            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL2", __func__);
     726            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     727            0 :             level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     728            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
     729            0 :             CHK_SMART_PTR_NULL(level2TempAlg);
     730            0 :             CHK_RET(level2TempAlg->Prepare(reduceAttr, false));
     731            0 :             if (algoAttr_.isSupportAtomicWrite) {
     732            0 :                 level2TempAlg->CloseBarrier();
     733              :             }
     734            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL2", __func__);
     735              :         } else {
     736            0 :             level2TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     737            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     738            0 :             CHK_SMART_PTR_NULL(level2TempAlg);
     739            0 :             CHK_RET(level2TempAlg->Prepare(reduceAttr));
     740            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL2", __func__);
     741              :         }
     742              : 
     743            0 :         CHK_RET(level2TempAlg->Prepare(execMem.inputMem, execMem.inputMem, execMem.scratchMem, execMem.count, dataType,
     744              :             param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, level2DataSegsSlice));
     745            0 :         CHK_RET(level2TempAlg->RegisterProfiler(
     746              :             (level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank,
     747              :             PROF_STAGE_2, HCCL_EXEC_STEP_NOT_SET, param.stream));
     748            0 :         CHK_RET(RunTemplate(level2TempAlg, level2CommInfo));
     749            0 :     }
     750              : 
     751           16 :     if (level1RankSize > 1 || level2RankSize > 1) {
     752              :         // 区分消减拷贝场景(消减拷贝数据需要拷贝到user output上)
     753           16 :         const u64 offset = CalcSrcMemOffset(execMem, param, perDataSize);
     754           16 :         DeviceMem srcMem = execMem.inputMem.range(offset, execMem.outputMem.size());
     755           16 :         if (opInfoPtr != nullptr) {
     756           16 :             DeviceMem dstMem = DeviceMem::create(static_cast<u8 *>(opInfoPtr->outputAddr), execMem.outputMem.size());
     757           16 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, const_cast<Stream&>(param.stream)));
     758           16 :         } else {
     759            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
     760              :         }
     761           16 :     }
     762              : 
     763           16 :     HCCL_INFO("ReduceScatter ring run success, rsv[%u]", isReduceScatterV_);
     764           16 :     return HCCL_SUCCESS;
     765           17 : }
     766              : 
     767            0 : HcclResult CollReduceScatterRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
     768              : {
     769            0 :     bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
     770            0 :         algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     771              : 
     772            0 :     if (isSelectAHC) {
     773            0 :         CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
     774            0 :         SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
     775              : 
     776            0 :         u32 commIndex = level0CommInfo.localRank;
     777              : 
     778            0 :         CommPlane commPlaneLevel1 = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
     779            0 :         CHK_RET(CheckCommSize(commPlaneLevel1, commIndex + 1));
     780            0 :         level1CommInfo = GetSubCommInfo(commPlaneLevel1, commIndex);
     781            0 :         return HCCL_SUCCESS;
     782            0 :     }
     783              : 
     784            0 :     if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
     785            0 :         return HCCL_E_UNAVAIL;
     786              :     }
     787            0 :     level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     788              : 
     789            0 :     return HCCL_SUCCESS;
     790              : }
     791              : 
     792            0 : HcclResult CollReduceScatterRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase> &level1TempAlg, u32 level1RankSize)
     793              : {
     794            0 :     bool isSelectAHC = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC ||
     795            0 :         algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     796            0 :     HCCL_DEBUG("[CollReduceScatterRingFor91093Executor]SelectTempAlg begins");
     797            0 :     if (isSelectAHC) {
     798            0 :         CommPlane commPlaneLevel1 = COMM_LEVEL1_AHC;
     799              :         // 获取通信域分组信息
     800            0 :         std::vector<std::vector<std::vector<u32>>> globalSubGroups;
     801            0 :         std::map<AHCConcOpType, TemplateType> ahcAlgOption;
     802            0 :         CHK_RET(topoMatcher_->GetGlobalSubGroups(commPlaneLevel1, globalSubGroups));
     803            0 :         topoMatcher_->GetAHCAlgOption(ahcAlgOption);
     804            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
     805            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC, dispatcher_);
     806            0 :             HCCL_INFO("reducescatter ring: using ahc algo inter-server.");
     807              :         } else {
     808            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_AHC_BROKE, dispatcher_);
     809            0 :             HCCL_INFO("reducescatter ring: using ahc-broke algo inter-server.");
     810              :         }
     811            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     812            0 :         CHK_RET(level1TempAlg->Prepare(NSLBDP_MIN_COUNT, globalSubGroups, ahcAlgOption));
     813            0 :         return HCCL_SUCCESS;
     814            0 :     }
     815            0 :     if (level1RankSize > 1) {
     816            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     817            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     818            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
     819            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NB in COMM_LEVEL2", __func__);
     820            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     821            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     822            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     823            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
     824            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_NHR in COMM_LEVEL2", __func__);
     825            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     826              :         } else {
     827            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     828            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     829            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_REDUCESCATTER_RING in COMM_LEVEL2", __func__);
     830            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     831              :         }
     832            0 :         return HCCL_SUCCESS;
     833              :     }
     834            0 :     return HCCL_E_UNAVAIL;
     835              : }
     836              : 
     837              : 
     838              : REGISTER_EXEC("ReduceScatterRingFor91093Executor", ReduceScatterRingFor91093, CollReduceScatterRingFor91093Executor);
     839              : }
        

Generated by: LCOV version 2.0-1