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