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