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 "endpoint_pair.h"
12 : #include "socket_config.h"
13 : #include "hcomm_c_adpt.h"
14 : #include "orion_adpt_utils.h"
15 : #include "channel_process.h"
16 : #include "comm_engine_utils.h"
17 :
18 : #include "hcom_common.h"
19 : #include "exception_handler.h"
20 :
21 : namespace hcomm {
22 :
23 19 : EndpointPair::~EndpointPair()
24 : {
25 36 : for (auto& channels : channelHandles_) {
26 17 : if (channels.second.empty()) {
27 3 : continue;
28 : }
29 14 : (void)ChannelProcess::ChannelDestroy(channels.second.data(), channels.second.size());
30 : }
31 19 : }
32 :
33 18 : HcclResult EndpointPair::Init()
34 : {
35 18 : EXCEPTION_CATCH(socketMgr_ = std::make_unique<SocketMgr>(), return HCCL_E_PTR);
36 18 : channelHandles_.clear();
37 : s32 devLogicId;
38 18 : CHK_RET(hrtGetDevice(&devLogicId));
39 18 : CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(devLogicId), devicePhyId_));
40 :
41 18 : return HCCL_SUCCESS;
42 : }
43 :
44 3 : HcclResult EndpointPair::GetHostSocketWithRank(
45 : const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, const uint32_t listenPort,
46 : u32 reuseIdx, Hccl::Socket*& socket)
47 : {
48 3 : uint32_t connectMode = 0;
49 3 : Hccl::LinkData linkData = BuildDefaultLinkData();
50 3 : CHK_RET(EndpointDescPairToLinkData(localEndpointDesc_, remoteEndpointDesc_, linkData, reuseIdx));
51 3 : std::string linkTag = socketTag;
52 3 : if (linkData.GetReuseIdx() != "0") {
53 0 : linkTag += ("_" + linkData.GetReuseIdx());
54 : }
55 :
56 : DevType devType;
57 3 : CHK_RET(hrtGetDeviceType(devType));
58 3 : if (devType == DevType::DEV_TYPE_910B && localEndpointDesc_.loc.locType != remoteEndpointDesc_.loc.locType) {
59 0 : connectMode = 1;
60 : }
61 :
62 : /* A2: host nic(cpu roce channel) -- device nic(transport ibv)时,两边ip地址格式不一样,判断大小算法不匹配
63 : * 修改成按照rank id大小判断server和client */
64 3 : Hccl::SocketConfig socketConfig = Hccl::SocketConfig(linkData, listenPort, linkTag, connectMode, myRank, rmtRank);
65 3 : CHK_RET(socketMgr_->GetHostSocket(socketConfig, socket));
66 3 : return HCCL_SUCCESS;
67 3 : }
68 :
69 20 : HcclResult EndpointPair::EnsureSocketMgrCompat(const uint32_t myRank, const std::string& socketTag)
70 : {
71 20 : if (!socketMgrCompat_) {
72 8 : int32_t devLogicId = HcclGetThreadDeviceId();
73 8 : uint32_t devPhyId{0};
74 8 : CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<uint32_t>(devLogicId), devPhyId));
75 8 : EXCEPTION_CATCH(
76 : socketMgrCompat_ = std::make_unique<Hccl::SocketManager>(myRank, devPhyId, devLogicId, socketTag),
77 : return HCCL_E_PTR);
78 8 : CHK_PTR_NULL(rankIpPortMap_);
79 8 : socketMgrCompat_->SetDeviceServerListenPortMap(*rankIpPortMap_);
80 : }
81 20 : return HCCL_SUCCESS;
82 : }
83 :
84 40 : Hccl::SocketConfig EndpointPair::BuildSocketConfig(const Hccl::LinkData& linkData, const std::string& socketTag)
85 : {
86 40 : std::string linkTag = socketTag;
87 40 : if (linkData.GetReuseIdx() != "0") {
88 6 : linkTag += ("_" + linkData.GetReuseIdx());
89 : }
90 80 : return Hccl::SocketConfig(linkData.GetRemoteRankId(), linkData, linkTag);
91 40 : }
92 :
93 23 : HcclResult EndpointPair::HandleHostSocketOrBuildLinkData(
94 : const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
95 : const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId,
96 : Hccl::LinkData& linkData, bool& isHost)
97 : {
98 23 : if (localEndpointDesc_.loc.locType == EndpointLocType::ENDPOINT_LOC_TYPE_HOST) {
99 3 : std::string socketTagPrefix = socketTag;
100 3 : if (myRank <= rmtRank) {
101 2 : socketTagPrefix += "_" + std::to_string(myRank) + "_" + std::to_string(rmtRank);
102 : } else {
103 1 : socketTagPrefix += "_" + std::to_string(rmtRank) + "_" + std::to_string(myRank);
104 : }
105 3 : CHK_RET(this->GetHostSocketWithRank(myRank, rmtRank, socketTagPrefix, listenPort, reuseIdx, socket));
106 3 : isHost = true;
107 3 : return HCCL_SUCCESS;
108 3 : }
109 20 : isHost = false;
110 20 : CHK_RET(EndpointDescPairToLinkDataWithRankIds(
111 : myRank, rmtRank, localEndpointDesc_, remoteEndpointDesc_, linkData, devicePhyId, remoteDevicePhyId, reuseIdx));
112 20 : return HCCL_SUCCESS;
113 : }
114 :
115 23 : HcclResult EndpointPair::GetSocketInternal(
116 : const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
117 : const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId,
118 : bool connectMode)
119 : {
120 23 : Hccl::LinkData linkData = BuildDefaultLinkData();
121 23 : bool isHost = false;
122 23 : CHK_RET(HandleHostSocketOrBuildLinkData(
123 : myRank, rmtRank, socketTag, reuseIdx, listenPort, socket, devicePhyId, remoteDevicePhyId, linkData, isHost));
124 23 : if (isHost) {
125 3 : return HCCL_SUCCESS;
126 : }
127 : EXCEPTION_HANDLE_BEGIN
128 20 : Hccl::SocketConfig socketConfig = BuildSocketConfig(linkData, socketTag);
129 20 : if (connectMode) {
130 20 : CHK_PTR_NULL(socketMgrCompat_);
131 20 : socketMgrCompat_->ConnectSockets(socketConfig);
132 : } else {
133 0 : CHK_RET(EnsureSocketMgrCompat(myRank, socketTag));
134 0 : socketMgrCompat_->BatchCreateSockets(socketConfig);
135 : }
136 20 : socket = socketMgrCompat_->GetConnectedSocket(socketConfig);
137 20 : CHK_PTR_NULL(socket);
138 20 : EXCEPTION_HANDLE_END
139 20 : return HCCL_SUCCESS;
140 : }
141 :
142 21 : HcclResult EndpointPair::ServerInit(
143 : const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx, uint32_t devicePhyId,
144 : uint32_t remoteDevicePhyId)
145 : {
146 21 : if (localEndpointDesc_.loc.locType == EndpointLocType::ENDPOINT_LOC_TYPE_HOST) {
147 : // host网卡不走device的socket监听
148 1 : return HCCL_SUCCESS;
149 : }
150 : // server监听
151 20 : Hccl::LinkData linkData = BuildDefaultLinkData();
152 20 : CHK_RET(EndpointDescPairToLinkDataWithRankIds(
153 : myRank, rmtRank, localEndpointDesc_, remoteEndpointDesc_, linkData, devicePhyId, remoteDevicePhyId, reuseIdx));
154 : EXCEPTION_HANDLE_BEGIN
155 20 : CHK_RET(EnsureSocketMgrCompat(myRank, socketTag));
156 20 : Hccl::SocketConfig socketConfig = BuildSocketConfig(linkData, socketTag);
157 : // 调用sock的server监听接口
158 20 : socketMgrCompat_->ServerListen(socketConfig);
159 20 : EXCEPTION_HANDLE_END
160 :
161 20 : return HCCL_SUCCESS;
162 : }
163 :
164 21 : HcclResult EndpointPair::GetConnectedSocket(
165 : const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
166 : const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId)
167 : {
168 : // 该接口内进行建链和获取socket
169 21 : return GetSocketInternal(
170 21 : myRank, rmtRank, socketTag, reuseIdx, listenPort, socket, devicePhyId, remoteDevicePhyId, true);
171 : }
172 :
173 2 : HcclResult EndpointPair::GetSocket(
174 : const uint32_t myRank, const uint32_t rmtRank, const std::string& socketTag, u32 reuseIdx,
175 : const uint32_t listenPort, Hccl::Socket*& socket, uint32_t devicePhyId, uint32_t remoteDevicePhyId)
176 : {
177 : // 临时方案:支持混跑新增,非Roce场景走orion socketMgr实现server socket复用
178 2 : return GetSocketInternal(
179 2 : myRank, rmtRank, socketTag, reuseIdx, listenPort, socket, devicePhyId, remoteDevicePhyId, false);
180 : }
181 :
182 34 : HcclResult EndpointPair::CreateChannel(
183 : EndpointHandle endpointHandle, CommEngine engine, u32 reuseIdx, HcommChannelDesc* channelDescs,
184 : ChannelHandle* channels)
185 : {
186 34 : if (channelHandles_.find(engine) == channelHandles_.end() || channelHandles_[engine].size() <= reuseIdx) {
187 24 : CHK_RET_UNAVAIL(
188 : static_cast<HcclResult>(HcommCollectiveChannelCreate(endpointHandle, engine, channelDescs, 1, channels)));
189 22 : channelHandles_[engine].push_back(channels[0]);
190 22 : return HCCL_SUCCESS;
191 : }
192 :
193 10 : channels[0] = channelHandles_[engine][reuseIdx];
194 10 : if (channelDescs->memHandleNum > 1) {
195 0 : CHK_RET(static_cast<HcclResult>(
196 : HcommChannelUpdateMemInfo(channelDescs->memHandles + 1, channelDescs->memHandleNum - 1, channels[0])));
197 : }
198 10 : return HCCL_SUCCESS;
199 : }
200 :
201 : // 找到对应的channelhandle,调用HcommChannelDestroy销毁平台层对象,并删除channelHandles_中的channelHandle元素
202 6 : HcclResult EndpointPair::DestroyChannel(CommEngine engine, u32 reuseIdx)
203 : {
204 6 : if (IsChannelNotExist(engine, reuseIdx)) {
205 1 : HCCL_WARNING(
206 : "EndpointPair::DestroyChannel: engine[%s] reuseIdx[%u], channelHandle size[%u],"
207 : "channel not found, skip destroy channel",
208 : GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), reuseIdx, channelHandles_[engine].size());
209 1 : return HCCL_SUCCESS;
210 : }
211 5 : HCCL_INFO(
212 : "EndpointPair::DestroyChannel: engine[%s] reuseIdx[%u], channelHandle size[%u],"
213 : "start destroy channel",
214 : GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), reuseIdx, channelHandles_[engine].size());
215 5 : ChannelHandle channelHandle = channelHandles_[engine][reuseIdx];
216 5 : CHK_RET(static_cast<HcclResult>(HcommChannelDestroy(&channelHandle, 1)));
217 : // 去掉channelHandles_中reuseIdx位置的channelHandle
218 5 : channelHandles_[engine].erase(channelHandles_[engine].begin() + reuseIdx);
219 5 : HCCL_INFO(
220 : "EndpointPair::DestroyChannel: engine[%s] reuseIdx[%u] destroy channel success,"
221 : "channelHandle size[%u]",
222 : GetEnumToString(GetCommEngineStatusStrMap(), engine).c_str(), reuseIdx, channelHandles_[engine].size());
223 5 : return HCCL_SUCCESS;
224 : }
225 :
226 : // 检查channel是否存在,channel不存在则返回true
227 40 : bool EndpointPair::IsChannelNotExist(CommEngine engine, u32 reuseIdx)
228 : {
229 40 : return channelHandles_.find(engine) == channelHandles_.end() || channelHandles_[engine].size() <= reuseIdx;
230 : }
231 :
232 1 : const std::unordered_map<CommEngine, std::vector<ChannelHandle>>& EndpointPair::GetChannelHandles()
233 : {
234 1 : return channelHandles_;
235 : }
236 :
237 : } // namespace hcomm
|