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

Generated by: LCOV version 2.0-1