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 "hccl_mem.h"
12 : #include "hccl_network.h"
13 : #include "remote_ipc_rma_buffer.h"
14 : #include "remote_rdma_rma_buffer.h"
15 :
16 : #ifdef __cplusplus
17 : extern "C" {
18 : #endif
19 : HcclResult __attribute__((weak)) HcclMemRegV2(HcclNetDev netDev, const HcclMem* mem, HcclBuf* buf);
20 : HcclResult __attribute__((weak)) HcclMemDeregV2(const HcclBuf* buf);
21 : HcclResult __attribute__((weak)) HcclMemExportV2(HcclBuf* buf, char** outDesc, uint64_t* outDescLen);
22 : HcclResult __attribute__((weak))
23 : HcclMemImportV2(const char* description, uint64_t descLen, bool isRemote, HcclBuf* outBuf, HcclNetDev netDev);
24 : HcclResult __attribute__((weak)) HcclMemCloseV2(HcclBuf* buf);
25 : #ifdef __cplusplus
26 : }
27 : #endif
28 :
29 : using namespace hccl;
30 :
31 : using LocalIpcRmaBufferMgr = NetDevContext::LocalIpcRmaBufferMgr;
32 : using LocalRdmaRmaBufferMgr = NetDevContext::LocalRdmaRmaBufferMgr;
33 :
34 0 : static HcclResult HcclMemRegIpc(NetDevContext* netDevCtx, const HcclMem* mem, HcclBuf* buf)
35 : {
36 0 : std::shared_ptr<LocalIpcRmaBufferMgr> localRmaBufferMgr = netDevCtx->GetlocalIpcRmaBufferMgr();
37 0 : if (!localRmaBufferMgr) {
38 0 : HCCL_ERROR("[HcclMemRegIpc]Can't get LocalIpcRmaBufferMgr");
39 0 : return HCCL_E_INTERNAL;
40 : }
41 :
42 0 : RmaMemType memType = static_cast<RmaMemType>(mem->type);
43 0 : u64 size = static_cast<u64>(mem->size);
44 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(mem->addr), size);
45 0 : std::shared_ptr<LocalIpcRmaBuffer> localbufferPtr = nullptr;
46 0 : EXCEPTION_CATCH(
47 : (localbufferPtr = std::make_shared<LocalIpcRmaBuffer>(netDevCtx, mem->addr, size, memType)), return HCCL_E_PTR);
48 0 : auto resultPair = localRmaBufferMgr->Add(tempKey, localbufferPtr);
49 0 : if (resultPair.first == localRmaBufferMgr->End()) {
50 : // 输入key是表中某一个最相近key的交集、子集。返回空迭代器
51 0 : HCCL_ERROR("[HcclMemRegIpc]The memory that is expected to be"
52 : " registered overlaps with the memory that has been registered, please check params");
53 0 : return HCCL_E_INTERNAL;
54 : }
55 : // 已注册:输入key是表中某一最相近key的全集。 返回添加该key的迭代器,及false
56 : // 未注册:输入key是表中某一最相近key的空集。 返回添加成功的迭代器,及true
57 0 : std::shared_ptr<LocalIpcRmaBuffer> localBuffer = resultPair.first->second.buffer;
58 0 : buf->addr = localBuffer->GetAddr();
59 0 : buf->len = localBuffer->GetSize();
60 0 : auto rmaBufferPtr = dynamic_cast<RmaBuffer*>(localBuffer.get());
61 0 : CHK_PTR_NULL(rmaBufferPtr);
62 0 : buf->handle = static_cast<void*>(rmaBufferPtr);
63 0 : if (resultPair.second) {
64 0 : HcclResult ret = localBuffer->Init();
65 0 : if (ret != HCCL_SUCCESS) {
66 : // 此分支中一定删除成功
67 0 : localRmaBufferMgr->Del(tempKey);
68 0 : HCCL_ERROR("[HcclMemRegRoce]localbuffer init failed %d.", ret);
69 0 : return ret;
70 : }
71 0 : HCCL_INFO("[HcclMemRegIpc]Register memory success! Add key {%p, %llu}", mem->addr, size);
72 0 : return HCCL_SUCCESS;
73 : } else { // 内存再次注册时
74 0 : HCCL_INFO(
75 : "[HcclMemRegIpc]Memory is already registered, just increase the reference count. Add key "
76 : "{%p, %llu}",
77 : mem->addr, size);
78 : ;
79 0 : return HCCL_E_AGAIN;
80 : }
81 0 : }
82 :
83 0 : static HcclResult HcclMemDeregIpc(NetDevContext* netDevCtx, const HcclBuf* buf)
84 : {
85 0 : std::shared_ptr<LocalIpcRmaBufferMgr> localRmaBufferMgr = netDevCtx->GetlocalIpcRmaBufferMgr();
86 0 : if (!localRmaBufferMgr) {
87 0 : HCCL_ERROR("[HcclMemDeregIpc]Can't get LocalIpcRmaBufferMgr");
88 0 : return HCCL_E_INTERNAL;
89 : }
90 :
91 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(buf->addr), buf->len);
92 0 : if (localRmaBufferMgr->Del(tempKey)) {
93 : // 删除成功:输入key是表中某一最相近key的全集,计数-1后为0,返回true
94 0 : HCCL_INFO("[HcclMemDeregIpc]Memory reference count is 0, deregister memory.");
95 0 : return HCCL_SUCCESS;
96 : } else {
97 : // 删除失败:输入key是表中某一最相近key的全集,计数不为0(存在其他remoteRank使用),返回false
98 0 : HCCL_INFO("[HcclMemDeregIpc]Memory reference count is larger than 0 "
99 : "(used by other RemoteRank), do not deregister memory.");
100 0 : return HCCL_E_AGAIN;
101 : }
102 0 : }
103 :
104 0 : static HcclResult HcclMemRegRoce(NetDevContext* netDevCtx, const HcclMem* mem, HcclBuf* buf)
105 : {
106 0 : std::shared_ptr<LocalRdmaRmaBufferMgr> localRmaBufferMgr = netDevCtx->GetlocalRdmaRmaBufferMgr();
107 0 : if (!localRmaBufferMgr) {
108 0 : HCCL_ERROR("[HcclMemRegRoce] can't get LocalRdmaRmaBufferMgr");
109 0 : return HCCL_E_INTERNAL;
110 : }
111 :
112 0 : RmaMemType memType = static_cast<RmaMemType>(mem->type);
113 0 : u64 size = static_cast<u64>(mem->size);
114 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(mem->addr), size);
115 0 : std::shared_ptr<LocalRdmaRmaBuffer> localbufferPtr = nullptr;
116 0 : EXCEPTION_CATCH(
117 : (localbufferPtr = std::make_shared<LocalRdmaRmaBuffer>(netDevCtx, mem->addr, size, memType)),
118 : return HCCL_E_PTR);
119 0 : auto resultPair = localRmaBufferMgr->Add(tempKey, localbufferPtr);
120 0 : if (resultPair.first == localRmaBufferMgr->End()) {
121 : // 输入key是表中某一个最相近key的交集、子集。返回空迭代器
122 0 : HCCL_ERROR("[HcclMemRegRoce]The memory that is expected to be"
123 : " registered overlaps with the memory that has been registered, please check params");
124 0 : return HCCL_E_INTERNAL;
125 : }
126 : // 已注册:输入key是表中某一最相近key的全集。 返回添加该key的迭代器,及false
127 : // 未注册:输入key是表中某一最相近key的空集。 返回添加成功的迭代器,及true
128 0 : std::shared_ptr<LocalRdmaRmaBuffer> localBuffer = resultPair.first->second.buffer;
129 0 : buf->addr = localBuffer->GetAddr();
130 0 : buf->len = localBuffer->GetSize();
131 0 : auto rmaBufferPtr = dynamic_cast<RmaBuffer*>(localBuffer.get());
132 0 : CHK_PTR_NULL(rmaBufferPtr);
133 0 : buf->handle = static_cast<void*>(rmaBufferPtr);
134 0 : if (resultPair.second) {
135 0 : HcclResult ret = localBuffer->Init();
136 0 : if (ret != HCCL_SUCCESS) {
137 : // 此分支中一定删除成功
138 0 : localRmaBufferMgr->Del(tempKey);
139 0 : HCCL_ERROR("[HcclMemRegRoce]localbuffer init failed %d.", ret);
140 0 : return ret;
141 : }
142 0 : HCCL_INFO("[HcclMemRegRoce]Register memory success! Add key {%p, %llu}", mem->addr, size);
143 0 : return HCCL_SUCCESS;
144 : } else { // 内存再次注册时
145 0 : HCCL_INFO(
146 : "[HcclMemRegRoce]Memory is already registered, just increase the reference count. Add key "
147 : "{%p, %llu}",
148 : mem->addr, size);
149 : ;
150 0 : return HCCL_E_AGAIN;
151 : }
152 0 : }
153 :
154 0 : static HcclResult HcclMemDeregRoce(NetDevContext* netDevCtx, const HcclBuf* buf)
155 : {
156 0 : std::shared_ptr<LocalRdmaRmaBufferMgr> localRmaBufferMgr = netDevCtx->GetlocalRdmaRmaBufferMgr();
157 0 : if (!localRmaBufferMgr) {
158 0 : HCCL_ERROR("[HcclMemDeregRoce]Can't get LocalRdmaRmaBufferMgr");
159 0 : return HCCL_E_INTERNAL;
160 : }
161 :
162 0 : BufferKey<uintptr_t, u64> tempKey(reinterpret_cast<uintptr_t>(buf->addr), buf->len);
163 0 : if (localRmaBufferMgr->Del(tempKey)) {
164 : // 删除成功:输入key是表中某一最相近key的全集,计数-1后为0,返回true
165 0 : HCCL_INFO("[HcclMemDeregRoce]Memory reference count is 0, deregister memory.");
166 0 : return HCCL_SUCCESS;
167 : } else {
168 : // 删除失败:输入key是表中某一最相近key的全集,计数不为0(存在其他remoteRank使用),返回false
169 0 : HCCL_INFO("[HcclMemDeregRoce]Memory reference count is larger than 0 "
170 : "(used by other RemoteRank), do not deregister memory.");
171 0 : return HCCL_E_AGAIN;
172 : }
173 0 : }
174 :
175 0 : static HcclResult HcclMemRempRoce(NetDevContext* netDevCtx, const HcclMem* memArray, u64 arraySize)
176 : {
177 0 : std::shared_ptr<LocalRdmaRmaBufferMgr> localRmaBufferMgr = netDevCtx->GetlocalRdmaRmaBufferMgr();
178 0 : if (!localRmaBufferMgr) {
179 0 : HCCL_ERROR("[HcclMemRempRoce]Can't get LocalRdmaRmaBufferMgr");
180 0 : return HCCL_E_INTERNAL;
181 : }
182 0 : HCCL_RUN_INFO("[HcclMemRempRoce] arraySize[%u]", arraySize);
183 0 : std::unordered_map<void*, bool> remapAddr;
184 0 : for (u64 i = 0; i < arraySize; i++) {
185 0 : const HcclMem& memInfo = memArray[i];
186 :
187 : // 检查地址和大小是否有效
188 0 : if (memInfo.addr == nullptr || memInfo.size <= 0 || memInfo.type != HcclMemType::HCCL_MEM_TYPE_DEVICE) {
189 0 : continue;
190 : }
191 :
192 : // 检查地址是否已经处理过
193 0 : if (remapAddr.find(memInfo.addr) != remapAddr.end()) {
194 0 : continue;
195 : }
196 :
197 : // 查找地址是否注册过
198 0 : BufferKey<uintptr_t, u64> searchKey(reinterpret_cast<uintptr_t>(memInfo.addr), 1U);
199 0 : auto bufferIter = localRmaBufferMgr->Find(searchKey);
200 0 : if (!bufferIter.first) {
201 0 : HCCL_ERROR(
202 : "[HcclMemRempRoce]Memory addr[%p] size[%llu] has not been registered.", memInfo.addr, memInfo.size);
203 0 : return HCCL_E_PARA;
204 : }
205 :
206 : // 计算需要注册的内存大小
207 0 : u64 size = std::min(static_cast<u64>(memInfo.size), bufferIter.second->GetSize());
208 :
209 : // 注册内存
210 0 : HCCL_RUN_INFO("[HcclMemRempRoce]Re-register memory addr[%p] size[%llu].", memInfo.addr, size);
211 0 : HcclResult ret = bufferIter.second->Remap(memInfo.addr, size);
212 0 : CHK_PRT_RET(
213 : ret != HCCL_SUCCESS,
214 : HCCL_ERROR("[HcclMemRempRoce]remap mem failed,addr[%p], size[%llu]", memInfo.addr, size), ret);
215 :
216 : // 标记地址已处理
217 0 : remapAddr.emplace(memInfo.addr, true);
218 0 : }
219 :
220 0 : return HCCL_SUCCESS;
221 0 : }
222 :
223 0 : HcclResult HcclMemReg(HcclNetDev netDev, const HcclMem* mem, HcclBuf* buf)
224 : {
225 0 : CHK_PTR_NULL(netDev);
226 0 : CHK_PTR_NULL(mem);
227 0 : CHK_PTR_NULL(buf);
228 0 : CHK_PTR_NULL(mem->addr);
229 0 : CHK_PRT_RET(
230 : (mem->type != HCCL_MEM_TYPE_DEVICE) && (mem->type != HCCL_MEM_TYPE_HOST),
231 : HCCL_ERROR("[HcclMemReg]memoryType[%d] must be device or host", mem->type), HCCL_E_PARA);
232 0 : CHK_PRT_RET(mem->size == 0, HCCL_ERROR("[HcclMemReg]memory size[%lld] is invalid", mem->size), HCCL_E_PARA);
233 :
234 : DevType devType;
235 0 : CHK_RET(hrtGetDeviceType(devType));
236 0 : if (devType == DevType::DEV_TYPE_950 || devType == DevType::DEV_TYPE_960) {
237 0 : return HcclMemRegV2(netDev, mem, buf);
238 : }
239 :
240 0 : NetDevContext* netDevCtx = static_cast<NetDevContext*>(netDev);
241 0 : if (netDevCtx->GetNicType() == NicType::VNIC_TYPE) {
242 0 : return HcclMemRegIpc(netDevCtx, mem, buf);
243 : } else {
244 0 : return HcclMemRegRoce(netDevCtx, mem, buf);
245 : }
246 : }
247 :
248 0 : HcclResult HcclMemDereg(const HcclBuf* buf)
249 : {
250 0 : CHK_PTR_NULL(buf);
251 0 : CHK_PTR_NULL(buf->addr);
252 0 : CHK_PTR_NULL(buf->handle);
253 0 : CHK_PRT_RET(buf->len == 0U, HCCL_ERROR("[HcclMemDereg]buf size[%llu] is invalid", buf->len), HCCL_E_PARA);
254 :
255 : DevType devType;
256 0 : CHK_RET(hrtGetDeviceType(devType));
257 0 : if (devType == DevType::DEV_TYPE_950 || devType == DevType::DEV_TYPE_960) {
258 0 : return HcclMemDeregV2(buf);
259 : }
260 :
261 0 : RmaBuffer* rmaBuffer = static_cast<RmaBuffer*>(buf->handle);
262 0 : NetDevContext* netDevCtx = static_cast<NetDevContext*>(const_cast<void*>(rmaBuffer->GetNetDevCtx()));
263 0 : if (netDevCtx->GetNicType() == NicType::VNIC_TYPE) {
264 0 : return HcclMemDeregIpc(netDevCtx, buf);
265 : } else {
266 0 : return HcclMemDeregRoce(netDevCtx, buf);
267 : }
268 : }
269 :
270 0 : HcclResult HcclMemRemap(HcclNetDev netDev, const HcclMem* memArray, uint64_t arraySize)
271 : {
272 0 : CHK_PTR_NULL(netDev);
273 0 : CHK_PTR_NULL(memArray);
274 0 : CHK_PRT_RET(arraySize == 0U, HCCL_ERROR("[HcclMemReMap]arraySize[%llu] is invalid", arraySize), HCCL_E_PARA);
275 :
276 0 : NetDevContext* netDevCtx = static_cast<NetDevContext*>(netDev);
277 0 : if (netDevCtx->GetNicType() == NicType::VNIC_TYPE) {
278 0 : HCCL_INFO("[HcclMemReMap][ReMapMemIpc] doesn't support ReMapMem");
279 0 : return HCCL_SUCCESS;
280 : } else {
281 0 : return HcclMemRempRoce(netDevCtx, memArray, arraySize);
282 : }
283 : }
284 :
285 0 : HcclResult HcclMemExport(HcclBuf* buf, char** outDesc, uint64_t* outDescLen)
286 : {
287 0 : CHK_PTR_NULL(buf);
288 0 : CHK_PTR_NULL(outDesc);
289 0 : CHK_PTR_NULL(outDescLen);
290 0 : CHK_PTR_NULL(buf->addr);
291 0 : CHK_PTR_NULL(buf->handle);
292 0 : CHK_PRT_RET(buf->len == 0U, HCCL_ERROR("[HcclMemExport]buf size[%llu] is invalid", buf->len), HCCL_E_PARA);
293 :
294 : DevType devType;
295 0 : CHK_RET(hrtGetDeviceType(devType));
296 0 : if (devType == DevType::DEV_TYPE_950 || devType == DevType::DEV_TYPE_960) {
297 0 : return HcclMemExportV2(buf, outDesc, outDescLen);
298 : }
299 :
300 0 : RmaBuffer* rmaBuffer = static_cast<RmaBuffer*>(buf->handle);
301 0 : if (rmaBuffer->GetRmaType() == RmaType::IPC_RMA) {
302 0 : LocalIpcRmaBuffer* localRmaBufer = dynamic_cast<LocalIpcRmaBuffer*>(rmaBuffer);
303 0 : CHK_PTR_NULL(localRmaBufer);
304 0 : std::string& tempLocalMemDesc = localRmaBufer->Serialize();
305 0 : if (tempLocalMemDesc.empty()) {
306 0 : HCCL_ERROR("[HcclMemExport][Ipc]tempLocalMemDesc is empty.");
307 0 : return HCCL_E_INTERNAL;
308 : }
309 :
310 0 : *outDesc = const_cast<char*>(tempLocalMemDesc.c_str());
311 0 : *outDescLen = tempLocalMemDesc.length();
312 : } else {
313 0 : LocalRdmaRmaBuffer* localRmaBufer = dynamic_cast<LocalRdmaRmaBuffer*>(rmaBuffer);
314 0 : CHK_PTR_NULL(localRmaBufer);
315 0 : std::string& tempLocalMemDesc = localRmaBufer->Serialize();
316 0 : if (tempLocalMemDesc.empty()) {
317 0 : HCCL_ERROR("[HcclMemExport][Roce]tempLocalMemDesc is empty.");
318 0 : return HCCL_E_INTERNAL;
319 : }
320 :
321 0 : *outDesc = const_cast<char*>(tempLocalMemDesc.c_str());
322 0 : *outDescLen = tempLocalMemDesc.length();
323 : }
324 0 : return HCCL_SUCCESS;
325 : }
326 :
327 0 : HcclResult HcclMemGrant(HcclBuf* localBuf, const HcclMemGrantInfo* remoteGrantInfo)
328 : {
329 0 : CHK_PTR_NULL(localBuf);
330 0 : CHK_PTR_NULL(remoteGrantInfo);
331 0 : CHK_PTR_NULL(localBuf->addr);
332 0 : CHK_PTR_NULL(localBuf->handle);
333 0 : CHK_PRT_RET(localBuf->len == 0U, HCCL_ERROR("[HcclMemGrant]buf size[%llu] is invalid", localBuf->len), HCCL_E_PARA);
334 0 : RmaBuffer* rmaBuffer = static_cast<RmaBuffer*>(localBuf->handle);
335 0 : if (rmaBuffer->GetRmaType() == RmaType::IPC_RMA) {
336 0 : LocalIpcRmaBuffer* localRmaBufer = dynamic_cast<LocalIpcRmaBuffer*>(rmaBuffer);
337 0 : CHK_PTR_NULL(localRmaBufer);
338 0 : HcclResult ret = localRmaBufer->Grant(remoteGrantInfo->remotePid, remoteGrantInfo->remoteSdid);
339 0 : CHK_PRT_RET((ret != HCCL_SUCCESS), HCCL_ERROR("[HcclMemGrant]Grant error"), ret);
340 : }
341 0 : return HCCL_SUCCESS;
342 : }
343 :
344 : HcclResult
345 0 : HcclMemImport(const char* description, uint32_t descLen, bool isRemote, HcclBuf* outBuf, HcclNetDevCtx netDevCtx)
346 : {
347 0 : CHK_PTR_NULL(netDevCtx);
348 0 : CHK_PTR_NULL(description);
349 0 : CHK_PTR_NULL(outBuf);
350 0 : CHK_PRT_RET(
351 : (descLen == 0), HCCL_ERROR("[HcclMemImport]input parameter is invalid descLen[%u] ", descLen), HCCL_E_PARA);
352 0 : CHK_PRT_RET(
353 : descLen > TRANSPORT_EMD_ESC_SIZE,
354 : HCCL_ERROR("[HcclMemImport]descLen[%u] is larger than limit[%u] ", descLen, TRANSPORT_EMD_ESC_SIZE),
355 : HCCL_E_PARA);
356 0 : if (isRemote == false) {
357 0 : HCCL_WARNING("[HcclMemImport]isRemote[%d] is invalid", isRemote);
358 : }
359 :
360 : DevType devType;
361 0 : CHK_RET(hrtGetDeviceType(devType));
362 0 : if (devType == DevType::DEV_TYPE_950 || devType == DevType::DEV_TYPE_960) {
363 0 : return HcclMemImportV2(description, descLen, isRemote, outBuf, netDevCtx);
364 : }
365 :
366 0 : std::string tempDesc = std::string(description, descLen);
367 0 : u8 rmaType = static_cast<unsigned char>(description[0]);
368 0 : switch (rmaType) {
369 0 : case static_cast<int>(RmaType::IPC_RMA): {
370 0 : RemoteIpcRmaBuffer* tempRemoteBufferPtr = new (std::nothrow) RemoteIpcRmaBuffer(netDevCtx);
371 0 : CHK_PTR_NULL(tempRemoteBufferPtr);
372 0 : HcclResult deRet = tempRemoteBufferPtr->Deserialize(tempDesc);
373 0 : HcclResult openRet = tempRemoteBufferPtr->Open();
374 0 : if (deRet != HCCL_SUCCESS || openRet != HCCL_SUCCESS) {
375 0 : delete tempRemoteBufferPtr;
376 0 : CHK_PRT_RET(
377 : deRet != HCCL_SUCCESS, HCCL_ERROR("[HcclMemImport]RemoteBuffer Deserialize failed."), deRet);
378 0 : CHK_PRT_RET(openRet != HCCL_SUCCESS, HCCL_ERROR("[HcclMemImport]RemoteBuffer Open failed."), openRet);
379 : }
380 0 : outBuf->addr = tempRemoteBufferPtr->GetAddr();
381 0 : outBuf->len = tempRemoteBufferPtr->GetSize();
382 0 : outBuf->handle = static_cast<void*>(tempRemoteBufferPtr);
383 0 : break;
384 : }
385 0 : case static_cast<int>(RmaType::RDMA_RMA): {
386 0 : RemoteRdmaRmaBuffer* tempRemoteBufferPtr = new (std::nothrow) RemoteRdmaRmaBuffer();
387 0 : CHK_PTR_NULL(tempRemoteBufferPtr);
388 0 : HcclResult ret = tempRemoteBufferPtr->Deserialize(tempDesc);
389 0 : if (ret != HCCL_SUCCESS) {
390 0 : delete tempRemoteBufferPtr;
391 0 : HCCL_ERROR("[HcclMemImport]RemoteBuffer Deserialize failed.");
392 0 : return ret;
393 : }
394 0 : outBuf->addr = tempRemoteBufferPtr->GetAddr();
395 0 : outBuf->len = tempRemoteBufferPtr->GetSize();
396 0 : outBuf->handle = static_cast<void*>(tempRemoteBufferPtr);
397 0 : break;
398 : }
399 0 : default: {
400 0 : HCCL_ERROR("[HcclMemImport]RmaType[%u] is invalid", rmaType);
401 0 : return HCCL_E_NOT_SUPPORT;
402 : }
403 : }
404 0 : return HCCL_SUCCESS;
405 0 : }
406 :
407 0 : HcclResult HcclMemClose(HcclBuf* buf)
408 : {
409 : // remoteIpcRmaBufferMgr_ 和 remoteRdmaRmaBufferMgr_ 要抽到HcclOneSidedConn里
410 0 : CHK_PTR_NULL(buf);
411 0 : CHK_PTR_NULL(buf->handle);
412 0 : RmaBuffer* rmaBuffer = static_cast<RmaBuffer*>(buf->handle);
413 :
414 : DevType devType;
415 0 : CHK_RET(hrtGetDeviceType(devType));
416 0 : if (devType == DevType::DEV_TYPE_950 || devType == DevType::DEV_TYPE_960) {
417 0 : return HcclMemCloseV2(buf);
418 : }
419 :
420 0 : if (rmaBuffer->GetRmaType() == RmaType::IPC_RMA) {
421 0 : HCCL_INFO("[HcclMemClose][Ipc] CloseMem");
422 0 : RemoteIpcRmaBuffer* tempRemoteBufferPtr = static_cast<RemoteIpcRmaBuffer*>(buf->handle);
423 0 : HcclResult ret = tempRemoteBufferPtr->Close();
424 0 : delete rmaBuffer;
425 0 : CHK_PRT_RET(ret != HCCL_SUCCESS, HCCL_ERROR("[HcclMemClose]RemoteBuffer Close failed"), ret);
426 0 : } else if (rmaBuffer->GetRmaType() == RmaType::RDMA_RMA) {
427 0 : HCCL_INFO("[HcclMemClose][Roce] CloseMem");
428 0 : delete rmaBuffer;
429 : } else {
430 0 : HCCL_ERROR("[HcclMemClose]RmaType[%d] is invalid", rmaBuffer->GetRmaType());
431 0 : return HCCL_E_INTERNAL;
432 : }
433 0 : return HCCL_SUCCESS;
434 : }
|