LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_reduce_scatter - coll_reduce_scatter_ring_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 51.7 % 240 124
Test Date: 2026-08-18 17:47:01 Functions: 66.7 % 12 8

            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_executor.h"
      12              : #include "alg_template_register.h"
      13              : 
      14              : namespace hccl {
      15              : 
      16            3 : CollReduceScatterRingExecutor::CollReduceScatterRingExecutor(
      17            3 :     const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
      18            3 :     : CollReduceScatterExecutor(dispatcher, topoMatcher)
      19              : {
      20              :     DMAReduceFlag_
      21            6 :         = (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE
      22            3 :            && topoAttr_.deviceType == DevType::DEV_TYPE_910_93);
      23            3 : }
      24              : 
      25            6 : void CollReduceScatterRingExecutor::ParseParam(const OpParam& param)
      26              : {
      27            6 :     tag_ = param.tag;
      28              : 
      29              :     // 是否需要scratch memory
      30           12 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE
      31            0 :         && (topoAttr_.deviceType == DevType::DEV_TYPE_910B || topoAttr_.deviceType == DevType::DEV_TYPE_910_93)
      32            6 :         && isSupportSDMAReduce_ && IsSupportRDMAReduce(param.DataDes.dataType, param.reduceType)) {
      33            0 :         scratchMemFlag_ = false;
      34              :     } else {
      35            6 :         scratchMemFlag_ = true;
      36              :     }
      37              : 
      38              :     // 记录图模式总数据量
      39            6 :     HCCL_DEBUG("[CollReduceScatterRingExecutor][ParseParam]scratchMemFlag is %d", scratchMemFlag_);
      40            6 :     totalSize_ = topoAttr_.userRankSize * param.DataDes.count * SIZE_TABLE[param.DataDes.dataType];
      41            6 :     aicpuUnfoldMode_ = param.aicpuUnfoldMode;
      42            6 : }
      43              : 
      44            3 : HcclResult CollReduceScatterRingExecutor::CalcScratchMemSize(u64& scratchMemSize)
      45              : {
      46            3 :     if (scratchMemFlag_) {
      47            3 :         if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      48            0 :             scratchMemSize = inCCLbufferSize_;
      49              :         } else {
      50            3 :             scratchMemSize = totalSize_;
      51              :         }
      52              :     } else {
      53            0 :         scratchMemSize = 0U;
      54              :     }
      55            3 :     HCCL_INFO(
      56              :         "[CollReduceScatterRingExecutor][CalcScratchMemSize] tag[%s] scratchMemSize[%llu]", tag_.c_str(),
      57              :         scratchMemSize);
      58            3 :     return HCCL_SUCCESS;
      59              : }
      60              : 
      61            3 : HcclResult CollReduceScatterRingExecutor::CalcStreamNum(u32& streamNum)
      62              : {
      63            3 :     u32 totalStreamNum = 1U;
      64            3 :     if (algType_.algoLevel0 == AlgTypeLevel0::ALG_LEVEL0_8P_RING) {
      65            0 :         totalStreamNum = LEVEL0_PLANE_NUM_IN_8PRING;
      66              :     }
      67            3 :     streamNum = totalStreamNum - 1;
      68            3 :     HCCL_INFO("[CollReduceScatterRingExecutor][CalcStreamNum] tag[%s] streamNum[%u]", tag_.c_str(), streamNum);
      69            3 :     return HCCL_SUCCESS;
      70              : }
      71              : 
      72            3 : HcclResult CollReduceScatterRingExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      73              : {
      74            3 :     TransportMemType inputType = TransportMemType::RESERVED;
      75            3 :     TransportMemType outputType = TransportMemType::RESERVED;
      76            3 :     CHK_RET(CalcTransportMemType(inputType, outputType));
      77            3 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
      78            3 :     CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
      79            3 :     return HCCL_SUCCESS;
      80              : }
      81              : 
      82              : HcclResult
      83            3 : CollReduceScatterRingExecutor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
      84              : {
      85            3 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      86            0 :         inputType = TransportMemType::CCL_INPUT;
      87            0 :         if (scratchMemFlag_) {
      88            0 :             outputType = TransportMemType::SCRATCH;
      89              :         } else {
      90            0 :             outputType = TransportMemType::CCL_OUTPUT;
      91              :         }
      92              :     } else {
      93            3 :         inputType = TransportMemType::PARAM_INPUT;
      94            3 :         if (scratchMemFlag_) {
      95            3 :             outputType = TransportMemType::SCRATCH;
      96              :         } else {
      97            0 :             outputType = TransportMemType::PARAM_OUTPUT;
      98              :         }
      99              :     }
     100            3 :     HCCL_INFO(
     101              :         "[CollReduceScatterRingExecutor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]", tag_.c_str(),
     102              :         inputType, outputType);
     103            3 :     return HCCL_SUCCESS;
     104              : }
     105              : 
     106            3 : HcclResult CollReduceScatterRingExecutor::CalcLevel0CommInfo(
     107              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
     108              : {
     109            3 :     HCCL_INFO("[CollReduceScatterRingExecutor][CalcLevel0CommInfo]tag[%s] start", tag_.c_str());
     110            3 :     CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
     111            3 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
     112            3 :     HCCL_INFO("[CollReduceScatterRingExecutor][CalcLevel0CommInfo]tag[%s] Calc RingComm finish", tag_.c_str());
     113            3 :     return HCCL_SUCCESS;
     114            3 : }
     115              : 
     116            0 : u64 CollReduceScatterRingExecutor::CalcLoopMaxCount(const u32 unitSize)
     117              : {
     118              :     // 中转内存单次最多能够接受的output count,放开ranksize限制
     119            0 :     u64 maxCountPerLoop = inCCLbufferSize_ / (topoAttr_.userRankSize * unitSize);
     120            0 :     return maxCountPerLoop;
     121              : }
     122              : 
     123            0 : bool CollReduceScatterRingExecutor::IsHugeData(const u64 curSize, [[maybe_unused]] OpParam* param)
     124              : {
     125              :     bool hugeData;
     126            0 :     if (DMAReduceFlag_) {
     127            0 :         hugeData = curSize > SDMA_SEND_MAX_SIZE;
     128              :     } else {
     129            0 :         hugeData = (curSize * topoAttr_.userRankSize / HCCL_INTERNODE_MAX_DATA_RATE > RDMA_SEND_MAX_SIZE)
     130            0 :                    || (curSize > SDMA_SEND_MAX_SIZE);
     131              :     }
     132              : 
     133            0 :     return hugeData;
     134              : }
     135              : 
     136            3 : HcclResult CollReduceScatterRingExecutor::KernelRun(const OpParam& param, ExecMem& execMem)
     137              : {
     138            3 :     HCCL_CONFIG_INFO(HCCL_ALG, "[CollReduceScatterRingExecutor][KernelRun] userRank[%u] starts.", topoAttr_.userRank);
     139            3 :     u32 perDataSize = 0;
     140            3 :     CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, perDataSize));
     141              : 
     142            3 :     CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
     143            3 :     SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
     144              : 
     145            3 :     u32 ringNum
     146            3 :         = (topoType_ == TopoType::TOPO_TYPE_8P_RING) ? LEVEL0_PLANE_NUM_IN_8PRING : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
     147              : 
     148            3 :     u32 commIndex = (ringNum == LEVEL0_PLANE_NUM_IN_8PRING) ? topoAttr_.devicePhyId : level0CommInfo.localRank;
     149              : 
     150            3 :     CHK_RET(CheckCommSize(COMM_LEVEL1, commIndex + 1));
     151            3 :     SubCommInfo level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
     152              : 
     153              :     /* ******************网口裁剪步骤: 节点内allreduce *******************************/
     154            3 :     std::vector<Slice> dataSegsSlice;                 // 数据分成ranksize份,每份的起始偏移和大小
     155            3 :     std::vector<std::vector<Slice>> multiStreamSlice; // 每个stream使用的数据基于用户buffer的偏移
     156            3 :     u32 sliceNum = level0CommInfo.localRankSize;
     157              :     // Slice sliceTemp;
     158            3 :     bool isMultiNic = topoType_ == TopoType::TOPO_TYPE_8P_RING && topoAttr_.nicList.size() != DEVICE_EIGHT;
     159            3 :     if (isMultiNic) {
     160            0 :         u64 inputDataCount = execMem.inputMem.size() / perDataSize;
     161            0 :         CHK_RET(AlgTemplateBase::PrepareSliceData(inputDataCount, perDataSize, sliceNum, 0, dataSegsSlice));
     162            0 :         multiStreamSlice = PrepareMultiRingSlice(dataSegsSlice, param.tag);
     163            0 :         CHK_PRT_RET(
     164              :             multiStreamSlice.size() != ringNum,
     165              :             HCCL_ERROR(
     166              :                 "[CollReduceScatterRingExecutor][KernelRun]ringNum[%u] != multiStreamSlice size[%zu]", ringNum,
     167              :                 multiStreamSlice.size()),
     168              :             HCCL_E_INTERNAL);
     169              : 
     170            0 :         CHK_RET(MultiRingAllReduce(
     171              :             param.tag, execMem.inputMem, execMem.scratchMem, inputDataCount, param.DataDes.dataType, param.reduceType,
     172              :             multiStreamSlice, param.stream, PROF_STAGE_0));
     173              : 
     174            0 :         CHK_RET(
     175              :             HcclD2DMemcpyAsync(dispatcher_, execMem.inputMem, execMem.scratchMem, const_cast<Stream&>(param.stream)));
     176              :     }
     177              : 
     178            3 :     std::vector<u32>& nicList = const_cast<std::vector<u32>&>(topoAttr_.nicList);
     179            3 :     std::vector<u32>::iterator iterNic = std::find(nicList.begin(), nicList.end(), topoAttr_.devicePhyId);
     180            3 :     bool innRunRet = isMultiNic && (iterNic == nicList.end());
     181            3 :     if (!innRunRet) { // 1. 8P ring的拓扑。2. 网口不满配。3. 当前device不出网口。 的情况下不进行节点间的reduce scatter
     182              :         /* ******************第一步: 节点间reducescatter *******************************/
     183            3 :         u32 level1RankSize = level1CommInfo.localRankSize;
     184            3 :         if (level1RankSize > 1) {
     185              :             u64 reduceAttr
     186            3 :                 = GetReduceAttr(execMem.inputMem, execMem.scratchMem, param.DataDes.dataType, param.reduceType);
     187            3 :             std::unique_ptr<AlgTemplateBase> level1TempAlg;
     188              : 
     189            3 :             if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     190            0 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     191            0 :                     TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     192            0 :                 HCCL_INFO("ReduceScatter ring: using ring algo inter-server.");
     193            0 :                 CHK_SMART_PTR_NULL(level1TempAlg);
     194            0 :                 CHK_RET(level1TempAlg->Prepare(reduceAttr));
     195              : 
     196            0 :                 u64 ringSize = execMem.inputMem.size() / level1RankSize;
     197            0 :                 u64 ringCount = ringSize / perDataSize;
     198              : 
     199            0 :                 CHK_RET(level1TempAlg->Prepare(
     200              :                     execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount, param.DataDes.dataType,
     201              :                     param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
     202            3 :             } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     203            0 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     204            0 :                     TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
     205            0 :                 HCCL_INFO("ReduceScatter ring: using nhr algo inter-server.");
     206            0 :                 CHK_SMART_PTR_NULL(level1TempAlg);
     207            0 :                 CHK_RET(level1TempAlg->Prepare(reduceAttr, false));
     208              : 
     209            0 :                 u64 ringSize = execMem.inputMem.size() / level1RankSize;
     210            0 :                 u64 ringCount = ringSize / perDataSize;
     211            0 :                 CHK_RET(level1TempAlg->Prepare(
     212              :                     execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount, param.DataDes.dataType,
     213              :                     param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
     214            0 :                 level1TempAlg->CloseBarrier();
     215            3 :             } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
     216            0 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     217            0 :                     TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, dispatcher_);
     218            0 :                 HCCL_INFO("ReduceScatter ring: using nhr_v1 algo inter-server.");
     219            0 :                 CHK_SMART_PTR_NULL(level1TempAlg);
     220            0 :                 CHK_RET(level1TempAlg->Prepare(reduceAttr));
     221              : 
     222            0 :                 u64 ringSize = execMem.inputMem.size() / level1RankSize;
     223            0 :                 u64 ringCount = ringSize / perDataSize;
     224            0 :                 CHK_RET(level1TempAlg->Prepare(
     225              :                     execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount, param.DataDes.dataType,
     226              :                     param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
     227            3 :             } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     228            0 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     229            0 :                     TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
     230            0 :                 HCCL_INFO("ReduceScatter ring: using nonuniform-bruck algo inter-server.");
     231            0 :                 CHK_SMART_PTR_NULL(level1TempAlg);
     232            0 :                 CHK_RET(level1TempAlg->Prepare(reduceAttr));
     233              : 
     234            0 :                 u64 ringSize = execMem.inputMem.size() / level1RankSize;
     235            0 :                 u64 ringCount = ringSize / perDataSize;
     236            0 :                 CHK_RET(level1TempAlg->Prepare(
     237              :                     execMem.inputMem, execMem.inputMem, execMem.scratchMem, ringCount, param.DataDes.dataType,
     238              :                     param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
     239              :             } else {
     240            6 :                 level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     241            3 :                     TemplateType::TEMPLATE_REDUCESCATTER_RECURSIVE_HD, dispatcher_);
     242            3 :                 HCCL_INFO("ReduceScatter ring: using halving-doubling algo inter-server.");
     243              : 
     244            3 :                 CHK_SMART_PTR_NULL(level1TempAlg);
     245            3 :                 CHK_RET(level1TempAlg->Prepare(reduceAttr));
     246            3 :                 u64 inputDataCount = execMem.inputMem.size() / perDataSize; // count是output的数据个数
     247           15 :                 CHK_RET(level1TempAlg->Prepare(
     248              :                     execMem.inputMem, execMem.inputMem, execMem.scratchMem, inputDataCount, param.DataDes.dataType,
     249              :                     param.stream, param.reduceType, LEVEL0_BRIDGE_RANK_ID, std::vector<Slice>(0)));
     250              :             }
     251            3 :             CHK_RET(level1TempAlg->RegisterProfiler(
     252              :                 (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1CommInfo.localRank, PROF_STAGE_0,
     253              :                 HCCL_EXEC_STEP_NOT_SET, param.stream));
     254            3 :             CHK_RET(RunTemplate(level1TempAlg, level1CommInfo));
     255            3 :         }
     256              :     }
     257              : 
     258              :     /* ***********第二步: 节点内reducescatter(正常场景), 节点内多根结点scatter(网口裁剪)*****************************/
     259            3 :     CHK_RET(ActiveSlaveStreams(param.stream));
     260              : 
     261            3 :     bool useInlineRduce = false;
     262            3 :     bool isInlineReduce = IsSupportSDMAReduce(
     263            3 :         execMem.inputMem.ptr(), execMem.scratchMem.ptr(), param.DataDes.dataType, param.reduceType);
     264            3 :     useInlineRduce = isInlineReduce && algoAttr_.inlineReduceSwitchOn;
     265              :     multiStreamSlice
     266            3 :         = ReduceScatterRingSlicePrepare(ringNum, sliceNum, useInlineRduce, execMem.outputMem, dataSegsSlice, param.tag);
     267            3 :     bool bRet = (multiStreamSlice.size() != ringNum);
     268            3 :     CHK_PRT_RET(
     269              :         bRet,
     270              :         HCCL_ERROR(
     271              :             "[CollReduceScatterRingExecutor][KernelRun]sliceNum-1[%u] != multiStreamSlice size[%zu]", sliceNum - 1,
     272              :             multiStreamSlice.size()),
     273              :         HCCL_E_INTERNAL);
     274              : 
     275            3 :     if (isMultiNic) { // 网口裁剪情况下需要改变slice最终在rank上位置
     276            0 :         PrepareMultiRingSlice(dataSegsSlice, param.tag, false, nicList); // 刷新多环ringRankList信息
     277            0 :         std::vector<std::vector<u32>> ringNics;
     278            0 :         CHK_RET(GetRingNics(param.tag, ringNics));
     279              : 
     280            0 :         for (u32 ringIdx = 0; ringIdx < ringNum; ringIdx++) { // 按第一个网口位置改变slice最终在rank上的位置
     281            0 :             u32 firstNicIdx = ringNics[ringIdx][0];
     282            0 :             std::rotate(
     283            0 :                 multiStreamSlice[ringIdx].begin(), multiStreamSlice[ringIdx].begin() + firstNicIdx,
     284            0 :                 multiStreamSlice[ringIdx].end());
     285              :         }
     286            0 :     }
     287              : 
     288            3 :     DeviceMem srcMem;
     289            3 :     if (isMultiNic) {
     290            0 :         u32 level1RankSize = topoAttr_.userRankSize / DEVICE_EIGHT; // currComm->commLevel0[0]->UserRankSize();
     291              :         // 每个server分配的slice大小
     292            0 :         CHK_PRT_RET(
     293              :             level1RankSize == 0, HCCL_ERROR("[CollReduceScatterRingExecutor][KernelRun]level1RankSize is illegal"),
     294              :             HCCL_E_PARA);
     295            0 :         u64 serverSliceSize = execMem.inputMem.size() / level1RankSize;
     296              :         // 每个服务器对应的偏移
     297            0 :         u32 serverIndex = level1CommInfo.localRank;
     298            0 :         CHK_PRT_RET(
     299              :             serverIndex == INVALID_VALUE_RANKID,
     300              :             HCCL_ERROR(
     301              :                 "[CollReduceScatterRingExecutor][KernelRun]get rank of "
     302              :                 "bridgeRank failed, commIdx[%u]",
     303              :                 commIndex),
     304              :             HCCL_E_PARA);
     305            0 :         u64 serverSliceOffset = serverSliceSize * serverIndex;
     306            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     307            0 :             CHK_RET(HcclD2DMemcpyAsync(
     308              :                 dispatcher_, execMem.scratchMem, execMem.inputMem, const_cast<Stream&>(param.stream)));
     309              :         }
     310            0 :         DeviceMem reduceScatterRingOutput = execMem.scratchMem.range(serverSliceOffset, serverSliceSize);
     311            0 :         CHK_SMART_PTR_NULL(reduceScatterRingOutput.ptr());
     312            0 :         u64 countLocal = serverSliceSize / perDataSize;
     313            0 :         CHK_RET(MultiRingMultiRootScatter(
     314              :             param.tag, reduceScatterRingOutput, reduceScatterRingOutput, countLocal, param.DataDes.dataType,
     315              :             multiStreamSlice, serverIndex * DEVICE_EIGHT, param.stream, serverSliceOffset));
     316              : 
     317              :         srcMem
     318            0 :             = reduceScatterRingOutput.range(dataSegsSlice[topoAttr_.devicePhyId].offset, execMem.count * perDataSize);
     319            0 :         CHK_SMART_PTR_NULL(srcMem.ptr());
     320            0 :     } else {
     321            3 :         u32 level1RankSize = level1CommInfo.localRankSize;
     322              :         // 每个server分配的slice大小
     323            3 :         u64 serverSliceSize = execMem.inputMem.size() / level1RankSize;
     324              :         // 每个服务器对应的偏移
     325            3 :         u32 serverIndex = level1CommInfo.localRank;
     326            3 :         u64 serverSliceOffset = serverSliceSize * serverIndex;
     327            3 :         HCCL_DEBUG(
     328              :             "inputMem.size=%llu, level0CommInfo.localRankSize=%u, serverSliceSize=%llu, serverSliceOffset=%llu "
     329              :             "commIndex=%u commLevel1[commIndex]->rank=%u",
     330              :             execMem.inputMem.size(), level0CommInfo.localRankSize, serverSliceSize, serverSliceOffset, commIndex,
     331              :             level1CommInfo.localRank);
     332            3 :         DeviceMem reduceScatterRingInput = execMem.inputMem.range(serverSliceOffset, serverSliceSize);
     333            3 :         CHK_SMART_PTR_NULL(reduceScatterRingInput.ptr());
     334            3 :         DeviceMem reduceScatterRingOutput = execMem.scratchMem.range(serverSliceOffset, serverSliceSize);
     335            3 :         CHK_SMART_PTR_NULL(reduceScatterRingOutput.ptr());
     336            3 :         u64 countLocal = serverSliceSize / perDataSize;
     337              : 
     338            3 :         HcomCollOpInfo opInfo = {"",
     339            3 :                                  execMem.inputPtr,
     340            3 :                                  execMem.outputPtr,
     341            3 :                                  param.DataDes.count,
     342            3 :                                  param.DataDes.dataType,
     343            3 :                                  param.root,
     344            3 :                                  param.reduceType,
     345            3 :                                  0};
     346            3 :         HcomCollOpInfo* opInfoPtr = nullptr;
     347            3 :         if (DMAReduceFlag_) {
     348            0 :             opInfoPtr = &opInfo;
     349              :         }
     350              : 
     351            9 :         CHK_RET(MultiRingReduceScatter(
     352              :             param.tag, reduceScatterRingInput, reduceScatterRingOutput, countLocal, param.DataDes.dataType,
     353              :             param.reduceType, multiStreamSlice, param.stream, PROF_STAGE_1, serverSliceOffset, opInfoPtr));
     354              : 
     355              :         srcMem
     356            3 :             = execMem.inputMem.range(serverSliceOffset + dataSegsSlice[commIndex].offset, execMem.count * perDataSize);
     357            3 :         CHK_SMART_PTR_NULL(srcMem.ptr());
     358            3 :     }
     359              : 
     360            3 :     CHK_RET(HcclD2DMemcpyAsync(dispatcher_, execMem.outputMem, srcMem, const_cast<Stream&>(param.stream)));
     361              : 
     362            3 :     return HCCL_SUCCESS;
     363            3 : }
     364              : 
     365            0 : HcclResult CollReduceScatterRingExecutor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
     366              : {
     367            0 :     if (CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
     368            0 :         return HCCL_E_UNAVAIL;
     369              :     }
     370            0 :     SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
     371            0 :     u32 ringNum
     372            0 :         = (topoType_ == TopoType::TOPO_TYPE_8P_RING) ? LEVEL0_PLANE_NUM_IN_8PRING : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
     373            0 :     u32 commIndex = (ringNum == LEVEL0_PLANE_NUM_IN_8PRING) ? topoAttr_.devicePhyId : level0CommInfo.localRank;
     374              : 
     375            0 :     if (CheckCommSize(COMM_LEVEL1, commIndex + 1) != HCCL_SUCCESS) {
     376            0 :         return HCCL_E_UNAVAIL;
     377              :     }
     378            0 :     level1CommInfo = GetSubCommInfo(COMM_LEVEL1, commIndex);
     379              : 
     380            0 :     return HCCL_SUCCESS;
     381            0 : }
     382              : 
     383              : HcclResult
     384            0 : CollReduceScatterRingExecutor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
     385              : {
     386            0 :     if (level1RankSize > 1) {
     387            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     388            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     389            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_RING, dispatcher_);
     390            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     391            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     392              :             level1TempAlg
     393            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NHR, dispatcher_);
     394            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     395            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
     396            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     397            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_NHR_V1, dispatcher_);
     398            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     399            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     400              :             level1TempAlg
     401            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_REDUCESCATTER_NB, dispatcher_);
     402            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     403              :         } else {
     404            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     405            0 :                 TemplateType::TEMPLATE_REDUCESCATTER_RECURSIVE_HD, dispatcher_);
     406            0 :             CHK_SMART_PTR_NULL(level1TempAlg);
     407              :         }
     408            0 :         return HCCL_SUCCESS;
     409              :     }
     410            0 :     return HCCL_E_UNAVAIL;
     411              : }
     412              : 
     413              : REGISTER_EXEC("ReduceScatterRingExecutor", ReduceScatterRing, CollReduceScatterRingExecutor);
     414              : 
     415              : } // namespace hccl
        

Generated by: LCOV version 2.0-1