LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_all_gather - coll_all_gather_ring_for_910_93_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 56.3 % 311 175
Test Date: 2026-08-18 17:47:01 Functions: 77.8 % 18 14

            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_all_gather_ring_for_910_93_executor.h"
      12              : 
      13              : namespace hccl {
      14            2 : CollAllGatherRingFor91093Executor::CollAllGatherRingFor91093Executor(
      15            2 :     const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
      16            2 :     : CollAllGatherExecutor(dispatcher, topoMatcher)
      17              : {
      18            2 :     DMAReduceFlag_ = workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE;
      19              :     desc_.level1SupportedAlgos
      20            2 :         = {AlgTypeLevel1::ALG_LEVEL1_NHR, AlgTypeLevel1::ALG_LEVEL1_NB, AlgTypeLevel1::ALG_LEVEL1_RING,
      21            2 :            AlgTypeLevel1::ALG_LEVEL1_AHC, AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE};
      22              :     desc_.level2SupportedAlgos
      23            2 :         = {AlgTypeLevel2::ALG_LEVEL2_NHR, AlgTypeLevel2::ALG_LEVEL2_NB, AlgTypeLevel2::ALG_LEVEL2_RING};
      24            2 : }
      25              : 
      26            2 : HcclResult CollAllGatherRingFor91093Executor::CalcStreamNum(u32& streamNum)
      27              : {
      28            2 :     u32 totalStreamNum
      29            2 :         = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING ? LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE :
      30              :                                                              LEVEL0_PLANE_NUM_IN_NPRING_SINGLE);
      31            2 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      32            2 :         totalStreamNum *= STREAM_NUM_FOR_DMAREDUCE_ONE_RING;
      33              :     }
      34              : 
      35            2 :     streamNum = totalStreamNum - 1;
      36            2 :     HCCL_INFO("[CollAllGatherRingFor91093Executor][CalcStreamNum] tag[%s] streamNum_[%u]", tag_.c_str(), streamNum);
      37            2 :     return HCCL_SUCCESS;
      38              : }
      39              : 
      40            2 : HcclResult CollAllGatherRingFor91093Executor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      41              : {
      42            2 :     TransportMemType inputType = TransportMemType::RESERVED;
      43            2 :     TransportMemType outputType = TransportMemType::RESERVED;
      44            2 :     CHK_RET(CalcTransportMemType(inputType, outputType));
      45            2 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
      46            2 :     CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
      47            2 :     CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
      48            2 :     return HCCL_SUCCESS;
      49              : }
      50              : 
      51              : HcclResult
      52            2 : CollAllGatherRingFor91093Executor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
      53              : {
      54            2 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      55            2 :         inputType = TransportMemType::CCL_INPUT;
      56            2 :         outputType = TransportMemType::CCL_OUTPUT;
      57              :     } else {
      58            0 :         inputType = TransportMemType::PARAM_INPUT;
      59            0 :         outputType = TransportMemType::PARAM_OUTPUT;
      60              :     }
      61            2 :     HCCL_INFO(
      62              :         "[CollAllGatherRingFor91093Executor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]", tag_.c_str(),
      63              :         inputType, outputType);
      64            2 :     return HCCL_SUCCESS;
      65              : }
      66              : 
      67            2 : HcclResult CollAllGatherRingFor91093Executor::CalcLevel0CommInfo(
      68              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      69              : {
      70            2 :     CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
      71            2 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
      72            2 :     return HCCL_SUCCESS;
      73            2 : }
      74              : 
      75            2 : HcclResult CollAllGatherRingFor91093Executor::CalcLevel2CommInfo(
      76              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      77              : {
      78            2 :     if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
      79            2 :         || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE) {
      80            0 :         HCCL_INFO("[CollAllGatherRingFor91093Executor][CalcLevel2CommInfo] select AHC bypass level2 comm calculate");
      81            0 :         return HCCL_SUCCESS;
      82              :     }
      83              : 
      84            2 :     CommParaInfo commParaLevel2(COMM_LEVEL2, CommType::COMM_TAG_MAX);
      85            2 :     HCCL_DEBUG("[CollAllGatherRingFor91093Executor][CalcLevel2CommInfo]Level2CommInfo start set");
      86            2 :     if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
      87            0 :         commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
      88            0 :         HCCL_INFO("[%s]Calc NHRCommInfo.", __func__);
      89            2 :     } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
      90            0 :         commParaLevel2.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
      91            0 :         HCCL_INFO("[%s]Calc NBCommInfo.", __func__);
      92              :     } else {
      93            2 :         commParaLevel2.commType = CommType::COMM_TAG_RING_INNER;
      94            2 :         HCCL_INFO("[%s]Calc RingCommInfo.", __func__);
      95              :     }
      96            2 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel2, opTransport[COMM_LEVEL2], inputType, outputType));
      97            2 :     return HCCL_SUCCESS;
      98            2 : }
      99              : 
     100            2 : u64 CollAllGatherRingFor91093Executor::CalcLoopMaxCount(const u64 cclBuffSize, const u32 unitSize)
     101              : {
     102            2 :     u64 maxCountPerLoop = cclBuffSize / topoAttr_.userRankSize / HCCL_MIN_SLICE_ALIGN * HCCL_MIN_SLICE_ALIGN / unitSize;
     103            2 :     return maxCountPerLoop;
     104              : }
     105              : 
     106            0 : HcclResult CollAllGatherRingFor91093Executor::RunIntraSeverAllGather(
     107              :     const std::string& tag, DeviceMem& inputMem, DeviceMem& outputMem, const u64 count, const HcclDataType& dataType,
     108              :     const std::vector<std::vector<Slice>>& multRingsSliceZero, const Stream& stream, s32 profStage,
     109              :     const u64 baseOffset, const HcomCollOpInfo* opInfo, const std::vector<std::vector<Slice>>& multRingsUserMemSlice)
     110              : {
     111            0 :     CHK_RET(MultiRingAllGather(
     112              :         tag, inputMem, outputMem, count, dataType, multRingsSliceZero, stream, profStage, baseOffset, opInfo,
     113              :         multRingsUserMemSlice, logicalLevel0plane_));
     114            0 :     return HCCL_SUCCESS;
     115              : }
     116              : 
     117           18 : u64 CollAllGatherRingFor91093Executor::CalcDstMemOffset(
     118              :     [[maybe_unused]] const OpParam& param, [[maybe_unused]] u32 perDataSize, u64 inputMemSize) const
     119              : {
     120           18 :     return topoAttr_.userRank * inputMemSize;
     121              : }
     122              : 
     123           18 : HcomCollOpInfo CollAllGatherRingFor91093Executor::GetHcomCollOpInfo(const OpParam& param, const ExecMem& execMem) const
     124              : {
     125           18 :     HcomCollOpInfo opInfo
     126           18 :         = {"", execMem.inputPtr,     execMem.outputPtr,        param.DataDes.count, param.DataDes.dataType,
     127           18 :            0,  HCCL_REDUCE_RESERVED, param.DataDes.strideCount};
     128           18 :     if (!DMAReduceFlag_ && (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING)) {
     129            0 :         opInfo.inputAddr = execMem.inputMem.ptr();
     130            0 :         opInfo.outputAddr = execMem.outputMem.ptr();
     131              :     }
     132           18 :     return opInfo;
     133              : }
     134              : 
     135           18 : HcclResult CollAllGatherRingFor91093Executor::PrepareSlicesL0(
     136              :     std::vector<std::vector<Slice>>& multRingsSlice, const OpParam& param, const SubCommInfo& level2CommInfo,
     137              :     const SubCommInfo& level1CommInfo, const SubCommInfo& level0CommInfo, [[maybe_unused]] u32 perDataSize,
     138              :     u64 inputMemSize)
     139              : {
     140           18 :     const u32 level0RankSize = level0CommInfo.localRankSize;
     141           18 :     const u32 level1RankSize = level1CommInfo.localRankSize;
     142           18 :     const u32 level2RankSize = level2CommInfo.localRankSize;
     143              : 
     144           18 :     std::vector<Slice> dataSegsSlice;
     145           18 :     CHK_RET(PrepareAllgatherSlice(level0RankSize, inputMemSize, dataSegsSlice));
     146              : 
     147              :     // 多环数据切分
     148           18 :     std::vector<std::vector<Slice>> multRingsSliceZero; // 数据基于该rank上环0的偏移
     149           18 :     bool ARSFlag = topoMatcher_->GetARSFlag();
     150           18 :     bool ARSDoubleRing = (ARSFlag && (level0RankSize > FACTOR_TWO) && topoAttr_.isARSDoubleRing);
     151              : 
     152           18 :     if (ARSDoubleRing) {
     153            0 :         std::vector<u32> mockNicList;
     154            0 :         mockNicList.reserve(level0RankSize);
     155            0 :         for (u32 rankIndex = 0; rankIndex < level0RankSize; rankIndex++) {
     156            0 :             mockNicList.push_back(rankIndex);
     157              :         }
     158            0 :         multRingsSliceZero = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, mockNicList);
     159            0 :     } else if (
     160           18 :         topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING
     161           18 :         && !IsSupportUnifiedMarch(param, topoType_, topoAttr_.serverNum, topoAttr_.superPodNum)) {
     162           18 :         multRingsSliceZero = PrepareMultiRingSlice(dataSegsSlice, param.tag, false, topoAttr_.nicList);
     163              :     } else {
     164            0 :         multRingsSliceZero.push_back(dataSegsSlice);
     165              :     }
     166           54 :     for (u32 ringIndex = 0; ringIndex < multRingsSliceZero.size(); ringIndex++) {
     167           36 :         std::vector<Slice> level2DataSlice;
     168           36 :         CHK_RET(CalculateLevel2AllgatherSlice(
     169              :             inputMemSize, level0RankSize, level1RankSize, level2RankSize, multRingsSliceZero, level2DataSlice,
     170              :             ringIndex));
     171           36 :         multRingsSlice.push_back(level2DataSlice);
     172           36 :     }
     173              : 
     174           18 :     return HCCL_SUCCESS;
     175           18 : }
     176              : 
     177           18 : std::vector<Slice> CollAllGatherRingFor91093Executor::PrepareSlicesL1(
     178              :     [[maybe_unused]] const OpParam& param, const SubCommInfo& level2CommInfo, const SubCommInfo& level1CommInfo,
     179              :     const SubCommInfo& level0CommInfo, [[maybe_unused]] u32 perDataSize, u64 inputMemSize) const
     180              : {
     181           18 :     const u32 level0RankSize = level0CommInfo.localRankSize;
     182           18 :     const u32 level0ServerIndex = level0CommInfo.localRank;
     183           18 :     const u32 level1RankSize = level1CommInfo.localRankSize;
     184           18 :     const u32 level2RankSize = level2CommInfo.localRankSize;
     185           18 :     std::vector<Slice> level1DataSegsSlice;
     186           54 :     for (u32 j = 0; j < level1RankSize; j++) {
     187           72 :         for (u32 i = 0; i < level2RankSize; i++) {
     188           36 :             Slice level1Slice;
     189           36 :             level1Slice.size = inputMemSize;
     190              :             level1Slice.offset
     191           36 :                 = inputMemSize * (i * level1RankSize * level0RankSize + j * level0RankSize + level0ServerIndex);
     192              : 
     193           36 :             HCCL_DEBUG(
     194              :                 "[CollAllGatherRingFor91093Executor][PrepareSlicesL1] rank[%u], level1index[%u], level2index[%u], "
     195              :                 "slices.offset=%llu, slices.size=%llu",
     196              :                 level0CommInfo.localRank, j, i, level1Slice.offset, level1Slice.size);
     197              : 
     198           36 :             level1DataSegsSlice.push_back(level1Slice);
     199              :         }
     200              :     }
     201           18 :     return level1DataSegsSlice;
     202            0 : }
     203              : 
     204            0 : std::vector<Slice> CollAllGatherRingFor91093Executor::PrepareSlicesL2(
     205              :     [[maybe_unused]] const OpParam& param, const SubCommInfo& level2CommInfo, const SubCommInfo& level1CommInfo,
     206              :     const SubCommInfo& level0CommInfo, [[maybe_unused]] u32 perDataSize, u64 inputMemSize) const
     207              : {
     208            0 :     const u32 level0RankSize = level0CommInfo.localRankSize;
     209            0 :     const u32 level0ServerIndex = level0CommInfo.localRank;
     210            0 :     const u32 level1RankSize = level1CommInfo.localRankSize;
     211            0 :     const u32 level1ServerIndex = level1CommInfo.localRank;
     212            0 :     const u32 level2RankSize = level2CommInfo.localRankSize;
     213            0 :     std::vector<Slice> level2DataSegsSlice;
     214            0 :     for (u32 i = 0; i < level2RankSize; i++) {
     215            0 :         Slice sliceTemp;
     216            0 :         sliceTemp.size = inputMemSize;
     217              :         sliceTemp.offset
     218            0 :             = inputMemSize
     219            0 :               * (i * level1RankSize * level0RankSize + level1ServerIndex * level0RankSize + level0ServerIndex);
     220            0 :         level2DataSegsSlice.push_back(sliceTemp);
     221              :     }
     222            0 :     return level2DataSegsSlice;
     223            0 : }
     224              : 
     225           18 : HcclResult CollAllGatherRingFor91093Executor::PrepareUserMemSlices(
     226              :     std::vector<std::vector<Slice>>& userMemSlices, const std::vector<std::vector<Slice>>& multRingsSlice,
     227              :     const OpParam& param, [[maybe_unused]] const SubCommInfo& level2CommInfo,
     228              :     [[maybe_unused]] const SubCommInfo& level1CommInfo, [[maybe_unused]] const SubCommInfo& level0CommInfo,
     229              :     u32 perDataSize, u64 inputMemSize)
     230              : {
     231           18 :     CHK_PRT_RET(
     232              :         0 < param.DataDes.strideCount && param.DataDes.strideCount < param.DataDes.count,
     233              :         HCCL_ERROR(
     234              :             "[CollAllGatherRingFor91093Executor][KernelRun]strideCount[%llu] is smaller than opCount[%llu]",
     235              :             param.DataDes.strideCount, param.DataDes.count),
     236              :         HCCL_E_PARA);
     237           18 :     HCCL_DEBUG(
     238              :         "[CollAllGatherRingFor91093Executor][KernelRun]strideCount[%llu], opCount[%llu]", param.DataDes.strideCount,
     239              :         param.DataDes.count);
     240              : 
     241           18 :     if (!DMAReduceFlag_) {
     242            0 :         userMemSlices = multRingsSlice;
     243              :         // 图模式,根据strideCount更新slice的offset
     244            0 :         if (param.DataDes.strideCount != 0) {
     245            0 :             CHK_RET(UpdateOffsetBasedOnStrideCount(param, userMemSlices));
     246              :         }
     247              :     } else {
     248           54 :         for (u32 ringIndex = 0; ringIndex < multRingsSlice.size(); ringIndex++) {
     249           36 :             std::vector<Slice> userMemSlice;
     250          180 :             for (const auto& cclSlice : multRingsSlice[ringIndex]) {
     251          144 :                 Slice tmpSlice;
     252          144 :                 u64 count = (param.DataDes.strideCount == 0) ? param.DataDes.count : param.DataDes.strideCount;
     253          144 :                 tmpSlice.size = cclSlice.size;
     254              :                 tmpSlice.offset
     255          144 :                     = (cclSlice.offset / inputMemSize) * count * perDataSize + multRingsSlice[ringIndex][0].offset;
     256          144 :                 userMemSlice.push_back(tmpSlice);
     257          144 :                 HCCL_DEBUG(
     258              :                     "rank[%u], ringIndex[%u], tmpSlice.offset=[%llu], size=[%llu]", topoAttr_.userRank, ringIndex,
     259              :                     tmpSlice.offset, tmpSlice.size);
     260              :             }
     261           36 :             userMemSlices.push_back(userMemSlice);
     262           36 :         }
     263              :     }
     264           18 :     return HCCL_SUCCESS;
     265              : }
     266              : 
     267           18 : HcclResult CollAllGatherRingFor91093Executor::GetLevelCommInfo()
     268              : {
     269           18 :     logicalLevel0plane_ = COMM_LEVEL0;
     270           18 :     CHK_RET(CheckCommSize(logicalLevel0plane_, COMM_INDEX_0 + 1));
     271           18 :     logicalLevel0CommInfo_ = GetSubCommInfo(logicalLevel0plane_, COMM_INDEX_0);
     272           18 :     u32 commIndex = logicalLevel0CommInfo_.localRank;
     273           18 :     bool isSelectAHC
     274           18 :         = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
     275           18 :            || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     276           18 :     logicalLevel1plane_ = isSelectAHC ? COMM_LEVEL1_AHC : COMM_LEVEL1;
     277           18 :     CHK_RET(CheckCommSize(logicalLevel1plane_, commIndex + 1));
     278           18 :     logicalLevel1CommInfo_ = GetSubCommInfo(logicalLevel1plane_, commIndex);
     279           18 :     return HCCL_SUCCESS;
     280              : }
     281              : 
     282           18 : HcclResult CollAllGatherRingFor91093Executor::KernelRun(const OpParam& param, ExecMem& execMem)
     283              : {
     284           18 :     HCCL_CONFIG_INFO(
     285              :         HCCL_ALG, "[%s] The AllGatherRingExecutor starts, topoType_[%u], agv[%u]", __func__, topoType_, isAllGatherV_);
     286           18 :     CHK_RET(GetLevelCommInfo()); // 设置逻辑通信域
     287           18 :     CHK_RET(ActiveSlaveStreams(param.stream));
     288           18 :     const HcclDataType dataType = param.GetDataType();
     289           18 :     u32 perDataSize = 0;
     290           18 :     CHK_RET(SalGetDataTypeSize(dataType, perDataSize));
     291           18 :     CHK_PRT_RET(
     292              :         perDataSize == 0,
     293              :         HCCL_ERROR(
     294              :             "[CollAllGatherRingFor91093Executor][KernelRun]errNo[0x%016llx] datatype[%s] is invalid",
     295              :             HCCL_ERROR_CODE(HCCL_E_PARA), GetDataTypeEnumStr(dataType).c_str()),
     296              :         HCCL_E_PARA);
     297              : 
     298           18 :     bool isSelectAHC
     299           18 :         = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
     300           18 :            || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     301              : 
     302           18 :     u32 level1RankSize = logicalLevel1CommInfo_.localRankSize;
     303              : 
     304           18 :     SubCommInfo level2CommInfo;
     305           18 :     if (isSelectAHC) {
     306            0 :         level2CommInfo = logicalLevel1CommInfo_;
     307            0 :         level2CommInfo.localRankSize = 1; // AHC bypass level2
     308              :     } else {
     309           18 :         CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
     310           18 :         level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     311              :     }
     312           18 :     const u32 level2RankSize = level2CommInfo.localRankSize;
     313              : 
     314              :     //  第一步,将数据从input内存拷贝到output内存的对应位置
     315           18 :     u64 inputMemSize = execMem.inputMem.size();
     316           18 :     u64 dstMemOffset = CalcDstMemOffset(param, perDataSize, inputMemSize);
     317           18 :     DeviceMem dstMem = execMem.outputMem.range(dstMemOffset, inputMemSize);
     318           18 :     CHK_SMART_PTR_NULL(dstMem);
     319              : 
     320           18 :     HcomCollOpInfo opInfo = GetHcomCollOpInfo(param, execMem);
     321           18 :     HcomCollOpInfo* opInfoPtr
     322           18 :         = (DMAReduceFlag_ || (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING)) ? &opInfo : nullptr;
     323              : 
     324              :     // 图模式opinfo不为空,但需要将数据从ccl input拷贝到ccl output上
     325           18 :     HcclResult ret = HCCL_SUCCESS;
     326           18 :     if (!DMAReduceFlag_) {
     327            0 :         ret = HcclD2DMemcpyAsync(dispatcher_, dstMem, execMem.inputMem, const_cast<Stream&>(param.stream));
     328            0 :         CHK_PRT_RET(
     329              :             ret != HCCL_SUCCESS,
     330              :             HCCL_ERROR(
     331              :                 "[CollAllGatherRingFor91093Executor][KernelRun]AllGather double "
     332              :                 "ring memcpy Failed, Offset[%llu], Size[%llu]",
     333              :                 dstMemOffset, inputMemSize),
     334              :             ret);
     335              :     } else {
     336              :         // 先做server间算法,带有消减拷贝场景数据需要从user input取,拷贝到ccl output上
     337           18 :         if (level1RankSize > 1 || level2RankSize > 1) {
     338           18 :             DeviceMem srcMem = DeviceMem::create(static_cast<u8*>(execMem.inputPtr), inputMemSize);
     339           18 :             ret = HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, const_cast<Stream&>(param.stream));
     340           18 :             CHK_PRT_RET(
     341              :                 ret != HCCL_SUCCESS,
     342              :                 HCCL_ERROR(
     343              :                     "[CollAllGatherRingFor91093Executor][KernelRun]AllGather double "
     344              :                     "ring user memcpy Failed, Offset[%llu], Size[%llu]",
     345              :                     dstMemOffset, inputMemSize),
     346              :                 ret);
     347           18 :         }
     348              :     }
     349           18 :     if (level2RankSize > 1) {
     350            0 :         std::unique_ptr<AlgTemplateBase> level2AGExecutor;
     351            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     352              :             level2AGExecutor
     353            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
     354            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL2", __func__);
     355            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     356              :             level2AGExecutor
     357            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
     358            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL2", __func__);
     359              :         } else {
     360              :             level2AGExecutor
     361            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     362            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL2", __func__);
     363              :         }
     364            0 :         CHK_SMART_PTR_NULL(level2AGExecutor);
     365              : 
     366            0 :         std::vector<Slice> level2DataSegsSlice = PrepareSlicesL2(
     367            0 :             param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_, perDataSize, inputMemSize);
     368            0 :         CHK_RET(level2AGExecutor->Prepare(
     369              :             execMem.outputMem, execMem.outputMem, execMem.inputMem, execMem.count, dataType, param.stream,
     370              :             HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, level2DataSegsSlice, 0));
     371              : 
     372            0 :         CHK_RET(level2AGExecutor->RegisterProfiler(
     373              :             (level2RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank, PROF_STAGE_0,
     374              :             HCCL_EXEC_STEP_NOT_SET, param.stream));
     375              : 
     376            0 :         CHK_RET(RunTemplate(level2AGExecutor, level2CommInfo));
     377            0 :         HCCL_INFO(
     378              :             "AllGather ring [superpod] level2 AllGather run successtopoType_[%u], agv[%u]", topoType_, isAllGatherV_);
     379            0 :     }
     380           18 :     if (level1RankSize > 1) {
     381              :         // 计算slice, 不同超节点相同slice
     382           18 :         std::vector<Slice> level1DataSegsSlice = PrepareSlicesL1(
     383           18 :             param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_, perDataSize, inputMemSize);
     384              : 
     385           18 :         std::unique_ptr<AlgTemplateBase> level1AGExecutor;
     386           18 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     387              :             level1AGExecutor
     388           18 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     389           18 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
     390            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     391              :             level1AGExecutor
     392            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
     393            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
     394            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     395              :             level1AGExecutor
     396            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
     397            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
     398            0 :         } else if (isSelectAHC) {
     399              :             // 获取通信域分组信息
     400            0 :             std::vector<std::vector<std::vector<u32>>> globalSubGroups;
     401            0 :             std::map<AHCConcOpType, TemplateType> ahcAlgOption;
     402            0 :             CHK_RET(topoMatcher_->GetGlobalSubGroups(logicalLevel1plane_, globalSubGroups));
     403            0 :             topoMatcher_->GetAHCAlgOption(ahcAlgOption);
     404            0 :             if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
     405            0 :                 level1AGExecutor = AlgTemplateRegistry::Instance().GetAlgTemplate(
     406            0 :                     TemplateType::TEMPLATE_ALL_GATHER_AHC, dispatcher_);
     407            0 :                 HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC in COMM_LEVEL1", __func__);
     408              :             } else {
     409            0 :                 level1AGExecutor = AlgTemplateRegistry::Instance().GetAlgTemplate(
     410            0 :                     TemplateType::TEMPLATE_ALL_GATHER_AHC_BROKE, dispatcher_);
     411            0 :                 HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC_BROKE in COMM_LEVEL1", __func__);
     412              :             }
     413            0 :             CHK_SMART_PTR_NULL(level1AGExecutor);
     414            0 :             CHK_RET(level1AGExecutor->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
     415            0 :         } else {
     416            0 :             HCCL_ERROR("AllGather ring: unsupported algtype [%s].", AlgTypeToStr(algType_).c_str());
     417            0 :             return HCCL_E_NOT_SUPPORT;
     418              :         }
     419           18 :         CHK_SMART_PTR_NULL(level1AGExecutor);
     420           54 :         CHK_RET(level1AGExecutor->Prepare(
     421              :             execMem.outputMem, execMem.outputMem, execMem.inputMem, execMem.count, dataType, param.stream,
     422              :             HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, level1DataSegsSlice, 0));
     423              : 
     424           18 :         CHK_RET(level1AGExecutor->RegisterProfiler(
     425              :             (level1RankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2CommInfo.localRank, PROF_STAGE_1,
     426              :             HCCL_EXEC_STEP_NOT_SET, param.stream));
     427              : 
     428           18 :         CHK_RET(RunTemplate(level1AGExecutor, logicalLevel1CommInfo_));
     429           18 :         HCCL_INFO(
     430              :             "AllGather ring [superpod] level1 AllGather run successtopoType_[%u], agv[%u]", topoType_, isAllGatherV_);
     431           18 :     }
     432              :     // 节点内做AllGather ring
     433           18 :     std::vector<std::vector<Slice>> multRingsSlice;
     434           18 :     CHK_RET(PrepareSlicesL0(
     435              :         multRingsSlice, param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_, perDataSize,
     436              :         inputMemSize));
     437              : 
     438           18 :     std::vector<std::vector<Slice>> multRingsUserMemSlice;
     439           18 :     CHK_RET(PrepareUserMemSlices(
     440              :         multRingsUserMemSlice, multRingsSlice, param, level2CommInfo, logicalLevel1CommInfo_, logicalLevel0CommInfo_,
     441              :         perDataSize, inputMemSize));
     442              : 
     443           18 :     if (DMAReduceFlag_ && (level1RankSize > 1 || level2RankSize > 1)) {
     444              :         // allgather输入放在CCL buffer上,通过设置nullptr指示要从CCL buffer获取输入
     445           18 :         opInfo.inputAddr = nullptr;
     446              :     }
     447           18 :     CHK_RET(RunIntraSeverAllGather(
     448              :         param.tag, execMem.inputMem, execMem.outputMem, execMem.count, dataType, multRingsSlice, param.stream,
     449              :         PROF_STAGE_2, 0, opInfoPtr, multRingsUserMemSlice));
     450           18 :     HCCL_INFO("AllGather ring run success. topoType_[%u], agv[%u]", topoType_, isAllGatherV_);
     451           18 :     return HCCL_SUCCESS;
     452           18 : }
     453              : 
     454            0 : HcclResult CollAllGatherRingFor91093Executor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
     455              : {
     456            0 :     HCCL_INFO("[CollAllGatherRingFor91093Executor][Getlevel1CommRank] Entry Getlevel1CommRank.");
     457            0 :     bool isSelectAHC
     458            0 :         = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
     459            0 :            || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     460            0 :     if (isSelectAHC) {
     461            0 :         SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
     462            0 :         u32 level0ServerIndex = level0CommInfo.localRank;
     463              : 
     464            0 :         CommPlane commPlaneLevel1 = COMM_LEVEL1;
     465            0 :         CHK_RET(CheckCommSize(commPlaneLevel1, level0ServerIndex + 1));
     466            0 :         level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0ServerIndex);
     467            0 :         u32 level1RankSize = level1CommInfo.localRankSize;
     468            0 :         HCCL_INFO("Getlevel1CommRank. level1RankSize[%u]", level1RankSize);
     469            0 :         return HCCL_SUCCESS;
     470            0 :     }
     471            0 :     if (CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1) != HCCL_SUCCESS) {
     472            0 :         HCCL_INFO("[nslbdp] Getlevel1CommRank size not match.");
     473            0 :         return HCCL_E_UNAVAIL;
     474              :     }
     475            0 :     level1CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     476              : 
     477            0 :     return HCCL_SUCCESS;
     478              : }
     479              : 
     480              : HcclResult
     481            0 : CollAllGatherRingFor91093Executor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
     482              : {
     483            0 :     HCCL_INFO("[nslbdp] Entry SelectTempAlg, level1RankSize = [%u].", level1RankSize);
     484            0 :     bool isSelectAHC
     485            0 :         = (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC
     486            0 :            || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE);
     487            0 :     if (isSelectAHC) {
     488            0 :         CommPlane commPlaneLevel1 = COMM_LEVEL1;
     489              :         // 获取通信域分组信息
     490            0 :         std::vector<std::vector<std::vector<u32>>> globalSubGroups;
     491            0 :         std::map<AHCConcOpType, TemplateType> ahcAlgOption;
     492            0 :         CHK_RET(topoMatcher_->GetGlobalSubGroups(commPlaneLevel1, globalSubGroups));
     493            0 :         topoMatcher_->GetAHCAlgOption(ahcAlgOption);
     494            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
     495              :             level1TempAlg
     496            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_AHC, dispatcher_);
     497            0 :             HCCL_INFO("allgather comm: using ahc algo inter-server.");
     498              :         } else {
     499            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     500            0 :                 TemplateType::TEMPLATE_ALL_GATHER_AHC_BROKE, dispatcher_);
     501            0 :             HCCL_INFO("allgather comm: using ahc-broke algo inter-server.");
     502              :         }
     503            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     504            0 :         CHK_RET(level1TempAlg->Prepare(NSLBDP_MIN_COUNT, globalSubGroups, ahcAlgOption));
     505            0 :         return HCCL_SUCCESS;
     506            0 :     }
     507            0 :     if (level1RankSize > 1) {
     508            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     509              :             level1TempAlg
     510            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
     511            0 :             HCCL_INFO("AllGather ring: using nonuniform-bruck algo inter-superPod.");
     512            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     513              :             level1TempAlg
     514            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
     515            0 :             HCCL_INFO("AllGather ring: using nonuniform-hierarchical-ring algo inter-superPod.");
     516              :         } else {
     517              :             level1TempAlg
     518            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     519            0 :             HCCL_INFO("AllGather ring: using ring algo inter-superPod.");
     520              :         }
     521            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     522            0 :         return HCCL_SUCCESS;
     523              :     }
     524            0 :     return HCCL_E_UNAVAIL;
     525              : }
     526              : 
     527              : REGISTER_EXEC("AllGatherRingFor91093Executor", AllGatherRingFor91093, CollAllGatherRingFor91093Executor);
     528              : 
     529              : } // namespace hccl
        

Generated by: LCOV version 2.0-1