LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_all_gather - coll_all_gather_ring_zerocopy_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 0.0 % 180 0
Test Date: 2026-08-04 10:52:23 Functions: 0.0 % 13 0

            Line data    Source code
       1              : /**
       2              :  * Copyright (c) 2025 Huawei Technologies Co., Ltd.
       3              :  * This program is free software, you can redistribute it and/or modify it under the terms and conditions of
       4              :  * CANN Open Software License Agreement Version 2.0 (the "License").
       5              :  * Please refer to the License for details. You may not use this file except in compliance with the License.
       6              :  * THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
       7              :  * INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
       8              :  * See LICENSE in the root of the software repository for the full text of the License.
       9              :  */
      10              : 
      11              : #include "coll_all_gather_ring_zerocopy_executor.h"
      12              : 
      13              : namespace hccl {
      14            0 : CollAllGatherRingZerocopyExecutor::CollAllGatherRingZerocopyExecutor(const HcclDispatcher dispatcher,
      15            0 :                                                                    std::unique_ptr<TopoMatcher> &topoMatcher)
      16            0 :     : CollAllGatherExecutor(dispatcher, topoMatcher)
      17              : {
      18            0 :     DMAReduceFlag_ = true;      // 设为true,以禁用RunLoop中的本地拷贝
      19            0 :     desc_.isZeroCopy = true;
      20            0 :     desc_.level1SupportedAlgos = {
      21              :         AlgTypeLevel1::ALG_LEVEL1_NHR,
      22              :         AlgTypeLevel1::ALG_LEVEL1_NB,
      23              :         AlgTypeLevel1::ALG_LEVEL1_RING,
      24              :         AlgTypeLevel1::ALG_LEVEL1_AHC,
      25              :         AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE
      26            0 :     };
      27            0 :     desc_.level2SupportedAlgos = {
      28              :         AlgTypeLevel2::ALG_LEVEL2_NHR,
      29              :         AlgTypeLevel2::ALG_LEVEL2_NB,
      30              :         AlgTypeLevel2::ALG_LEVEL2_RING
      31            0 :     };
      32            0 : }
      33              : 
      34            0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcStreamNum(u32& streamNum)
      35              : {
      36            0 :     u32 totalStreamNum = (topoType_ == TopoType::TOPO_TYPE_NP_DOUBLE_RING) ?
      37              :                          (LEVEL0_PLANE_NUM_IN_NPRING_DOUBLE + 1) : LEVEL0_PLANE_NUM_IN_NPRING_SINGLE;
      38            0 :     streamNum = totalStreamNum - 1;
      39            0 :     HCCL_INFO("[%s] tag[%s] streamNum_[%u]", __func__, tag_.c_str(), streamNum);
      40            0 :     return HCCL_SUCCESS;
      41              : }
      42              : 
      43            0 : void CollAllGatherRingZerocopyExecutor::ParseParam(const OpParam& param)
      44              : {
      45            0 :     tag_ = param.tag;
      46            0 : }
      47              : 
      48            0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      49              : {
      50            0 :     TransportMemType inputType = TransportMemType::RESERVED;
      51            0 :     TransportMemType outputType = TransportMemType::RESERVED;
      52            0 :     CHK_RET(CalcTransportMemType(inputType, outputType));
      53            0 :     CHK_RET(CalcLevel0CommInfo(inputType, outputType, opTransport));
      54            0 :     CHK_RET(CalcLevel1CommInfo(inputType, outputType, opTransport));
      55            0 :     CHK_RET(CalcLevel2CommInfo(inputType, outputType, opTransport));
      56            0 :     return HCCL_SUCCESS;
      57              : }
      58              : 
      59            0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcTransportMemType(TransportMemType &inputType,
      60              :     TransportMemType &outputType)
      61              : {
      62            0 :     inputType = TransportMemType::CCL_INPUT;
      63            0 :     outputType = TransportMemType::CCL_OUTPUT;
      64            0 :     return HCCL_SUCCESS;
      65              : }
      66              : 
      67            0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcLevel0CommInfo(TransportMemType inputType,
      68              :     TransportMemType outputType,
      69              :     std::vector<LevelNSubCommTransport>& opTransport)
      70              : {
      71            0 :     CommParaInfo commParaLevel0(COMM_LEVEL0, CommType::COMM_TAG_RING_INNER);
      72            0 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaLevel0, opTransport[COMM_LEVEL0], inputType, outputType));
      73            0 :     LevelNSubCommTransport &commTransportLevel0 = opTransport[COMM_LEVEL0];
      74            0 :     for (u32 subCommIndex = 0; subCommIndex < commTransportLevel0.size(); subCommIndex++) {
      75            0 :         commTransportLevel0[subCommIndex].isZeroCopy = true;
      76              :     }
      77            0 :     return HCCL_SUCCESS;
      78            0 : }
      79              : 
      80            0 : u64 CollAllGatherRingZerocopyExecutor::CalcLoopMaxCount(const u64 cclBuffSize, const u32 unitSize)
      81              : {
      82            0 :     u64 maxCountPerLoop = cclBuffSize / topoAttr_.serverNum / HCCL_MIN_SLICE_ALIGN
      83            0 :         * HCCL_MIN_SLICE_ALIGN / unitSize;
      84            0 :     return maxCountPerLoop;
      85              : }
      86              : 
      87            0 : HcclResult CollAllGatherRingZerocopyExecutor::SemiRingAllGather(
      88              :     const std::string &tag, DeviceMem &inputMem, DeviceMem &outputMem,
      89              :     const u64 count, const HcclDataType &dataType, const std::vector<std::vector<Slice>> &multRingsSliceZero,
      90              :     const Stream &stream, s32 profStage, const u64 baseOffset, const HcomCollOpInfo *opInfo,
      91              :     const std::vector<std::vector<Slice>> &multRingsUserMemSlice)
      92              : {    
      93              :     (void) multRingsSliceZero;
      94              :     (void) tag;
      95              :     (void) baseOffset;
      96              :     (void) opInfo;
      97            0 :     CHK_RET(CheckCommSize(COMM_LEVEL0, COMM_INDEX_0 + 1));
      98            0 :     SubCommInfo level0CommInfo = GetSubCommInfo(COMM_LEVEL0, COMM_INDEX_0);
      99              : 
     100              :     // 执行
     101            0 :     std::unique_ptr<AlgTemplateBase> executor = AlgTemplateRegistry::Instance().GetAlgTemplate(
     102            0 :         TemplateType::TEMPLATE_ALL_GATHER_UNIFIED_MARCH, dispatcher_);
     103            0 :     HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_UNIFIED_MARCH in COMM_LEVEL0", __func__);
     104            0 :     CHK_SMART_PTR_NULL(executor);
     105              : 
     106            0 :     CHK_RET(executor->Prepare(stream, level0CommInfo, algResResp_->paramInputMem, algResResp_->paramOutputMem,
     107              :         inputMem, outputMem, count * SIZE_TABLE[dataType], algResResp_->slaveStreams, algResResp_->notifiesMain,
     108              :         algResResp_->notifiesAux, multRingsUserMemSlice));
     109            0 :     HcclResult ret = executor->RegisterProfiler(
     110              :         ((COMM_INDEX_0 + 1) << PROF_RINGINDEX_OFFSET_OF_PLANEID) +
     111            0 :         (level0CommInfo.localRankSize << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level0CommInfo.localRank,
     112              :         profStage, HCCL_EXEC_STEP_NOT_SET, stream);
     113            0 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     114              :         HCCL_ERROR("[CollAllGatherRingZerocopyExecutor][SemiRingAllGather]Double ring "
     115              :         "AllGather failed, return[%d]", ret), ret);
     116            0 :     CHK_RET(executor->RunAsync());
     117            0 :     return ret;
     118            0 : }
     119              : 
     120            0 : HcclResult CollAllGatherRingZerocopyExecutor::KernelRunIntraServerPost(const OpParam &param, ExecMem &execMem)
     121              : {
     122            0 :     bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
     123            0 :     CHK_RET(GetCommRankInfoNormal(level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, isAHCAlgo));
     124              :     
     125              :     // 计算slice信息
     126            0 :     std::vector<Slice> dataSegsSlice;
     127            0 :     CHK_RET(CalcLevel0DataSlices(param, execMem, dataSegsSlice));
     128              :     // 执行AllGather
     129            0 :     u64 level0Count = (dataSegsSlice.size() > level0RankSize_) ?     // 如果是非连续数据通信
     130            0 :                       (execMem.count) : (execMem.count * level1RankSize_ * level2RankSize_);
     131            0 :     std::vector<std::vector<Slice>> multRingsUserMemSlice = {dataSegsSlice};
     132            0 :     if (topoType_ == TopoType::TOPO_TYPE_NP_SINGLE_RING) {
     133            0 :         HCCL_INFO("[%s] single ring AllGather", __func__);
     134            0 :         CHK_RET(MultiRingAllGather(param.tag, execMem.inputMem, execMem.outputMem, level0Count, param.DataDes.dataType,
     135              :             multRingsUserMemSlice, param.stream, PROF_STAGE_0, 0, nullptr, multRingsUserMemSlice));
     136              :     } else {
     137            0 :         CHK_PRT_RET(topoType_ != TopoType::TOPO_TYPE_NP_DOUBLE_RING,
     138              :             HCCL_ERROR("[%s] unknown topoType: %u", __func__, topoType_), HCCL_E_NOT_SUPPORT);
     139            0 :         HCCL_INFO("[%s] semi ring AllGather", __func__);
     140            0 :         CHK_RET(SemiRingAllGather(param.tag, execMem.inputMem, execMem.outputMem, level0Count, param.DataDes.dataType,
     141              :             multRingsUserMemSlice, param.stream, PROF_STAGE_0, 0, nullptr,  multRingsUserMemSlice));
     142              :     }
     143              : 
     144            0 :     return HCCL_SUCCESS;
     145            0 : }
     146              : 
     147            0 : HcclResult CollAllGatherRingZerocopyExecutor::KernelRunInterServerPreProcess(const OpParam &param, const ExecMem &execMem)
     148              : {
     149              :     // 将数据从User Input拷到CCL Output
     150            0 :     u32 dataIndex = level1Rank_ * level2RankSize_ + level2Rank_;
     151            0 :     u64 curSize = execMem.inputMem.size();
     152            0 :     DeviceMem dstMem = execMem.outputMem.range(curSize * dataIndex, curSize);
     153            0 :     DeviceMem srcMem = DeviceMem::create(static_cast<u8 *>(execMem.inputPtr), curSize);
     154            0 :     Stream stream = param.stream;
     155            0 :     return HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream);
     156            0 : }
     157              : 
     158            0 : HcclResult CollAllGatherRingZerocopyExecutor::KernelRunInterServer(const OpParam &param, ExecMem &execMem)
     159              : {
     160            0 :     HCCL_CONFIG_INFO(HCCL_ALG,
     161              :         "[CollAllGatherRingZerocopyExecutor][KernelRunInterServer] The AllGatherDoubleRingExecutor starts");
     162            0 :     bool isAHCAlgo = algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC || algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC_BROKE;
     163            0 :     CHK_RET(GetCommRankInfoNormal(level0Rank_, level0RankSize_, level1Rank_, level1RankSize_, level2Rank_, level2RankSize_, isAHCAlgo));
     164              : 
     165              :     // 前处理
     166            0 :     CHK_RET(KernelRunInterServerPreProcess(param, execMem));
     167              : 
     168              :     // 计算slice
     169            0 :     std::vector<Slice> level1DataSegsSlice;
     170            0 :     CalcLevel1DataSlices(execMem.inputMem.size(), level1RankSize_, level2RankSize_, level1DataSegsSlice);
     171              : 
     172              :     // 超节点间通信
     173            0 :     if (level2RankSize_ > 1 && !isAHCAlgo) {
     174              :         // 获取对应算法的Template
     175            0 :         std::unique_ptr<AlgTemplateBase> level2AGTemplage;
     176            0 :         if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NB) {
     177            0 :             level2AGTemplage = AlgTemplateRegistry::Instance().GetAlgTemplate(
     178            0 :                 TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
     179            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL2", __func__);
     180            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_NHR) {
     181            0 :             level2AGTemplage = AlgTemplateRegistry::Instance().GetAlgTemplate(
     182            0 :                 TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
     183            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL2", __func__);
     184            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_RING){
     185            0 :             level2AGTemplage = AlgTemplateRegistry::Instance().GetAlgTemplate(
     186            0 :                 TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     187            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL2", __func__);
     188            0 :         } else if (algType_.algoLevel2 == AlgTypeLevel2::ALG_LEVEL2_HD) {
     189            0 :             level2AGTemplage = AlgTemplateRegistry::Instance().GetAlgTemplate(
     190            0 :                     TemplateType::TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING, dispatcher_);
     191            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING in COMM_LEVEL2", __func__);
     192              :         } else {
     193            0 :             HCCL_ERROR("AllGather ring: unsupported level2 algtype [%s]", AlgTypeToStr(algType_).c_str());
     194            0 :             return HCCL_E_NOT_SUPPORT;
     195              :         }
     196            0 :         CHK_SMART_PTR_NULL(level2AGTemplage);
     197              :         // 执行算法编排
     198            0 :         DeviceMem level2OutputMem = execMem.outputMem.range(level1DataSegsSlice[level1Rank_].offset,
     199            0 :                                                             level1DataSegsSlice[level1Rank_].size);
     200            0 :         CHK_RET(level2AGTemplage->Prepare(level2OutputMem, level2OutputMem, execMem.inputMem, execMem.count,
     201              :             param.DataDes.dataType, param.stream, HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID,
     202              :             std::vector<Slice>(0), level1DataSegsSlice[level1Rank_].offset));
     203            0 :         CHK_RET(level2AGTemplage->RegisterProfiler((
     204              :             level2RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level2Rank_,
     205              :             PROF_STAGE_0, HCCL_EXEC_STEP_NOT_SET, param.stream));
     206            0 :         CHK_RET(CheckCommSize(COMM_LEVEL2, COMM_INDEX_0 + 1));
     207            0 :         SubCommInfo level2CommInfo = GetSubCommInfo(COMM_LEVEL2, COMM_INDEX_0);
     208            0 :         CHK_RET(RunTemplate(level2AGTemplage, level2CommInfo));
     209            0 :         HCCL_INFO("AllGather double ring [superpod] level2 AllGather run success");
     210            0 :     }
     211              : 
     212              :     // 超节点内、节点间通信
     213            0 :     if (level1RankSize_ > 1) {
     214              :         // 获取对应算法的Template
     215            0 :         std::unique_ptr<AlgTemplateBase> level1AGTemplate;
     216            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_RING) {
     217            0 :             level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(
     218            0 :                 TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     219            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_LEVEL1", __func__);
     220            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     221            0 :             level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(
     222            0 :                 TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
     223            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_LEVEL1", __func__);
     224            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     225            0 :             level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(
     226            0 :                 TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
     227            0 :             HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_LEVEL1", __func__);
     228            0 :         } else if (isAHCAlgo) {
     229              :             // 获取通信域分组信息
     230            0 :             std::vector<std::vector<std::vector<u32>>> globalSubGroups;
     231            0 :             std::map<AHCConcOpType, TemplateType> ahcAlgOption;
     232            0 :             CHK_RET(topoMatcher_->GetGlobalSubGroups(COMM_LEVEL1_AHC, globalSubGroups));
     233            0 :             topoMatcher_->GetAHCAlgOption(ahcAlgOption);
     234            0 :             if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_AHC) {
     235            0 :                 level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_AHC, dispatcher_);
     236            0 :                 HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC in COMM_LEVEL1", __func__);
     237              :             } else {
     238            0 :                 level1AGTemplate = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_AHC_BROKE, dispatcher_);
     239            0 :                 HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_AHC_BROKE in COMM_LEVEL1", __func__);
     240              :             }
     241            0 :             CHK_SMART_PTR_NULL(level1AGTemplate);
     242            0 :             CHK_RET(level1AGTemplate->Prepare(execMem.count, globalSubGroups, ahcAlgOption));
     243            0 :         } else {
     244            0 :             HCCL_ERROR("AllGather ring: unsupported level1 algtype [%s]", AlgTypeToStr(algType_).c_str());
     245            0 :             return HCCL_E_NOT_SUPPORT;
     246              :         }
     247            0 :         CHK_SMART_PTR_NULL(level1AGTemplate);
     248              :         // 执行算法编排
     249            0 :         CHK_RET(level1AGTemplate->Prepare(execMem.outputMem, execMem.outputMem, execMem.inputMem, INVALID_U64,
     250              :             param.DataDes.dataType, param.stream,
     251              :             HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID, level1DataSegsSlice));
     252            0 :         CHK_RET(level1AGTemplate->RegisterProfiler((
     253              :             level1RankSize_ << PROF_RANKSIZE_OFFSET_OF_PLANEID) + level1Rank_,
     254              :             PROF_STAGE_1, HCCL_EXEC_STEP_NOT_SET, param.stream));
     255            0 :         CommPlane commPlaneLevel1 = isAHCAlgo ? COMM_LEVEL1_AHC : COMM_LEVEL1;
     256            0 :         CHK_RET(CheckCommSize(commPlaneLevel1, level0Rank_ + 1));
     257            0 :         SubCommInfo level1CommInfo = GetSubCommInfo(commPlaneLevel1, level0Rank_);
     258            0 :         CHK_RET(RunTemplate(level1AGTemplate, level1CommInfo));
     259            0 :         HCCL_INFO("AllGather double ring [superpod] level1 AllGather run success");
     260            0 :     }
     261              : 
     262              :     // 后处理
     263            0 :     CHK_RET(KernelRunInterServerPostProcess(param, execMem));
     264              : 
     265            0 :     return HCCL_SUCCESS;
     266            0 : }
     267              : 
     268            0 : HcclResult CollAllGatherRingZerocopyExecutor::KernelRunInterServerPostProcess(const OpParam &param, const ExecMem &execMem)
     269              : {
     270            0 :     u32 unitSize = 0;
     271            0 :     CHK_RET(SalGetDataTypeSize(param.DataDes.dataType, unitSize));
     272              : 
     273            0 :     DeviceMem dstMem;
     274            0 :     DeviceMem srcMem;
     275            0 :     u64 curSize = execMem.inputMem.size();
     276            0 :     Stream stream = param.stream;
     277            0 :     for (u32 i = 0; i < level1RankSize_; i++) {
     278            0 :         for (u32 j = 0; j < level2RankSize_; j++) {
     279              :             // 拷贝input上每个slice的数据到中转内存,源端每个slice的size固定为output的size
     280            0 :             u32 dstIndex = i * level2RankSize_ + j;
     281            0 :             u32 srcIndex = j * level1RankSize_ + i;
     282            0 :             srcMem = execMem.outputMem.range(dstIndex * curSize, curSize);
     283            0 :             dstMem = DeviceMem::create(static_cast<u8 *>(execMem.outputPtr)
     284            0 :                     + param.DataDes.count * unitSize * level0RankSize_ * srcIndex
     285            0 :                     + param.DataDes.count * unitSize * level0Rank_,
     286            0 :                     curSize);
     287            0 :             CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, srcMem, stream));
     288              :         }
     289              :     }
     290            0 :     return HCCL_SUCCESS;
     291            0 : }
     292              : 
     293            0 : HcclResult CollAllGatherRingZerocopyExecutor::CalcLevel0DataSlices(const OpParam &param, const ExecMem &execMem,
     294              :     std::vector<Slice> &dataSegsSlice)
     295              : {
     296            0 :     return CalcIntraServerDataSlicesDiscontinuous(param, execMem,
     297            0 :         level0RankSize_, level1RankSize_, level2RankSize_, dataSegsSlice);
     298              : }
     299              : 
     300              : REGISTER_EXEC("AllGatherRingZerocopyExecutor", AllGatherRingZerocopy, CollAllGatherRingZerocopyExecutor);
     301              : 
     302              : } // namespace hccl
        

Generated by: LCOV version 2.0-1