LCOV - code coverage report
Current view: top level - coll_communicator_mgr/api_c_adpt - coll_comm_res_c_adpt.cc (source / functions) Coverage Total Hit
Test: coverage.info Lines: 66.9 % 354 237
Test Date: 2026-08-04 10:52:23 Functions: 65.4 % 26 17

            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 "my_rank.h"
      12              : #include <algorithm>
      13              : #include <limits>
      14              : #include "hccl_comm_pub.h"
      15              : #include "exception_handler.h"
      16              : #include "config_log.h"
      17              : #include "config/env_config.h"
      18              : #include "env_config/env_config.h"
      19              : 
      20              : #include "coll_comm_mgr.h"
      21              : #include "hcclCommOp.h"
      22              : #include "channel_process.h"
      23              : #include "aicpu_ts_roce_channel_v2.h"
      24              : #include "aiv_urma_channel.h"
      25              : #include "hccl_group.h"
      26              : #include "../resource_mgr/local/my_rank/comm_engine/kernel_launch/hccl_kernel_launch_aicpu.h"
      27              : #include "param_check_basic_v2.h"
      28              : #include "comm_engine_utils.h"
      29              : #include "rank_consistency_checker_v2.h"
      30              : #include "rank_table_crc_bridge.h"
      31              : #include "hccl/hccl_types.h"
      32              : #include "tp_qos.h"
      33              : 
      34              : using namespace hccl;
      35              : /**
      36              :  * @note 职责:集合通信的通信域资源管理的C接口的C到C++适配
      37              :  */
      38              : 
      39              : /**
      40              :  * @note C接口适配参考示例
      41              :  * @code {.c}
      42              :  * HcclResult HcclThreadAcquire(HcclComm comm, CommEngine engine, uint32_t threadNum,
      43              :  *     uint32_t notifyNumPerThread, ThreadHandle *threads) {
      44              :  *     return HCCL_SUCCESS;
      45              :  * }
      46              :  * @endcode
      47              :  */
      48              : 
      49              : constexpr uint32_t HCCL_CHANNEL_VERSION_ONE = 1;
      50              : constexpr uint32_t MULTIPLE = 4;               // 用于A5判断TC是否为4的倍数
      51              : constexpr uint32_t TC_MAX = 255;               // TC的最大值(不区分芯片类型)
      52              : constexpr uint32_t RETRY_INTERVAL_MIN = 5u;    // retryInterval范围的最小值(不区分芯片类型)
      53              : constexpr uint32_t A5_RETRY_INTERVAL_MAX = 24u;// A5的retryInterval范围的最大值
      54              : constexpr uint32_t RETRY_CNT_MIN = 1u;         // retryCnt范围的最小值(不区分芯片类型)
      55              : constexpr uint32_t RETRY_CNT_MAX = 7u;         // retryCnt范围的最大值(不区分芯片类型)
      56              : constexpr uint32_t SL_MAX = 7u;                // sl范围的最大值,sl即serviceLevel(不区分芯片类型)
      57              : constexpr uint32_t TC_DEFAULT = 0xFFFFFFFFu;   // TC的默认值(不区分芯片类型)
      58              : constexpr uint32_t SL_DEFAULT = 0xFFFFFFFFu;   // SL的默认值(不区分芯片类型)
      59              : constexpr uint32_t kDscpToRoceTcShift = 2U;    // RoCE TC = DSCP << 2(DiffServ 高 6 位为 DSCP)
      60              : 
      61            3 : static uint32_t ResolveRoceDevPhyId(const HcclChannelDesc &channelDesc)
      62              : {
      63            3 :     if (channelDesc.localEndpoint.loc.locType == ENDPOINT_LOC_TYPE_DEVICE) {
      64            3 :         return channelDesc.localEndpoint.loc.device.devPhyId;
      65              :     }
      66            0 :     s32 deviceLogicId = 0;
      67            0 :     u32 devicePhyId = 0U;
      68            0 :     if (hrtGetDevice(&deviceLogicId) != HCCL_SUCCESS) {
      69            0 :         return 0U;
      70              :     }
      71            0 :     if (hrtGetDevicePhyIdByIndex(static_cast<u32>(deviceLogicId), devicePhyId) != HCCL_SUCCESS) {
      72            0 :         return 0U;
      73              :     }
      74            0 :     return devicePhyId;
      75              : }
      76              : 
      77            8 : static void FillRoceQos(const hccl::CommConfig &commConfig, const Hccl::EnvRdmaConfig &rdmaConfig,
      78              :     const HcclChannelDesc &channelDesc, uint8_t &slOut, uint8_t &tcOut)
      79              : {
      80            8 :     const uint32_t hcclQos = commConfig.GetConfigHcclQos();
      81            8 :     if (hcclQos == HCCL_COMM_QOS_CONFIG_NOT_SET) {
      82            5 :         tcOut = static_cast<uint8_t>((commConfig.GetConfigTrafficClass() == INVALID_UINT) ?
      83            5 :             rdmaConfig.GetRdmaTrafficClass() : commConfig.GetConfigTrafficClass());
      84            5 :         slOut = static_cast<uint8_t>((commConfig.GetConfigServiceLevel() == INVALID_UINT) ?
      85            5 :             rdmaConfig.GetRdmaServerLevel() : commConfig.GetConfigServiceLevel());
      86            5 :         return;
      87              :     }
      88              : 
      89            3 :     slOut = static_cast<uint8_t>(hcclQos & 0xFFU);
      90            3 :     const uint32_t devPhyId = ResolveRoceDevPhyId(channelDesc);
      91            3 :     uint8_t dscp = Hccl::kUboeDefaultDscp;
      92            3 :     (void)Hccl::TpQosGetDscpByQosFromHccnCfg(devPhyId, slOut, dscp);
      93            3 :     tcOut = static_cast<uint8_t>((static_cast<uint32_t>(dscp) << kDscpToRoceTcShift) & 0xFFU);
      94            3 :     HCCL_INFO("[FillRoceQos] hcclQos compat: hcclQos[%u] devPhyId[%u] dscp[%u] sl[%u] tc[%u].",
      95              :         hcclQos, devPhyId, static_cast<unsigned>(dscp), static_cast<unsigned>(slOut),
      96              :         static_cast<unsigned>(tcOut));
      97              : }
      98              : 
      99            8 : static void FillChannelDescFinal(hccl::CommConfig commConfig, const HcclChannelDesc &channelDesc, HcclChannelDesc &channelDescFinal, bool isCommunicatorV2)
     100              : {
     101            8 :     if (isCommunicatorV2) { // A5
     102            8 :         auto& rdmaConfig = Hccl::EnvConfig::GetInstance().GetRdmaConfig();
     103            8 :         channelDescFinal.roceAttr.retryCnt = (channelDesc.roceAttr.retryCnt == INVALID_UINT) ? rdmaConfig.GetRdmaRetryCnt() : channelDesc.roceAttr.retryCnt;
     104            8 :         channelDescFinal.roceAttr.retryInterval = (channelDesc.roceAttr.retryInterval == INVALID_UINT) ? rdmaConfig.GetRdmaTimeOut() : channelDesc.roceAttr.retryInterval;
     105            8 :         FillRoceQos(commConfig, rdmaConfig, channelDesc, channelDescFinal.roceAttr.sl, channelDescFinal.roceAttr.tc);
     106            8 :         channelDescFinal.roceAttr.queueNum = (channelDesc.roceAttr.queueNum == INVALID_UINT) ? rdmaConfig.GetRdmaQueueNum() : channelDesc.roceAttr.queueNum;
     107              :     } else {
     108            0 :         channelDescFinal.roceAttr.retryCnt = (channelDesc.roceAttr.retryCnt == INVALID_UINT) ? EnvConfig::GetExternalInputRdmaRetryCnt() : channelDesc.roceAttr.retryCnt;
     109            0 :         channelDescFinal.roceAttr.retryInterval = (channelDesc.roceAttr.retryInterval == INVALID_UINT) ? EnvConfig::GetExternalInputRdmaTimeOut() : channelDesc.roceAttr.retryInterval;
     110            0 :         channelDescFinal.roceAttr.tc = (channelDesc.roceAttr.tc == 0xFF) ? EnvConfig::GetExternalInputRdmaTrafficClass() : channelDesc.roceAttr.tc;
     111            0 :         channelDescFinal.roceAttr.sl = (channelDesc.roceAttr.sl == 0xFF) ? EnvConfig::GetExternalInputRdmaServerLevel() : channelDesc.roceAttr.sl;
     112            0 :         channelDescFinal.roceAttr.queueNum = (channelDesc.roceAttr.queueNum == INVALID_UINT) ? GetExternalInputQpsPerConnection() : channelDesc.roceAttr.queueNum;
     113              :     }
     114            8 : }
     115              : 
     116           12 : static HcclResult CheckA5Config(hccl::CommConfig commConfig, const HcclChannelDesc &channelDesc)
     117              : {
     118           12 :     u32 tc = commConfig.GetConfigTrafficClass();
     119           12 :     CHK_PRT_RET((tc != TC_DEFAULT) && (tc > TC_MAX || (tc % MULTIPLE != 0)),
     120              :         HCCL_ERROR("[ProcessRoceChannelDesc]errNo[0x%016llx] invalid hcclRdmaTrafficClass[%u], must be 0xFFFFFFFF or in [0,255] and a multiple of 4",
     121              :             static_cast<unsigned long long>(HCCL_ERROR_CODE(HCCL_E_PARA)), tc),
     122              :         HCCL_E_PARA);
     123              : 
     124           11 :     u32 sl = commConfig.GetConfigServiceLevel();
     125           11 :     CHK_PRT_RET((sl != SL_DEFAULT) && (sl > SL_MAX),
     126              :         HCCL_ERROR("[ProcessRoceChannelDesc]errNo[0x%016llx] invalid hcclRdmaServiceLevel[%u], must be 0xFFFFFFFF or in [0,7]",
     127              :             static_cast<unsigned long long>(HCCL_ERROR_CODE(HCCL_E_PARA)), sl),
     128              :         HCCL_E_PARA);
     129              : 
     130           10 :     u32 retryInterval = channelDesc.roceAttr.retryInterval;
     131           10 :     CHK_PRT_RET((retryInterval != INVALID_UINT) && (retryInterval < RETRY_INTERVAL_MIN || retryInterval > A5_RETRY_INTERVAL_MAX),
     132              :         HCCL_ERROR("[ProcessRoceChannelDesc]errNo[0x%016llx] invalid hcclRdmaRetryInterval[%u], must be 0xFFFFFFFF or in [5,24]",
     133              :         static_cast<unsigned long long>(HCCL_ERROR_CODE(HCCL_E_PARA)), retryInterval),
     134              :         HCCL_E_PARA);
     135              : 
     136            9 :     u32 retryCnt = channelDesc.roceAttr.retryCnt;
     137            9 :     CHK_PRT_RET((retryCnt != INVALID_UINT) && (retryCnt < RETRY_CNT_MIN || retryCnt > RETRY_CNT_MAX),
     138              :         HCCL_ERROR("[ProcessRoceChannelDesc]errNo[0x%016llx] invalid hcclRdmaRetryCnt[%u], must be 0xFFFFFFFF or in [1,7]",
     139              :         static_cast<unsigned long long>(HCCL_ERROR_CODE(HCCL_E_PARA)), retryCnt),
     140              :         HCCL_E_PARA);
     141            8 :     return HCCL_SUCCESS;
     142              : }
     143              : 
     144           12 : HcclResult ProcessRoceChannelDesc(const HcclChannelDesc &channelDesc, HcclChannelDesc &channelDescFinal, hccl::hcclComm *hcclComm)
     145              : {
     146           12 :     bool isCommunicatorV2 = hcclComm->IsCommunicatorV2();
     147           12 :     hccl::CommConfig commConfig{}; // A5使用
     148           12 :     if (isCommunicatorV2) { // A5
     149           12 :         hccl::CollComm* collComm = hcclComm->GetCollComm();
     150           12 :         CHK_PTR_NULL(collComm);
     151           12 :         commConfig = collComm->GetCommConfig();
     152           12 :         CHK_RET(CheckA5Config(commConfig, channelDesc));
     153              :     }
     154            8 :     FillChannelDescFinal(commConfig, channelDesc, channelDescFinal, isCommunicatorV2);
     155            8 :     HCCL_INFO("[%s]queueNum[%u], retryCnt[%u], retryInterval[%u], tc[%u], sl[%u]", __func__,
     156              :         channelDescFinal.roceAttr.queueNum, channelDescFinal.roceAttr.retryCnt, channelDescFinal.roceAttr.retryInterval,
     157              :         channelDescFinal.roceAttr.tc, channelDescFinal.roceAttr.sl);
     158            8 :     return HCCL_SUCCESS;
     159           12 : }
     160              : 
     161            9 : HcclResult ProcessUbChannelDesc(const HcclChannelDesc &channelDesc, const HcclChannelDesc &channelDescFinal,
     162              :     const hccl::hcclComm *hcclComm)
     163              : {
     164              :     (void)channelDescFinal;
     165              :     (void)hcclComm;
     166              : 
     167            9 :     if (channelDesc.channelProtocol != COMM_PROTOCOL_UBC_CTP &&
     168            7 :         channelDesc.channelProtocol != COMM_PROTOCOL_UBC_TP &&
     169            6 :         channelDesc.channelProtocol != COMM_PROTOCOL_UBOE &&
     170            3 :         channelDesc.channelProtocol != COMM_PROTOCOL_UBG) {
     171            2 :         HCCL_ERROR("[%s] unexpected channelProtocol[%d], expect UBC_CTP/UBC_TP/UBOE/UBG", __func__,
     172              :             static_cast<int>(channelDesc.channelProtocol));
     173            2 :         return HCCL_E_PARA;
     174              :     }
     175            7 :     HCCL_INFO("[%s] channelProtocol[%d] ub comm-domain qos applied in HcommChannelDesc::qos when converting (HcclChannelDesc has no qos field)",
     176              :         __func__, static_cast<int>(channelDesc.channelProtocol));
     177            7 :     return HCCL_SUCCESS;
     178              : }
     179              : 
     180           16 : HcclResult ProcessHcclChannelDesc(const HcclChannelDesc &channelDesc, HcclChannelDesc &channelDescFinal, hccl::hcclComm *hcclComm)
     181              : {
     182           16 :     channelDescFinal.remoteRank = channelDesc.remoteRank;
     183           16 :     channelDescFinal.channelProtocol   = channelDesc.channelProtocol;
     184           16 :     channelDescFinal.localEndpoint  = channelDesc.localEndpoint;
     185           16 :     channelDescFinal.remoteEndpoint  = channelDesc.remoteEndpoint;
     186           16 :     channelDescFinal.notifyNum  = channelDesc.notifyNum;
     187           16 :     channelDescFinal.memHandles  = channelDesc.memHandles;
     188           16 :     channelDescFinal.memHandleNum  = channelDesc.memHandleNum;
     189              : 
     190              :      // 根据协议类型拷贝union中的相应成员
     191           16 :     switch (channelDesc.channelProtocol) {
     192            1 :         case COMM_PROTOCOL_HCCS:
     193              :         case COMM_PROTOCOL_HCCS_ONLY:
     194              :         case COMM_PROTOCOL_PCIE:
     195              :         case COMM_PROTOCOL_SIO:
     196              :         case COMM_PROTOCOL_UB_MEM:
     197            1 :             break;
     198            3 :         case COMM_PROTOCOL_UBC_CTP:
     199              :         case COMM_PROTOCOL_UBC_TP:
     200              :         case COMM_PROTOCOL_UBOE:
     201              :         case COMM_PROTOCOL_UBG:
     202            3 :             return ProcessUbChannelDesc(channelDesc, channelDescFinal, hcclComm);
     203           12 :         case COMM_PROTOCOL_ROCE:
     204           12 :             return ProcessRoceChannelDesc(channelDesc, channelDescFinal, hcclComm);
     205            0 :         default: {
     206            0 :             auto ProtocolToString = [](const CommProtocol proto) -> const char* {
     207            0 :                 switch (proto) {
     208            0 :                     case COMM_PROTOCOL_HCCS:    return "COMM_PROTOCOL_HCCS";
     209            0 :                     case COMM_PROTOCOL_PCIE:    return "COMM_PROTOCOL_PCIE";
     210            0 :                     case COMM_PROTOCOL_SIO:     return "COMM_PROTOCOL_SIO";
     211            0 :                     case COMM_PROTOCOL_UBC_CTP: return "COMM_PROTOCOL_UBC_CTP";
     212            0 :                     case COMM_PROTOCOL_UB_MEM:  return "COMM_PROTOCOL_UB_MEM";
     213            0 :                     case COMM_PROTOCOL_ROCE:    return "COMM_PROTOCOL_ROCE";
     214            0 :                     case COMM_PROTOCOL_UBC_TP:  return "COMM_PROTOCOL_UBC_TP";
     215            0 :                     case COMM_PROTOCOL_UBOE:    return "COMM_PROTOCOL_UBOE";
     216            0 :                     case COMM_PROTOCOL_UBG:     return "COMM_PROTOCOL_UBG";
     217            0 :                     case COMM_PROTOCOL_HCCS_ONLY:   return "COMM_PROTOCOL_HCCS_ONLY";
     218            0 :                     default:                    return "UNKNOWN_PROTOCOL";
     219              :                 }
     220              :             };
     221            0 :             HCCL_ERROR("[%s] Unsupported protocol[%s] found in HcclChannelDesc.",
     222              :                        __func__, ProtocolToString(channelDesc.channelProtocol));
     223            0 :             return HCCL_E_PARA;
     224              :         }
     225              :     }
     226            1 :     return HCCL_SUCCESS;
     227              : }
     228              : 
     229           16 : HcclResult ProcessHcclResPackReq(const HcclChannelDesc &channelDesc, HcclChannelDesc &channelDescFinal, hccl::hcclComm *hcclComm)
     230              : {
     231           16 :     if (channelDesc.header.size < channelDescFinal.header.size) {
     232              :         // 需要前向兼容HcclChannelDesc,末尾部分字段不支持处理
     233           16 :     } else if (channelDesc.header.size > channelDescFinal.header.size) {
     234              :         // 需要后向向兼容HcclChannelDesc,末尾部分字段会被忽略
     235              :     }
     236              :  
     237           16 :     if (channelDesc.header.magicWord != channelDescFinal.header.magicWord) {
     238            0 :         HCCL_ERROR("[%s]channelDescFinal.header.magicWord[%u] not equal to channelDesc.header.magicWord[%u]",
     239              :             __func__, channelDescFinal.header.magicWord, channelDesc.header.magicWord);
     240            0 :         return HCCL_E_PARA;
     241              :     }
     242              :  
     243           16 :     uint32_t copySize = (channelDescFinal.header.size < channelDesc.header.size ?
     244           16 :         channelDescFinal.header.size : channelDesc.header.size) - sizeof(CommAbiHeader);
     245           16 :     CHK_SAFETY_FUNC_RET(memcpy_s(reinterpret_cast<uint8_t *>(&channelDescFinal) + sizeof(CommAbiHeader), copySize,
     246              :         reinterpret_cast<const uint8_t *>(&channelDesc) + sizeof(CommAbiHeader), copySize));
     247              :  
     248           16 :     if (channelDesc.header.version >= HCCL_CHANNEL_VERSION_ONE) {
     249           16 :         CHK_RET(ProcessHcclChannelDesc(channelDesc, channelDescFinal, hcclComm));
     250              :     }
     251              :  
     252           12 :     if (channelDesc.header.version > HCCL_CHANNEL_VERSION) {
     253              :         // 传入的版本高于当前版本,警告不支持的配置项将被忽略
     254            0 :         HCCL_WARNING("The version of provided [%u] is higher than the current version[%u], "
     255              :             "unsupported configuration will be ignored.",
     256              :             channelDesc.header.version, HCCL_CHANNEL_VERSION);
     257           12 :     } else if (channelDesc.header.version < HCCL_CHANNEL_VERSION) {
     258              :         // 传入的版本低于当前版本,警告高版本支持的配置项将被忽略
     259            0 :         HCCL_WARNING("The version of provided [%u] is lower than the current version[%u], "
     260              :             "configurations supported by later versions will be ignored.",
     261              :             channelDesc.header.version, HCCL_CHANNEL_VERSION);
     262              :     }
     263              :  
     264              :     // 如果扩展到version=2后
     265              :     // 1) 在底层为新的结构体和版本(version为2)上,会正常执行下面的判断处理逻辑;
     266              :     // 2) 在底层为旧的结构体和版本(version为1)上,下面的逻辑没有,version的2 > 1的部分会被忽略掉;
     267           12 :     if (channelDesc.header.version >= 2) {
     268              :     }
     269              :  
     270           12 :     return HCCL_SUCCESS;
     271              : }
     272              : 
     273            1 : static HcclResult BuildAivDeviceChannelEntity(const HcclChannelDesc &channelDesc, ChannelHandle hostChannel,
     274              :     ChannelHandle &deviceChannel)
     275              : {
     276            1 :     void *channel = nullptr;
     277            1 :     CHK_RET(hcomm::ChannelProcess::ChannelGet(hostChannel, &channel));
     278            1 :     hcomm::Channel *baseChannel = static_cast<hcomm::Channel *>(channel);
     279            1 :     CHK_PTR_NULL(baseChannel);
     280              : 
     281            1 :     if (channelDesc.channelProtocol == COMM_PROTOCOL_ROCE) {
     282            0 :         auto *aicpuTsRoceChannelV2 = dynamic_cast<hcomm::AicpuTsRoceChannelV2 *>(baseChannel);
     283            0 :         CHK_PTR_NULL(aicpuTsRoceChannelV2);
     284            0 :         HCCL_INFO("[%s] build AIV direct device channel by AICPU+Host RoCE flow, protocol[%d], "
     285              :             "hostHandle[0x%llx]", __func__, channelDesc.channelProtocol,
     286              :             static_cast<unsigned long long>(hostChannel));
     287            0 :         CHK_RET(aicpuTsRoceChannelV2->BuildAndGetDevChannelEntity(&deviceChannel));
     288            0 :         return HCCL_SUCCESS;
     289              :     }
     290              : 
     291            1 :     if (channelDesc.channelProtocol == COMM_PROTOCOL_UBC_CTP ||
     292            1 :         channelDesc.channelProtocol == COMM_PROTOCOL_UBC_TP ||
     293            1 :         channelDesc.channelProtocol == COMM_PROTOCOL_UBG) {
     294            1 :         auto *aivUrmaChannel = dynamic_cast<hcomm::AivUrmaChannel *>(baseChannel);
     295            1 :         CHK_PTR_NULL(aivUrmaChannel);
     296            1 :         HCCL_INFO("[%s] build AIV direct device channel by AIV+URMA flow, protocol[%d], "
     297              :             "hostHandle[0x%llx]", __func__, channelDesc.channelProtocol,
     298              :             static_cast<unsigned long long>(hostChannel));
     299            1 :         void *devChannelEntity = nullptr;
     300            1 :         CHK_RET(aivUrmaChannel->BuildChannelEntityToDevice(&devChannelEntity));
     301            1 :         CHK_PTR_NULL(devChannelEntity);
     302            1 :         deviceChannel = static_cast<ChannelHandle>(reinterpret_cast<uintptr_t>(devChannelEntity));
     303            1 :         return HCCL_SUCCESS;
     304              :     }
     305              : 
     306            0 :     HCCL_ERROR("[%s] protocol[%d] is not AIV direct channel protocol", __func__, channelDesc.channelProtocol);
     307            0 :     return HCCL_E_PARA;
     308              : }
     309              : 
     310            4 : static HcclResult ConvertAivChannelHandlesToDevicePtrs(CommEngine engine, const HcclChannelDesc *channelDescs,
     311              :     uint32_t channelNum, ChannelHandle *channels)
     312              : {
     313            4 :     if (engine != COMM_ENGINE_AIV) {
     314            3 :         return HCCL_SUCCESS;
     315              :     }
     316              : 
     317            1 :     std::vector<ChannelHandle> hostChannels(channels, channels + channelNum);
     318            1 :     std::vector<ChannelHandle> deviceChannels(hostChannels);
     319            1 :     std::vector<ChannelHandle> mappedDeviceChannels;
     320            1 :     std::vector<ChannelHandle> mappedHostChannels;
     321            2 :     for (uint32_t idx = 0; idx < channelNum; ++idx) {
     322            1 :         if (channelDescs[idx].channelProtocol != COMM_PROTOCOL_ROCE &&
     323            1 :             channelDescs[idx].channelProtocol != COMM_PROTOCOL_UBC_CTP &&
     324            1 :             channelDescs[idx].channelProtocol != COMM_PROTOCOL_UBC_TP &&
     325            1 :             channelDescs[idx].channelProtocol != COMM_PROTOCOL_UBG) {
     326            0 :             continue;
     327              :         }
     328            1 :         CHK_RET(BuildAivDeviceChannelEntity(channelDescs[idx], hostChannels[idx], deviceChannels[idx]));
     329            1 :         mappedDeviceChannels.emplace_back(deviceChannels[idx]);
     330            1 :         mappedHostChannels.emplace_back(hostChannels[idx]);
     331            1 :         HCCL_INFO("[%s] convert AIV channel success, idx[%u], protocol[%d], hostHandle[0x%llx], devEntity[0x%llx]",
     332              :             __func__, idx, channelDescs[idx].channelProtocol, static_cast<unsigned long long>(hostChannels[idx]),
     333              :             static_cast<unsigned long long>(deviceChannels[idx]));
     334              :     }
     335              : 
     336            1 :     if (!mappedDeviceChannels.empty()) {
     337            1 :         CHK_RET(hcomm::ChannelProcess::RegisterChannelD2HMap(mappedDeviceChannels.data(), mappedHostChannels.data(),
     338              :             static_cast<uint32_t>(mappedDeviceChannels.size())));
     339              :     }
     340              : 
     341            2 :     for (uint32_t idx = 0; idx < channelNum; ++idx) {
     342            1 :         channels[idx] = deviceChannels[idx];
     343              :     }
     344            1 :     return HCCL_SUCCESS;
     345            1 : }
     346            2 : static bool IsUbUrmaChannelProtocol(CommProtocol protocol)
     347              : {
     348            2 :     return protocol == COMM_PROTOCOL_UBC_CTP || protocol == COMM_PROTOCOL_UBC_TP || protocol == COMM_PROTOCOL_UBOE
     349            4 :         || protocol == COMM_PROTOCOL_UBG;
     350              : }
     351              : 
     352            2 : static bool HasUbUrmaChannel(const std::vector<HcclChannelDesc> &channelDescFinals)
     353              : {
     354            3 :     for (const HcclChannelDesc &channelDesc : channelDescFinals) {
     355            2 :         if (IsUbUrmaChannelProtocol(channelDesc.channelProtocol)) {
     356            1 :             return true;
     357              :         }
     358              :     }
     359            1 :     return false;
     360              : }
     361              : 
     362            0 : static void AppendUniqueMemHandle(std::vector<HcclMemHandle> &mergedHandles, HcclMemHandle memHandle)
     363              : {
     364            0 :     if (memHandle == nullptr) {
     365            0 :         return;
     366              :     }
     367            0 :     if (std::find(mergedHandles.begin(), mergedHandles.end(), memHandle) == mergedHandles.end()) {
     368            0 :         mergedHandles.emplace_back(memHandle);
     369              :     }
     370              : }
     371              : 
     372            0 : static HcclResult MergeSymmetricMemHandles(HcclChannelDesc &channelDesc,
     373              :     const std::vector<HcclMemHandle> &symmetricMemHandles, std::vector<HcclMemHandle> &mergedHandles)
     374              : {
     375            0 :     if (!IsUbUrmaChannelProtocol(channelDesc.channelProtocol)) {
     376            0 :         return HCCL_SUCCESS;
     377              :     }
     378            0 :     if (channelDesc.memHandleNum != 0) {
     379            0 :         CHK_PTR_NULL(channelDesc.memHandles);
     380            0 :         for (uint32_t handleIdx = 0; handleIdx < channelDesc.memHandleNum; ++handleIdx) {
     381            0 :             AppendUniqueMemHandle(mergedHandles, channelDesc.memHandles[handleIdx]);
     382              :         }
     383              :     }
     384            0 :     for (HcclMemHandle memHandle : symmetricMemHandles) {
     385            0 :         AppendUniqueMemHandle(mergedHandles, memHandle);
     386              :     }
     387            0 :     CHK_PRT_RET(mergedHandles.size() > static_cast<size_t>(std::numeric_limits<uint32_t>::max()),
     388              :         HCCL_ERROR("[MergeSymmetricMemHandles] merged memHandleNum[%zu] exceeds uint32 max.",
     389              :             mergedHandles.size()), HCCL_E_PARA);
     390            0 :     channelDesc.memHandles = mergedHandles.data();
     391            0 :     channelDesc.memHandleNum = static_cast<uint32_t>(mergedHandles.size());
     392            0 :     return HCCL_SUCCESS;
     393              : }
     394              : 
     395            2 : static HcclResult AppendSymmetricMemHandles(hccl::CollComm *collComm,
     396              :     std::vector<HcclChannelDesc> &channelDescFinals,
     397              :     std::vector<std::vector<HcclMemHandle>> &mergedMemHandles,
     398              :     bool &hasSymmetricMemHandles)
     399              : {
     400            2 :     CHK_PTR_NULL(collComm);
     401            2 :     hasSymmetricMemHandles = false;
     402            2 :     if (!HasUbUrmaChannel(channelDescFinals)) {
     403            1 :         return HCCL_SUCCESS;
     404              :     }
     405              :     // 只有UB/URMA类channel需要追加symmetric memHandle参与建链交换。
     406            1 :     std::vector<HcclMemHandle> symmetricMemHandles;
     407            1 :     CHK_RET(collComm->RegisterPendingSymmetricMemHandles(symmetricMemHandles));
     408            1 :     if (symmetricMemHandles.empty()) {
     409            1 :         return HCCL_SUCCESS;
     410              :     }
     411            0 :     hasSymmetricMemHandles = true;
     412              : 
     413            0 :     mergedMemHandles.clear();
     414            0 :     mergedMemHandles.resize(channelDescFinals.size());
     415            0 :     for (size_t idx = 0; idx < channelDescFinals.size(); ++idx) {
     416            0 :         CHK_RET(MergeSymmetricMemHandles(channelDescFinals[idx], symmetricMemHandles, mergedMemHandles[idx]));
     417              :     }
     418            0 :     HCCL_INFO("[AppendSymmetricMemHandles] append symmetric memHandles success, channelNum[%zu], symMemHandleNum[%zu], "
     419              :         "protocols[UBC_CTP/UBC_TP/UBOE].",
     420              :         channelDescFinals.size(), symmetricMemHandles.size());
     421            0 :     return HCCL_SUCCESS;
     422            1 : }
     423              : 
     424            0 : static HcclResult UpdateSymmetricRemoteMems(hccl::CollComm *collComm, const hccl::MyRank *myRank,
     425              :     const std::vector<HcclChannelDesc> &channelDescFinals, const ChannelHandle *channels, uint32_t channelNum)
     426              : {
     427            0 :     CHK_PTR_NULL(collComm);
     428            0 :     CHK_PTR_NULL(myRank);
     429            0 :     CHK_PTR_NULL(channels);
     430            0 :     for (uint32_t idx = 0; idx < channelNum; ++idx) {
     431            0 :         const HcclChannelDesc &channelDesc = channelDescFinals[idx];
     432            0 :         if (!IsUbUrmaChannelProtocol(channelDesc.channelProtocol)) {
     433            0 :             continue;
     434              :         }
     435            0 :         CommMem *remoteMems = nullptr;
     436            0 :         uint32_t memNum = 0;
     437            0 :         std::vector<std::string> memTags;
     438              :         // CreateChannels完成后,从channel取回交换到的remoteMem/memTag并回填window。
     439            0 :         CHK_RET(myRank->ChannelGetRemoteMems(channels[idx], &memNum, &remoteMems, memTags));
     440            0 :         if (memNum == 0) {
     441            0 :             continue;
     442              :         }
     443            0 :         CHK_RET(collComm->UpdateSymmetricRemoteMem(channelDesc.remoteRank, remoteMems, memTags));
     444            0 :     }
     445            0 :     return HCCL_SUCCESS;
     446              : }
     447              : 
     448            7 : bool CheckCommEngine(const CommEngine engine, const uint32_t opExpansionMode)
     449              : {
     450            7 :     constexpr uint32_t DEFAULT_MODE = 0;
     451            7 :     constexpr uint32_t CCU_MS_MODE = 5;
     452            7 :     constexpr uint32_t CCU_SCHE_MODE = 6;
     453            7 :     if (engine == CommEngine::COMM_ENGINE_CCU) {
     454              :         return opExpansionMode == DEFAULT_MODE
     455            0 :             || opExpansionMode == CCU_MS_MODE
     456            0 :             || opExpansionMode == CCU_SCHE_MODE;
     457              :     }
     458              : 
     459            7 :     return true;
     460              : }
     461              : 
     462            9 : static bool IsAicpuEngine(CommEngine engine)
     463              : {
     464            9 :     return engine == COMM_ENGINE_AICPU || engine == COMM_ENGINE_AICPU_TS;
     465              : }
     466              : 
     467              : constexpr uint32_t CHANNEL_NUM_MAX = 1024 * 1024;  // channel的默认限制最大为1024 * 1024
     468              : 
     469            5 : HcclResult RegisterToClusterMonitor(HcclComm comm)
     470              : {
     471            5 :     HCCL_INFO("[%s] START, comm[%p].", __func__, comm);
     472            5 :     CHK_PRT_RET(comm == nullptr,  HCCL_ERROR("[%s] comm is null", __func__), HCCL_E_PTR);
     473            5 :     auto* hcclComm = static_cast<hccl::hcclComm*>(comm);
     474            5 :     CHK_PTR_NULL(hcclComm);
     475            5 :     if (!hcclComm->IsCommunicatorV2()) {
     476            0 :         HCCL_ERROR("[%s] comm is not support", __func__);
     477            0 :         return HCCL_E_NOT_SUPPORT;
     478              :     }
     479            5 :     hccl::CollComm* collComm = hcclComm->GetCollComm();
     480            5 :     CHK_PTR_NULL(collComm);
     481            5 :     CHK_RET(CollCommMgr::GetInstance()->GetClusterMonitor(collComm->GetDeviceLogicId()).RegisterToClusterMonitor(comm));
     482            3 :     HCCL_INFO("%s Success", __func__);
     483            3 :     return HCCL_SUCCESS;
     484              : }
     485              : 
     486           13 : HcclResult HcclChannelAcquire(HcclComm comm, CommEngine engine, 
     487              :     const HcclChannelDesc* channelDescs, uint32_t channelNum, ChannelHandle* channels)
     488              : {
     489           13 :     HcclUs startut = TIME_NOW();
     490           13 :     u64 beginTime =  Hccl::DlProfFunction::GetInstance().dlMsprofSysCycleTime();
     491              :     EXCEPTION_HANDLE_BEGIN
     492              : 
     493              :     // 入参校验
     494           21 :     CHK_PTR_NULL(comm);
     495           12 :     CHK_PTR_NULL(channelDescs);
     496           12 :     CHK_PTR_NULL(channels);
     497           12 :     CHK_PRT_RET(
     498              :         (channelNum == 0 || channelNum > CHANNEL_NUM_MAX), 
     499              :         HCCL_ERROR("[%s]Invalid channelNum, channelNum[%u], max channel num[%u]",
     500              :         __func__, channelNum, CHANNEL_NUM_MAX), HCCL_E_PARA
     501              :     );
     502              :  
     503           12 :     HcclResult ret = HCCL_SUCCESS;
     504           12 :     hccl::hcclComm *hcclComm = static_cast<hccl::hcclComm *>(comm);
     505           12 :     HCCL_RUN_INFO("Entry-%s channelNum[%u], engine[%s] group[%s]", __func__, channelNum, GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), hcclComm->GetIdentifier().c_str());
     506           12 :     std::vector<HcclChannelDesc> channelDescFinals;
     507           12 :     std::vector<std::vector<HcclMemHandle>> mergedMemHandles;
     508           20 :     for (uint32_t idx = 0; idx < channelNum; idx++) {
     509              :         HcclChannelDesc channelDescFinal;
     510           12 :         HcclChannelDescInit(&channelDescFinal, 1);
     511           12 :         ret = ProcessHcclResPackReq(channelDescs[idx], channelDescFinal, hcclComm);
     512           12 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     513              :             HCCL_ERROR("ProcessHcclResPackReq failed. channelDesc idx[%u], group[%s], engine[%s] channelNum[%u], ret[%d]", idx, hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, ret), ret);
     514            8 :         channelDescFinals.push_back(channelDescFinal);
     515              :     }
     516              :  
     517            8 :     if (hcclComm->IsCommunicatorV2()) {  // A5
     518            7 :         hccl::CollComm* collComm = hcclComm->GetCollComm();
     519           10 :         CHK_PTR_NULL(collComm);
     520            7 :         const std::string &commTag = hcclComm->GetIdentifier();
     521            7 :         hccl::MyRank* myRank = collComm->GetMyRank();
     522            7 :         CHK_PTR_NULL(myRank);
     523              : 
     524            7 :         s32 deviceLogicId = 0;
     525            7 :         (void)hrtGetDeviceRefresh(&deviceLogicId);
     526            7 :         u32 rankTableCrc = RankTableCrcBridge::GetInstance().ConsumeRankTableJsonCrc(deviceLogicId);
     527            7 :         if (rankTableCrc != 0) {
     528            0 :             CHK_RET(RankConsistencyCheckerV2::GetInstance(deviceLogicId).RecordRankTableCrcV2(rankTableCrc));
     529              :         }
     530            7 :         char hcommPkgName[] = "hcomm";
     531            7 :         char hcommVersionStr[CANN_VERSION_MAX_LEN + 1] = {0};
     532            7 :         aclError aclRet = aclsysGetVersionStr(hcommPkgName, hcommVersionStr);
     533            7 :         CHK_PRT_RET(aclRet != ACL_SUCCESS,
     534              :             HCCL_ERROR("[HcclChannelAcquire] aclsysGetVersionStr failed, aclRet[%d].", aclRet), HCCL_E_INTERNAL);
     535            7 :         std::string curVersion(hcommVersionStr);
     536            7 :         CHK_RET(RankConsistencyCheckerV2::GetInstance(deviceLogicId).RecordCannVersionV2(curVersion));
     537              :  
     538            7 :         const uint32_t opExpansionMode = myRank->GetOpExpansionMode();
     539            7 :         if (!CheckCommEngine(engine, opExpansionMode)) {
     540            0 :             HCCL_ERROR("[%s] failed, coll comm[%p] group[%s] opExpansionMode[%u] is not supported by CCU engine[%s].", 
     541              :                 __func__, hcclComm, hcclComm->GetIdentifier().c_str(), opExpansionMode, GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str());
     542            0 :             return HcclResult::HCCL_E_PARA;
     543              :         }
     544              : 
     545            7 :         if (engine != CommEngine::COMM_ENGINE_CPU) { // host dpu场景暂不支持cluster monitor
     546            5 :             ret = RegisterToClusterMonitor(comm);
     547            5 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
     548              :                 HCCL_ERROR("RegisterToClusterMonitor failed. group[%s], engine[%s], channelNum[%u], ret[%d]", hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, ret), ret);
     549              :         }
     550              : 
     551            5 :         if (!GetDebugConfigInited()) {
     552            1 :             InitDebugConfigByEnv();
     553              :         }
     554              : 
     555            5 :         bool hasSymmetricMemHandles = false;
     556            5 :         if (IsAicpuEngine(engine)) {
     557            2 :             CHK_RET(AppendSymmetricMemHandles(collComm, channelDescFinals, mergedMemHandles, hasSymmetricMemHandles));
     558              :         }
     559            5 :         HCCL_INFO("[HcclChannelAcquire] AppendSymmetricMemHandles done, group[%s], engine[%d], channelNum[%u], "
     560              :             "hasSymmetricMemHandles[%d], mergedMemHandleGroups[%zu].",
     561              :             commTag.c_str(), engine, channelNum, hasSymmetricMemHandles, mergedMemHandles.size());
     562            5 :         ret = myRank->CreateChannels(engine, commTag, channelDescFinals.data(), channelNum, channels);
     563            5 :         CHK_PRT_RET((ret == HCCL_E_AGAIN || ret == HCCL_E_UNAVAIL),
     564              :             HCCL_WARNING("CreateChannels group[%s], engine[%s] ret[%d]", commTag.c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), ret), ret);
     565            5 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     566              :             HCCL_ERROR("CreateChannels failed. group[%s], engine[%s] ret[%d]", commTag.c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), ret), ret);
     567            4 :         if (hasSymmetricMemHandles) {
     568            0 :             CHK_RET(UpdateSymmetricRemoteMems(collComm, myRank, channelDescFinals, channels, channelNum));
     569              :         }
     570            4 :         if (engine == COMM_ENGINE_CPU) {
     571            2 :             HcclCommDfx* hcclCommDfx = collComm->GetHcclCommDfx();
     572            2 :             CHK_PTR_NULL(hcclCommDfx);
     573            2 :             auto callback = hcclCommDfx->GetDpuCallback();
     574            4 :             for (uint32_t idx = 0; idx < channelNum; idx++) {
     575            2 :                 int32_t ret = HcommDpuChannelRegisterDfx(channels[idx], callback);
     576            2 :                 CHK_PRT_RET(ret != HCCL_SUCCESS,
     577              :                     HCCL_ERROR("[HcclChannelAcquire] group[%s] Failed to register DFX callback for channel[%u], ret[%d]", commTag.c_str(), idx, ret),
     578              :                     static_cast<HcclResult>(ret));
     579              :             }
     580            2 :             HCCL_INFO("[HcclChannelAcquire] group[%s] channelNum[%u] Register DFX callback for CPU channels success", commTag.c_str(), channelNum);
     581            2 :         }
     582            4 :         if (IsAicpuEngine(engine)) {
     583            1 :             HCCL_INFO("[HcclChannelAcquire] ReportChannelAicpuKernel start");
     584            1 :             HcclCommDfx* hcclCommDfx = collComm->GetHcclCommDfx();
     585            1 :             CHK_PTR_NULL(hcclCommDfx);
     586            1 :             std::string kernelName = "RunAicpuIndOpChannelInitV2";
     587              :             // 还是kernel的当前无法判断
     588            1 :             ret = hcclCommDfx->ReportKernel(beginTime, commTag, kernelName, SalGetTid(), false);
     589            1 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
     590              :                 HCCL_ERROR("[HcclChannelAcquire] group[%s] Failed to report kernel for kernelName[%s], tid[%d], ret[%d]", commTag.c_str(), kernelName.c_str(), SalGetTid(), ret), ret);
     591            1 :         }
     592           10 :     } else {
     593            1 :         hccl::CollComm* collComm = hcclComm->GetCollComm();
     594            1 :         if (collComm != nullptr) {
     595            0 :             hccl::MyRank *myRank = collComm->GetMyRank();
     596            0 :             if (hcclComm->GetConnectMode() != 0 && engine == COMM_ENGINE_CPU && myRank != nullptr) {
     597            0 :                 const std::string &commTag = hcclComm->GetIdentifier();
     598            0 :                 ret = myRank->CreateChannels(engine, commTag, channelDescFinals.data(), channelNum, channels);
     599            0 :             } else {
     600            0 :                 auto& channelMgr = hcclComm->GetIndependentOp().GetChannelManager();
     601            0 :                 ret = channelMgr.ChannelCommCreate(hcclComm->GetIdentifier(), engine,
     602            0 :                     channelDescFinals.data(), channelNum, channels);
     603              :             }
     604              :         } else {
     605            1 :             auto& channelMgr = hcclComm->GetIndependentOp().GetChannelManager();
     606            1 :             ret = channelMgr.ChannelCommCreate(hcclComm->GetIdentifier(), engine,
     607            1 :                 channelDescFinals.data(), channelNum, channels);
     608              :         }
     609              :     }
     610              :  
     611            5 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     612              :         HCCL_ERROR("[%s] Failed to acquire channel, group[%s], engine[%s], channelNum[%u], ret[%d]", __func__, hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, ret), ret);
     613              : 
     614            4 :     CHK_RET(ConvertAivChannelHandlesToDevicePtrs(engine, channelDescFinals.data(), channelNum, channels));
     615              :  
     616            4 :     HCCL_RUN_INFO("[%s] acquire channel success, group[%s], engine[%s], channelNum[%u], take time [%lld]us.", __func__, hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, DURATION_US(TIME_NOW() - startut).count());
     617           20 :     EXCEPTION_HANDLE_END
     618            4 :     return HCCL_SUCCESS;
     619              : }
     620              : 
     621            0 : HcclResult HcclGroupStart()
     622              : {
     623            0 :     return HcclLegacyGroupStart();
     624              : }
     625              : 
     626            0 : HcclResult HcclGroupEndV2()
     627              : {
     628            0 :     CHK_RET(groupLaunchA5());
     629            0 :     HCCL_INFO("[GroupEnd] to the end");
     630            0 :     return HCCL_SUCCESS;
     631              : }
     632              : 
     633            0 : HcclResult HcclGroupEnd()
     634              : {
     635            0 :     if (hcclGroupDepth == 0) {
     636            0 :         HCCL_ERROR("HcclGroupEnd: not in a group call. Didn't call HcclGroupStart before.");
     637            0 :         return HCCL_E_NOT_SUPPORT;
     638              :     }
     639            0 :     if (--hcclGroupDepth > 0) {
     640            0 :         return HCCL_SUCCESS;
     641              :     }
     642              : 
     643            0 :     HCCL_INFO("[HcclGroupEnd] hcclGroupDepth=[%d]", hcclGroupDepth);
     644              :     /*遇到最后一个HcclGroupEnd才处理group内的所有任务*/
     645            0 :     HCCLV2_FUNC_RUN([&]() -> HcclResult {
     646              :         CHK_RET(HcclLegacyAsyncJobLaunch());
     647              :         return HcclGroupEndV2();
     648              :     }());
     649            0 :     return HcclLegacyGroupEnd();
     650              : }
     651              : 
     652            0 : HcclResult HcclGroupStatusGet(bool *isGroupEnabled)
     653              : {
     654            0 :     CHK_PTR_NULL(isGroupEnabled);
     655            0 :     *isGroupEnabled = (hcclGroupDepth > 0);
     656            0 :     return HCCL_SUCCESS;
     657              : }
        

Generated by: LCOV version 2.0-1