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 "comm_manager.h"
12 :
13 : #include <list>
14 : #include <mutex>
15 : #include <vector>
16 : #include <string>
17 : #include <fstream>
18 : #include <sys/stat.h>
19 : #include <algorithm>
20 : #include <securec.h>
21 : #include <linux/limits.h>
22 : #include <sstream>
23 :
24 : #include "log.h"
25 : #include "hccl/base.h"
26 : #include "hccl_common_v2.h"
27 : #include "hccl/hccl_types.h"
28 : #include "orion_adapter_rts.h"
29 :
30 : #include "tp_manager.h"
31 : #include "inner_net_dev_manager.h"
32 : #include "hccp_hdc_manager.h"
33 : #include "hccp_peer_manager.h"
34 : #include "hccp_tlv_hdc_manager.h"
35 : #include "ccu_driver_handle.h"
36 : #include "rdma_handle_manager.h"
37 : #include "socket_handle_manager.h"
38 : #include "host_socket_handle_manager.h"
39 :
40 : #include "ccu_context_mgr_imp.h"
41 : #include "ccu_res_batch_allocator_legacy.h"
42 : #include "ccu_component.h"
43 : #include "communicator_callback.h"
44 : #include "types.h"
45 :
46 : using namespace std;
47 : using namespace Hccl;
48 :
49 : std::mutex g_commInfoV2CtxMutex;
50 :
51 2 : u64 GetFileSize(const std::string& path)
52 : {
53 : struct stat fileStat;
54 2 : if (stat(path.c_str(), &fileStat) != 0) {
55 3 : HCCL_ERROR("[GetFileSize] Get file stat failed , file path:%s", path.c_str());
56 1 : return 0;
57 : }
58 1 : return static_cast<u64>(fileStat.st_size);
59 : }
60 :
61 21 : HcclResult CcuResAllocAndCtxMgrInit(s32 deviceLogicId)
62 : {
63 : try {
64 21 : CcuComponent::GetInstance(deviceLogicId);
65 21 : CcuResBatchAllocator::GetInstance(deviceLogicId);
66 21 : CtxMgrImp::GetInstance(deviceLogicId);
67 0 : } catch (HcclException& e) {
68 0 : HCCL_ERROR(e.what());
69 0 : return e.GetErrorCode();
70 0 : } catch (exception& e) {
71 0 : HCCL_ERROR(e.what());
72 0 : return HcclResult::HCCL_E_INTERNAL;
73 0 : } catch (...) {
74 0 : HCCL_ERROR("Unknown error occurs!");
75 0 : return HcclResult::HCCL_E_INTERNAL;
76 0 : }
77 21 : return HcclResult::HCCL_SUCCESS;
78 : }
79 :
80 : // 规避默认析构顺序导致单例调用接口时序错误,架构优化后必须删除
81 : // 析构含时序要求接口的单例应在此声明
82 : // 声明顺序与期望析构顺序相反
83 23 : HcclResult CallSingletons()
84 : {
85 23 : s32 deviceLogicId = 0;
86 : try {
87 23 : deviceLogicId = HrtGetDevice();
88 : // 避免设备粒度单例访问错误设备
89 22 : if (deviceLogicId < 0 || static_cast<uint32_t>(deviceLogicId) >= ::MAX_MODULE_DEVICE_NUM) {
90 3 : HCCL_WARNING("[CallSingletons] deviceLogicId[%d] may not have device, passed.", deviceLogicId);
91 1 : return HCCL_E_RUNTIME;
92 : }
93 :
94 : // 不同通信域初始化方式时序不同,hdc manager 重复 init 内部会跳过
95 21 : HccpHdcManager::GetInstance();
96 21 : HccpPeerManager::GetInstance(); // host网卡需要拉起peer模式hccp
97 21 : HccpTlvHdcManager::GetInstance();
98 21 : RdmaHandleManager::GetInstance();
99 21 : InnerNetDevManager::GetInstance();
100 21 : SocketHandleManager::GetInstance();
101 21 : HostSocketHandleManager::GetInstance(); // host网卡需要
102 21 : TpManager::GetInstance(deviceLogicId);
103 1 : } catch (HcclException& e) {
104 3 : HCCL_ERROR(e.what());
105 1 : return e.GetErrorCode();
106 1 : } catch (exception& e) {
107 0 : HCCL_ERROR(e.what());
108 0 : return HcclResult::HCCL_E_INTERNAL;
109 0 : } catch (...) {
110 0 : HCCL_ERROR("Unknown error occurs!");
111 0 : return HcclResult::HCCL_E_INTERNAL;
112 0 : }
113 :
114 21 : if (CcuResAllocAndCtxMgrInit(deviceLogicId) != HCCL_SUCCESS) {
115 : // 遗留问题,处理ccu资源申请失败,走aicpu流程
116 0 : HCCL_ERROR("Ccu res batch allocator or ctx mgr init failed.");
117 0 : return HcclResult::HCCL_E_INTERNAL;
118 : }
119 21 : return HcclResult::HCCL_SUCCESS;
120 : }
121 :
122 777 : CommManager& CommManager::GetInstance(s32 deviceLogicId)
123 : {
124 : // 预留额外一个作为兜底通信域
125 843 : static CommManager commManager[::MAX_MODULE_DEVICE_NUM + 1]; // 使用全局命名空间变量
126 :
127 777 : if (deviceLogicId < 0 || static_cast<uint32_t>(deviceLogicId) > ::MAX_MODULE_DEVICE_NUM) {
128 9 : HCCL_WARNING("[GetInstance] deviceLogicId[%d] is invalid, use backup comm instead.", deviceLogicId);
129 3 : deviceLogicId = ::MAX_MODULE_DEVICE_NUM;
130 : }
131 777 : commManager[deviceLogicId].deviceLogicId = deviceLogicId;
132 777 : return commManager[deviceLogicId];
133 : };
134 :
135 485 : HcclCommInfoV2& CommManager::GetCommInfoV2() { return commInfoV2; }
136 :
137 3 : void CommManager::PrintChannelInfo()
138 : {
139 3 : std::lock_guard<std::mutex> lock(commInfoV2.groupParamsLock);
140 3 : u32 channelNum = 0;
141 3 : s32 logicDevId = HrtGetDevice();
142 9 : HCCL_INFO("[CommManager][PrintChannelInfo]devId[%d].", logicDevId);
143 7 : for (u32 dieId = 0; dieId < MAX_CCU_IODIE_NUM; dieId++) {
144 5 : auto ret = CcuGetChannelSpecNum(logicDevId, dieId, channelNum);
145 5 : if (ret != HCCL_SUCCESS) {
146 3 : HCCL_WARNING(
147 : "[CommManager][PrintChannelInfo]Get channel num failed, devId[%d], dieId[%u]", logicDevId, dieId);
148 1 : return;
149 : }
150 12 : HCCL_RUN_INFO(
151 : "[CommManager][PrintChannelInfo]devId[%d], dieId[%u], Channel num[%u].", logicDevId, dieId, channelNum);
152 : }
153 :
154 4 : for (const auto& group : commInfoV2.hcclGroupMap) {
155 6 : for (u32 dieId = 0; dieId < MAX_CCU_IODIE_NUM; dieId++) {
156 4 : u32 channelCount = group.second.pComm->GetUsedChannelCount(dieId);
157 4 : if (channelCount != 0) {
158 6 : HCCL_RUN_INFO(
159 : "[CommManager][PrintChannelInfo]group[%s], dieId[%u], used channel count[%u].", group.first.c_str(),
160 : dieId, channelCount);
161 : }
162 : }
163 : }
164 3 : }
165 :
166 13 : std::function<void()> CommManager::GetPrintChannelInfoCallback()
167 : {
168 0 : auto callBack = [this]() {
169 0 : PrintChannelInfo();
170 13 : };
171 13 : return callBack;
172 : }
173 :
174 168 : HcclCommInfoV2& GetCommInfoV2(void)
175 : {
176 168 : std::lock_guard<std::mutex> lock(g_commInfoV2CtxMutex);
177 168 : s32 logicDevId = 0;
178 168 : aclError ret = aclrtGetDevice(&logicDevId);
179 168 : if (ret == ACL_SUCCESS && (static_cast<u32>(logicDevId) < ::MAX_MODULE_DEVICE_NUM)) {
180 : /* 当前线程获取到deviceId, 如果是首次使用该deviceId的HcomInfo, 先判断之前是否已经配置过 */
181 167 : HcclCommInfoV2& commInfoV2 = CommManager::GetInstance(logicDevId).GetCommInfoV2();
182 167 : if (!commInfoV2.isUsed) {
183 3 : HCCL_WARNING("[GetCommInfoV2] logicDevId[%d] is not Used.", logicDevId);
184 :
185 1 : HcclCommInfoV2& backupCommInfoV2 = CommManager::GetInstance(::MAX_MODULE_DEVICE_NUM).GetCommInfoV2();
186 1 : if (backupCommInfoV2.isUsed) {
187 0 : return backupCommInfoV2;
188 : }
189 : }
190 167 : commInfoV2.isUsed = true;
191 167 : return commInfoV2;
192 : }
193 :
194 : /* 当前线程没有获取到deviceId, 查找是否有使用过的Ctx */
195 1 : for (u32 i = 0; i <= ::MAX_MODULE_DEVICE_NUM; i++) {
196 1 : HcclCommInfoV2& commInfoV2 = CommManager::GetInstance(i).GetCommInfoV2();
197 1 : if (commInfoV2.isUsed) {
198 3 : HCCL_WARNING("[GetCommInfoV2] no set device Used logicDevId[%u].", i);
199 1 : return commInfoV2;
200 : }
201 : }
202 :
203 0 : HCCL_WARNING("[GetCommInfoV2] HrtGetDevice fail.");
204 : /* 当前线程没有获取到deviceId, 使用兜底Ctx */
205 0 : HcclCommInfoV2& backupCommInfoV2 = CommManager::GetInstance(::MAX_MODULE_DEVICE_NUM).GetCommInfoV2();
206 0 : backupCommInfoV2.isUsed = true;
207 0 : return backupCommInfoV2;
208 168 : }
209 :
210 2 : static HcclResult ParseRankIdsV2(u32 rankNum, const u32* rankIds, std::vector<u32>& groupRanks)
211 : {
212 2 : unordered_set<uint32_t> rankIdSet;
213 2 : std::ostringstream printRankIds;
214 2 : printRankIds << "input rankIds: ";
215 10 : for (u32 i = 0; i < rankNum; i++) {
216 8 : CHK_PTR_NULL(rankIds + i);
217 8 : printRankIds << "rank[";
218 8 : printRankIds << i;
219 8 : printRankIds << "] = ";
220 8 : printRankIds << rankIds[i];
221 8 : if (i < rankNum - 1) {
222 6 : printRankIds << ", ";
223 : }
224 8 : CHK_PRT_RET(
225 : rankIdSet.find(rankIds[i]) != rankIdSet.end(),
226 : HCCL_ERROR(
227 : "[ParseRankIdsV2]errNo[0x%016llx], "
228 : "duplicated rankId[%u] in rankIds.",
229 : HCCL_ERROR_CODE(HCCL_E_PARA), rankIds[i]),
230 : HCCL_E_PARA);
231 8 : rankIdSet.insert(rankIds[i]);
232 8 : groupRanks.push_back(rankIds[i]);
233 : }
234 6 : HCCL_RUN_INFO("Entry-%s: %s", __func__, printRankIds.str().c_str());
235 2 : return HCCL_SUCCESS;
236 2 : }
237 :
238 2 : HcclResult GetHcomRankListV2(u32 rankNum, const u32* rankIds, HcclGroupParamsV2& params, HcclComm globalComm)
239 : {
240 2 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
241 :
242 2 : u32 worldRank = hcomCommInfoV2.commParams.myRank;
243 2 : u32 worldRankSize = hcomCommInfoV2.commParams.rankSize;
244 2 : if (globalComm != nullptr) {
245 0 : Hccl::HcclCommunicator* globalCommunicator = static_cast<Hccl::HcclCommunicator*>(globalComm);
246 0 : CHK_RET(globalCommunicator->GetRankId(worldRank));
247 0 : CHK_RET(globalCommunicator->GetRankSize(&worldRankSize));
248 : }
249 :
250 2 : params.totalRanks = rankNum;
251 2 : params.worldRank = worldRank;
252 2 : params.groupRank = INVALID_VALUE_RANKID;
253 :
254 2 : CHK_RET(ParseRankIdsV2(rankNum, rankIds, params.groupRanks));
255 :
256 2 : std::vector<u32> rankListCopy = params.groupRanks;
257 2 : auto maxIt = std::max_element(rankListCopy.begin(), rankListCopy.end());
258 2 : if (*maxIt >= worldRankSize) {
259 0 : HCCL_ERROR(
260 : "[get][RankList]errNo[0x%016llx] maxRank[%u] is invalid, worldRankSize[%u]", HCOM_ERROR_CODE(HCCL_E_PARA),
261 : *maxIt, worldRankSize);
262 0 : return HCCL_E_PARA;
263 : }
264 :
265 2 : for (u32 i = 0; i < rankNum; i++) {
266 2 : if (params.groupRanks[i] == params.worldRank) {
267 2 : params.groupRank = i;
268 2 : break;
269 : }
270 : }
271 :
272 2 : u32 serverNum = 1; // severNum初始值应为1,代表groupId为0的serverId;
273 2 : params.serverNum = serverNum;
274 :
275 2 : return HCCL_SUCCESS;
276 2 : }
277 :
278 : // 图模式 创建子通信域 V2
279 4 : HcclResult HcomCreateGroupImplV2(const std::string& group, u32 rankNum, const std::vector<u32>& rankIds)
280 : {
281 4 : HcclUs startut = TIME_NOW();
282 : /* 接口交互信息日志 */
283 4 : rankNum = rankIds.size();
284 4 : std::string rankId = "";
285 20 : for (u32 i = 0; i < rankNum; i++) {
286 16 : rankId += std::to_string(rankIds[i]);
287 16 : if (i < rankNum - 1) {
288 12 : rankId += ',';
289 : }
290 : }
291 12 : HCCL_RUN_INFO("Entry-HcomCreateGroup:group[%s], rankNum[%u], rankIds[%s]", group.c_str(), rankNum, rankId.c_str());
292 :
293 4 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
294 4 : CHK_PRT_RET(
295 : hcomCommInfoV2.pComm == nullptr,
296 : HCCL_ERROR("[Create][Group]hcomCommInfoV2.pComm is null, please check if the initialize process is called."),
297 : HCCL_E_PTR);
298 :
299 : /* 已经存在的group不允许再次创建 */
300 4 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
301 4 : if (hcomCommInfoV2.hcclGroupMap.find(group) != hcomCommInfoV2.hcclGroupMap.end()) {
302 6 : HCCL_ERROR(
303 : "[Create][Group]errNo[0x%016llx] group[%s] is already exist", HCOM_ERROR_CODE(HCCL_E_PARA), group.c_str());
304 2 : return HCCL_E_PARA;
305 : }
306 2 : groupParaLock.unlock();
307 :
308 : /* 创建groupParamsV2Tem */
309 2 : HcclGroupParamsV2 groupParamsV2Tem;
310 2 : CHK_RET(GetHcomRankListV2(rankNum, rankIds.data(), groupParamsV2Tem, nullptr));
311 :
312 : /* 如果是groupRank = INVALID_VALUE_RANKID,即本rank不参与create group */
313 2 : if (groupParamsV2Tem.groupRank == INVALID_VALUE_RANKID) {
314 0 : HCCL_ERROR(
315 : "[Create][Group]errNo[0x%016llx] confirm groupRank from worldRank[%d] error",
316 : HCOM_ERROR_CODE(HCCL_E_NOT_FOUND), hcomCommInfoV2.commParams.myRank);
317 0 : return HCCL_E_NOT_FOUND;
318 : }
319 :
320 : /* 创建子通信域 */
321 : Hccl::CommParams subCommParams{
322 2 : group, static_cast<Hccl::RankId>(groupParamsV2Tem.groupRank), rankNum,
323 2 : static_cast<Hccl::RankId>(groupParamsV2Tem.worldRank), hcomCommInfoV2.commParams.devType};
324 2 : auto ret = hcomCommInfoV2.pComm->CreateSubComm(subCommParams, groupParamsV2Tem.groupRanks, groupParamsV2Tem.pComm);
325 2 : CHK_PRT_RET(
326 : ret != HCCL_SUCCESS, HCCL_ERROR("[Create][Group]errNo[0x%016llx] create group failed.", HCOM_ERROR_CODE(ret)),
327 : ret);
328 :
329 2 : CHK_SMART_PTR_NULL(groupParamsV2Tem.pComm);
330 2 : groupParamsV2Tem.pComm->RegisterAcceStateCallBack(CommunicatorCallback());
331 2 : s32 logicDevId = HrtGetDevice();
332 2 : CHK_RET(CommManager::GetInstance(logicDevId)
333 : .SetCommAcceleratorV2(groupParamsV2Tem.pComm.get(), 0)); // 子通信域创建,设置默认accelerator
334 :
335 2 : groupParaLock.lock();
336 2 : hcomCommInfoV2.hcclGroupMap.insert(std::make_pair(group, groupParamsV2Tem));
337 2 : groupParaLock.unlock();
338 :
339 4 : groupParamsV2Tem.pComm->RegisterPrintChannelInfoCallback(
340 4 : CommManager::GetInstance(logicDevId).GetPrintChannelInfoCallback());
341 6 : HCCL_RUN_INFO(
342 : "hcom create group[%s] success, take time [%lld]us", group.c_str(), DURATION_US(TIME_NOW() - startut));
343 :
344 2 : return HCCL_SUCCESS;
345 4 : }
346 :
347 2 : HcclResult HcomDestroyGroupImplV2(const std::string& group)
348 : {
349 2 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
350 :
351 : /* 接口交互信息日志 */
352 6 : HCCL_RUN_INFO("Entry-HcomDestroyGroup:group[%s]", group.c_str());
353 :
354 2 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
355 2 : auto iter = hcomCommInfoV2.hcclGroupMap.find(group);
356 2 : if (iter == hcomCommInfoV2.hcclGroupMap.end()) {
357 0 : HCCL_ERROR(
358 : "[Destroy][Group]errNo[0x%016llx] group[%s] is not exist", HCOM_ERROR_CODE(HCCL_E_PARA), group.c_str());
359 0 : return HCCL_E_PARA;
360 : }
361 2 : hcomCommInfoV2.hcclGroupMap.erase(group);
362 : // 通信域销毁,更新ccu使用情况
363 2 : hcomCommInfoV2.ccuStatus.RemoveCommId(group);
364 :
365 2 : groupParaLock.unlock();
366 :
367 6 : HCCL_RUN_INFO("hcom destroy group[%s] success.", group.c_str());
368 2 : return HCCL_SUCCESS;
369 2 : }
370 :
371 1 : HcclResult HcomGetWorldRankFromGroupRankV2(const char* group, u32 groupRank, u32* worldRank)
372 : {
373 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
374 : // 校验通信域非空
375 1 : CHK_PRT_RET(
376 : hcomCommInfoV2.pComm == nullptr,
377 : HCCL_ERROR("[Get][WorldRank]hcomCommInfoV2.pComm is null, "
378 : "please check if the initialize process is called."),
379 : HCCL_E_PTR);
380 : // 获取group
381 1 : std::string strGroup = (group == nullptr) ? HCCL_WORLD_GROUP : group;
382 1 : if (strGroup == HCCL_WORLD_GROUP) {
383 1 : *worldRank = hcomCommInfoV2.commParams.myRank;
384 3 : HCCL_INFO(
385 : "hcom get world rank success, group[%s], groupRank[%u], worldRank[%u]", strGroup.c_str(), groupRank,
386 : *worldRank);
387 1 : return HCCL_SUCCESS;
388 : }
389 0 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
390 0 : auto iter = hcomCommInfoV2.hcclGroupMap.find(strGroup);
391 0 : if (iter == hcomCommInfoV2.hcclGroupMap.end()) {
392 0 : HCCL_ERROR(
393 : "[Get][WorldRank]errNo[0x%016llx] group[%s] is not exist", HCOM_ERROR_CODE(HCCL_E_PARA), strGroup.c_str());
394 0 : return HCCL_E_PARA;
395 : }
396 :
397 : // groupRanks判空
398 0 : CHK_PRT_RET(
399 : (iter->second).groupRanks.empty(),
400 : HCCL_ERROR(
401 : "[Get][WorldRank]errNo[0x%016llx] group[%s] ranks is empty", HCOM_ERROR_CODE(HCCL_E_INTERNAL),
402 : strGroup.c_str()),
403 : HCCL_E_INTERNAL);
404 :
405 : // 校验groupRank合法性
406 0 : if (groupRank >= (iter->second).totalRanks) {
407 0 : HCCL_ERROR(
408 : "[Get][WorldRank]errNo[0x%016llx] group[%s] groupRank[%u] is invalid", HCOM_ERROR_CODE(HCCL_E_PARA),
409 : strGroup.c_str(), groupRank);
410 0 : return HCCL_E_PARA;
411 : }
412 0 : *worldRank = (iter->second).groupRanks[groupRank];
413 :
414 0 : HCCL_INFO(
415 : "hcom get world rank success, group[%s], groupRank[%u], worldRank[%u]", strGroup.c_str(), groupRank,
416 : *worldRank);
417 0 : return HCCL_SUCCESS;
418 1 : }
419 :
420 1 : HcclResult HcomGetGroupRankFromWorldRankV2(u32 worldRank, const char* group, u32* groupRank)
421 : {
422 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
423 : // 校验通信域非空
424 1 : CHK_PRT_RET(
425 : hcomCommInfoV2.pComm == nullptr,
426 : HCCL_ERROR("[Get][GroupRank]hcomCommInfoV2.pComm is null, "
427 : "please check if the initialize process is called."),
428 : HCCL_E_PTR);
429 : // 校验worldRank合法性
430 1 : if (worldRank >= hcomCommInfoV2.commParams.rankSize) {
431 0 : HCCL_ERROR(
432 : "[Get][GroupRank]errNo[0x%016llx] world[%u] rank is invalid", HCOM_ERROR_CODE(HCCL_E_PARA), worldRank);
433 0 : return HCCL_E_PARA;
434 : }
435 : // 获取group
436 1 : std::string strGroup = (group == nullptr) ? HCCL_WORLD_GROUP : group;
437 1 : if (strGroup == HCCL_WORLD_GROUP) {
438 1 : *groupRank = hcomCommInfoV2.commParams.myRank;
439 3 : HCCL_INFO(
440 : "hcom get group rank success, group[%s], worldRank[%u], groupRank[%u]", strGroup.c_str(), worldRank,
441 : *groupRank);
442 1 : return HCCL_SUCCESS;
443 : }
444 0 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
445 0 : auto iter = hcomCommInfoV2.hcclGroupMap.find(strGroup);
446 0 : if (iter == hcomCommInfoV2.hcclGroupMap.end()) {
447 0 : HCCL_ERROR(
448 : "[Get][GroupRank]errNo[0x%016llx] group[%s] is not exist", HCOM_ERROR_CODE(HCCL_E_PARA), strGroup.c_str());
449 0 : return HCCL_E_PARA;
450 : }
451 :
452 : // groupRanks判空
453 0 : CHK_PRT_RET(
454 : (iter->second).groupRanks.empty(),
455 : HCCL_ERROR(
456 : "[Get][GroupRank]errNo[0x%016llx] group[%s] ranks is empty", HCOM_ERROR_CODE(HCCL_E_INTERNAL),
457 : strGroup.c_str()),
458 : HCCL_E_INTERNAL);
459 :
460 : // 获取groupRank
461 0 : for (u32 rank = 0; rank < (iter->second).totalRanks; rank++) {
462 0 : if (worldRank == (iter->second).groupRanks[rank]) {
463 0 : *groupRank = rank;
464 :
465 0 : HCCL_INFO(
466 : "hcom get group rank success, group[%s], worldRank[%u], groupRank[%u]", strGroup.c_str(), worldRank,
467 : *groupRank);
468 0 : return HCCL_SUCCESS;
469 : }
470 : }
471 0 : return HCCL_E_PARA;
472 1 : }
473 :
474 1 : HcclResult HcomGetRankSizeV2(const char* group, u32* rankSize)
475 : {
476 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
477 :
478 : // 获取group
479 1 : std::string strGroup = (group == nullptr) ? HCCL_WORLD_GROUP : group;
480 1 : if (strGroup == HCCL_WORLD_GROUP) {
481 1 : *rankSize = hcomCommInfoV2.commParams.rankSize;
482 3 : HCCL_INFO("hcom get world rank size success, rankSize[%u]", *rankSize);
483 1 : return HCCL_SUCCESS;
484 : }
485 0 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
486 0 : auto iter = hcomCommInfoV2.hcclGroupMap.find(strGroup);
487 0 : if (iter == hcomCommInfoV2.hcclGroupMap.end()) {
488 0 : HCCL_ERROR(
489 : "[Get][RankSize]errNo[0x%016llx] group[%s] is not exist", HCOM_ERROR_CODE(HCCL_E_PARA), strGroup.c_str());
490 0 : return HCCL_E_PARA;
491 : }
492 0 : CHK_SMART_PTR_NULL((iter->second).pComm);
493 0 : HcclResult ret = (iter->second).pComm->GetRankSize(rankSize);
494 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Get][RankSize]GetRankSize failed."), HCCL_E_PTR);
495 0 : groupParaLock.unlock();
496 :
497 0 : HCCL_INFO("hcom get rank size success, group[%s], rankSize[%u]", strGroup.c_str(), *rankSize);
498 0 : return HCCL_SUCCESS;
499 1 : }
500 :
501 1 : HcclResult HcomGetCommV2(void** commV2)
502 : {
503 1 : CHK_PTR_NULL(commV2);
504 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
505 1 : CHK_PTR_NULL(hcomCommInfoV2.pComm);
506 1 : *commV2 = static_cast<void*>(hcomCommInfoV2.pComm.get());
507 3 : HCCL_INFO("[HcomGetCommV2] success.");
508 1 : return HCCL_SUCCESS;
509 : }
510 :
511 1 : HcclResult HcomGetGroupParamsV2(const char* group, void* groupParams, void** commV2)
512 : {
513 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
514 2 : auto iter = hcomCommInfoV2.hcclGroupMap.find(group);
515 1 : if (iter == hcomCommInfoV2.hcclGroupMap.end()) {
516 0 : HCCL_ERROR("[HcomGetGroupParamsV2] group[%s] not found", group);
517 0 : return HCCL_E_PARA;
518 : }
519 1 : HcclGroupParamsV2& groupParamsV2 = iter->second;
520 1 : HcclGroupParamsV2* groupParamsTem = static_cast<HcclGroupParamsV2*>(groupParams);
521 1 : *groupParamsTem = groupParamsV2;
522 1 : CHK_PTR_NULL(groupParamsV2.pComm);
523 1 : *commV2 = static_cast<Hccl::HcclCommunicator*>(groupParamsV2.pComm.get());
524 3 : HCCL_INFO("[HcomGetGroupParamsV2] success. group[%s]", group);
525 1 : return HCCL_SUCCESS;
526 : }
527 :
528 1 : HcclResult HcomDestroyV2(void)
529 : {
530 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
531 1 : if (hcomCommInfoV2.pComm != nullptr) {
532 1 : bool isSetDeviceByHcomm = false;
533 1 : if (hcomCommInfoV2.devId != HOST_DEVICE_ID) {
534 1 : s32 logicDevId = 0;
535 1 : aclError ret = aclrtGetDevice(&logicDevId);
536 1 : if (ret == ACL_ERROR_RT_CONTEXT_NULL) {
537 : // 若当前线程获取不到context,则由hccl进行setDevice,并在通信域析构完成后resetDevice
538 0 : HrtSetDevice(hcomCommInfoV2.devId);
539 0 : isSetDeviceByHcomm = true;
540 1 : } else if (ret != ACL_SUCCESS) {
541 0 : HCCL_ERROR("[HcomDestroyV2] get device failed, ret[%d]", ret);
542 : }
543 : }
544 : // 通信域销毁,更新ccu使用情况
545 1 : hcomCommInfoV2.ccuStatus.RemoveCommId(hcomCommInfoV2.pComm->GetId());
546 1 : hcomCommInfoV2.pComm = nullptr;
547 1 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
548 :
549 : // 通信域销毁,更新子通信域ccu使用情况
550 2 : for (auto iterGroup : hcomCommInfoV2.hcclGroupMap) {
551 1 : hcomCommInfoV2.ccuStatus.RemoveCommId(iterGroup.first);
552 1 : }
553 1 : hcomCommInfoV2.hcclGroupMap.clear();
554 1 : if (isSetDeviceByHcomm) {
555 0 : HrtResetDevice(hcomCommInfoV2.devId);
556 : }
557 1 : }
558 1 : return HCCL_SUCCESS;
559 : }
560 :
561 1 : static HcclResult GetRankTableInfo(const char* rankTablePath, std::string& ranktableInfo)
562 : {
563 : // 校验文件是否存在
564 1 : char resolvedPath[PATH_MAX] = {0};
565 1 : if (realpath(rankTablePath, resolvedPath) == nullptr) {
566 0 : HCCL_ERROR("RanktableRealPath: %s is not a valid real path", rankTablePath);
567 0 : return HCCL_E_INTERNAL;
568 : }
569 :
570 3 : HCCL_INFO("waiting for json file load complete");
571 1 : u64 ranktableFileSize = GetFileSize(resolvedPath);
572 1 : if (ranktableFileSize > RANKTABLE_FILE_MAX_SIZE || ranktableFileSize <= 0) {
573 3 : HCCL_ERROR(
574 : "[GetRankTableInfo] ranktablefile size: %u, ranktable must be greater than 0 and less than %u",
575 : ranktableFileSize, RANKTABLE_FILE_MAX_SIZE);
576 1 : return HCCL_E_OPEN_FILE_FAILURE;
577 : }
578 :
579 0 : std::ifstream infoFile(resolvedPath, std::ifstream::in);
580 0 : if (!infoFile) {
581 0 : HCCL_ERROR("open file %s failed", resolvedPath);
582 0 : return HCCL_E_INTERNAL;
583 : }
584 :
585 0 : std::stringstream rankTableStr;
586 0 : rankTableStr << infoFile.rdbuf();
587 0 : ranktableInfo = rankTableStr.str();
588 :
589 0 : return HCCL_SUCCESS;
590 0 : }
591 :
592 : // 图模式 创建全局通信域 V2
593 1 : HcclResult HcomInitByFileV2(const char* rankTablePath, const char* identify)
594 : {
595 : // 待解决:目前主要为了芯片验证,非最终版本
596 3 : HCCL_RUN_INFO("Entry-HcomInitByFile V950, ranktable[%s], identify[%s]", rankTablePath, identify);
597 :
598 : // 解析myRank
599 : s32 myRank;
600 : try {
601 1 : myRank = std::atoi(identify);
602 : } catch (...) {
603 : HCCL_ERROR("atoi(identify) failed!");
604 : return HCCL_E_INTERNAL;
605 : }
606 :
607 1 : CallSingletons(); // 临时规避,在初始化通信域前声明单例保证时序
608 :
609 : // 防止重复调用初始化
610 1 : string commId(HCCL_WORLD_GROUP);
611 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
612 1 : CHK_PRT_RET(
613 : hcomCommInfoV2.hcclGroupMap.find(commId) != hcomCommInfoV2.hcclGroupMap.end(),
614 : HCCL_ERROR(
615 : "[Init][CheckOpBasedHcom]errNo[0x%016llx] The comm name[%s] already exists in Group2Comm map.",
616 : HCCL_ERROR_CODE(HCCL_E_PARA), commId.c_str()),
617 : HCCL_E_PARA);
618 :
619 : // 解析ranktable
620 1 : std::string ranktableInfo;
621 1 : HcclResult ret = GetRankTableInfo(rankTablePath, ranktableInfo);
622 4 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[HcomInitByFile] get ranktable info failed"), ret);
623 :
624 0 : bool devUsed = false;
625 0 : bool isWorldGroup = true;
626 : // 临时修改这个为4,后边要改掉这个,在初始流程中,解析完虚拟拓扑后,添加ranksize
627 : Hccl::CommParams commParams{commId,
628 : static_cast<Hccl::RankId>(myRank),
629 : 0,
630 : static_cast<Hccl::RankId>(myRank),
631 0 : Hccl::HrtGetDeviceType(),
632 : devUsed,
633 0 : isWorldGroup};
634 0 : hcomCommInfoV2.pComm.reset(new (std::nothrow) Hccl::HcclCommunicator(commParams));
635 0 : CHK_PTR_NULL(hcomCommInfoV2.pComm);
636 0 : auto res = hcomCommInfoV2.pComm->Init(ranktableInfo);
637 0 : CHK_PRT_RET(
638 : res != HcclResult::HCCL_SUCCESS, HCCL_ERROR("[HcomInitByFile] Hccl::Communicator Init failed, res %d", res),
639 : HCCL_E_INTERNAL);
640 :
641 0 : hcomCommInfoV2.pComm->RegisterAcceStateCallBack(CommunicatorCallback());
642 0 : s32 logicDevId = HrtGetDevice();
643 0 : CHK_RET(CommManager::GetInstance(logicDevId)
644 : .SetCommAcceleratorV2(hcomCommInfoV2.pComm.get(), 0)); // 全局通信域创建,设置默认accelerator
645 :
646 0 : res = hcomCommInfoV2.pComm->GetRankSize(&commParams.rankSize);
647 0 : CHK_PRT_RET(
648 : res != HCCL_SUCCESS,
649 : HCCL_ERROR("[HcomInitByFile] Hccl::Communicator GetRankSize failed, rankSize = %u", commParams.rankSize), res);
650 0 : hcomCommInfoV2.commParams = commParams;
651 :
652 0 : HcclGroupParamsV2 params;
653 0 : params.pComm = hcomCommInfoV2.pComm;
654 0 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
655 0 : hcomCommInfoV2.hcclGroupMap[commId] = params;
656 0 : groupParaLock.unlock();
657 :
658 0 : hcomCommInfoV2.pComm->RegisterPrintChannelInfoCallback(
659 0 : CommManager::GetInstance(logicDevId).GetPrintChannelInfoCallback());
660 :
661 0 : HCCL_INFO(
662 : "[HcomInitByFile] HcomInitByFile success! logicDevId[%d], commId[%s]", logicDevId, commParams.commId.c_str());
663 :
664 0 : return HCCL_SUCCESS;
665 1 : }
666 :
667 1 : HcclResult HcomInitByStringV2(const char* rankTableM, const char* identify)
668 : {
669 : // 待解决:目前主要为了芯片验证,非最终版本
670 3 : HCCL_RUN_INFO("Entry-HcomInitByString V950, rankTableM[%s], identify[%s]", rankTableM, identify);
671 :
672 : // 解析myRank
673 : s32 myRank;
674 : try {
675 1 : myRank = std::atoi(identify);
676 : } catch (...) {
677 : HCCL_ERROR("atoi(identify) failed!");
678 : return HCCL_E_INTERNAL;
679 : }
680 :
681 1 : CallSingletons(); // 临时规避,在初始化通信域前声明单例保证时序
682 :
683 : // 防止重复调用初始化
684 1 : string commId(HCCL_WORLD_GROUP);
685 1 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
686 1 : CHK_PRT_RET(
687 : hcomCommInfoV2.hcclGroupMap.find(commId) != hcomCommInfoV2.hcclGroupMap.end(),
688 : HCCL_ERROR(
689 : "[Init][CheckOpBasedHcom]errNo[0x%016llx] The comm name[%s] already exists in Group2Comm map.",
690 : HCCL_ERROR_CODE(HCCL_E_PARA), commId.c_str()),
691 : HCCL_E_PARA);
692 :
693 1 : bool devUsed = false;
694 1 : bool isWorldGroup = true;
695 : // 临时修改这个为4,后边要改掉这个,在初始流程中,解析完虚拟拓扑后,添加ranksize
696 : Hccl::CommParams commParams{commId,
697 : static_cast<Hccl::RankId>(myRank),
698 : 0,
699 : static_cast<Hccl::RankId>(myRank),
700 0 : Hccl::HrtGetDeviceType(),
701 : devUsed,
702 1 : isWorldGroup};
703 1 : hcomCommInfoV2.pComm.reset(new (std::nothrow) Hccl::HcclCommunicator(commParams));
704 1 : CHK_PTR_NULL(hcomCommInfoV2.pComm);
705 2 : auto res = hcomCommInfoV2.pComm->Init(rankTableM);
706 1 : CHK_PRT_RET(
707 : res != HcclResult::HCCL_SUCCESS, HCCL_ERROR("[HcomInitByString] Hccl::Communicator Init failed, res %d", res),
708 : HCCL_E_INTERNAL);
709 :
710 1 : hcomCommInfoV2.pComm->RegisterAcceStateCallBack(CommunicatorCallback());
711 1 : s32 logicDevId = HrtGetDevice();
712 1 : hcomCommInfoV2.devId = logicDevId;
713 1 : CHK_RET(CommManager::GetInstance(logicDevId)
714 : .SetCommAcceleratorV2(hcomCommInfoV2.pComm.get(), 0)); // 全局通信域创建,设置默认accelerator
715 :
716 1 : res = hcomCommInfoV2.pComm->GetRankSize(&commParams.rankSize);
717 1 : CHK_PRT_RET(
718 : res != HCCL_SUCCESS,
719 : HCCL_ERROR("[HcomInitByString] Hccl::Communicator GetRankSize failed, rankSize = %u", commParams.rankSize),
720 : res);
721 1 : hcomCommInfoV2.commParams = commParams;
722 :
723 1 : HcclGroupParamsV2 params;
724 1 : params.pComm = hcomCommInfoV2.pComm;
725 1 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
726 1 : hcomCommInfoV2.hcclGroupMap[commId] = params;
727 1 : groupParaLock.unlock();
728 :
729 2 : hcomCommInfoV2.pComm->RegisterPrintChannelInfoCallback(
730 2 : CommManager::GetInstance(logicDevId).GetPrintChannelInfoCallback());
731 :
732 3 : HCCL_INFO(
733 : "[HcomInitByString] HcomInitByString success! logicDevId[%d], commId[%s]", logicDevId,
734 : commParams.commId.c_str());
735 :
736 1 : return HCCL_SUCCESS;
737 1 : }
738 :
739 17 : void CcuStatus::RemoveCommId(const std::string& commId)
740 : {
741 17 : auto itMs = std::find(useMsCommIds.begin(), useMsCommIds.end(), commId);
742 17 : if (itMs != useMsCommIds.end()) {
743 0 : HCCL_DEBUG("[CcuStatus][%s] commId[%s] used ccu ms, removed", __func__, commId.c_str());
744 0 : useMsCommIds.erase(itMs);
745 : }
746 :
747 17 : auto itSched = std::find(useSchedCommIds.begin(), useSchedCommIds.end(), commId);
748 17 : if (itSched != useSchedCommIds.end()) {
749 3 : HCCL_DEBUG("[CcuStatus][%s] commId[%s] used ccu sched, removed", __func__, commId.c_str());
750 1 : useSchedCommIds.erase(itSched);
751 : }
752 17 : }
753 :
754 16 : bool CcuStatus::IsMsAvailable(const std::string& commId) const
755 : {
756 16 : auto itMs = std::find(useMsCommIds.begin(), useMsCommIds.end(), commId);
757 : // ms没有通信域使用,或者就是传入通信域在使用,则可用
758 16 : return (useMsCommIds.size() < MAX_NUM_COMM_USING_MS) || (itMs != useMsCommIds.end());
759 : }
760 :
761 6 : HcclResult CcuStatus::InsertCommId(const std::string& commId, bool isUsingCcuMs, bool isUsingCcuSched)
762 : {
763 : // 先删再加,避免重复添加到两种模式
764 6 : RemoveCommId(commId);
765 : // ccu ms 没有被使用过,则将ccu ms 标记为已使用
766 6 : if (isUsingCcuMs) {
767 5 : CHK_RET(InsertMsCommId(commId));
768 4 : } else if (isUsingCcuSched) {
769 2 : InsertSchedCommId(commId);
770 : } else {
771 6 : HCCL_DEBUG("NotUsingCcu comm [%s]", commId.c_str());
772 : }
773 5 : return HCCL_SUCCESS;
774 : }
775 :
776 2 : HcclResult CcuStatus::InsertMsCommId(const std::string& commId)
777 : {
778 2 : if (!IsMsAvailable(commId)) {
779 3 : HCCL_WARNING(
780 : "[%s] ccu ms has been used by comm [%s], no more than 2 comms can use ccu ms at the same time.", __func__,
781 : (*(useMsCommIds.begin())).c_str());
782 1 : return HCCL_E_INTERNAL;
783 : }
784 3 : HCCL_DEBUG("[%s] UsingCcuMs comm [%s]", __func__, commId.c_str());
785 1 : useMsCommIds.push_back(commId);
786 1 : return HCCL_SUCCESS;
787 : }
788 :
789 2 : void CcuStatus::InsertSchedCommId(const std::string& commId)
790 : {
791 6 : HCCL_DEBUG("[%s] UsingCcuSched comm [%s]", __func__, commId.c_str());
792 2 : useSchedCommIds.push_back(commId);
793 2 : }
794 :
795 14 : HcclResult CommManager::SetCommAcceleratorV2(Hccl::HcclCommunicator* communicator, int32_t accelerator)
796 : {
797 14 : CHK_PTR_NULL(communicator);
798 14 : if (accelerator < static_cast<int32_t>(HcclAccelerator::DEFAULT)
799 14 : || accelerator > static_cast<int32_t>(HcclAccelerator::AICPU)) {
800 0 : HCCL_ERROR("[SetCommAcceleratorV2] Invalid accelerator value [%d], valid range is [0,7]", accelerator);
801 0 : return HCCL_E_NOT_SUPPORT;
802 : }
803 14 : HcclAccelerator hcclAccelerator = static_cast<HcclAccelerator::Value>(accelerator);
804 :
805 14 : HcclCommInfoV2& opbasedCommInfoV2 = GetCommInfoV2();
806 : // 通过进程锁看护,避免多个通信域同时占用CCU_MS
807 14 : std::unique_lock<std::mutex> lock(opbasedCommInfoV2.groupParamsLock);
808 28 : if ((hcclAccelerator == HcclAccelerator::CCU_MS || hcclAccelerator == HcclAccelerator::CCU_SCHED)
809 28 : && !isCcuAvailable) {
810 0 : HCCL_WARNING("CCU not support reuse in single device multi-precess services, accelerator fallback AICPU_TS");
811 0 : hcclAccelerator = HcclAccelerator::AICPU_TS;
812 : }
813 14 : bool isMsAvailable = opbasedCommInfoV2.ccuStatus.IsMsAvailable(communicator->GetId());
814 42 : HCCL_INFO(
815 : "[CommManager][%s] hcclAccelerator is [%s], isMsAvailable is [%d]", __func__,
816 : hcclAccelerator.Describe().c_str(), isMsAvailable);
817 14 : CHK_RET(communicator->SetAccelerator(hcclAccelerator, isMsAvailable));
818 14 : return HCCL_SUCCESS;
819 14 : }
820 :
821 0 : std::shared_ptr<Hccl::CcuDriverHandle> CommManager::GetCcuDriver()
822 : {
823 0 : if (isCcuAvailable == true && ccuDriverHandle == nullptr) {
824 0 : ccuDriverHandle = std::make_shared<Hccl::CcuDriverHandle>(deviceLogicId);
825 0 : if (ccuDriverHandle->Init() == HCCL_E_UNAVAIL) {
826 0 : isCcuAvailable = false;
827 0 : ccuDriverHandle = nullptr;
828 0 : HCCL_WARNING("[CommManager::GetCcuDriver]Tlv already open, isCcuAvailable updated to false");
829 : }
830 : }
831 0 : return ccuDriverHandle;
832 : }
833 :
834 273 : void CommManager::DeinitCcuDriver()
835 : {
836 273 : if (ccuDriverHandle.use_count() == 1) {
837 0 : ccuDriverHandle = nullptr;
838 : }
839 273 : }
840 :
841 3 : HcclResult HcomGetCcuTaskInfo(const std::string& group, void* tilingData, void* ccuTaskGroup)
842 : {
843 3 : CHK_PTR_NULL(tilingData);
844 3 : CHK_PTR_NULL(ccuTaskGroup);
845 3 : CHK_PRT_RET(group.empty(), HCCL_ERROR("[HcomGetCcuTaskInfo] group is null"), HCCL_E_PARA);
846 :
847 : /* 接口交互信息日志 */
848 9 : HCCL_RUN_INFO("HcomGetCcuTaskInfo:group[%s]", group.c_str());
849 :
850 3 : HcclCommInfoV2& hcomCommInfoV2 = GetCommInfoV2();
851 :
852 3 : std::unique_lock<std::mutex> groupParaLock(hcomCommInfoV2.groupParamsLock);
853 3 : auto iter = hcomCommInfoV2.hcclGroupMap.find(group);
854 3 : if (iter == hcomCommInfoV2.hcclGroupMap.end()) {
855 3 : HCCL_ERROR(
856 : "[HcomGetCcuTaskInfo]errNo[0x%016llx] group[%s] is not exist", HCOM_ERROR_CODE(HCCL_E_PARA), group.c_str());
857 1 : return HCCL_E_PARA;
858 : }
859 2 : HcclGroupParamsV2& groupParam = iter->second;
860 2 : Hccl::HcclCommunicator* comm = static_cast<Hccl::HcclCommunicator*>(groupParam.pComm.get());
861 2 : CHK_PTR_NULL(comm);
862 2 : auto ret = comm->GetCcuTaskInfo(tilingData, ccuTaskGroup);
863 2 : if (ret != HCCL_SUCCESS) {
864 3 : HCCL_ERROR("[HcomGetCcuTaskInfo] GetCcuTaskInfo failed.");
865 1 : return HCCL_E_INTERNAL;
866 : }
867 :
868 3 : HCCL_RUN_INFO("HcomGetCcuTaskInfo success group[%s]", group.c_str());
869 1 : return HCCL_SUCCESS;
870 3 : }
|