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

Generated by: LCOV version 2.0-1