LCOV - code coverage report
Current view: top level - legacy/ascend910/algorithm/impl/coll_executor/coll_all_gather - coll_all_gather_comm_executor.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 48.0 % 98 47
Test Date: 2026-08-18 17:47:01 Functions: 71.4 % 7 5

            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_comm_executor.h"
      12              : 
      13              : namespace hccl {
      14            5 : CollAllGatherCommExecutor::CollAllGatherCommExecutor(
      15            5 :     const HcclDispatcher dispatcher, std::unique_ptr<TopoMatcher>& topoMatcher)
      16            5 :     : CollAllGatherExecutor(dispatcher, topoMatcher)
      17              : {
      18            5 :     DMAReduceFlag_ = false;
      19            5 : }
      20              : 
      21            5 : HcclResult CollAllGatherCommExecutor::CalcCommInfo(std::vector<LevelNSubCommTransport>& opTransport)
      22              : {
      23            5 :     TransportMemType inputType = TransportMemType::RESERVED;
      24            5 :     TransportMemType outputType = TransportMemType::RESERVED;
      25            5 :     CHK_RET(CalcTransportMemType(inputType, outputType));
      26            5 :     CHK_RET(CalcCombinedCommInfo(inputType, outputType, opTransport));
      27            5 :     return HCCL_SUCCESS;
      28              : }
      29              : 
      30            5 : HcclResult CollAllGatherCommExecutor::CalcTransportMemType(TransportMemType& inputType, TransportMemType& outputType)
      31              : {
      32            5 :     if (workflowMode_ == HcclWorkflowMode::HCCL_WORKFLOW_MODE_OP_BASE) {
      33            0 :         inputType = TransportMemType::CCL_INPUT;
      34            0 :         outputType = TransportMemType::CCL_OUTPUT;
      35              :     } else {
      36            5 :         inputType = TransportMemType::PARAM_INPUT;
      37            5 :         outputType = TransportMemType::PARAM_OUTPUT;
      38              :     }
      39            5 :     HCCL_INFO(
      40              :         "[CollAllGatherCommExecutor][CalcTransportMemType] tag[%s] inputType[%d], outputType[%d]", tag_.c_str(),
      41              :         inputType, outputType);
      42            5 :     return HCCL_SUCCESS;
      43              : }
      44              : 
      45            5 : HcclResult CollAllGatherCommExecutor::CalcCombinedCommInfo(
      46              :     TransportMemType inputType, TransportMemType outputType, std::vector<LevelNSubCommTransport>& opTransport)
      47              : {
      48            5 :     CommPlane commPlane = COMM_COMBINE;
      49            5 :     if (topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
      50            0 :         commPlane = COMM_COMBINE_ORDER;
      51              :     }
      52              : 
      53            5 :     CommParaInfo commParaInfo(commPlane, CommType::COMM_TAG_MAX);
      54            5 :     if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
      55            0 :         commParaInfo.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING;
      56            5 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
      57            0 :         commParaInfo.commType = CommType::COMM_TAG_NONUNIFORM_HIERARCHICAL_RING_V1;
      58            5 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
      59            0 :         commParaInfo.commType = CommType::COMM_TAG_NONUNIFORM_BRUCK;
      60            5 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_HD) {
      61            0 :         commParaInfo.commType = CommType::COMM_TAG_HALVING_DOUBLING;
      62              :     } else {
      63            5 :         commParaInfo.commType = CommType::COMM_TAG_RING_INNER;
      64              :     }
      65            5 :     CHK_RET(CalcCommPlaneInfo(tag_, commParaInfo, opTransport[commPlane], inputType, outputType));
      66              : 
      67            5 :     return HCCL_SUCCESS;
      68            5 : }
      69              : 
      70            5 : HcclResult CollAllGatherCommExecutor::KernelRun(const OpParam& param, ExecMem& execMem)
      71              : {
      72            5 :     HCCL_CONFIG_INFO(HCCL_ALG, "[CollAllGatherCommExecutor][KernelRun]AllGather enter");
      73            5 :     CommPlane commPlane = COMM_COMBINE;
      74            5 :     if (topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
      75            0 :         commPlane = COMM_COMBINE_ORDER;
      76              :     }
      77              : 
      78            5 :     CHK_RET(CheckCommSize(commPlane, COMM_INDEX_0 + 1));
      79            5 :     SubCommInfo combinedCommInfo = GetSubCommInfo(commPlane, COMM_INDEX_0);
      80              : 
      81              :     // 构造ring algorithm对应的all_gather实例
      82            5 :     std::unique_ptr<AlgTemplateBase> tempAlg;
      83            5 :     if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
      84            0 :         tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
      85            0 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHR in COMM_COMBINE_ORDER", __func__);
      86            5 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
      87            0 :         tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHRV1, dispatcher_);
      88            0 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NHRV1 in COMM_COMBINE_ORDER", __func__);
      89            5 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
      90            0 :         tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
      91            0 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_NB in COMM_COMBINE_ORDER", __func__);
      92            5 :     } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_HD) {
      93            0 :         u8* dstPtr = static_cast<u8*>(execMem.outputMem.ptr()) + execMem.inputMem.size() * combinedCommInfo.localRank;
      94            0 :         DeviceMem dstMem = DeviceMem::create(dstPtr, execMem.inputMem.size());
      95            0 :         CHK_RET(HcclD2DMemcpyAsync(dispatcher_, dstMem, execMem.inputMem, const_cast<Stream&>(param.stream)));
      96            0 :         tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
      97            0 :             TemplateType::TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING, dispatcher_);
      98            0 :         HCCL_CONFIG_INFO(
      99              :             HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING in COMM_COMBINE_ORDER", __func__);
     100            0 :     } else {
     101            5 :         tempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     102            5 :         HCCL_CONFIG_INFO(HCCL_ALG, "[%s] Run TEMPLATE_ALL_GATHER_RING in COMM_COMBINE_ORDER", __func__);
     103              :     }
     104            5 :     CHK_SMART_PTR_NULL(tempAlg);
     105              : 
     106           25 :     CHK_RET(tempAlg->Prepare(
     107              :         execMem.inputMem, execMem.outputMem, execMem.outputMem, execMem.count, param.DataDes.dataType, param.stream,
     108              :         HCCL_REDUCE_RESERVED, INVALID_VALUE_RANKID));
     109              : 
     110            5 :     CHK_RET(RunTemplate(tempAlg, combinedCommInfo));
     111              : 
     112            5 :     return HCCL_SUCCESS;
     113            5 : }
     114            0 : HcclResult CollAllGatherCommExecutor::Getlevel1CommRank(SubCommInfo& level1CommInfo)
     115              : {
     116            0 :     HCCL_INFO("[CollAllGatherCommExecutor] Entry Getlevel1CommRank.");
     117            0 :     CommPlane commPlane = COMM_COMBINE;
     118            0 :     if (topoAttr_.deviceType == DevType::DEV_TYPE_910_93) {
     119            0 :         commPlane = COMM_COMBINE_ORDER;
     120            0 :         HCCL_INFO("nslbdp AllGather comm: Getlevel1CommRank.");
     121              :     }
     122            0 :     CHK_RET(CheckCommSize(commPlane, COMM_INDEX_0 + 1));
     123            0 :     level1CommInfo = GetSubCommInfo(commPlane, COMM_INDEX_0);
     124              : 
     125            0 :     return HCCL_SUCCESS;
     126              : }
     127              : 
     128            0 : HcclResult CollAllGatherCommExecutor::SelectTempAlg(std::unique_ptr<AlgTemplateBase>& level1TempAlg, u32 level1RankSize)
     129              : {
     130            0 :     HCCL_INFO("[CollAllGatherCommExecutor] Entry SelectTempAlg, algoLevel1 = [%u].", algType_.algoLevel1);
     131            0 :     if (level1RankSize > 1) {
     132            0 :         if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR) {
     133              :             level1TempAlg
     134            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHR, dispatcher_);
     135            0 :             HCCL_INFO("allgather comm: using nhr algo inter-server.");
     136            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NHR_V1) {
     137              :             level1TempAlg
     138            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NHRV1, dispatcher_);
     139            0 :             HCCL_INFO("allgather comm: using nhr_v1 algo inter-server.");
     140            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_NB) {
     141              :             level1TempAlg
     142            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_NB, dispatcher_);
     143            0 :             HCCL_INFO("allgather comm: using nonuniform-bruck algo inter-server.");
     144            0 :         } else if (algType_.algoLevel1 == AlgTypeLevel1::ALG_LEVEL1_HD) {
     145            0 :             level1TempAlg = AlgTemplateRegistry::Instance().GetAlgTemplate(
     146            0 :                 TemplateType::TEMPLATE_ALL_GATHER_RECURSIVE_HALVING_DOUBLING, dispatcher_);
     147            0 :             HCCL_INFO("allgather comm: using halving-doubling algo inter-server.");
     148              :         } else {
     149              :             level1TempAlg
     150            0 :                 = AlgTemplateRegistry::Instance().GetAlgTemplate(TemplateType::TEMPLATE_ALL_GATHER_RING, dispatcher_);
     151            0 :             HCCL_INFO("allgather comm: ring algo inter-server.");
     152              :         }
     153            0 :         CHK_SMART_PTR_NULL(level1TempAlg);
     154            0 :         return HCCL_SUCCESS;
     155              :     }
     156            0 :     return HCCL_E_UNAVAIL;
     157              : }
     158              : REGISTER_EXEC("AllGatherComm", AllGatherComm, CollAllGatherCommExecutor);
     159              : 
     160              : } // namespace hccl
        

Generated by: LCOV version 2.0-1