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_base.h"
12 : #include <arpa/inet.h>
13 : #include <securec.h>
14 :
15 : #include "externalinput_pub.h"
16 : #include "hccl_common.h"
17 : #include "device_capacity.h"
18 : #include "p2p_mgmt_pub.h"
19 : #include "rank_consistentcy_checker.h"
20 :
21 : namespace hccl {
22 : constexpr s32 HCCL_DEFAULT_INITIAL_VALUE = -1;
23 28 : CommBase::CommBase(const std::string &collectiveId, const u32 userRank, const u32 userRankSize,
24 : const u32 rank, const u32 rankSize, const std::vector<RankInfo> paraVector, const TopoType topoFlag,
25 : const HcclDispatcher dispatcher, const std::unique_ptr<NotifyPool> ¬ifyPool,
26 : std::map<HcclIpAddress, HcclNetDevCtx> &netDevCtxMap,
27 : const IntraExchanger &exchanger, const DeviceMem &inputMem, const DeviceMem &outputMem,
28 : const bool isUsedRdmaLevel0,
29 : const std::string &tag, const NICDeployment nicDeployInner,
30 : bool isAlltoAllCommMesh, const bool useOneDoorbell, const bool isAicpuModeEn, const u32 rankRoot,
31 28 : const bool isHaveCpuRank, const bool useSuperPodMode, DeviceMem expMem)
32 28 : : linkDummy_(nullptr), collectiveId_(collectiveId), userRank_(userRank),
33 28 : userRankSize_(userRankSize), rank_(rank), rankSize_(rankSize), paraVector_(paraVector),
34 56 : transportType_(rankSize, TransportType::TRANS_TYPE_RESERVED),
35 28 : deviceLogicId_(HCCL_DEFAULT_INITIAL_VALUE), devicePhyId_(INVALID_UINT),
36 112 : topoFlag_(topoFlag), tag_(tag), transportInfo_(rankSize), rankMap_(userRankSize, INVALID_VALUE_RANKID),
37 56 : userRankMap_(rankSize, INVALID_VALUE_RANKID), dispatcher_(dispatcher), notifyPool_(notifyPool),
38 28 : netDevCtxMap_(netDevCtxMap), exchanger_(exchanger), inputMem_(inputMem), outputMem_(outputMem),
39 28 : isUsedRdmaLevel0_(isUsedRdmaLevel0),
40 28 : dstInterServerMap_(), dstInterClientMap_(), dstIntraServerVec_(), dstIntraClientVec_(),
41 28 : linkThreads_(), threadsRapplyNum_(0),
42 28 : shmDev_(0), isAlltoAllCommMesh_(isAlltoAllCommMesh),
43 56 : nicDeployInner_(nicDeployInner), isNeedHeterogP2P_(false),
44 56 : useOneDoorbell_(useOneDoorbell), isAicpuModeEn_(isAicpuModeEn), subUserRankRoot_(rankRoot),
45 140 : isHaveCpuRank_(isHaveCpuRank), useSuperPodMode_(useSuperPodMode), expMem_(expMem)
46 : {
47 28 : }
48 :
49 181 : CommBase::~CommBase()
50 : {
51 28 : (void)DeInit();
52 41 : }
53 :
54 28 : HcclResult CommBase::DeInit()
55 : {
56 28 : for (u32 index = 0; index < linkThreads_.size(); index++) {
57 0 : if (linkThreads_[index]) {
58 0 : if (linkThreads_[index]->joinable()) {
59 0 : HCCL_DEBUG("Joining Link Thread[%u]", index);
60 0 : linkThreads_[index]->join(); // 等待线程执行后释放资源
61 : }
62 :
63 0 : HcclResult ret = hrtResetDevice(deviceLogicId_); // 防止线程里面异常退出,在进程中reset
64 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[CommBase][DeInit]comm base reset device[%d] failed",
65 : deviceLogicId_), ret);
66 : }
67 : }
68 28 : linkThreads_.clear();
69 :
70 57 : for (u32 i = 0; i < transportInfo_.size(); i++) {
71 29 : if (transportInfo_[i]) { // 使用对应类型的port销毁
72 0 : CHK_RET(transportInfo_[i]->DeInit());
73 : }
74 : }
75 :
76 28 : return HCCL_SUCCESS;
77 : }
78 :
79 18 : HcclResult CommBase::Init()
80 : {
81 : // 获取rank->userrank以及userrank->rank的映射关系
82 18 : CHK_RET(SetRankMap());
83 :
84 18 : if (!IsGeneralServer()) {
85 : // 获取当前线程操作的设备ID
86 18 : CHK_RET(hrtGetDevice(&deviceLogicId_));
87 18 : CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(deviceLogicId_), devicePhyId_));
88 : }
89 :
90 18 : intraSocketsMap_.insert(exchanger_.socketsMap.begin(), exchanger_.socketsMap.end());
91 :
92 : // 创建当前rank与其他rank之间的link(RDMA异步、TCP)
93 18 : CHK_RET(CreateLinks());
94 :
95 : // 校验当前rank与dst rank建链的链路有效性
96 18 : CHK_RET(CheckLinks());
97 :
98 : // task多线程并行下发,根据当前transport创建vtransport信息
99 18 : CHK_RET(CreateVirturalTransport());
100 :
101 18 : return HCCL_SUCCESS;
102 : }
103 :
104 1 : std::shared_ptr<Transport> &CommBase::GetTransportByRank(const u32 dstRank)
105 : {
106 1 : if (transportInfo_.size() <= dstRank) {
107 1 : HCCL_ERROR("[Get][TransportByRank]dstRank[%u] is bigger than link size[%llu]", dstRank, transportInfo_.size());
108 1 : return linkDummy_;
109 : }
110 :
111 0 : return transportInfo_[dstRank];
112 : }
113 :
114 6 : HcclResult CommBase::GetRankByUserRank(const u32 userRank, u32 &rank) const
115 : {
116 6 : if (rankMap_.size() > userRank) {
117 4 : rank = rankMap_[userRank];
118 4 : if (rank == INVALID_VALUE_RANKID) {
119 2 : HCCL_INFO("This userRank[%u] is not in this sub communication. ", userRank);
120 2 : return HCCL_E_NOT_FOUND;
121 : }
122 2 : return HCCL_SUCCESS;
123 : }
124 :
125 2 : HCCL_ERROR("[Get][RankByUserRank]This userRank[%u] is invalid. ", userRank);
126 2 : rank = INVALID_VALUE_RANKID;
127 2 : return HCCL_E_PARA;
128 : }
129 :
130 6 : HcclResult CommBase::GetUserRankByRank(const u32 rank, u32 &userRank) const
131 : {
132 6 : if (userRankMap_.size() > rank) {
133 2 : if (userRankMap_[rank] == INVALID_VALUE_RANKID) {
134 0 : HCCL_INFO("This rank[%u] is not in this sub communication.", rank);
135 0 : userRank = INVALID_VALUE_RANKID;
136 0 : return HCCL_E_NOT_FOUND;
137 : }
138 :
139 2 : userRank = userRankMap_[rank];
140 2 : return HCCL_SUCCESS;
141 : }
142 :
143 4 : HCCL_ERROR("[Get][UserRankByRank]This rank[%u] is invalid.", rank);
144 4 : userRank = INVALID_VALUE_RANKID;
145 4 : return HCCL_E_PARA;
146 : }
147 :
148 18 : HcclResult CommBase::CreateLinks()
149 : {
150 18 : HCCL_DEBUG("[CreateLinks] [comm_base] rankSize_[%u]", rankSize_);
151 18 : if (rankSize_ == HCCL_RANK_SIZE_EQ_ONE) {
152 18 : HCCL_INFO("comm base needn't to create links, rankSize_[%u].", rankSize_);
153 18 : return HCCL_SUCCESS;
154 : }
155 :
156 0 : CHK_RET(CalcLink());
157 0 : u32 threadsNum = dstInterClientMap_.size() + dstIntraClientVec_.size() +
158 0 : dstInterServerMap_.size() + dstIntraServerVec_.size();
159 0 : CHK_PRT_RET((threadsNum == 0), HCCL_ERROR("[Create][Links]no link to create, threadsNum[%u]", threadsNum),
160 : HCCL_E_INTERNAL);
161 :
162 0 : linkThreads_.resize(threadsNum);
163 0 : HCCL_INFO("comm base threads info:link threads size[%llu], dst inter client map size[%llu], " \
164 : "dst intra client vec size[%llu], dst inter server map size[%llu], dst intra server vec size[%llu]",
165 : linkThreads_.size(), dstInterClientMap_.size(), dstIntraClientVec_.size(),
166 : dstInterServerMap_.size(), dstIntraServerVec_.size());
167 :
168 0 : CHK_RET(CreateExchangerNetwork());
169 :
170 0 : CHK_RET(CreateIntraLinks());
171 :
172 0 : CHK_RET(CreateInterLinks());
173 :
174 0 : bool check = (threadsRapplyNum_ != linkThreads_.size());
175 0 : CHK_PRT_RET(check, HCCL_ERROR("[Create][Links]comm base rapply num[%u] is not equal to link threads[%llu]",
176 : threadsRapplyNum_, linkThreads_.size()), HCCL_E_INTERNAL);
177 :
178 0 : for (u32 index = 0; index < linkThreads_.size(); index++) {
179 0 : if (linkThreads_[index] == nullptr) {
180 0 : continue;
181 : }
182 0 : if (linkThreads_[index]->joinable()) {
183 0 : HCCL_DEBUG("Joining Link Thread[%u]", index);
184 0 : linkThreads_[index]->join(); // 等待线程执行完毕
185 : }
186 0 : if (!IsGeneralServer()) {
187 0 : CHK_RET(hrtResetDevice(deviceLogicId_)); // 防止线程里面异常退出,在进程中reset
188 : }
189 : }
190 0 : linkThreads_.clear();
191 : // 建链结束立即释放socket资源(添加判断host网卡走的是TCP就不释放资源)
192 0 : if (pyhIdResourseSockets_.size()) {
193 0 : for (auto &iter : pyhIdResourseSockets_) {
194 0 : iter.second->DestroySockets(tag_);
195 : }
196 0 : } else if (!GetExternalInputHcclIsTcpMode() && interSocketManager_ != nullptr) {
197 : // 建链结束,关闭socket
198 0 : interSocketManager_->DestroySockets(tag_);
199 : }
200 0 : return HCCL_SUCCESS;
201 : }
202 :
203 0 : HcclResult CommBase::CalcLink()
204 : {
205 0 : return HCCL_SUCCESS;
206 : }
207 :
208 0 : u32 CommBase::GetSocketsPerLink()
209 : {
210 0 : return 1;
211 : }
212 :
213 2 : bool CommBase::NeedDataReceivedAck()
214 : {
215 2 : return false;
216 : }
217 :
218 : // 获取rank间的link type
219 1 : HcclResult CommBase::SetTransportType(const u32 dstRank)
220 : {
221 1 : LinkTypeInServer linkType = LinkTypeInServer::RESERVED_LINK_TYPE;
222 :
223 : // 适配910_93的RDMA+SIO ring,创建RDMA类型下的SIO连接
224 1 : if (linkType == LinkTypeInServer::SIO_TYPE && paraVector_[rank_].deviceType == DevType::DEV_TYPE_910_93) {
225 0 : transportType_[dstRank] = TransportType::TRANS_TYPE_P2P;
226 1 : } else if (paraVector_[rank_].serverId == paraVector_[dstRank].serverId) { // 判断是否在同一个server
227 : // Server内判断是否使用rdma
228 0 : if (isUsedRdmaLevel0_ || isAlltoAllCommMesh_ ||
229 0 : (paraVector_[rank_].deviceType != DevType::DEV_TYPE_310P3 &&
230 0 : (paraVector_[rank_].devicePhyId / HCCL_AISERVER_DEVICE_NUM !=
231 0 : paraVector_[dstRank].devicePhyId / HCCL_AISERVER_DEVICE_NUM))) {
232 0 : transportType_[dstRank] = TransportType::TRANS_TYPE_IBV_EXP;
233 : } else {
234 0 : transportType_[dstRank] = TransportType::TRANS_TYPE_P2P;
235 : }
236 : } else { // server间
237 1 : if (IsSupportInterHccs(dstRank)) {
238 : // 超节点内节点间走HCCS通信
239 0 : transportType_[dstRank] = TransportType::TRANS_TYPE_P2P;
240 : } else {
241 1 : transportType_[dstRank] = TransportType::TRANS_TYPE_IBV_EXP;
242 : }
243 : }
244 :
245 1 : HCCL_INFO("SetTransportType: dstRank[%u] transport_type[%d]", dstRank, transportType_[dstRank]);
246 1 : return HCCL_SUCCESS;
247 : }
248 :
249 0 : HcclResult CommBase::RunTemplateAlg(const std::unique_ptr<AlgTemplateBase> &tempAlg)
250 : {
251 0 : HcclResult ret = tempAlg->RunAsync(Rank(), RankSize(), transportInfo_);
252 0 : CHK_PRT_RET(ret == HCCL_E_AGAIN, HCCL_WARNING("[Run][AlgTemplateBase]group has been destroyed. Break!"), ret);
253 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][AlgTemplateBase]comm base run tempAlg "\
254 : "rank[%u] rank size[%u] failed", rank_, rankSize_), ret);
255 0 : return HCCL_SUCCESS;
256 : }
257 :
258 0 : HcclResult CommBase::RunTemplateAlgStaged(const std::unique_ptr<AlgTemplateBase> &tempAlg, const RunStage &stage)
259 : {
260 0 : HcclResult ret = tempAlg->RunAsyncStaged(Rank(), RankSize(), transportInfo_, stage);
261 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Run][RunTemplateAlgStaged]comm base run tempAlg staged "\
262 : "rank[%u] rank size[%u] failed", rank_, rankSize_), ret);
263 0 : return HCCL_SUCCESS;
264 : }
265 :
266 18 : HcclResult CommBase::SetRankMap()
267 : {
268 : // 参数有效性校验
269 18 : if ((userRankSize_ <= userRank_) || (rankSize_ <= rank_)) {
270 0 : HCCL_ERROR("[Set][RankMap]invalid:userRankSize_[%u] userRank_[%u] rankSize_[%u] rank_[%u].",
271 : userRankSize_, userRank_, rankSize_, rank_);
272 0 : return HCCL_E_PARA;
273 : }
274 :
275 36 : for (u32 index = 0; index < rankSize_; index++) {
276 18 : userRankMap_[index] = paraVector_[index].userRank;
277 :
278 18 : if (userRankSize_ > userRankMap_[index]) {
279 18 : rankMap_[userRankMap_[index]] = index;
280 : }
281 18 : HCCL_INFO("userRankMap: [%u] -> [%u]", index, paraVector_[index].worldRank);
282 : }
283 :
284 18 : return HCCL_SUCCESS;
285 : }
286 :
287 18 : HcclResult CommBase::CheckLinks() const
288 : {
289 36 : for (u32 index = 0; index < transportType_.size(); index++) {
290 18 : bool check = (transportInfo_.size() <= index);
291 18 : CHK_PRT_RET(check, HCCL_ERROR("[Check][Links]index[%u] is bigger than link size[%llu]",
292 : index, transportInfo_.size()), HCCL_E_INTERNAL);
293 18 : if ((transportType_[index] != TransportType::TRANS_TYPE_RESERVED) && !transportInfo_[index]) {
294 0 : HCCL_ERROR("[Check][Links]there is no effective link(type[%d]) between rank[%u] and dst rank[%u]!",
295 : transportType_[index], rank_, index);
296 0 : return HCCL_E_NOT_FOUND;
297 : }
298 : }
299 :
300 18 : return HCCL_SUCCESS;
301 : }
302 :
303 : // 只有节点内,采用虚拟网卡时,才会进入该函数,有且仅有一个nic_ip
304 0 : HcclResult CommBase::CreateIntraLinks()
305 : {
306 0 : HcclUs startut = TIME_NOW();
307 0 : HcclResult ret = HCCL_SUCCESS;
308 :
309 0 : auto socketsMap = intraSocketsMap_;
310 0 : HCCL_DEBUG("[Create][IntraLinks] dstIntraServerVec size[%u].", dstIntraServerVec_.size());
311 0 : for (auto &rank : dstIntraServerVec_) {
312 0 : HCCL_DEBUG("[Create][IntraLinks] localrank[%u] remoterank[%u].", userRank_, rank);
313 : // 与当前Inter的Socket不同, 在Intra Socket创建时, 使用的userRank, 所以这里需要使用 userRank 为 key
314 0 : auto item = socketsMap.find(paraVector_[rank].userRank);
315 0 : if (item != socketsMap.end()) {
316 0 : ret = CreateIntraThread(CLIENT_ROLE_SOCKET, rank, item->second);
317 : } else {
318 0 : HCCL_INFO("[Create][IntraLinks] remoterank[%u] socket item not find.", rank);
319 : // 异构场景下, 当前上层 CreateCommP2PAsync 并不会创建 IntraExchanger, 后继的
320 : // TransportHeterogP2P 场景使用原有的 Socket 逻辑, 所以在这里没有找到 socket item 时,
321 : // 使用一个空的 sockets 作为参数, 创建后继处理线程.
322 0 : std::vector<std::shared_ptr<HcclSocket> > sockets;
323 0 : ret = CreateIntraThread(CLIENT_ROLE_SOCKET, rank, sockets);
324 0 : }
325 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
326 : HCCL_ERROR("[Create][IntraLinks] create intra thread failed, socket role is client"), ret);
327 : }
328 :
329 0 : HCCL_DEBUG("[Create][IntraLinks] dstIntraClientVec size[%u].", dstIntraClientVec_.size());
330 0 : for (auto &rank : dstIntraClientVec_) {
331 0 : HCCL_DEBUG("[Create][IntraLinks] localrank[%u] remoterank[%u].", userRank_, rank);
332 0 : auto item = socketsMap.find(paraVector_[rank].userRank);
333 0 : if (item != socketsMap.end()) {
334 0 : ret = CreateIntraThread(SERVER_ROLE_SOCKET, rank, item->second);
335 : } else {
336 0 : HCCL_INFO("[Create][IntraLinks] remoterank[%u] socket item not find.", rank);
337 : // 异构场景下, 当前上层 CreateCommP2PAsync 并不会创建 IntraExchanger, 后继的
338 : // TransportHeterogP2P 场景使用原有的 Socket 逻辑, 所以在这里没有找到 socket item 时,
339 : // 使用一个空的 sockets 作为参数, 创建后继处理线程.
340 0 : std::vector<std::shared_ptr<HcclSocket> > sockets;
341 0 : ret = CreateIntraThread(SERVER_ROLE_SOCKET, rank, sockets);
342 0 : }
343 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
344 : HCCL_ERROR("[Create][IntraLinks] create intra thread failed, socket role is server"), ret);
345 : }
346 :
347 0 : HCCL_DEBUG("[Create][IntraLinks] create intra link used time:%lld us.", DURATION_US(TIME_NOW() - startut));
348 0 : return ret;
349 0 : }
350 :
351 0 : HcclResult CommBase::CreateIntraThread(const u32 role, u32 dstRank,
352 : const std::vector<std::shared_ptr<HcclSocket> > &sockets)
353 : {
354 0 : if (threadsRapplyNum_ >= linkThreads_.size()) {
355 0 : HCCL_ERROR("[Create][InterThread] threadsRapplyNum_[%u] is bigger than link threads size[%llu] ",
356 : threadsRapplyNum_, linkThreads_.size());
357 0 : return HCCL_E_INTERNAL;
358 : }
359 :
360 : // 线程命名,TraL_ 代表Intra Link
361 0 : std::string threadStr = "HcclTraL_" + std::to_string(threadsRapplyNum_);
362 :
363 : // 创建新线程前更新一下最新的workflowMode
364 0 : workflowMode_ = GetWorkflowMode();
365 0 : if (role == SERVER_ROLE_SOCKET) {
366 0 : linkThreads_[threadsRapplyNum_].reset(
367 0 : new (std::nothrow) std::thread(&CommBase::CreateDestLink, this, hrtErrMGetErrorContextPub(),
368 0 : MachineType::MACHINE_SERVER_TYPE, paraVector_[rank_].serverId, dstRank, threadStr, sockets));
369 : }
370 :
371 0 : HCCL_DEBUG("[CommBase][CreateIntraThread]role is %u", role);
372 0 : if (role == CLIENT_ROLE_SOCKET) {
373 0 : linkThreads_[threadsRapplyNum_].reset(
374 0 : new (std::nothrow) std::thread(&CommBase::CreateDestLink, this, hrtErrMGetErrorContextPub(),
375 0 : MachineType::MACHINE_CLIENT_TYPE, paraVector_[rank_].serverId, dstRank, threadStr, sockets));
376 : }
377 :
378 0 : if (!linkThreads_[threadsRapplyNum_]) {
379 0 : HCCL_ERROR("[Create][IntraThread] link threads[%u] reset failed.", threadsRapplyNum_);
380 0 : return HCCL_E_INTERNAL;
381 : }
382 0 : threadsRapplyNum_++;
383 :
384 0 : HCCL_DEBUG("[Create][IntraThread] role[%u], dstRank[%u], sockets size[%u], threadsRapplyNum[%u]",
385 : role, dstRank, sockets.size(), threadsRapplyNum_);
386 0 : return HCCL_SUCCESS;
387 0 : }
388 :
389 0 : void CommBase::PrintCreateInterLinksInfo()
390 : {
391 0 : HCCL_RUN_INFO("[PrintCreateInterLinksInfo] dstInterServerMap size[%llu], dstInterClientMap size[%llu]",
392 : dstInterServerMap_.size(), dstInterClientMap_.size());
393 :
394 : // 维护建链输出的信息
395 0 : std::string outLogInfo = "";
396 0 : for (auto iter = dstInterServerMap_.begin(); iter != dstInterServerMap_.end(); iter++) {
397 0 : outLogInfo.append(std::to_string(paraVector_[iter->first].userRank));
398 0 : outLogInfo.append("/");
399 0 : outLogInfo.append(paraVector_[iter->first].serverId);
400 0 : outLogInfo.append("/");
401 0 : outLogInfo.append(std::to_string(paraVector_[iter->first].devicePhyId));
402 0 : outLogInfo.append("; ");
403 : }
404 :
405 0 : for (auto iter = dstInterClientMap_.begin(); iter != dstInterClientMap_.end(); iter++) {
406 0 : outLogInfo.append(std::to_string(paraVector_[iter->first].userRank));
407 0 : outLogInfo.append("/");
408 0 : outLogInfo.append(paraVector_[iter->first].serverId);
409 0 : outLogInfo.append("/");
410 0 : outLogInfo.append(std::to_string(paraVector_[iter->first].devicePhyId));
411 0 : outLogInfo.append("; ");
412 : }
413 :
414 0 : HCCL_RUN_INFO("serverInterConnectInfo:tag[%s], userRank/serverIp/devicePhyId:[%u/%s/%d], connectRankInfo[%s]",
415 : tag_.c_str(), userRank_, paraVector_[rank_].serverId.c_str(), paraVector_[rank_].devicePhyId,
416 : outLogInfo.c_str());
417 0 : }
418 :
419 0 : HcclResult CommBase::CreateInterLinks()
420 : {
421 0 : interSocketManager_.reset(
422 0 : new (std::nothrow) HcclSocketManager(nicDeployInner_, deviceLogicId_, devicePhyId_, userRank_));
423 0 : CHK_PTR_NULL(interSocketManager_);
424 :
425 0 : if (dstInterServerMap_.size() + dstInterClientMap_.size() == 0) {
426 0 : HCCL_DEBUG("[Create][InterLinks] do not need create links.");
427 0 : return HCCL_SUCCESS;
428 : }
429 :
430 0 : PrintCreateInterLinksInfo();
431 :
432 0 : HcclUs startut = TIME_NOW();
433 0 : HcclResult ret = HCCL_SUCCESS;
434 0 : std::map <u32, std::vector<std::shared_ptr<HcclSocket> > > serverSocketsMap;
435 0 : std::map <u32, std::vector<std::shared_ptr<HcclSocket> > > clientSocketsMap;
436 0 : ret = interSocketManager_->CreateSockets(tag_, true, netDevCtxMap_[paraVector_[rank_].nicIp[0]],
437 0 : dstInterServerMap_, dstInterClientMap_,
438 : serverSocketsMap, clientSocketsMap);
439 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
440 : HCCL_ERROR("[Create][InterLinks] socket manager create connections failed, ret[%u]", ret), ret);
441 :
442 0 : for (auto &sockets : clientSocketsMap) {
443 0 : ret = CreateInterThread(CLIENT_ROLE_SOCKET, sockets.first, sockets.second);
444 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
445 : HCCL_ERROR("[Create][InterLinks] create inter thread failed, socket role[CLIENT_ROLE_SOCKET] "),
446 : ret);
447 : }
448 0 : HCCL_DEBUG("[CommBase][CreateInterLinks]create inter thread success");
449 0 : for (auto &sockets : serverSocketsMap) {
450 0 : ret = CreateInterThread(SERVER_ROLE_SOCKET, sockets.first, sockets.second);
451 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
452 : HCCL_ERROR("[Create][InterLinks] create inter thread failed, socket role[SERVER_ROLE_SOCKET] "),
453 : ret);
454 : }
455 :
456 0 : HCCL_DEBUG("[Create][InterLinks] create inter link used time:%lld us", DURATION_US(TIME_NOW() - startut));
457 0 : return ret;
458 0 : }
459 :
460 0 : HcclResult CommBase::CreateInterThread(const u32 role, u32 dstRank,
461 : const std::vector<std::shared_ptr<HcclSocket> > &sockets)
462 : {
463 0 : if (sockets.empty()) {
464 0 : HCCL_ERROR("[Create][InterThread] create inter link failed, rank's sockets is empty");
465 0 : return HCCL_E_INTERNAL;
466 : }
467 :
468 0 : if (threadsRapplyNum_ >= linkThreads_.size()) {
469 0 : HCCL_ERROR("[Create][InterThread] threadsRapplyNum_[%u] is bigger than link threads size[%llu] ",
470 : threadsRapplyNum_, linkThreads_.size());
471 0 : return HCCL_E_INTERNAL;
472 : }
473 :
474 : // 线程命名,TerL代表Inter Link
475 0 : std::string threadStr = "HcclTerL_" + std::to_string(threadsRapplyNum_);
476 :
477 : // 创建新线程前更新一下最新的workflowMode
478 0 : workflowMode_ = GetWorkflowMode();
479 0 : if (role == SERVER_ROLE_SOCKET) {
480 0 : linkThreads_[threadsRapplyNum_].reset(
481 0 : new (std::nothrow) std::thread(&CommBase::CreateDestLink, this, hrtErrMGetErrorContextPub(),
482 0 : MachineType::MACHINE_SERVER_TYPE, paraVector_[rank_].serverId, dstRank, threadStr, sockets));
483 : }
484 :
485 0 : if (role == CLIENT_ROLE_SOCKET) {
486 0 : linkThreads_[threadsRapplyNum_].reset(
487 0 : new (std::nothrow) std::thread(&CommBase::CreateDestLink, this, hrtErrMGetErrorContextPub(),
488 0 : MachineType::MACHINE_CLIENT_TYPE, paraVector_[rank_].serverId, dstRank, threadStr, sockets));
489 : }
490 :
491 0 : if (!linkThreads_[threadsRapplyNum_]) {
492 0 : HCCL_ERROR("[Create][InterThread] link threads[%u] reset failed.", threadsRapplyNum_);
493 0 : return HCCL_E_INTERNAL;
494 : }
495 0 : threadsRapplyNum_++;
496 :
497 0 : HCCL_DEBUG("[Create][InterThread] role[%u], dstRank[%u], sockets size[%u], threadsRapplyNum[%u]",
498 : role, dstRank, sockets.size(), threadsRapplyNum_);
499 :
500 0 : return HCCL_SUCCESS;
501 0 : }
502 :
503 0 : u32 CommBase::GetInterRemotePort(s32 devicePhyId, u32 dstUserRank)
504 : {
505 0 : if (isUseRankPort_ && dstUserRank < ranksPort_.size() && ranksPort_[dstUserRank] != HCCL_INVALID_PORT) {
506 0 : HCCL_INFO("[GetInterRemotePort] port[%u] from ranks port", ranksPort_[dstUserRank]);
507 0 : return ranksPort_[dstUserRank];
508 0 : } else if (!isUseRankPort_ && !Is310PDevice()) {
509 0 : HCCL_INFO("[GetInterRemotePort] port[%u]", HETEROG_CCL_PORT);
510 0 : return HETEROG_CCL_PORT;
511 0 : } else if (GetExternalInputHcclIfBasePort() == HCCL_INVALID_PORT) {
512 0 : return (devicePhyId + HOST_PARA_BASE_PORT);
513 : } else {
514 0 : return (devicePhyId + GetExternalInputHcclIfBasePort() + HCCL_AISERVER_DEVICE_NUM);
515 : }
516 : }
517 :
518 0 : HcclResult CommBase::CalcLinksNum(const MachineType machineType, const u32 dstRank)
519 : {
520 0 : bool check = (paraVector_.size() <= dstRank) || (paraVector_.size() <= rank_);
521 0 : CHK_PRT_RET(check, HCCL_ERROR("[Calc][LinksNum]para check failed, para vector size[%llu], dstRank[%u], rank[%u] ",
522 : paraVector_.size(), dstRank, rank_), HCCL_E_INTERNAL);
523 : // 节点间或者是节点内采用RDMA通信的, 放至dst_inter_client_map_,采用rdma建链
524 0 : bool isInterRdma = paraVector_[rank_].serverId != paraVector_[dstRank].serverId ||
525 0 : isUsedRdmaLevel0_ || isAlltoAllCommMesh_;
526 :
527 0 : bool isInterHccs = IsSupportInterHccs(dstRank);
528 :
529 0 : HCCL_DEBUG("[Calc][LinksNum]rank[%u], dstRank[%u], isInterRdma[%d], isInterHccs[%d], machineType[%d]",
530 : rank_, dstRank, isInterRdma, isInterHccs, machineType);
531 :
532 0 : auto dstRankInfo = paraVector_[dstRank];
533 0 : if (machineType == MachineType::MACHINE_SERVER_TYPE) {
534 0 : CHK_RET(MakeClientInfo(dstRank, dstRankInfo, isInterRdma, isInterHccs));
535 : }
536 :
537 0 : if (machineType == MachineType::MACHINE_CLIENT_TYPE) {
538 0 : CHK_RET(MakeServerInfo(dstRank, dstRankInfo, isInterRdma, isInterHccs));
539 : }
540 :
541 0 : return HCCL_SUCCESS;
542 0 : }
543 :
544 0 : HcclResult CommBase::MakeClientInfo(const u32 dstRank, RankInfo &dstRankInfo, bool isInterRdma, bool isInterHccs)
545 : {
546 0 : if (isInterRdma && !isInterHccs) {
547 0 : HcclRankLinkInfo tempLinkInfo {};
548 0 : tempLinkInfo.userRank = dstRankInfo.userRank;
549 0 : tempLinkInfo.devicePhyId = dstRankInfo.devicePhyId;
550 :
551 0 : tempLinkInfo.ip = dstRankInfo.nicIp[0];
552 0 : tempLinkInfo.port = GetInterRemotePort(tempLinkInfo.devicePhyId, dstRankInfo.userRank);
553 0 : tempLinkInfo.socketsPerLink = GetSocketsPerLink();
554 :
555 0 : auto iter = dstInterClientMap_.find(dstRank);
556 0 : bool check = (iter != dstInterClientMap_.end());
557 0 : CHK_PRT_RET(check, HCCL_ERROR("[Make][ClientInfo]dstRank[%u] already exists in dst inter client map. ",
558 : dstRank), HCCL_E_PARA);
559 0 : dstInterClientMap_.insert(std::make_pair(dstRank, tempLinkInfo));
560 0 : } else {
561 0 : dstIntraClientVec_.push_back(dstRank);
562 : }
563 0 : return HCCL_SUCCESS;
564 : }
565 :
566 0 : HcclResult CommBase::MakeServerInfo(const u32 dstRank, RankInfo &dstRankInfo, bool isInterRdma, bool isInterHccs)
567 : {
568 : // 节点间或者是节点内采用RDMA通信的,放至dst_inter_client_map_,采用rdma建链
569 0 : if (isInterRdma && !isInterHccs) {
570 0 : HcclRankLinkInfo tempLinkInfo {};
571 0 : tempLinkInfo.userRank = dstRankInfo.userRank;
572 0 : tempLinkInfo.devicePhyId = dstRankInfo.devicePhyId;
573 :
574 0 : HCCL_INFO("dstRank = %u, useRank = %u, ip = %s",
575 : dstRank, dstRankInfo.userRank, dstRankInfo.nicIp[0].GetReadableAddress());
576 :
577 0 : tempLinkInfo.ip = dstRankInfo.nicIp[0];
578 0 : tempLinkInfo.port = GetInterRemotePort(tempLinkInfo.devicePhyId, dstRankInfo.userRank);
579 0 : tempLinkInfo.socketsPerLink = GetSocketsPerLink();
580 :
581 0 : auto iter = dstInterServerMap_.find(dstRank);
582 0 : bool check = (iter != dstInterServerMap_.end());
583 0 : CHK_PRT_RET(check, HCCL_ERROR("[Make][ServerInfo]dstRank[%u] already exists in dst inter server map",
584 : dstRank), HCCL_E_PARA);
585 0 : dstInterServerMap_.insert(std::make_pair(dstRank, tempLinkInfo));
586 0 : } else {
587 0 : dstIntraServerVec_.push_back(dstRank);
588 : }
589 0 : return HCCL_SUCCESS;
590 : }
591 :
592 1 : HcclResult CommBase::CreateDestLink(const ErrContextPub &error_context, const MachineType machineType,
593 : const std::string &serverId, const u32 dstRank, const std::string &threadStr,
594 : const std::vector<std::shared_ptr<HcclSocket> > &sockets)
595 : {
596 1 : hrtErrMSetErrorContextPub(error_context);
597 : // 给当前线程添加名字
598 1 : SetThreadName(threadStr);
599 1 : if (!IsGeneralServer()) {
600 1 : CHK_RET(hrtSetDevice(deviceLogicId_));
601 1 : SetWorkflowMode(workflowMode_); // 新的线程,更新workflowMode
602 : }
603 :
604 2 : bool check = (paraVector_.size() <= dstRank) || (paraVector_.size() <= rank_) ||
605 2 : (transportInfo_.size() <= dstRank) || (transportType_.size() <= dstRank);
606 1 : CHK_PRT_RET(check, HCCL_ERROR("[Create][DestLink]paraCheck failed, paraVector size[%llu], linkInfo size[%llu], "
607 : "linkType size[%llu], dstRank[%u], rank[%u] ", paraVector_.size(), transportInfo_.size(),
608 : transportType_.size(), dstRank, rank_), HCCL_E_INTERNAL);
609 :
610 1 : MachinePara machinePara;
611 1 : CHK_RET(SetMachinePara(machineType, serverId, dstRank, sockets, machinePara));
612 1 : HCCL_INFO("[creakLink para]rank[%u]-localUserrank[%u]-localIpAddr[%s], linkMode[%d] "
613 : "dst_rank[%u]-remoteUserrank[%u]-remote_ip_addr[%s], machineType[%d], serverId[%s], nicDeploy[%d] ",
614 : rank_, paraVector_[rank_].worldRank, paraVector_[rank_].serverId.c_str(), machinePara.linkMode,
615 : dstRank, paraVector_[dstRank].worldRank, paraVector_[dstRank].serverId.c_str(), machinePara.machineType,
616 : machinePara.serverId.c_str(), machinePara.nicDeploy);
617 :
618 : // transport初始化
619 1 : HcclResult ret = TransportInit(dstRank, machinePara);
620 1 : if (ret != HCCL_SUCCESS) {
621 1 : transportInfo_[dstRank] = nullptr;
622 1 : if (ret == HCCL_E_MEMORY) {
623 : std::string err_str = "[Create][DestLink]Transport init error! IPC memory allocation failed due to "
624 : "possible memory limit exceeded. Suggested solution: Use 3TB / (ranksize * 2) as the upper limit of "
625 1 : "HCCL_BUFFSIZE.";
626 1 : HCCL_ERROR("%s", err_str.c_str());
627 1 : }
628 1 : const std::string CREATE_LINK_ERR = "[Create][DestLink]Create Dest error! creakLink para:rank[" + \
629 3 : std::to_string(rank_) + "]-localUserrank[" + std::to_string(paraVector_[rank_].worldRank) + \
630 3 : "]-localIpAddr[" + paraVector_[rank_].serverId.c_str() + "], dst_rank[" + \
631 4 : std::to_string(dstRank) + "]-remoteUserrank[" + std::to_string(paraVector_[dstRank].worldRank) + \
632 2 : "]-remote_ip_addr[" + paraVector_[dstRank].serverId.c_str() + "]";
633 :
634 1 : HCCL_ERROR("[Create][DestLink]Transport init error! creakLink para:rank[%u]-localUserrank[%u]-localIpAddr[%s], "
635 : "dst_rank[%u]-remoteUserrank[%u]-remote_ip_addr[%s], machineType[%d], serverId[%s], linkMode[%d], "
636 : "shmDev_[%u], tag[%s]",
637 : rank_, paraVector_[rank_].worldRank, paraVector_[rank_].serverId.c_str(), dstRank,
638 : paraVector_[dstRank].worldRank, paraVector_[dstRank].serverId.c_str(),
639 : machinePara.machineType, machinePara.serverId.c_str(), machinePara.linkMode, shmDev_,
640 : machinePara.tag.c_str());
641 1 : return ret;
642 1 : }
643 0 : HCCL_INFO("[creakLink success]:rank[%u]-localUserrank[%u]-localIpAddr[%s], " \
644 : "dst_rank[%u]-remoteUserrank[%u]-remote_ip_addr[%s], transportType_[%d], tag[%s]", rank_,
645 : paraVector_[rank_].worldRank, paraVector_[rank_].serverId.c_str(), dstRank, paraVector_[dstRank].worldRank,
646 : paraVector_[dstRank].serverId.c_str(), transportType_[dstRank], machinePara.tag.c_str());
647 :
648 0 : return HCCL_SUCCESS;
649 1 : }
650 :
651 2 : void CommBase::SetTransportParam(TransportPara ¶, MachinePara &machinePara)
652 : {
653 0 : std::chrono::milliseconds kdefaultTimeout = std::chrono::seconds(
654 2 : GetExternalInputHcclLinkTimeOut());
655 2 : para.isRootRank = subUserRankRoot_ == rank_ ? true : false;
656 2 : para.timeout = kdefaultTimeout;
657 2 : para.virtualFlag = false;
658 2 : }
659 :
660 2 : HcclResult CommBase::TransportInit(const u32 dstRank, MachinePara &machinePara)
661 : {
662 2 : CHK_PRT_RET(dstRank >= transportInfo_.size(),
663 : HCCL_ERROR("[TransportQuerry] Transport[%u] is invalid, should init before query it.", dstRank), HCCL_E_PARA);
664 : // 实例化TransportBase
665 2 : CHK_RET(SetTransportType(dstRank));
666 2 : TransportPara para{};
667 2 : SetTransportParam(para, machinePara);
668 :
669 2 : TransportType type = transportType_[dstRank];
670 2 : if (type == TransportType::TRANS_TYPE_P2P) {
671 0 : transportInfo_[dstRank].reset(new (std::nothrow) Transport(type, para, dispatcher_, notifyPool_, machinePara));
672 2 : } else if (type == TransportType::TRANS_TYPE_IBV_EXP) {
673 0 : transportInfo_[dstRank].reset(new (std::nothrow) Transport(type, para, dispatcher_, notifyPool_, machinePara));
674 : } else {
675 2 : HCCL_ERROR("[Init][Transport]not supported transport type");
676 2 : return HCCL_E_NOT_SUPPORT;
677 : }
678 :
679 0 : CHK_PRT_RET(!transportInfo_[dstRank], HCCL_ERROR("[Init][Transport]In create link, new link failed"), HCCL_E_PTR);
680 :
681 0 : if (useOneDoorbell_) {
682 0 : transportInfo_[dstRank]->EnableUseOneDoorbell();
683 : }
684 :
685 0 : CHK_RET(transportInfo_[dstRank]->Init());
686 :
687 0 : CHK_RET(CheckExchangeInfo(transportInfo_[dstRank], machinePara.localDeviceId));
688 :
689 0 : return HCCL_SUCCESS;
690 : }
691 :
692 2 : HcclResult CommBase::SetMachinePara(MachineType machineType, const std::string &serverId, u32 dstRank,
693 : const std::vector<std::shared_ptr<HcclSocket> > &socketList, MachinePara &machinePara)
694 : {
695 2 : SetMachineLinkMode(machinePara);
696 2 : HCCL_INFO("[Set][MachinePara]rankSize %u, linkMode %d", rankSize_, machinePara.linkMode);
697 :
698 2 : machinePara.machineType = machineType;
699 2 : machinePara.serverId = serverId;
700 2 : machinePara.localIpAddr = paraVector_[rank_].nicIp[0];
701 2 : machinePara.remoteIpAddr = paraVector_[dstRank].nicIp[0];
702 2 : machinePara.localUserrank = paraVector_[rank_].userRank;
703 2 : machinePara.remoteUserrank = paraVector_[dstRank].userRank;
704 2 : machinePara.localWorldRank = paraVector_[rank_].worldRank;
705 2 : machinePara.remoteWorldRank = paraVector_[dstRank].worldRank;
706 2 : machinePara.collectiveId = collectiveId_;
707 2 : machinePara.localDeviceId = paraVector_[rank_].devicePhyId;
708 2 : machinePara.remoteDeviceId = paraVector_[dstRank].devicePhyId;
709 2 : machinePara.deviceType = static_cast<DevType>(paraVector_[dstRank].deviceType);
710 2 : machinePara.inputMem = inputMem_;
711 2 : machinePara.outputMem = outputMem_;
712 2 : if(expMem_.ptr() != nullptr){
713 0 : machinePara.mem.push_back(expMem_);
714 : } else {
715 2 : machinePara.mem.clear();
716 : }
717 2 : machinePara.linkAttribute = 0x03; /* 0x03同时支持目的端和源端发起 */
718 2 : machinePara.tag = tag_;
719 :
720 : // MoE算子优化,MC2 多机场景使用普通QP模式
721 2 : const std::string &suffix = HCCL_MC2_MULTISERVER_SUFFIX;
722 3 : if (tag_.size() > suffix.size() &&
723 1 : tag_.compare(tag_.size() - suffix.size(), suffix.size(), suffix) == 0) {
724 1 : bool isSupportNormalQP{false};
725 1 : CHK_RET(IsSupportAicpuNormalQP(paraVector_[rank_].devicePhyId, isSupportNormalQP));
726 1 : if (isSupportNormalQP) {
727 1 : HCCL_INFO("[Set][MachinePara] Set machinePara.qpMode to [NORMAL]");
728 1 : machinePara.qpMode = QPMode::NORMAL;
729 : }
730 : }
731 :
732 : // 把原来的两层vector变成一层, 方便后继调用
733 2 : for (u32 i = 0; i < socketList.size(); i++) {
734 0 : machinePara.sockets.push_back(socketList[i]);
735 : }
736 : u64 rankConsistentDataLength =
737 2 : RankConsistentcyChecker::GetInstance(machinePara.localDeviceId).GetRankConsistentDataLength();
738 2 : machinePara.exchangeInfo.resize(rankConsistentDataLength);
739 2 : CHK_RET(RankConsistentcyChecker::GetInstance(machinePara.localDeviceId).GetCheckFrame(&machinePara.exchangeInfo[0],
740 : rankConsistentDataLength, tag_));
741 2 : machinePara.supportDataReceivedAck = NeedDataReceivedAck();
742 2 : machinePara.nicDeploy = nicDeployInner_;
743 2 : machinePara.localSocketPort = paraVector_[rank_].hostPort;
744 2 : machinePara.remoteSocketPort = paraVector_[dstRank].hostPort;
745 2 : machinePara.isAicpuModeEn = isAicpuModeEn_;
746 2 : machinePara.deviceLogicId = deviceLogicId_;
747 2 : machinePara.srcPorts = std::vector<std::uint16_t>(1, 0); /* 默认填充一个元素,0代表默认不配置 */
748 2 : return HCCL_SUCCESS;
749 : }
750 :
751 18 : HcclResult CommBase::CreateVirturalTransport()
752 : {
753 18 : MachinePara machinePara;
754 0 : std::chrono::milliseconds kdefaultTimeout = std::chrono::seconds(
755 18 : GetExternalInputHcclLinkTimeOut());
756 :
757 18 : vTransportInfo_.resize(transportInfo_.size());
758 36 : for (u32 i = 0; i < transportInfo_.size(); i++) {
759 18 : TransportPara para {};
760 18 : para.virtualFlag = true;
761 18 : para.timeout = kdefaultTimeout;
762 18 : para.index = i;
763 36 : vTransportInfo_[i].reset(new (std::nothrow) Transport(TransportType::TRANS_TYPE_RESERVED, para, dispatcher_,
764 36 : notifyPool_, machinePara));
765 :
766 18 : CHK_PRT_RET(!vTransportInfo_[i], HCCL_ERROR("[CreateVirturalTransport]In create link, new link failed"),
767 : HCCL_E_PTR);
768 : }
769 :
770 18 : return HCCL_SUCCESS;
771 18 : }
772 :
773 0 : std::shared_ptr<Transport> &CommBase::GetTrasportInfoByVTransportInfoIndex(u32 index)
774 : {
775 0 : if (vTransportInfo_.size() <= index) {
776 0 : HCCL_ERROR("[GetTrasportInfoByVTransportInfoIndex]index[%u] is bigger than vlink size[%llu]", index,
777 : vTransportInfo_.size());
778 0 : return linkDummy_;
779 : }
780 :
781 0 : if (transportInfo_.size() <= index) {
782 0 : HCCL_ERROR("[GetTrasportInfoByVTransportInfoIndex]index[%u] is bigger than link size[%llu]", index,
783 : transportInfo_.size());
784 0 : return linkDummy_;
785 : }
786 0 : return transportInfo_[index];
787 : }
788 0 : HcclResult CommBase::BuildAsync(u32& status)
789 : {
790 0 : transportStatus_.resize(rankSize_, 1);
791 0 : checkStatus_.resize(rankSize_, false);
792 :
793 : // 获取rank->userrank以及userrank->rank的映射关系
794 0 : CHK_RET(SetRankMap());
795 :
796 : // 获取当前线程操作的设备ID
797 0 : deviceLogicId_ = 0;
798 0 : if (paraVector_[rank_].devicePhyId != HOST_DEVICE_ID) {
799 0 : CHK_RET(hrtGetDevice(&deviceLogicId_));
800 : }
801 :
802 0 : if (rankSize_ == HCCL_RANK_SIZE_EQ_ONE) {
803 0 : HCCL_INFO("comm base needn't to create links, rankSize_[%u].", rankSize_);
804 0 : status = 0;
805 0 : return HCCL_SUCCESS;
806 : }
807 0 : CHK_RET(CalcLink());
808 :
809 : // 当前rank作为client端角色
810 0 : u32 dstIntraServerNum = dstIntraServerVec_.size();
811 0 : std::vector<std::shared_ptr<HcclSocket> > sockets;
812 0 : for (u32 intraIndex = 0; intraIndex < dstIntraServerNum; intraIndex++) {
813 0 : HcclResult ret = TransportBuildAsync(MachineType::MACHINE_CLIENT_TYPE, paraVector_[rank_].serverId,
814 0 : dstIntraServerVec_[intraIndex], sockets, transportStatus_[dstIntraServerVec_[intraIndex]]);
815 0 : CHK_PRT_RET(ret, HCCL_ERROR("[BuildAsync] transport build async failed, self rank[%u], peer rank[%u]",
816 : paraVector_[rank_].worldRank, paraVector_[dstIntraServerVec_[intraIndex]].worldRank),
817 : HCCL_E_INTERNAL);
818 : }
819 :
820 : // 当前rank作为server端角色
821 0 : u32 dstIntraClientNum = dstIntraClientVec_.size();
822 0 : for (u32 intraIndex = 0; intraIndex < dstIntraClientNum; intraIndex++) {
823 0 : HcclResult ret = TransportBuildAsync(MachineType::MACHINE_SERVER_TYPE, paraVector_[rank_].serverId,
824 0 : dstIntraClientVec_[intraIndex], sockets, transportStatus_[dstIntraClientVec_[intraIndex]]);
825 0 : CHK_PRT_RET(ret, HCCL_ERROR("[BuildAsync] transport build async failed, self rank[%u], peer rank[%u]",
826 : paraVector_[rank_].worldRank, paraVector_[dstIntraClientVec_[intraIndex]].worldRank),
827 : HCCL_E_INTERNAL);
828 : }
829 :
830 : // 暂不支持 跨node通信
831 :
832 0 : CHK_RET(GetBuildStatus(status));
833 0 : return HCCL_SUCCESS;
834 0 : }
835 :
836 0 : HcclResult CommBase::BuildQuerry(u32& status)
837 : {
838 0 : for (u32 i = 0; i < transportStatus_.size(); i++) {
839 0 : if (transportStatus_[i] == HETEROG_P2P_WAIT) {
840 0 : CHK_RET(TransportBuildQuerry(i, transportStatus_[i]));
841 : }
842 : }
843 0 : CHK_RET(GetBuildStatus(status));
844 0 : HCCL_DEBUG("BuildQuerry: %u", status);
845 0 : return HCCL_SUCCESS;
846 : }
847 :
848 0 : HcclResult CommBase::GetBuildStatus(u32& status)
849 : {
850 0 : u32 transportDoneNum = 0;
851 0 : u32 transportErrorNum = 0;
852 0 : for (u32 i = 0; i < transportStatus_.size(); i++) {
853 0 : if (transportStatus_[i] == HETEROG_P2P_SUCCESS) {
854 0 : transportDoneNum++;
855 0 : } else if (transportStatus_[i] == HETEROG_P2P_FAILED) {
856 0 : transportErrorNum++;
857 : }
858 : }
859 0 : u32 transportNum = dstInterClientMap_.size() + dstIntraClientVec_.size() + dstInterServerMap_.size() +
860 0 : dstIntraServerVec_.size();
861 0 : if (transportErrorNum > 0) {
862 0 : status = HETEROG_P2P_FAILED;
863 0 : HCCL_ERROR("transport error num[%u].", transportErrorNum);
864 0 : return HCCL_E_INTERNAL;
865 0 : } else if (transportDoneNum == transportNum) {
866 0 : status = HETEROG_P2P_SUCCESS;
867 0 : HCCL_INFO("CommBase connect complete.");
868 0 : } else if (transportDoneNum < transportNum) {
869 0 : status = HETEROG_P2P_WAIT;
870 : } else {
871 0 : status = HETEROG_P2P_FAILED;
872 0 : HCCL_ERROR("transport done num[%u] invalid, expect[%u].", transportDoneNum, transportNum);
873 0 : return HCCL_E_INTERNAL;
874 : }
875 0 : return HCCL_SUCCESS;
876 : }
877 :
878 0 : HcclResult CommBase::TransportBuildAsync(const MachineType machineType, const std::string &serverId, u32 dstRank,
879 : const std::vector<std::shared_ptr<HcclSocket> > &sockets, u32& status)
880 : {
881 0 : CHK_PRT_RET(dstRank >= transportInfo_.size(),
882 : HCCL_ERROR("[TransportQuerry] Transport[%u] is invalid, should init before query it.", dstRank), HCCL_E_PARA);
883 0 : MachinePara machinePara;
884 0 : CHK_RET(SetMachinePara(machineType, serverId, dstRank, sockets, machinePara));
885 : // 实例化TransportBase
886 0 : CHK_RET(SetTransportType(dstRank));
887 0 : std::chrono::milliseconds kdefaultTimeout = std::chrono::seconds(
888 0 : GetExternalInputHcclLinkTimeOut());
889 :
890 0 : TransportPara para {};
891 0 : para.timeout = kdefaultTimeout;
892 0 : para.virtualFlag = false;
893 0 : transportInfo_[dstRank].reset(new (std::nothrow) Transport(TransportType::TRANS_TYPE_RESERVED, para,
894 0 : dispatcher_, notifyPool_, machinePara));
895 0 : CHK_PRT_RET(!transportInfo_[dstRank], HCCL_ERROR("[Init][Transport]In create link, new link failed"),
896 : HCCL_E_PTR);
897 :
898 0 : CHK_RET(transportInfo_[dstRank]->ConnectAsync(status));
899 0 : if (status == HETEROG_P2P_SUCCESS && checkStatus_[dstRank] == false) {
900 0 : checkStatus_[dstRank] = true;
901 0 : CHK_RET(CheckExchangeInfo(transportInfo_[dstRank], machinePara.localDeviceId));
902 : }
903 0 : HCCL_DEBUG("TransportBuildAsync[%u] %u", dstRank, status);
904 0 : return HCCL_SUCCESS;
905 0 : }
906 :
907 0 : HcclResult CommBase::TransportBuildQuerry(u32 dstRank, u32& status)
908 : {
909 0 : CHK_PRT_RET(dstRank >= transportInfo_.size(),
910 : HCCL_ERROR("[TransportQuerry] Transport[%u] is invalid, should init before query it.", dstRank), HCCL_E_PARA);
911 0 : if (transportInfo_[dstRank]) {
912 0 : CHK_RET(transportInfo_[dstRank]->ConnectQuerry(status));
913 0 : if (status == HETEROG_P2P_SUCCESS && checkStatus_[dstRank] == false) {
914 0 : checkStatus_[dstRank] = true;
915 0 : CHK_RET(CheckExchangeInfo(transportInfo_[dstRank], paraVector_[rank_].devicePhyId));
916 : }
917 : } else {
918 0 : status = HETEROG_P2P_WAIT;
919 : }
920 0 : HCCL_DEBUG("TransportBuildQuerry[%u] %u", dstRank, status);
921 0 : return HCCL_SUCCESS;
922 : }
923 :
924 0 : HcclResult CommBase::CreateExchangerNetwork()
925 : {
926 0 : CHK_PRT_RET(dstIntraServerVec_.empty() && dstIntraClientVec_.empty(),
927 : HCCL_DEBUG("[Create][ExchangerNetwork]dstIntraServerVec and dstIntraClientVec is empty, do nothing."),
928 : HCCL_SUCCESS);
929 :
930 0 : bool isInterServer = false; // 是否跨server
931 0 : bool isInterHccs = true; // 是否超节点模式
932 0 : std::map<u32, HcclSocketRole> rankRole;
933 0 : CHK_RET(GetRankLinkInfo(isInterServer, isInterHccs, rankRole));
934 :
935 : // 保持原逻辑不变,将rank&deviceIP信息,构造成 std::map<u32, std::vector<HcclIpAddress> >,使用 SocketManager创建链接
936 0 : std::string commTag = (Is310PDevice() || isHaveCpuRank_) ? tag_ : collectiveId_;
937 0 : HcclIpAddress localIP;
938 0 : std::map<u32, HcclRankLinkInfo> dstServerMap;
939 0 : std::map<u32, HcclRankLinkInfo> dstClientMap;
940 0 : bool isSupportReuse = false;
941 :
942 0 : CHK_RET(GetRankIPInfo(isInterServer, isInterHccs, isSupportReuse, rankRole, localIP,
943 : dstServerMap, dstClientMap, socketManager_));
944 :
945 0 : std::map <u32, std::vector<std::shared_ptr<HcclSocket> > > serverSocketsMap;
946 0 : std::map <u32, std::vector<std::shared_ptr<HcclSocket> > > clientSocketsMap;
947 :
948 0 : HcclUs startut = TIME_NOW();
949 0 : HcclResult ret = socketManager_->CreateSockets(commTag, false,
950 0 : netDevCtxMap_[localIP], dstServerMap, dstClientMap, serverSocketsMap, clientSocketsMap, isSupportReuse);
951 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
952 : HCCL_ERROR("[Create][ExchangerNetwork]sync create connections Failed, ret[%u].", ret), ret);
953 :
954 0 : HCCL_DEBUG("[Create][Exchanger] serverSocketsMap size[%u], clientSocketsMap size[%u]",
955 : serverSocketsMap.size(), clientSocketsMap.size());
956 0 : intraSocketsMap_.insert(serverSocketsMap.begin(), serverSocketsMap.end());
957 0 : intraSocketsMap_.insert(clientSocketsMap.begin(), clientSocketsMap.end());
958 :
959 0 : HCCL_INFO("[Create][ExchangerNetwork]create connections duration time:%lld us.", DURATION_US(TIME_NOW() - startut));
960 0 : return HCCL_SUCCESS;
961 0 : }
962 :
963 0 : HcclResult CommBase::GetRankIPInfo(bool isInterServer, bool isInterHccs, bool &isSupportReuse,
964 : std::map<u32, HcclSocketRole> &rankRole, HcclIpAddress &localIP,
965 : std::map<u32, HcclRankLinkInfo> &dstServerMap,
966 : std::map<u32, HcclRankLinkInfo> &dstClientMap,
967 : std::shared_ptr<HcclSocketManager> &socketManager)
968 : {
969 0 : if (Is310PDevice() || isHaveCpuRank_) {
970 : // 310P和异构场景
971 0 : std::vector<u32> dstIntraVec;
972 0 : for (auto it = rankRole.begin(); it != rankRole.end(); ++it) {
973 0 : dstIntraVec.push_back(it->first);
974 : }
975 0 : CHK_RET(GetIntraRankIPInfo(dstIntraVec, localIP, dstServerMap, dstClientMap));
976 :
977 : // 不复用,每次都创建
978 0 : isSupportReuse = false;
979 0 : socketManager.reset(new (std::nothrow) HcclSocketManager(nicDeployInner_, deviceLogicId_,
980 0 : devicePhyId_, userRank_));
981 0 : CHK_PTR_NULL(socketManager);
982 0 : } else if (isInterServer && isInterHccs) {
983 : // 超节点间Hccs模式
984 0 : CHK_RET(GetSuperNodeIntraRankIPInfo(rankRole, localIP, dstServerMap, dstClientMap));
985 0 : isSupportReuse = true;
986 0 : socketManager = exchanger_.socketManager;
987 0 : CHK_PTR_NULL(socketManager);
988 0 : } else if (isInterServer == false) {
989 : // server内模式
990 0 : CHK_RET(GetIntraRankIPInfo(rankRole, localIP, dstServerMap, dstClientMap));
991 0 : isSupportReuse = true;
992 0 : socketManager = exchanger_.socketManager;
993 0 : CHK_PTR_NULL(socketManager);
994 : } else {
995 0 : HCCL_ERROR("[Create][ExchangerNetwork]isInterServer[%d] and isInterHccs[%d] is not support",
996 : isInterServer, isInterHccs);
997 0 : return HCCL_E_INTERNAL;
998 : }
999 0 : return HCCL_SUCCESS;
1000 : }
1001 :
1002 0 : HcclResult CommBase::GetRankLinkInfo(bool &isInterServer, bool &isInterHccs, std::map<u32, HcclSocketRole> &rankRole)
1003 : {
1004 0 : std::vector<u32> devicePhyIds;
1005 0 : for (u32 dstRank : dstIntraServerVec_) {
1006 0 : isInterServer |= (paraVector_[dstRank].serverId != paraVector_[rank_].serverId);
1007 0 : isInterHccs &= IsSupportInterHccs(dstRank);
1008 0 : rankRole.insert(std::make_pair(dstRank, HcclSocketRole::SOCKET_ROLE_SERVER));
1009 0 : devicePhyIds.push_back(paraVector_[dstRank].devicePhyId);
1010 : }
1011 0 : for (u32 dstRank : dstIntraClientVec_) {
1012 0 : isInterServer |= (paraVector_[dstRank].serverId != paraVector_[rank_].serverId);
1013 0 : isInterHccs &= IsSupportInterHccs(dstRank);
1014 0 : rankRole.insert(std::make_pair(dstRank, HcclSocketRole::SOCKET_ROLE_CLIENT));
1015 0 : devicePhyIds.push_back(paraVector_[dstRank].devicePhyId);
1016 : }
1017 0 : rankRole.insert(std::make_pair(rank_, HcclSocketRole::SOCKET_ROLE_RESERVED));
1018 0 : devicePhyIds.push_back(paraVector_[rank_].devicePhyId);
1019 :
1020 0 : if (paraVector_[rank_].deviceType == DevType::DEV_TYPE_310P3) {
1021 0 : HcclResult ret = P2PMgmtPub::EnableP2P(devicePhyIds);
1022 0 : CHK_PRT_RET(ret != HCCL_SUCCESS,
1023 : HCCL_ERROR("[Get][RankLinkInfo]Enable P2P Failed, devicePhyId[%d], ret[%u]",
1024 : paraVector_[rank_].devicePhyId, ret), ret);
1025 : }
1026 : // server内非异构场景,使能P2P
1027 : // 心跳需要单独WaitP2PEnabled?
1028 0 : if (!isInterServer && !isHaveCpuRank_) {
1029 0 : HcclResult ret = P2PMgmtPub::WaitP2PEnabled(devicePhyIds);
1030 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[Get][RankLinkInfo]Enable P2P Failed, ret[%u]", ret), ret);
1031 : }
1032 0 : return HCCL_SUCCESS;
1033 0 : }
1034 :
1035 0 : HcclResult CommBase::GetIntraRankIPInfo(std::map<u32, HcclSocketRole> &rankRole,
1036 : HcclIpAddress &localIP,
1037 : std::map<u32, HcclRankLinkInfo> &dstServerMap,
1038 : std::map<u32, HcclRankLinkInfo> &dstClientMap)
1039 : {
1040 0 : u32 userRankSize = rankRole.size();
1041 0 : for (auto rankIter = rankRole.begin(); rankIter != rankRole.end(); rankIter++) {
1042 0 : u32 dstRank = rankIter->first;
1043 0 : s32 dstDeviceId = useSuperPodMode_ ?
1044 0 : static_cast<s32>(paraVector_[dstRank].superDeviceId) : paraVector_[dstRank].devicePhyId;
1045 0 : HcclSocketRole localRole = rankIter->second;
1046 :
1047 : // Rank devicePhyId 作为地址
1048 0 : HcclRankLinkInfo linkInfo {};
1049 0 : linkInfo.userRank = paraVector_[dstRank].userRank;
1050 0 : linkInfo.devicePhyId = dstDeviceId;
1051 0 : if (vnicRanksPort_.empty() || (userRankSize > 1 && IsSupportMC2(tag_) >= MC2_PLANE_MODE_COMBINE)) {
1052 0 : linkInfo.port = GetNicPort(paraVector_[dstRank].devicePhyId, ranksPort_,
1053 0 : linkInfo.userRank, isUseRankPort_);
1054 : } else {
1055 0 : linkInfo.port = GetNicPort(paraVector_[dstRank].devicePhyId, vnicRanksPort_,
1056 0 : linkInfo.userRank, isUseRankPort_);
1057 : }
1058 0 : HcclIpAddress ipAddress(linkInfo.devicePhyId);
1059 0 : DeviceIdType deviceidType =
1060 0 : useSuperPodMode_ ? (DeviceIdType::DEVICE_ID_TYPE_SDID) : (DeviceIdType::DEVICE_ID_TYPE_PHY_ID);
1061 : // rank个数小于等于1时,没有初始化ra资源,无法调用device侧hccp接口
1062 0 : if (userRankSize > 1) {
1063 0 : if(IsSupportMC2(tag_) >= MC2_PLANE_MODE_COMBINE) {
1064 0 : ipAddress = paraVector_[dstRank].nicIp.front();
1065 0 : CHK_PRT_RET(ipAddress.IsInvalid(),
1066 : HCCL_ERROR("[Get][IntraRankIPInfo] ipAddress is invalid when NIC, check the ip configuration for "
1067 : "dstRank[%u]",
1068 : dstRank),
1069 : HCCL_E_PARA);
1070 : } else {
1071 0 : CHK_RET(hrtRaGetSingleSocketVnicIpInfo(
1072 : paraVector_[rank_].devicePhyId, deviceidType, linkInfo.devicePhyId, ipAddress));
1073 : }
1074 : }
1075 0 : linkInfo.ip = ipAddress;
1076 0 : linkInfo.socketsPerLink = 1;
1077 :
1078 0 : HCCL_DEBUG("[Get][IntraRankIPInfo] tag[%s], userRank[%u], destRank[%u], localRole[%d], port[%u], ip[%s], "
1079 : "devicePhyId[%u]",
1080 : tag_.c_str(),
1081 : rank_,
1082 : linkInfo.userRank,
1083 : localRole,
1084 : linkInfo.port,
1085 : linkInfo.ip.GetReadableAddress(),
1086 : paraVector_[rank_].devicePhyId);
1087 :
1088 0 : if (localRole == HcclSocketRole::SOCKET_ROLE_CLIENT) {
1089 0 : dstServerMap.insert(std::make_pair(linkInfo.userRank, linkInfo));
1090 0 : } else if (localRole == HcclSocketRole::SOCKET_ROLE_SERVER) {
1091 0 : dstClientMap.insert(std::make_pair(linkInfo.userRank, linkInfo));
1092 : } else {
1093 : // 当前上层逻辑,保证 userRank_(当前 Rank) 在 userRanks 中
1094 0 : localIP = linkInfo.ip;
1095 : }
1096 0 : }
1097 0 : return HCCL_SUCCESS;
1098 : }
1099 :
1100 1 : HcclResult CommBase::GetIntraRankIPInfo(std::vector<u32> &dstIntraVec,
1101 : HcclIpAddress &localIP,
1102 : std::map<u32, HcclRankLinkInfo> &dstServerMap,
1103 : std::map<u32, HcclRankLinkInfo> &dstClientMap)
1104 : {
1105 2 : for (u32 dstRank : dstIntraVec) {
1106 1 : auto &rankInfo = paraVector_[dstRank];
1107 1 : HcclRankLinkInfo linkInfo {};
1108 1 : linkInfo.userRank = rankInfo.userRank;
1109 1 : linkInfo.devicePhyId = rankInfo.devicePhyId;
1110 1 : linkInfo.ip = isHaveCpuRank_ ? rankInfo.hostIp : rankInfo.nicIp[0];
1111 1 : if (!vnicRanksPort_.empty()) {
1112 0 : linkInfo.port = GetNicPort(linkInfo.devicePhyId, vnicRanksPort_, linkInfo.userRank, isUseRankPort_);
1113 : } else {
1114 1 : linkInfo.port = GetNicPort(linkInfo.devicePhyId, ranksPort_, linkInfo.userRank, isUseRankPort_);
1115 : }
1116 1 : linkInfo.socketsPerLink = 1;
1117 :
1118 : HcclSocketRole localRole;
1119 1 : if (paraVector_[rank_].userRank < linkInfo.userRank) {
1120 0 : dstClientMap.insert(std::make_pair(linkInfo.userRank, linkInfo));
1121 0 : localRole = HcclSocketRole::SOCKET_ROLE_CLIENT;
1122 1 : } else if (paraVector_[rank_].userRank > linkInfo.userRank) {
1123 0 : dstServerMap.insert(std::make_pair(linkInfo.userRank, linkInfo));
1124 0 : localRole = HcclSocketRole::SOCKET_ROLE_SERVER;
1125 : } else {
1126 1 : localIP = linkInfo.ip;
1127 1 : localRole = HcclSocketRole::SOCKET_ROLE_RESERVED;
1128 : }
1129 1 : HCCL_DEBUG("[Get][IntraRankIPInfo] userRank[%u], destRank[%u], localRole[%d], port[%u], ip[%s]",
1130 : userRank_, linkInfo.userRank, localRole, linkInfo.port, linkInfo.ip.GetReadableAddress());
1131 1 : }
1132 1 : return HCCL_SUCCESS;
1133 : }
1134 :
1135 0 : HcclResult CommBase::GetSuperNodeIntraRankIPInfo(std::map<u32, HcclSocketRole> &rankRole,
1136 : HcclIpAddress &localIP,
1137 : std::map<u32, HcclRankLinkInfo> &dstServerMap,
1138 : std::map<u32, HcclRankLinkInfo> &dstClientMap)
1139 : {
1140 0 : for (auto rankIter = rankRole.begin(); rankIter != rankRole.end(); rankIter++) {
1141 0 : u32 dstUserRank = paraVector_[rankIter->first].userRank;
1142 0 : HcclIpAddress ipAddr(paraVector_[rankIter->first].nicIp[0]);
1143 0 : if (!GetExternalInputInterHccsDisable()) {
1144 0 : CHK_RET(hrtRaGetSingleSocketVnicIpInfo(devicePhyId_, DeviceIdType::DEVICE_ID_TYPE_SDID,
1145 : paraVector_[rankIter->first].superDeviceId, ipAddr));
1146 : }
1147 0 : HcclSocketRole localRole = rankIter->second;
1148 :
1149 0 : HcclRankLinkInfo linkInfo {};
1150 0 : u32 dstRank = INVALID_VALUE_RANKID;
1151 0 : linkInfo.userRank = dstUserRank;
1152 0 : linkInfo.devicePhyId = -1;
1153 0 : linkInfo.ip = ipAddr;
1154 0 : CHK_RET(GetRankByUserRank(linkInfo.userRank, dstRank));
1155 0 : if (!vnicRanksPort_.empty()) {
1156 0 : linkInfo.port = GetNicPort(paraVector_[dstRank].devicePhyId, vnicRanksPort_,
1157 0 : linkInfo.userRank, isUseRankPort_);
1158 : } else {
1159 0 : linkInfo.port = GetNicPort(paraVector_[dstRank].devicePhyId, ranksPort_,
1160 0 : linkInfo.userRank, isUseRankPort_);
1161 : }
1162 0 : linkInfo.socketsPerLink = 1;
1163 :
1164 0 : HCCL_DEBUG("[Get][SuperNodeIntraRankIPInfo] userRank[%u], destRank[%u], localRole[%d], port[%u], ip[%s]",
1165 : userRank_, dstUserRank, localRole, linkInfo.port, linkInfo.ip.GetReadableAddress());
1166 :
1167 0 : if (localRole == HcclSocketRole::SOCKET_ROLE_CLIENT) {
1168 0 : dstServerMap.insert(std::make_pair(dstUserRank, linkInfo));
1169 0 : } else if (localRole == HcclSocketRole::SOCKET_ROLE_SERVER) {
1170 0 : dstClientMap.insert(std::make_pair(dstUserRank, linkInfo));
1171 : } else {
1172 : // 当前上层逻辑,保证 userRank_(当前 Rank) 在 userRanks 中
1173 0 : localIP = linkInfo.ip;
1174 : }
1175 0 : }
1176 0 : return HCCL_SUCCESS;
1177 : }
1178 :
1179 1 : bool CommBase::IsSupportInterHccs(const u32 dstRank)
1180 : {
1181 : // 仅判断超节点内, 兼容打平通信域同时有server内和server间, 因此不判断server_id
1182 1 : bool isInterHccs = GetExternalInputInterHccsDisable() == false &&
1183 1 : paraVector_[rank_].deviceType == DevType::DEV_TYPE_910_93 &&
1184 2 : paraVector_[rank_].superPodId.empty() == false &&
1185 0 : paraVector_[rank_].superPodId == paraVector_[dstRank].superPodId;
1186 :
1187 1 : HCCL_INFO("[IsSupportInterHccs]rank[%u], superPodId[%s], dstRank[%u], dstSuperPodId[%s], isInterHccs[%d]",
1188 : rank_, paraVector_[rank_].superPodId.c_str(), dstRank, paraVector_[dstRank].superPodId.c_str(), isInterHccs);
1189 1 : return isInterHccs;
1190 : }
1191 :
1192 2 : void CommBase::SetMachineLinkMode(MachinePara &machinePara)
1193 : {
1194 2 : machinePara.linkMode = LinkMode::LINK_DUPLEX_MODE;
1195 2 : }
1196 :
1197 10 : HcclResult CommBase::SetHDCModeInfo(
1198 : std::unordered_map<std::string, std::map<u32, HcclIpAddress>> &rankDevicePhyIdNicInfoMap,
1199 : std::vector<u32> &ranksPort, std::vector<u32> &vnicRanksPort, bool isSetHDCModeInfo, bool isUseRankPort)
1200 : {
1201 10 : rankDevicePhyIdNicInfoMap_ = rankDevicePhyIdNicInfoMap;
1202 10 : ranksPort_ = ranksPort;
1203 10 : vnicRanksPort_ = vnicRanksPort;
1204 10 : isSetHDCModeInfo_ = isSetHDCModeInfo;
1205 10 : isUseRankPort_ = isUseRankPort;
1206 10 : return HCCL_SUCCESS;
1207 : }
1208 :
1209 0 : HcclResult CommBase::CheckExchangeInfo(const std::shared_ptr<Transport> &link, const s32 deviceId)
1210 : {
1211 : // 算子一致性校验
1212 0 : u64 exchangeInfoLength = RankConsistentcyChecker::GetInstance(deviceId).GetRankConsistentDataLength();
1213 0 : std::vector<u8> recvData = link->GetExchangeInfo();
1214 0 : if (recvData.size() != 0) {
1215 0 : CHK_PRT_RET(recvData.size() != exchangeInfoLength,
1216 : HCCL_ERROR("[Check][ExchangeInfo]remote exchangInfo size[%zu], local exchangeInfo size[%llu]",
1217 : recvData.size(), exchangeInfoLength), HCCL_E_INTERNAL);
1218 0 : CHK_RET(RankConsistentcyChecker::GetInstance(deviceId).CheckFrameRecv(&recvData[0],
1219 : recvData.size(), tag_.c_str()));
1220 : }
1221 :
1222 0 : return HCCL_SUCCESS;
1223 0 : }
1224 :
1225 0 : u32 CommBase::IsSupportMC2(const std::string &tag)
1226 : {
1227 0 : const std::string &suffix = HCCL_MC2_MULTISERVER_SUFFIX;
1228 0 : u32 mc2MultiServerType = MC2_PLANE_MODE_HOST; // 非直驱场景
1229 0 : if (tag.size() > suffix.size() &&
1230 0 : tag.compare(tag.size() - suffix.size(), suffix.size(), suffix) == 0) {
1231 0 : mc2MultiServerType = MC2_PLANE_MODE_COMBINE; // 非分层建链场景
1232 0 : if(tag.find("_HIE") != std::string::npos) {
1233 0 : mc2MultiServerType = MC2_PLANE_MODE_HIERARCHY; // 分层建链场景
1234 : }
1235 : }
1236 0 : return mc2MultiServerType;
1237 : }
1238 : } // namespace hccl
|