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 : }
|