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 "aicpu_ts_urma_channel.h"
12 : #include "endpoint.h"
13 : #include "orion_adpt_utils.h"
14 : #include "hcomm_c_adpt.h"
15 : #include "config_log.h"
16 :
17 : // Orion
18 : #include "adapter_rts_common.h"
19 : #include "topo_common_types.h"
20 : #include "virtual_topo.h"
21 : #include "aicpu_res_package_helper.h"
22 : #include "makebufs_helper.h"
23 :
24 : namespace hcomm {
25 : constexpr uint16_t DEFAULT_LISTENING_PORT = 60001;
26 :
27 9 : AicpuTsUrmaChannel::AicpuTsUrmaChannel(EndpointHandle endpointHandle, const HcommChannelDesc &channelDesc):
28 9 : endpointHandle_(endpointHandle), channelDesc_(channelDesc) {}
29 :
30 10 : AicpuTsUrmaChannel::~AicpuTsUrmaChannel()
31 : {
32 9 : if (channelDesc_.socket == nullptr && socket_ != nullptr) {
33 0 : SocketMgr::GetInstance(devicePhyId_).PutSocket(socketConfig_, socket_);
34 0 : socket_ = nullptr;
35 : }
36 10 : }
37 :
38 2 : HcclResult AicpuTsUrmaChannel::ParseInputParam()
39 : {
40 : // 1. 从 endpointHandle_,获得 localEp_ 和 rdmaHandle_
41 : // TODO: 使用 HcommEndpointGet
42 2 : Endpoint* localEpPtr = reinterpret_cast<Endpoint*>(endpointHandle_);
43 2 : CHK_PTR_NULL(localEpPtr);
44 2 : localEp_ = localEpPtr->GetEndpointDesc();
45 2 : rdmaHandle_ = localEpPtr->GetRdmaHandle();
46 :
47 2 : HCCL_INFO("[AicpuTsUrmaChannel][%s] localProtocol[%d]", __func__, localEp_.protocol);
48 :
49 : // 2. 从 channelDesc_,获得 remoteEp_, socket_ 和 notifyNum
50 2 : remoteEp_ = channelDesc_.remoteEndpoint;
51 2 : socket_ = reinterpret_cast<Hccl::Socket*>(channelDesc_.socket);
52 2 : notifyNum_ = channelDesc_.notifyNum;
53 2 : commonRes_.bufferVec.clear();
54 :
55 2 : if (channelDesc_.exchangeAllMems) {
56 : // 3. Get memHandles from endpoint
57 1 : HCCL_INFO("[AicpuTsUrmaChannel][%s] exchangeAllMems == True. Get memHandles from endpoint.", __func__);
58 1 : std::shared_ptr<Hccl::LocalUbRmaBuffer> *memHandles = nullptr;
59 1 : uint32_t memHandleNum = 0;
60 1 : CHK_RET(static_cast<HcclResult>(HcommMemGetAllMemHandles(
61 : endpointHandle_, reinterpret_cast<void**>(&memHandles), &memHandleNum)));
62 1 : HCCL_INFO("[AicpuTsUrmaChannel][%s] Got memHandleNum[%u].", __func__, memHandleNum);
63 2 : for (uint32_t i = 0; i < memHandleNum; ++i) {
64 1 : std::shared_ptr<Hccl::LocalUbRmaBuffer> &localUbRmaBuffer = memHandles[i];
65 1 : HCCL_INFO("[AicpuTsUrmaChannel][%s] Got memHandle No.%u: addr[0x%llx], size[0x%llx], type[%d], memInfo[%s].",
66 : __func__, i, static_cast<unsigned long long>(localUbRmaBuffer->GetAddr()),
67 : static_cast<unsigned long long>(localUbRmaBuffer->GetSize()),
68 : static_cast<int>(localUbRmaBuffer->GetBuf()->GetMemType()),
69 : localUbRmaBuffer->GetBuf()->GetMemInfo().c_str());
70 1 : commonRes_.bufferVec.push_back(localUbRmaBuffer.get());
71 : }
72 : } else {
73 : // 3. 从 channelDesc 的 memHandle,获得 bufs_
74 1 : HCCL_INFO("[AicpuTsUrmaChannel][%s] exchangeAllMems == false. Get memHandles from channelDesc.", __func__);
75 1 : CHK_RET(MakeRmaBufferVecFromMemHandles(
76 : channelDesc_.memHandles, channelDesc_.memHandleNum, commonRes_.bufferVec, "AicpuTsUrmaChannel"));
77 : }
78 :
79 2 : return HCCL_SUCCESS;
80 : }
81 :
82 0 : HcclResult AicpuTsUrmaChannel::BuildAttr()
83 : {
84 0 : attr_.devicePhyId = localEp_.loc.device.devPhyId;
85 0 : attr_.opMode = Hccl::OpMode::OPBASE;
86 0 : return HCCL_SUCCESS;
87 : }
88 :
89 0 : HcclResult AicpuTsUrmaChannel::BuildConnection()
90 : {
91 0 : UbConnBuildContext ctx;
92 0 : CHK_RET(PrepareUbConnBuildContext(localEp_, remoteEp_, channelDesc_.qos, ctx));
93 :
94 0 : Hccl::OpMode opMode = Hccl::OpMode::OPBASE;
95 0 : bool devUsed = true; // aicpu 为 true
96 0 : std::unique_ptr<Hccl::DevUbConnection> ubConn = nullptr;
97 0 : switch (ctx.protocol) {
98 0 : case Hccl::LinkProtocol::UB_TP:
99 0 : EXCEPTION_CATCH(
100 : ubConn = std::make_unique<Hccl::DevUbTpConnection>(rdmaHandle_, ctx.locAddr, ctx.rmtAddr, opMode,
101 : devUsed, Hccl::HrtUbJfcMode::STARS_POLL, Hccl::IpAddress(), Hccl::IpAddress(), ctx.qosPre),
102 : return HCCL_E_PTR
103 : );
104 0 : break;
105 0 : case Hccl::LinkProtocol::UB_CTP:
106 0 : EXCEPTION_CATCH(
107 : ubConn = std::make_unique<Hccl::DevUbCtpConnection>(rdmaHandle_, ctx.locAddr, ctx.rmtAddr, opMode,
108 : devUsed, Hccl::HrtUbJfcMode::STARS_POLL, Hccl::IpAddress(), Hccl::IpAddress(), ctx.qosPre),
109 : return HCCL_E_PTR
110 : );
111 0 : break;
112 0 : default:
113 0 : HCCL_ERROR("%s No LinkProtocol to match", __func__);
114 0 : break;
115 : }
116 0 : CHK_SMART_PTR_NULL(ubConn);
117 :
118 0 : if (devBaseAttr_.maxReadSize == 0 || devBaseAttr_.maxWriteSize == 0) {
119 0 : HCCL_ERROR("%s maxReadSize[%u] or maxWriteSize[%u] must not be zero", __func__, devBaseAttr_.maxReadSize,
120 : devBaseAttr_.maxWriteSize);
121 0 : return HCCL_E_PARA;
122 : }
123 0 : ubConn->SetMaxReadSize(devBaseAttr_.maxReadSize);
124 0 : ubConn->SetMaxWriteSize(devBaseAttr_.maxWriteSize);
125 0 : HCCL_INFO("%s maxReadSize[%u], maxWriteSize[%u]", __func__, devBaseAttr_.maxReadSize, devBaseAttr_.maxWriteSize);
126 :
127 0 : commonRes_.connVec.clear();
128 0 : commonRes_.connVec.emplace_back(ubConn.get());
129 0 : connections_.clear();
130 0 : connections_.push_back(std::move(ubConn));
131 0 : return HCCL_SUCCESS;
132 0 : }
133 :
134 0 : HcclResult AicpuTsUrmaChannel::BuildNotify()
135 : {
136 0 : localNotifies_.clear();
137 0 : commonRes_.notifyVec.clear();
138 0 : bool devUsed = true;
139 0 : for (uint32_t i = 0; i < notifyNum_; ++i) {
140 0 : std::unique_ptr<Hccl::UbLocalNotify> notifyPtr = nullptr;
141 0 : EXCEPTION_CATCH(
142 : notifyPtr = std::make_unique<Hccl::UbLocalNotify>(rdmaHandle_, devUsed),
143 : return HCCL_E_PTR
144 : );
145 0 : commonRes_.notifyVec.push_back(notifyPtr.get());
146 0 : localNotifies_.push_back(std::move(notifyPtr));
147 0 : }
148 0 : return HCCL_SUCCESS;
149 : }
150 :
151 0 : HcclResult AicpuTsUrmaChannel::BuildUbMemTransport()
152 : {
153 0 : Hccl::BaseMemTransport::LocCntNotifyRes locCntNotifyRes{};
154 0 : locCntNotifyRes.vec.clear();
155 0 : locCntNotifyRes.desc.clear();
156 0 : const Hccl::Socket &socket = *socket_;
157 :
158 0 : Hccl::LinkData linkData = BuildDefaultLinkData();
159 0 : CHK_RET(EndpointDescPairToLinkData(localEp_, remoteEp_, linkData));
160 :
161 0 : bool isRecvFirst = socket.GetRole() == Hccl::SocketRole::CLIENT ? true : false;
162 :
163 : // make_unique / make_shared / release 包一层抛异常的宏
164 0 : EXCEPTION_CATCH(
165 : memTransport_ = std::make_unique<Hccl::UbMemTransport>(
166 : commonRes_, attr_, linkData, socket, rdmaHandle_, locCntNotifyRes, isRecvFirst
167 : ),
168 : return HCCL_E_PTR
169 : );
170 0 : return HCCL_SUCCESS;
171 0 : }
172 :
173 0 : HcclResult AicpuTsUrmaChannel::BuildSocket()
174 : {
175 0 : if (socket_ != nullptr) {
176 0 : return HCCL_SUCCESS;
177 : }
178 0 : HCCL_INFO("[AicpuTsUrmaChannel][%s] socket ptr is NULL, rebuildSocket", __func__);
179 0 : Hccl::IpAddress ipaddr{};
180 0 : CHK_RET(CommAddrToIpAddress(localEp_.commAddr, ipaddr));
181 0 : Hccl::DevNetPortType type = Hccl::DevNetPortType(Hccl::ConnectProtoType::UB);
182 0 : Hccl::PortData localPort = Hccl::PortData(static_cast<Hccl::RankId>(localEp_.loc.device.devPhyId), type, 0, ipaddr);
183 0 : if (channelDesc_.role == HCOMM_SOCKET_ROLE_RESERVED) {
184 : Hccl::SocketHandle socketHandle
185 0 : = Hccl::SocketHandleManager::GetInstance().Create(localEp_.loc.device.devPhyId, localPort);
186 0 : EXCEPTION_CATCH(serverSocket_ = std::make_unique<Hccl::Socket>(socketHandle, ipaddr, DEFAULT_LISTENING_PORT, ipaddr, "server",
187 : Hccl::SocketRole::SERVER, Hccl::NicType::DEVICE_NIC_TYPE),
188 : return HCCL_E_PARA);
189 0 : HCCL_INFO("[AicpuTsUrmaChannel][%s] listen_socket_info[%s]", __func__, serverSocket_->Describe().c_str());
190 0 : EXCEPTION_CATCH(serverSocket_->Listen(), return HCCL_E_INTERNAL);
191 0 : Hccl::LinkData linkData = BuildDefaultLinkData();
192 0 : CHK_RET(EndpointDescPairToLinkData(localEp_, remoteEp_, linkData));
193 0 : HCCL_INFO("[AicpuTsUrmaChannel][%s] built linkData: %s", __func__, linkData.Describe().c_str());
194 0 : std::string socketTag = (channelDesc_.channelName != nullptr)
195 0 : ? std::string(channelDesc_.channelName) : "AUTOMATIC_SOCKET_TAG";
196 0 : bool noRankId = true;
197 0 : Hccl::SocketConfig socketConfig = Hccl::SocketConfig(linkData, socketTag, noRankId);
198 0 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).GetSocket(socketConfig, socket_));
199 0 : } else {
200 0 : uint16_t port = channelDesc_.port;
201 0 : if (port == 0) {
202 0 : port = DEFAULT_LISTENING_PORT;
203 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] channelDesc port is 0, use default port [%u]", __func__, port);
204 : }
205 0 : Hccl::LinkData linkData = BuildDefaultLinkData();
206 0 : CHK_RET(EndpointDescPairToLinkData(localEp_, remoteEp_, linkData));
207 0 : HCCL_INFO("[AicpuTsUrmaChannel][%s] built linkData: %s", __func__, linkData.Describe().c_str());
208 0 : std::string socketTag = (channelDesc_.channelName != nullptr)
209 0 : ? std::string(channelDesc_.channelName) : "AUTOMATIC_SOCKET_TAG";
210 0 : bool isServer = (channelDesc_.role == HCOMM_SOCKET_ROLE_SERVER);
211 0 : Hccl::SocketConfig socketConfig = Hccl::SocketConfig(linkData, port, socketTag, isServer);
212 0 : CHK_RET(SocketMgr::GetInstance(devicePhyId_).GetSocket(socketConfig, socket_));
213 0 : }
214 0 : return HCCL_SUCCESS;
215 : }
216 :
217 0 : HcclResult AicpuTsUrmaChannel::Init()
218 : {
219 : /*
220 : Argue result: make_unique 配合一场捕获的宏 EXCEPTION CATCH
221 : Attention: const 和引用
222 : */
223 : // TODO: 处理抛异常
224 : s32 devLogicId;
225 0 : CHK_RET(ParseInputParam());
226 0 : CHK_RET(hrtGetDevice(&devLogicId));
227 0 : CHK_RET(hrtGetDevicePhyIdByIndex(static_cast<u32>(devLogicId), devicePhyId_));
228 0 : CHK_RET(StartListen());
229 0 : CHK_RET(BuildSocket());
230 0 : CHK_RET(BuildAttr());
231 : /*
232 : HccpRaGetDevBaseAttr
233 : 获取urma read/write 单个wr的最大传输数据大小
234 : 调用前,rdmaHandle_要在ParseInputParam中被赋值好,之后BuildConnection会使用获取的属性
235 : */
236 0 : CHK_RET(HccpRaGetDevBaseAttr(rdmaHandle_, &devBaseAttr_));
237 0 : CHK_RET(BuildConnection());
238 0 : CHK_RET(BuildNotify());
239 0 : CHK_RET(BuildUbMemTransport());
240 :
241 0 : return HCCL_SUCCESS;
242 : }
243 :
244 0 : HcclResult AicpuTsUrmaChannel::GetNotifyNum(uint32_t *notifyNum) const
245 : {
246 0 : *notifyNum = this->notifyNum_;
247 0 : return HCCL_SUCCESS;
248 : }
249 :
250 0 : HcclResult AicpuTsUrmaChannel::GetRemoteMems(uint32_t *memNum, CommMem **remoteMem, char ***memInfos)
251 : {
252 0 : return memTransport_->GetRemoteMems(memNum, remoteMem, memInfos);
253 : }
254 :
255 2 : ChannelStatus AicpuTsUrmaChannel::GetStatus()
256 : {
257 2 : ChannelStatus out = Channel::TransportStatusToChannelStatus(memTransport_->GetStatus());
258 :
259 2 : if (isFirstPrintChannelInfo_ && out == ChannelStatus::READY) {
260 2 : std::string channelInfo = "create channel info:channel handle[";
261 2 : channelInfo.append(std::to_string(reinterpret_cast<uint64_t>(this)));
262 2 : channelInfo.append("] ");
263 2 : HcclResult ret = memTransport_->Describe(channelInfo);
264 2 : if (ret != HCCL_SUCCESS) {
265 1 : HCCL_ERROR("[AicpuTsUrmaChannel][%s] Describe channel info failed, ret=%d", __func__, ret);
266 1 : out = ChannelStatus::FAILED;
267 : } else {
268 1 : channelInfo.append(" TA[RM]"); // 目前TA只支持RM
269 1 : HCCL_CONFIG_DEBUG(hccl::HCCL_RES, "%s", channelInfo.c_str());
270 : }
271 2 : isFirstPrintChannelInfo_ = false;
272 2 : }
273 2 : return out;
274 : }
275 :
276 0 : HcclResult SetModuleDataName(Hccl::ModuleData &module, const std::string &name)
277 : {
278 0 : int ret = strcpy_s(module.name, sizeof(module.name), name.c_str());
279 0 : if (ret != 0) {
280 0 : HCCL_ERROR("[SetModuleDataName] strcpy_s name %s failed", name.c_str());
281 0 : return HCCL_E_INTERNAL;
282 : }
283 :
284 0 : return HCCL_SUCCESS;
285 : }
286 :
287 0 : HcclResult AicpuTsUrmaChannel::PackOpData(std::vector<char> &data)
288 : {
289 0 : std::vector<Hccl::ModuleData> dataVec;
290 0 : dataVec.resize(Hccl::AicpuResMgrType::__COUNT__);
291 :
292 0 : Hccl::AicpuResMgrType resType = Hccl::AicpuResMgrType::STREAM;
293 0 : CHK_RET(SetModuleDataName(dataVec[resType], "UbMemTransport"));
294 :
295 0 : std::vector<char> result;
296 0 : Hccl::BinaryStream binaryStream;
297 0 : binaryStream << memTransport_->GetUniqueIdV2();
298 :
299 0 : binaryStream.Dump(result);
300 :
301 0 : dataVec[resType].data = result;
302 :
303 : Hccl::AicpuResPackageHelper helper;
304 0 : data = helper.GetPackedData(dataVec);
305 :
306 0 : return HCCL_SUCCESS;
307 0 : }
308 :
309 0 : HcclResult AicpuTsUrmaChannel::H2DResPack(std::vector<char>& buffer)
310 : {
311 0 : CHK_RET(PackOpData(buffer));
312 0 : HCCL_INFO("[AicpuTsUrmaChannel][%s] Pack Buffer data[%p], Pack Buffer size[%zu].",
313 : __func__, buffer.data(), buffer.size());
314 0 : return HCCL_SUCCESS;
315 : }
316 :
317 1 : HcclResult AicpuTsUrmaChannel::Clean()
318 : {
319 1 : memTransport_.reset();
320 1 : return HCCL_SUCCESS;
321 : }
322 :
323 1 : HcclResult AicpuTsUrmaChannel::Resume()
324 : {
325 1 : BuildSocket();
326 1 : BuildConnection();
327 1 : BuildUbMemTransport();
328 1 : return HCCL_SUCCESS;
329 : }
330 :
331 2 : HcclResult AicpuTsUrmaChannel::UpdateMemInfo(HcommMemHandle *memHandles, uint32_t memHandleNum)
332 : {
333 2 : std::vector<Hccl::LocalRmaBuffer *> bufferVecTemp;
334 2 : CHK_RET(MakeRmaBufferVecFromMemHandles(memHandles, memHandleNum, bufferVecTemp, "AicpuTsUrmaChannel"));
335 1 : CHK_RET(memTransport_->UpdateMemInfo(bufferVecTemp));
336 1 : commonRes_.bufferVec.insert(commonRes_.bufferVec.end(), bufferVecTemp.begin(), bufferVecTemp.end());
337 1 : return HCCL_SUCCESS;
338 2 : }
339 :
340 : // 返回当前 channel 类型,供上层区分不同 channel 的能力和行为
341 0 : HcommChannelKind AicpuTsUrmaChannel::GetChannelKind() const
342 : {
343 0 : return HcommChannelKind::AICPU_TS_URMA;
344 : }
345 :
346 0 : HcclResult AicpuTsUrmaChannel::NotifyRecord(const uint32_t remoteNotifyIdx)
347 : {
348 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] not supported yet.", __func__);
349 0 : return HCCL_E_NOT_SUPPORT;
350 : }
351 :
352 0 : HcclResult AicpuTsUrmaChannel::NotifyWait(const uint32_t localNotifyIdx, const uint32_t timeout)
353 : {
354 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] not supported yet.", __func__);
355 0 : return HCCL_E_NOT_SUPPORT;
356 : }
357 :
358 0 : HcclResult AicpuTsUrmaChannel::WriteWithNotify(void *dst, const void *src, const uint64_t len, uint32_t remoteNotifyIdx)
359 : {
360 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] not supported yet.", __func__);
361 0 : return HCCL_E_NOT_SUPPORT;
362 : }
363 :
364 0 : HcclResult AicpuTsUrmaChannel::Write(void *dst, const void *src, uint64_t len)
365 : {
366 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] not supported yet.", __func__);
367 0 : return HCCL_E_NOT_SUPPORT;
368 : }
369 :
370 0 : HcclResult AicpuTsUrmaChannel::Read(void *dst, const void *src, uint64_t len)
371 : {
372 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] not supported yet.", __func__);
373 0 : return HCCL_E_NOT_SUPPORT;
374 : }
375 :
376 0 : HcclResult AicpuTsUrmaChannel::ChannelFence()
377 : {
378 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] not supported yet.", __func__);
379 0 : return HCCL_E_NOT_SUPPORT;
380 : }
381 :
382 1 : HcclResult AicpuTsUrmaChannel::StartListen()
383 : {
384 1 : if (channelDesc_.role != HCOMM_SOCKET_ROLE_SERVER) {
385 1 : return HCCL_SUCCESS;
386 : }
387 :
388 0 : uint16_t port = channelDesc_.port;
389 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] Start. EndpointHandle[0x%llx], port[%u]", __func__, reinterpret_cast<uint64_t>(endpointHandle_), port);
390 0 : if (port == 0) {
391 0 : port = DEFAULT_LISTENING_PORT;
392 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] channelDesc port is 0, use default port [%u]", __func__, port);
393 : }
394 0 : CHK_RET(static_cast<HcclResult>(HcommEndpointStartListen(endpointHandle_, port, nullptr)));
395 0 : HCCL_INFO("[AicpuTsUrmaChannel::%s] SUCCESS. port[%u].", __func__, port);
396 0 : return HCCL_SUCCESS;
397 : }
398 :
399 : } // namespace hcomm
|