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 "aiv_ub_mem_transport.h"
12 : #include "exception_handler.h"
13 : #include "../../../../../../legacy/ascend950/unified_platform/resource/socket/socket.h"
14 : #include "../../../../../../legacy/ascend950/unified_platform/resource/buffer/exchange_ipc_buffer_dto.h"
15 : #include "../../../../../../legacy/ascend950/unified_platform/resource/mem/user_remote_mem_getter.h"
16 : #include "env_config/env_config.h"
17 :
18 : namespace hcomm {
19 :
20 9 : AivUbMemTransport::AivUbMemTransport(Hccl::Socket *socket, HcommChannelDesc &channelDesc) : socket_(socket),
21 9 : channelDesc_(channelDesc) {}
22 :
23 2 : HcclResult AivUbMemTransport::FillBufferVec(HcommMemHandle *memHandles, uint32_t bufferNum,
24 : std::vector<Hccl::LocalIpcRmaBuffer *> &bufferVec)
25 : {
26 2 : uint32_t totalBufferNum = localRmaBufferVec_.size() + bufferNum;
27 2 : if (UNLIKELY(totalBufferNum > MAX_BUFFER_NUM)) {
28 0 : HCCL_ERROR("[AivUbMemTransport][FillBufferVec] totalBufferNum[%u] exceeds limit[%u]", totalBufferNum, MAX_BUFFER_NUM);
29 0 : return HCCL_E_PARA;
30 : }
31 5 : for (uint32_t i = 0; i < bufferNum; ++i) {
32 3 : auto localIpcRmaBuffer = reinterpret_cast<Hccl::LocalIpcRmaBuffer *>(memHandles[i]);
33 3 : CHK_PTR_NULL(localIpcRmaBuffer);
34 3 : auto buf = localIpcRmaBuffer->GetBuf();
35 3 : CHK_PTR_NULL(buf);
36 3 : bufferVec.push_back(localIpcRmaBuffer);
37 3 : HCCL_INFO("[AivUbMemTransport][FillBufferVec] memHandleNum[%u] buffer[%s]", i, localIpcRmaBuffer->Describe().data());
38 : }
39 2 : return HCCL_SUCCESS;
40 : }
41 :
42 1 : HcclResult AivUbMemTransport::Init()
43 : {
44 1 : uint32_t bufferNum = channelDesc_.memHandleNum;
45 1 : if (bufferNum == 0) {
46 1 : HCCL_ERROR("[AivUbMemTransport][Init] bufferNum is 0.");
47 1 : return HCCL_E_PARA;
48 : }
49 0 : HCCL_INFO("[AivUbMemTransport][Init] channelDesc_.memHandleNum: %u", bufferNum);
50 0 : CHK_RET(FillBufferVec(channelDesc_.memHandles, bufferNum, localRmaBufferVec_));
51 :
52 0 : baseStatus_ = Hccl::TransportStatus::INIT;
53 0 : return HCCL_SUCCESS;
54 : }
55 :
56 8 : HcclResult AivUbMemTransport::IsSocketReady(bool &isReady)
57 : {
58 8 : CHK_PTR_NULL(socket_);
59 : EXCEPTION_HANDLE_BEGIN
60 8 : Hccl::SocketStatus socketStatus = socket_->GetAsyncStatus();
61 7 : if (socketStatus == Hccl::SocketStatus::OK) {
62 7 : baseStatus_ = Hccl::TransportStatus::SOCKET_OK;
63 7 : isReady = true;
64 0 : } else if (socketStatus == Hccl::SocketStatus::TIMEOUT) {
65 0 : baseStatus_ = Hccl::TransportStatus::SOCKET_TIMEOUT;
66 0 : isReady = false;
67 : }
68 1 : EXCEPTION_HANDLE_END
69 7 : return HCCL_SUCCESS;
70 : }
71 :
72 6 : void AivUbMemTransport::CheckStatusFuncResult(std::string funcName, HcclResult ret)
73 : {
74 6 : if (UNLIKELY(ret != HCCL_SUCCESS)) {
75 1 : HCCL_ERROR("[%s] fail ret[%d], aivUbStatus_[%d], baseStatus_[%d]",
76 : funcName.c_str(), ret, aivUbStatus_, baseStatus_);
77 1 : baseStatus_ = Hccl::TransportStatus::INVALID;
78 : }
79 6 : }
80 :
81 8 : Hccl::TransportStatus AivUbMemTransport::GetStatus()
82 : {
83 8 : if (baseStatus_ == Hccl::TransportStatus::READY || baseStatus_ == Hccl::TransportStatus::INVALID) {
84 0 : return baseStatus_;
85 8 : } else if (baseStatus_ == Hccl::TransportStatus::INIT) {
86 0 : aivUbStatus_ = AivUbMemTransportStatus::INIT;
87 : }
88 :
89 8 : bool isReady = false;
90 8 : if (UNLIKELY(IsSocketReady(isReady) != HCCL_SUCCESS)) {
91 1 : HCCL_ERROR("[%s] IsSocketReady fail, aivUbStatus_[%d], baseStatus_[%d]", __func__, aivUbStatus_, baseStatus_);
92 1 : baseStatus_ = Hccl::TransportStatus::INVALID;
93 1 : return baseStatus_;
94 : }
95 7 : if (!isReady) {
96 0 : return baseStatus_;
97 : }
98 7 : return UpdateStatus();
99 : }
100 :
101 7 : Hccl::TransportStatus AivUbMemTransport::UpdateStatus()
102 : {
103 7 : HCCL_INFO("%s aivUbStatus_[%d], baseStatus_[%d] start, aivUbStatus_::SOCKET_OK[%d]",
104 : __func__, aivUbStatus_, baseStatus_, AivUbMemTransportStatus::SOCKET_OK);
105 : HcclResult ret;
106 7 : switch (aivUbStatus_) {
107 0 : case AivUbMemTransportStatus::INIT:
108 0 : aivUbStatus_ = AivUbMemTransportStatus::SOCKET_OK;
109 0 : baseStatus_ = Hccl::TransportStatus::SOCKET_OK;
110 0 : break;
111 1 : case AivUbMemTransportStatus::SOCKET_OK:
112 1 : ret = SendDataSize();
113 1 : CheckStatusFuncResult("SendDataSize", ret);
114 1 : aivUbStatus_ = AivUbMemTransportStatus::SEND_DATA_SIZE;
115 1 : break;
116 1 : case AivUbMemTransportStatus::SEND_DATA_SIZE:
117 1 : ret = RecvDataSize();
118 1 : CheckStatusFuncResult("RecvDataSize", ret);
119 1 : aivUbStatus_ = AivUbMemTransportStatus::RECV_DATA_SIZE;
120 1 : break;
121 1 : case AivUbMemTransportStatus::RECV_DATA_SIZE:
122 1 : ret = SendMemInfo();
123 1 : CheckStatusFuncResult("SendMemInfo", ret);
124 1 : aivUbStatus_ = AivUbMemTransportStatus::SEND_MEM_INFO;
125 1 : break;
126 1 : case AivUbMemTransportStatus::SEND_MEM_INFO:
127 1 : ret = RecvMemInfo();
128 1 : CheckStatusFuncResult("RecvMemInfo", ret);
129 1 : aivUbStatus_ = AivUbMemTransportStatus::RECV_MEM_INFO;
130 1 : break;
131 2 : case AivUbMemTransportStatus::RECV_MEM_INFO:
132 2 : ret = RecvDataProcess();
133 2 : CheckStatusFuncResult("RecvDataProcess", ret);
134 2 : aivUbStatus_ = AivUbMemTransportStatus::RECV_MEM_FIN;
135 2 : break;
136 1 : case AivUbMemTransportStatus::RECV_MEM_FIN:
137 1 : aivUbStatus_ = AivUbMemTransportStatus::READY;
138 1 : baseStatus_ = Hccl::TransportStatus::READY;
139 1 : break;
140 0 : default:
141 0 : break;
142 : }
143 7 : HCCL_INFO("%s aivUbStatus_[%d], baseStatus_[%d]", __func__, aivUbStatus_, baseStatus_);
144 7 : return baseStatus_;
145 : }
146 :
147 1 : HcclResult AivUbMemTransport::SendDataSize()
148 : {
149 1 : HCCL_INFO("[%s] start", __func__);
150 :
151 1 : Hccl::BinaryStream binaryStream;
152 1 : CHK_RET(BufferPack(binaryStream, localRmaBufferVec_));
153 :
154 1 : binaryStream.Dump(sendData_);
155 1 : u32 sendSize = sendData_.size();
156 : EXCEPTION_HANDLE_BEGIN
157 1 : socket_->SendAsync(&sendSize, sizeof(sendSize));
158 0 : EXCEPTION_HANDLE_END
159 1 : HCCL_INFO("[%s] finished", __func__);
160 1 : return HCCL_SUCCESS;
161 1 : }
162 :
163 2 : HcclResult AivUbMemTransport::RecvDataSize()
164 : {
165 2 : HCCL_INFO("[%s] start", __func__);
166 :
167 : EXCEPTION_HANDLE_BEGIN
168 2 : socket_->RecvAsync(reinterpret_cast<u8 *>(&exchangeDataSize_), sizeof(exchangeDataSize_));
169 0 : EXCEPTION_HANDLE_END
170 2 : HCCL_INFO("[%s] finished", __func__);
171 2 : return HCCL_SUCCESS;
172 : }
173 :
174 2 : HcclResult AivUbMemTransport::SendMemInfo()
175 : {
176 2 : HCCL_INFO("[%s] start", __func__);
177 :
178 : EXCEPTION_HANDLE_BEGIN
179 2 : socket_->SendAsync(&sendData_[0], sendData_.size());
180 0 : EXCEPTION_HANDLE_END
181 2 : HCCL_INFO("[%s] finished", __func__);
182 2 : return HCCL_SUCCESS;
183 : }
184 :
185 3 : HcclResult AivUbMemTransport::BufferPack(Hccl::BinaryStream &binaryStream, std::vector<Hccl::LocalIpcRmaBuffer *> &bufferVec)
186 : {
187 3 : u32 vecSize = bufferVec.size();
188 3 : binaryStream << vecSize;
189 3 : HCCL_INFO("BufferPack vecSize=%u", vecSize);
190 :
191 6 : for (uint32_t i = 0; i < vecSize; ++i) {
192 3 : std::unique_ptr<Hccl::Serializable> dto = bufferVec[i]->GetExchangeDto();
193 3 : CHK_PTR_NULL(dto);
194 3 : dto->Serialize(binaryStream);
195 3 : HCCL_INFO("[%s] dto[%s]", __func__, dto->Describe().c_str());
196 3 : }
197 3 : return HCCL_SUCCESS;
198 : }
199 :
200 2 : HcclResult AivUbMemTransport::RecvMemInfo()
201 : {
202 2 : recvData_.resize(exchangeDataSize_);
203 : EXCEPTION_HANDLE_BEGIN
204 2 : socket_->RecvAsync(reinterpret_cast<u8 *>(&recvData_[0]), recvData_.size());
205 0 : EXCEPTION_HANDLE_END
206 : // HCCL_INFO("recv data, size=%llu, data=%s", data.size(), Hccl::Bytes2hex(data.data(), data.size()).c_str());
207 2 : return HCCL_SUCCESS;
208 : }
209 :
210 2 : HcclResult AivUbMemTransport::RecvDataProcess()
211 : {
212 2 : Hccl::BinaryStream binaryStream(recvData_);
213 2 : rmtBufferVec_.clear();
214 2 : rmtRmaBufferVec_.clear();
215 : EXCEPTION_HANDLE_BEGIN
216 2 : RmtBufferUnpackProc(binaryStream);
217 1 : EXCEPTION_HANDLE_END
218 1 : return HCCL_SUCCESS;
219 2 : }
220 :
221 1 : void AivUbMemTransport::RmtBufferUnpackProc(Hccl::BinaryStream &binaryStream)
222 : {
223 1 : u32 vecSize{0};
224 1 : binaryStream >> vecSize;
225 1 : HCCL_INFO("vecSize=%u", vecSize);
226 1 : uint32_t totalBufferNum = rmtBufferVec_.size() + vecSize;
227 1 : if (UNLIKELY(totalBufferNum > MAX_BUFFER_NUM)) {
228 0 : EXCEPTION_THROW_IF_ERR(HCCL_E_PARA, "[AivUbMemTransport][RmtBufferUnpackProc] vecSize exceeds limit.");
229 : }
230 :
231 1 : for (u32 pos = 0; pos < vecSize; ++pos) {
232 0 : Hccl::ExchangeIpcBufferDto dto;
233 0 : dto.Deserialize(binaryStream);
234 0 : HCCL_INFO("[%s] dto[%s]", __func__, dto.Describe().c_str());
235 0 : if (dto.size == 0) { // size为0,则为 remote 空buffer
236 0 : HCCL_INFO("unpack nullptr, pos=%u", pos);
237 0 : rmtBufferVec_.push_back(nullptr);
238 0 : rmtRmaBufferVec_.push_back(nullptr);
239 : } else { // size非0,则构造一个remote buffer
240 0 : HCCL_INFO("[AivUbMemTransport][RmtBufferUnpackProc] unpack buffer memInfo[%s]", dto.memInfo.c_str());
241 0 : rmtBufferVec_.push_back(std::make_unique<Hccl::RemoteIpcRmaBuffer>(dto));
242 0 : rmtRmaBufferVec_.push_back(rmtBufferVec_.back().get());
243 : }
244 0 : }
245 1 : }
246 :
247 2 : HcclResult AivUbMemTransport::GetRemoteMems(uint32_t *memNum, CommMem **remoteMem, char ***memInfos)
248 : {
249 2 : std::lock_guard<std::mutex> lock(remoteMemsMutex_);
250 2 : Hccl::RemoteMemCtx<std::unique_ptr<Hccl::RemoteIpcRmaBuffer>> remoteMemCtx{cacheValid_, rmtBufferVec_,
251 2 : remoteUserMems_, memInfoCopies_, memInfoPointers_, remoteMem, memInfos, memNum};
252 2 : CHK_RET(GetRemoteUserMems(remoteMemCtx));
253 2 : return HCCL_SUCCESS;
254 2 : }
255 :
256 5 : HcclResult AivUbMemTransport::CheckSocketStatus(std::string socketOperator)
257 : {
258 5 : CHK_PTR_NULL(socket_);
259 5 : auto timeout = std::chrono::seconds(Hccl::EnvConfig::GetInstance().GetSocketConfig().GetLinkTimeOut());
260 5 : auto startTime = std::chrono::steady_clock::now();
261 5 : uint32_t retryCount = 0;
262 : while(true) {
263 : EXCEPTION_HANDLE_BEGIN
264 5 : Hccl::SocketStatus socketStatus = socket_->GetAsyncStatus();
265 5 : if (socketStatus == Hccl::SocketStatus::OK) {
266 4 : auto elapsed = std::chrono::duration_cast<std::chrono::milliseconds>(
267 8 : std::chrono::steady_clock::now() - startTime).count();
268 4 : HCCL_INFO("[AivUbMemTransport][%s] socket transport operation[%s] success, elapsed[%lld]ms, retryCount[%u]",
269 : __func__, socketOperator.c_str(), elapsed, retryCount);
270 4 : break;
271 : }
272 2 : if ((std::chrono::steady_clock::now() - startTime) >= timeout ||
273 1 : socketStatus == Hccl::SocketStatus::TIMEOUT) {
274 1 : auto elapsed = std::chrono::duration_cast<std::chrono::milliseconds>(
275 2 : std::chrono::steady_clock::now() - startTime).count();
276 1 : HCCL_ERROR("[AivUbMemTransport][%s] socket transport operation[%s] timeout after %lld sec, elapsed[%lld]ms, retryCount[%u]",
277 : __func__, socketOperator.c_str(), timeout, elapsed, retryCount);
278 1 : return HCCL_E_TIMEOUT;
279 : }
280 0 : EXCEPTION_HANDLE_END
281 0 : retryCount++;
282 0 : }
283 4 : return HCCL_SUCCESS;
284 : }
285 :
286 3 : HcclResult AivUbMemTransport::UpdateMemInfo(HcommMemHandle *memHandles, uint32_t memHandleNum)
287 : {
288 3 : if (memHandleNum == 0) {
289 1 : HCCL_WARNING("[AivUbMemTransport][UpdateMemInfo] bufferNum is 0.");
290 1 : return HCCL_SUCCESS;
291 : }
292 2 : locMemTemp_.clear();
293 2 : CHK_RET(FillBufferVec(memHandles, memHandleNum, locMemTemp_));
294 2 : HCCL_INFO("[AivUbMemTransport][UpdateMemInfo] bufferNum[%zu]", locMemTemp_.size());
295 2 : sendData_.clear();
296 2 : Hccl::BinaryStream sendStream;
297 2 : CHK_RET(BufferPack(sendStream, locMemTemp_));
298 2 : sendStream.Dump(sendData_);
299 2 : u32 sendSize = sendData_.size();
300 : EXCEPTION_HANDLE_BEGIN
301 2 : socket_->SendAsync(&sendSize, sizeof(sendSize));
302 0 : EXCEPTION_HANDLE_END
303 4 : CHK_RET(CheckSocketStatus("SendDataSize"));
304 1 : CHK_RET(RecvDataSize());
305 2 : CHK_RET(CheckSocketStatus("RecvDataSize"));
306 1 : CHK_RET(SendMemInfo());
307 2 : CHK_RET(CheckSocketStatus("SendMemInfo"));
308 1 : CHK_RET(RecvMemInfo());
309 2 : CHK_RET(CheckSocketStatus("RecvMemInfo"));
310 1 : Hccl::BinaryStream recvStream(recvData_);
311 : EXCEPTION_HANDLE_BEGIN
312 1 : RmtBufferUnpackProc(recvStream);
313 0 : EXCEPTION_HANDLE_END
314 1 : localRmaBufferVec_.insert(localRmaBufferVec_.end(), locMemTemp_.begin(), locMemTemp_.end());
315 : // 流程中已有新增内存数量判断,故执行到此位置一定存在新增内存,需要将标识置位false,使得再次调用GetRemoteMems时重新构造缓存
316 1 : cacheValid_ = false;
317 1 : return HCCL_SUCCESS;
318 2 : }
319 : }
|