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: 57.1 % 352 201
Test Date: 2026-07-28 12:11:00 Functions: 61.5 % 26 16

            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_reses/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              :             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              :             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              :         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              :         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            8 : HcclResult ProcessUbChannelDesc(const HcclChannelDesc &channelDesc, const HcclChannelDesc &channelDescFinal,
     162              :     const hccl::hcclComm *hcclComm)
     163              : {
     164              :     (void)channelDescFinal;
     165              :     (void)hcclComm;
     166              : 
     167            8 :     if (channelDesc.channelProtocol != COMM_PROTOCOL_UBC_CTP &&
     168            6 :         channelDesc.channelProtocol != COMM_PROTOCOL_UBC_TP &&
     169            5 :         channelDesc.channelProtocol != COMM_PROTOCOL_UBOE &&
     170            2 :         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            6 :     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            6 :     return HCCL_SUCCESS;
     178              : }
     179              : 
     180           15 : HcclResult ProcessHcclChannelDesc(const HcclChannelDesc &channelDesc, HcclChannelDesc &channelDescFinal, hccl::hcclComm *hcclComm)
     181              : {
     182           15 :     channelDescFinal.remoteRank = channelDesc.remoteRank;
     183           15 :     channelDescFinal.channelProtocol   = channelDesc.channelProtocol;
     184           15 :     channelDescFinal.localEndpoint  = channelDesc.localEndpoint;
     185           15 :     channelDescFinal.remoteEndpoint  = channelDesc.remoteEndpoint;
     186           15 :     channelDescFinal.notifyNum  = channelDesc.notifyNum;
     187           15 :     channelDescFinal.memHandles  = channelDesc.memHandles;
     188           15 :     channelDescFinal.memHandleNum  = channelDesc.memHandleNum;
     189              : 
     190              :      // 根据协议类型拷贝union中的相应成员
     191           15 :     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            2 :         case COMM_PROTOCOL_UBC_CTP:
     199              :         case COMM_PROTOCOL_UBC_TP:
     200              :         case COMM_PROTOCOL_UBOE:
     201              :         case COMM_PROTOCOL_UBG:
     202            2 :             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           15 : HcclResult ProcessHcclResPackReq(const HcclChannelDesc &channelDesc, HcclChannelDesc &channelDescFinal, hccl::hcclComm *hcclComm)
     230              : {
     231           15 :     if (channelDesc.header.size < channelDescFinal.header.size) {
     232              :         // 需要前向兼容HcclChannelDesc,末尾部分字段不支持处理
     233           15 :     } else if (channelDesc.header.size > channelDescFinal.header.size) {
     234              :         // 需要后向向兼容HcclChannelDesc,末尾部分字段会被忽略
     235              :     }
     236              :  
     237           15 :     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           15 :     uint32_t copySize = (channelDescFinal.header.size < channelDesc.header.size ?
     244           15 :         channelDescFinal.header.size : channelDesc.header.size) - sizeof(CommAbiHeader);
     245           15 :     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           15 :     if (channelDesc.header.version >= HCCL_CHANNEL_VERSION_ONE) {
     249           15 :         CHK_RET(ProcessHcclChannelDesc(channelDesc, channelDescFinal, hcclComm));
     250              :     }
     251              :  
     252           11 :     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           11 :     } 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           11 :     if (channelDesc.header.version >= 2) {
     268              :     }
     269              :  
     270           11 :     return HCCL_SUCCESS;
     271              : }
     272              : 
     273            0 : static HcclResult BuildAivDeviceChannelEntity(const HcclChannelDesc &channelDesc, ChannelHandle hostChannel,
     274              :     ChannelHandle &deviceChannel)
     275              : {
     276            0 :     void *channel = nullptr;
     277            0 :     CHK_RET(hcomm::ChannelProcess::ChannelGet(hostChannel, &channel));
     278            0 :     hcomm::Channel *baseChannel = static_cast<hcomm::Channel *>(channel);
     279            0 :     CHK_PTR_NULL(baseChannel);
     280              : 
     281            0 :     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            0 :     if (channelDesc.channelProtocol == COMM_PROTOCOL_UBC_CTP ||
     292            0 :         channelDesc.channelProtocol == COMM_PROTOCOL_UBC_TP) {
     293            0 :         auto *aivUrmaChannel = dynamic_cast<hcomm::AivUrmaChannel *>(baseChannel);
     294            0 :         CHK_PTR_NULL(aivUrmaChannel);
     295            0 :         HCCL_INFO("[%s] build AIV direct device channel by AIV+URMA flow, protocol[%d], "
     296              :             "hostHandle[0x%llx]", __func__, channelDesc.channelProtocol,
     297              :             static_cast<unsigned long long>(hostChannel));
     298            0 :         void *devChannelEntity = nullptr;
     299            0 :         CHK_RET(aivUrmaChannel->BuildChannelEntityToDevice(&devChannelEntity));
     300            0 :         CHK_PTR_NULL(devChannelEntity);
     301            0 :         deviceChannel = static_cast<ChannelHandle>(reinterpret_cast<uintptr_t>(devChannelEntity));
     302            0 :         return HCCL_SUCCESS;
     303              :     }
     304              : 
     305            0 :     HCCL_ERROR("[%s] protocol[%d] is not AIV direct channel protocol", __func__, channelDesc.channelProtocol);
     306            0 :     return HCCL_E_PARA;
     307              : }
     308              : 
     309            3 : static HcclResult ConvertAivChannelHandlesToDevicePtrs(CommEngine engine, const HcclChannelDesc *channelDescs,
     310              :     uint32_t channelNum, ChannelHandle *channels)
     311              : {
     312            3 :     if (engine != COMM_ENGINE_AIV) {
     313            3 :         return HCCL_SUCCESS;
     314              :     }
     315              : 
     316            0 :     std::vector<ChannelHandle> hostChannels(channels, channels + channelNum);
     317            0 :     std::vector<ChannelHandle> deviceChannels(hostChannels);
     318            0 :     std::vector<ChannelHandle> mappedDeviceChannels;
     319            0 :     std::vector<ChannelHandle> mappedHostChannels;
     320            0 :     for (uint32_t idx = 0; idx < channelNum; ++idx) {
     321            0 :         if (channelDescs[idx].channelProtocol != COMM_PROTOCOL_ROCE &&
     322            0 :             channelDescs[idx].channelProtocol != COMM_PROTOCOL_UBC_CTP &&
     323            0 :             channelDescs[idx].channelProtocol != COMM_PROTOCOL_UBC_TP) {
     324            0 :             continue;
     325              :         }
     326            0 :         CHK_RET(BuildAivDeviceChannelEntity(channelDescs[idx], hostChannels[idx], deviceChannels[idx]));
     327            0 :         mappedDeviceChannels.emplace_back(deviceChannels[idx]);
     328            0 :         mappedHostChannels.emplace_back(hostChannels[idx]);
     329            0 :         HCCL_INFO("[%s] convert AIV channel success, idx[%u], protocol[%d], hostHandle[0x%llx], devEntity[0x%llx]",
     330              :             __func__, idx, channelDescs[idx].channelProtocol, static_cast<unsigned long long>(hostChannels[idx]),
     331              :             static_cast<unsigned long long>(deviceChannels[idx]));
     332              :     }
     333              : 
     334            0 :     if (!mappedDeviceChannels.empty()) {
     335            0 :         CHK_RET(hcomm::ChannelProcess::RegisterChannelD2HMap(mappedDeviceChannels.data(), mappedHostChannels.data(),
     336              :             static_cast<uint32_t>(mappedDeviceChannels.size())));
     337              :     }
     338              : 
     339            0 :     for (uint32_t idx = 0; idx < channelNum; ++idx) {
     340            0 :         channels[idx] = deviceChannels[idx];
     341              :     }
     342            0 :     return HCCL_SUCCESS;
     343            0 : }
     344            2 : static bool IsUbUrmaChannelProtocol(CommProtocol protocol)
     345              : {
     346            2 :     return protocol == COMM_PROTOCOL_UBC_CTP || protocol == COMM_PROTOCOL_UBC_TP || protocol == COMM_PROTOCOL_UBOE
     347            4 :         || protocol == COMM_PROTOCOL_UBG;
     348              : }
     349              : 
     350            2 : static bool HasUbUrmaChannel(const std::vector<HcclChannelDesc> &channelDescFinals)
     351              : {
     352            3 :     for (const HcclChannelDesc &channelDesc : channelDescFinals) {
     353            2 :         if (IsUbUrmaChannelProtocol(channelDesc.channelProtocol)) {
     354            1 :             return true;
     355              :         }
     356              :     }
     357            1 :     return false;
     358              : }
     359              : 
     360            0 : static void AppendUniqueMemHandle(std::vector<HcclMemHandle> &mergedHandles, HcclMemHandle memHandle)
     361              : {
     362            0 :     if (memHandle == nullptr) {
     363            0 :         return;
     364              :     }
     365            0 :     if (std::find(mergedHandles.begin(), mergedHandles.end(), memHandle) == mergedHandles.end()) {
     366            0 :         mergedHandles.emplace_back(memHandle);
     367              :     }
     368              : }
     369              : 
     370            0 : static HcclResult MergeSymmetricMemHandles(HcclChannelDesc &channelDesc,
     371              :     const std::vector<HcclMemHandle> &symmetricMemHandles, std::vector<HcclMemHandle> &mergedHandles)
     372              : {
     373            0 :     if (!IsUbUrmaChannelProtocol(channelDesc.channelProtocol)) {
     374            0 :         return HCCL_SUCCESS;
     375              :     }
     376            0 :     if (channelDesc.memHandleNum != 0) {
     377            0 :         CHK_PTR_NULL(channelDesc.memHandles);
     378            0 :         for (uint32_t handleIdx = 0; handleIdx < channelDesc.memHandleNum; ++handleIdx) {
     379            0 :             AppendUniqueMemHandle(mergedHandles, channelDesc.memHandles[handleIdx]);
     380              :         }
     381              :     }
     382            0 :     for (HcclMemHandle memHandle : symmetricMemHandles) {
     383            0 :         AppendUniqueMemHandle(mergedHandles, memHandle);
     384              :     }
     385            0 :     CHK_PRT_RET(mergedHandles.size() > static_cast<size_t>(std::numeric_limits<uint32_t>::max()),
     386              :         HCCL_ERROR("[MergeSymmetricMemHandles] merged memHandleNum[%zu] exceeds uint32 max.",
     387              :             mergedHandles.size()), HCCL_E_PARA);
     388            0 :     channelDesc.memHandles = mergedHandles.data();
     389            0 :     channelDesc.memHandleNum = static_cast<uint32_t>(mergedHandles.size());
     390            0 :     return HCCL_SUCCESS;
     391              : }
     392              : 
     393            2 : static HcclResult AppendSymmetricMemHandles(hccl::CollComm *collComm,
     394              :     std::vector<HcclChannelDesc> &channelDescFinals,
     395              :     std::vector<std::vector<HcclMemHandle>> &mergedMemHandles,
     396              :     bool &hasSymmetricMemHandles)
     397              : {
     398            2 :     CHK_PTR_NULL(collComm);
     399            2 :     hasSymmetricMemHandles = false;
     400            2 :     if (!HasUbUrmaChannel(channelDescFinals)) {
     401            1 :         return HCCL_SUCCESS;
     402              :     }
     403              :     // 只有UB/URMA类channel需要追加symmetric memHandle参与建链交换。
     404            1 :     std::vector<HcclMemHandle> symmetricMemHandles;
     405            1 :     CHK_RET(collComm->RegisterPendingSymmetricMemHandles(symmetricMemHandles));
     406            1 :     if (symmetricMemHandles.empty()) {
     407            1 :         return HCCL_SUCCESS;
     408              :     }
     409            0 :     hasSymmetricMemHandles = true;
     410              : 
     411            0 :     mergedMemHandles.clear();
     412            0 :     mergedMemHandles.resize(channelDescFinals.size());
     413            0 :     for (size_t idx = 0; idx < channelDescFinals.size(); ++idx) {
     414            0 :         CHK_RET(MergeSymmetricMemHandles(channelDescFinals[idx], symmetricMemHandles, mergedMemHandles[idx]));
     415              :     }
     416            0 :     HCCL_INFO("[AppendSymmetricMemHandles] append symmetric memHandles success, channelNum[%zu], symMemHandleNum[%zu], "
     417              :         "protocols[UBC_CTP/UBC_TP/UBOE].",
     418              :         channelDescFinals.size(), symmetricMemHandles.size());
     419            0 :     return HCCL_SUCCESS;
     420            1 : }
     421              : 
     422            0 : static HcclResult UpdateSymmetricRemoteMems(hccl::CollComm *collComm, const hccl::MyRank *myRank,
     423              :     const std::vector<HcclChannelDesc> &channelDescFinals, const ChannelHandle *channels, uint32_t channelNum)
     424              : {
     425            0 :     CHK_PTR_NULL(collComm);
     426            0 :     CHK_PTR_NULL(myRank);
     427            0 :     CHK_PTR_NULL(channels);
     428            0 :     for (uint32_t idx = 0; idx < channelNum; ++idx) {
     429            0 :         const HcclChannelDesc &channelDesc = channelDescFinals[idx];
     430            0 :         if (!IsUbUrmaChannelProtocol(channelDesc.channelProtocol)) {
     431            0 :             continue;
     432              :         }
     433            0 :         CommMem *remoteMems = nullptr;
     434            0 :         uint32_t memNum = 0;
     435            0 :         std::vector<std::string> memTags;
     436              :         // CreateChannels完成后,从channel取回交换到的remoteMem/memTag并回填window。
     437            0 :         CHK_RET(myRank->ChannelGetRemoteMems(channels[idx], &memNum, &remoteMems, memTags));
     438            0 :         if (memNum == 0) {
     439            0 :             continue;
     440              :         }
     441            0 :         CHK_RET(collComm->UpdateSymmetricRemoteMem(channelDesc.remoteRank, remoteMems, memTags));
     442            0 :     }
     443            0 :     return HCCL_SUCCESS;
     444              : }
     445              : 
     446            6 : bool CheckCommEngine(const CommEngine engine, const uint32_t opExpansionMode)
     447              : {
     448            6 :     constexpr uint32_t DEFAULT_MODE = 0;
     449            6 :     constexpr uint32_t CCU_MS_MODE = 5;
     450            6 :     constexpr uint32_t CCU_SCHE_MODE = 6;
     451            6 :     if (engine == CommEngine::COMM_ENGINE_CCU) {
     452              :         return opExpansionMode == DEFAULT_MODE
     453            0 :             || opExpansionMode == CCU_MS_MODE
     454            0 :             || opExpansionMode == CCU_SCHE_MODE;
     455              :     }
     456              : 
     457            6 :     return true;
     458              : }
     459              : 
     460            7 : static bool IsAicpuEngine(CommEngine engine)
     461              : {
     462            7 :     return engine == COMM_ENGINE_AICPU || engine == COMM_ENGINE_AICPU_TS;
     463              : }
     464              : 
     465              : constexpr uint32_t CHANNEL_NUM_MAX = 1024 * 1024;  // channel的默认限制最大为1024 * 1024
     466              : 
     467            4 : HcclResult RegisterToClusterMonitor(HcclComm comm)
     468              : {
     469            4 :     HCCL_INFO("[%s] START, comm[%p].", __func__, comm);
     470            4 :     CHK_PRT_RET(comm == nullptr,  HCCL_ERROR("[%s] comm is null", __func__), HCCL_E_PTR);
     471            4 :     auto* hcclComm = static_cast<hccl::hcclComm*>(comm);
     472            4 :     CHK_PTR_NULL(hcclComm);
     473            4 :     if (!hcclComm->IsCommunicatorV2()) {
     474            0 :         HCCL_ERROR("comm is not support [%s]", __func__);
     475            0 :         return HCCL_E_NOT_SUPPORT;
     476              :     }
     477            4 :     hccl::CollComm* collComm = hcclComm->GetCollComm();
     478            4 :     CHK_PTR_NULL(collComm);
     479            4 :     CHK_RET(CollCommMgr::GetInstance()->GetClusterMonitor(collComm->GetDeviceLogicId()).RegisterToClusterMonitor(comm));
     480            2 :     HCCL_INFO("%s Success", __func__);
     481            2 :     return HCCL_SUCCESS;
     482              : }
     483              : 
     484           12 : HcclResult HcclChannelAcquire(HcclComm comm, CommEngine engine, 
     485              :     const HcclChannelDesc* channelDescs, uint32_t channelNum, ChannelHandle* channels)
     486              : {
     487           12 :     HcclUs startut = TIME_NOW();
     488           12 :     u64 beginTime =  Hccl::DlProfFunction::GetInstance().dlMsprofSysCycleTime();
     489              :     EXCEPTION_HANDLE_BEGIN
     490              : 
     491              :     // 入参校验
     492           20 :     CHK_PTR_NULL(comm);
     493           11 :     CHK_PTR_NULL(channelDescs);
     494           11 :     CHK_PTR_NULL(channels);
     495           11 :     CHK_PRT_RET(
     496              :         (channelNum == 0 || channelNum > CHANNEL_NUM_MAX), 
     497              :         HCCL_ERROR("[%s]Invalid channelNum, channelNum[%u], max channel num[%u]",
     498              :         __func__, channelNum, CHANNEL_NUM_MAX), HCCL_E_PARA
     499              :     );
     500              :  
     501           11 :     HcclResult ret = HCCL_SUCCESS;
     502           11 :     hccl::hcclComm *hcclComm = static_cast<hccl::hcclComm *>(comm);
     503           11 :     HCCL_RUN_INFO("Entry-%s channelNum[%u], engine[%s] group[%s]", __func__, channelNum, GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), hcclComm->GetIdentifier().c_str());
     504           11 :     std::vector<HcclChannelDesc> channelDescFinals;
     505           11 :     std::vector<std::vector<HcclMemHandle>> mergedMemHandles;
     506           18 :     for (uint32_t idx = 0; idx < channelNum; idx++) {
     507              :         HcclChannelDesc channelDescFinal;
     508           11 :         HcclChannelDescInit(&channelDescFinal, 1);
     509           11 :         ret = ProcessHcclResPackReq(channelDescs[idx], channelDescFinal, hcclComm);
     510           11 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     511              :             HCCL_ERROR("ProcessHcclResPackReq failed. channelDesc idx[%u], group[%s], engine[%s] channelNum[%llu], ret[%d]", idx, hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, ret), ret);
     512            7 :         channelDescFinals.push_back(channelDescFinal);
     513              :     }
     514              :  
     515            7 :     if (hcclComm->IsCommunicatorV2()) {  // A5
     516            6 :         hccl::CollComm* collComm = hcclComm->GetCollComm();
     517            9 :         CHK_PTR_NULL(collComm);
     518            6 :         const std::string &commTag = hcclComm->GetIdentifier();
     519            6 :         hccl::MyRank* myRank = collComm->GetMyRank();
     520            6 :         CHK_PTR_NULL(myRank);
     521              : 
     522            6 :         s32 deviceLogicId = 0;
     523            6 :         (void)hrtGetDeviceRefresh(&deviceLogicId);
     524            6 :         u32 rankTableCrc = RankTableCrcBridge::GetInstance().ConsumeRankTableJsonCrc(deviceLogicId);
     525            6 :         if (rankTableCrc != 0) {
     526            0 :             CHK_RET(RankConsistencyCheckerV2::GetInstance(deviceLogicId).RecordRankTableCrcV2(rankTableCrc));
     527              :         }
     528            6 :         char hcommPkgName[] = "hcomm";
     529            6 :         int hcommVersion = 0;
     530            6 :         aclError aclRet = aclsysGetVersionNum(hcommPkgName, &hcommVersion);
     531            6 :         CHK_PRT_RET(aclRet != ACL_SUCCESS,
     532              :             HCCL_ERROR("[HcclChannelAcquire] aclsysGetVersionNum failed, aclRet[%d].", aclRet), HCCL_E_INTERNAL);
     533            6 :         std::string curVersion = std::to_string(hcommVersion);
     534            6 :         CHK_RET(RankConsistencyCheckerV2::GetInstance(deviceLogicId).RecordCannVersionV2(curVersion));
     535              :  
     536            6 :         const uint32_t opExpansionMode = myRank->GetOpExpansionMode();
     537            6 :         if (!CheckCommEngine(engine, opExpansionMode)) {
     538            0 :             HCCL_ERROR("[%s] failed, coll comm[%p] group[%s] opExpansionMode[%d] is not supported by CCU engine[%s].", 
     539              :                 __func__, hcclComm, hcclComm->GetIdentifier().c_str(), opExpansionMode, GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str());
     540            0 :             return HcclResult::HCCL_E_PARA;
     541              :         }
     542              : 
     543            6 :         if (engine != CommEngine::COMM_ENGINE_CPU) { // host dpu场景暂不支持cluster monitor
     544            4 :             ret = RegisterToClusterMonitor(comm);
     545            4 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
     546              :                 HCCL_ERROR("RegisterToClusterMonitor failed. group[%s], engine[%s], channelNum[%llu], ret[%d]", hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, ret), ret);
     547              :         }
     548              : 
     549            4 :         if (!GetDebugConfigInited()) {
     550            1 :             InitDebugConfigByEnv();
     551              :         }
     552              : 
     553            4 :         bool hasSymmetricMemHandles = false;
     554            4 :         if (IsAicpuEngine(engine)) {
     555            2 :             CHK_RET(AppendSymmetricMemHandles(collComm, channelDescFinals, mergedMemHandles, hasSymmetricMemHandles));
     556              :         }
     557            4 :         HCCL_INFO("[HcclChannelAcquire] AppendSymmetricMemHandles done, group[%s], engine[%d], channelNum[%u], "
     558              :             "hasSymmetricMemHandles[%d], mergedMemHandleGroups[%zu].",
     559              :             commTag.c_str(), engine, channelNum, hasSymmetricMemHandles, mergedMemHandles.size());
     560            4 :         ret = myRank->CreateChannels(engine, commTag, channelDescFinals.data(), channelNum, channels);
     561            4 :         CHK_PRT_RET((ret == HCCL_E_AGAIN || ret == HCCL_E_UNAVAIL),
     562              :             HCCL_WARNING("CreateChannels group[%s], engine[%s] ret[%d]", commTag.c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), ret), ret);
     563            4 :         CHK_PRT_RET(ret != HCCL_SUCCESS,
     564              :             HCCL_ERROR("CreateChannels failed. group[%s], engine[%s] ret[%d]", commTag.c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), ret), ret);
     565            3 :         if (hasSymmetricMemHandles) {
     566            0 :             CHK_RET(UpdateSymmetricRemoteMems(collComm, myRank, channelDescFinals, channels, channelNum));
     567              :         }
     568            3 :         if (engine == COMM_ENGINE_CPU) {
     569            2 :             HcclCommDfx* hcclCommDfx = collComm->GetHcclCommDfx();
     570            2 :             CHK_PTR_NULL(hcclCommDfx);
     571            2 :             auto callback = hcclCommDfx->GetDpuCallback();
     572            4 :             for (uint32_t idx = 0; idx < channelNum; idx++) {
     573            2 :                 int32_t ret = HcommDpuChannelRegisterDfx(channels[idx], callback);
     574            2 :                 CHK_PRT_RET(ret != HCCL_SUCCESS,
     575              :                     HCCL_ERROR("[HcclChannelAcquire] group[%s] Failed to register DFX callback for channel[%u], ret[%d]", commTag.c_str(), idx, ret),
     576              :                     static_cast<HcclResult>(ret));
     577              :             }
     578            2 :             HCCL_INFO("[HcclChannelAcquire] group[%s] channelNum[%u] Register DFX callback for CPU channels success", commTag.c_str(), channelNum);
     579            2 :         }
     580            3 :         if (IsAicpuEngine(engine)) {
     581            1 :             HCCL_INFO("[HcclChannelAcquire] ReportChannelAicpuKernel start");
     582            1 :             HcclCommDfx* hcclCommDfx = collComm->GetHcclCommDfx();
     583            1 :             CHK_PTR_NULL(hcclCommDfx);
     584            1 :             std::string kernelName = "RunAicpuIndOpChannelInitV2";
     585              :             // 还是kernel的当前无法判断
     586            1 :             ret = hcclCommDfx->ReportKernel(beginTime, commTag, kernelName, SalGetTid(), false);
     587            1 :             CHK_PRT_RET(ret != HCCL_SUCCESS,
     588              :                 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);
     589            1 :         }
     590            9 :     } else {
     591            1 :         hccl::CollComm* collComm = hcclComm->GetCollComm();
     592            1 :         if (collComm != nullptr) {
     593            0 :             hccl::MyRank *myRank = collComm->GetMyRank();
     594            0 :             if (hcclComm->GetConnectMode() != 0 && engine == COMM_ENGINE_CPU && myRank != nullptr) {
     595            0 :                 const std::string &commTag = hcclComm->GetIdentifier();
     596            0 :                 ret = myRank->CreateChannels(engine, commTag, channelDescFinals.data(), channelNum, channels);
     597            0 :             } else {
     598            0 :                 auto& channelMgr = hcclComm->GetIndependentOp().GetChannelManager();
     599            0 :                 ret = channelMgr.ChannelCommCreate(hcclComm->GetIdentifier(), engine,
     600            0 :                     channelDescFinals.data(), channelNum, channels);
     601              :             }
     602              :         } else {
     603            1 :             auto& channelMgr = hcclComm->GetIndependentOp().GetChannelManager();
     604            1 :             ret = channelMgr.ChannelCommCreate(hcclComm->GetIdentifier(), engine,
     605            1 :                 channelDescFinals.data(), channelNum, channels);
     606              :         }
     607              :     }
     608              :  
     609            4 :     CHK_PRT_RET(ret != HCCL_SUCCESS,
     610              :         HCCL_ERROR("[%s] Failed to acquire channel, group[%s], engine[%s], channelNum[%llu], ret[%d]", __func__, hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, ret), ret);
     611              : 
     612            3 :     CHK_RET(ConvertAivChannelHandlesToDevicePtrs(engine, channelDescFinals.data(), channelNum, channels));
     613              :  
     614            3 :     HCCL_RUN_INFO("[%s] acquire channel success, group[%s], engine[%s], channelNum[%llu], take time [%lld]us.", __func__, hcclComm->GetIdentifier().c_str(), GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), channelNum, DURATION_US(TIME_NOW() - startut));
     615           19 :     EXCEPTION_HANDLE_END
     616            3 :     return HCCL_SUCCESS;
     617              : }
     618              : 
     619            0 : HcclResult HcclGroupStart()
     620              : {
     621            0 :     return HcclLegacyGroupStart();
     622              : }
     623              : 
     624            0 : HcclResult HcclGroupEndV2()
     625              : {
     626            0 :     CHK_RET(groupLaunchA5());
     627            0 :     HCCL_INFO("[GroupEnd] to the end");
     628            0 :     return HCCL_SUCCESS;
     629              : }
     630              : 
     631            0 : HcclResult HcclGroupEnd()
     632              : {
     633            0 :     if (hcclGroupDepth == 0) {
     634            0 :         HCCL_ERROR("HcclGroupEnd: not in a group call. Didn't call HcclGroupStart before.");
     635            0 :         return HCCL_E_NOT_SUPPORT;
     636              :     }
     637            0 :     if (--hcclGroupDepth > 0) {
     638            0 :         return HCCL_SUCCESS;
     639              :     }
     640              : 
     641            0 :     HCCL_INFO("[HcclGroupEnd] hcclGroupDepth=[%d]", hcclGroupDepth);
     642              :     /*遇到最后一个HcclGroupEnd才处理group内的所有任务*/
     643            0 :     HCCLV2_FUNC_RUN([&]() -> HcclResult {
     644              :         CHK_RET(HcclLegacyAsyncJobLaunch());
     645              :         return HcclGroupEndV2();
     646              :     }());
     647            0 :     return HcclLegacyGroupEnd();
     648              : }
     649              : 
     650            0 : HcclResult HcclGroupStatusGet(bool *isGroupEnabled)
     651              : {
     652            0 :     CHK_PTR_NULL(isGroupEnabled);
     653            0 :     *isGroupEnabled = (hcclGroupDepth > 0);
     654            0 :     return HCCL_SUCCESS;
     655              : }
        

Generated by: LCOV version 2.0-1