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 : #include "endpoint_mgr.h"
11 : #include "hccl_common.h"
12 : #include "urma_mem.h"
13 : #include <algorithm>
14 : #include "log.h"
15 : #include "hccl/hccl_res.h"
16 : #include "hccl_mem_v2.h"
17 : #include "exchange_ub_buffer_dto.h"
18 : #include "local_ub_rma_buffer_manager.h"
19 : #include "remote_rma_buffer.h"
20 : #include "local_ub_rma_buffer.h"
21 :
22 : namespace hcomm {
23 :
24 51 : UbRegedMemMgr::UbRegedMemMgr() { localUbRmaBufferMgr_ = std::make_unique<LocalUbRmaBufferMgr>(); }
25 :
26 22 : HcclResult UbRegedMemMgr::RegisterMemory(HcommMem mem, const char* memTag, void** memHandle)
27 : {
28 22 : HCCL_INFO("[%s] Begin", __FUNCTION__);
29 22 : CHK_PTR_NULL(localUbRmaBufferMgr_);
30 22 : std::lock_guard<std::mutex> lock(memMtx_);
31 44 : return RegisterMemoryImpl(
32 22 : mem, memTag, memHandle, localUbRmaBufferMgr_, allRegisteredBuffers_, &handlesRecords_, "UbRegedMemMgr",
33 0 : [&](auto& bufPtr, auto& parent) {
34 6 : return std::make_shared<Hccl::LocalUbRmaBuffer>(bufPtr, rdmaHandle_, *parent);
35 : },
36 22 : [&](auto& bufPtr) {
37 10 : return std::make_shared<Hccl::LocalUbRmaBuffer>(bufPtr, rdmaHandle_);
38 22 : });
39 22 : }
40 :
41 20 : HcclResult UbRegedMemMgr::UnregisterMemory(void* memHandle)
42 : {
43 20 : HCCL_INFO("[%s] Begin", __FUNCTION__);
44 20 : CHK_PTR_NULL(localUbRmaBufferMgr_);
45 20 : std::lock_guard<std::mutex> lock(memMtx_);
46 40 : return UnregisterMemoryImpl(
47 20 : memHandle, localUbRmaBufferMgr_, allRegisteredBuffers_, &handlesRecords_,
48 0 : [](auto* b) {
49 14 : return b->GetMemRegOutParam();
50 : },
51 0 : [](const void* lhs, const void* rhs) {
52 7 : return Hccl::LocalUbRmaBuffer::IsSameMemRegOutParam(lhs, rhs);
53 20 : });
54 20 : }
55 :
56 0 : HcclResult UbRegedMemMgr::GetMemDesc(const EndpointDesc endpointDesc, Hccl::LocalUbRmaBuffer* localUbRmaBuffer)
57 : {
58 0 : auto dto = localUbRmaBuffer->GetExchangeDto();
59 0 : Hccl::BinaryStream localUbRmaBufferStream;
60 0 : dto->Serialize(localUbRmaBufferStream);
61 0 : std::vector<char> tempLocalMemDesc;
62 0 : localUbRmaBufferStream.Dump(tempLocalMemDesc);
63 0 : HCCL_DEBUG("[UbRegedMemMgr][GetMemDesc] [%s] dump data size [%u]", __func__, tempLocalMemDesc.size());
64 : // 判断内存描述符是否正确导出
65 0 : if (tempLocalMemDesc.empty()) {
66 0 : HCCL_ERROR("[UbRegedMemMgr][GetMemDesc] [%s] tempLocalMemDesc export failed.", __func__);
67 0 : return HCCL_E_INTERNAL;
68 : }
69 :
70 0 : std::vector<char> tempLocalEndpointDesc;
71 0 : tempLocalEndpointDesc.resize(sizeof(EndpointDesc));
72 0 : if (memcpy_s(tempLocalEndpointDesc.data(), sizeof(EndpointDesc), &endpointDesc, sizeof(EndpointDesc)) != EOK) {
73 0 : HCCL_ERROR("[UbRegedMemMgr][GetMemDesc] [%s] endpointDesc memcpy_s failed.", __func__);
74 0 : return HCCL_E_INTERNAL;
75 : }
76 :
77 0 : tempLocalMemDesc.insert(tempLocalMemDesc.end(), tempLocalEndpointDesc.begin(), tempLocalEndpointDesc.end());
78 :
79 : // 内存描述符拷贝
80 0 : localUbRmaBuffer->Desc = std::move(tempLocalMemDesc);
81 0 : return HCCL_SUCCESS;
82 0 : }
83 :
84 : HcclResult
85 1 : UbRegedMemMgr::MemoryExport(const EndpointDesc endpointDesc, void* memHandle, void** memDesc, uint32_t* memDescLen)
86 : {
87 1 : HCCL_INFO("[%s] Begin", __FUNCTION__);
88 :
89 1 : CHK_PTR_NULL(memHandle);
90 1 : CHK_PTR_NULL(memDesc);
91 1 : CHK_PTR_NULL(memDescLen);
92 1 : std::lock_guard<std::mutex> lock(memMtx_);
93 :
94 1 : Hccl::LocalUbRmaBuffer* localUbRmaBuffer = nullptr;
95 1 : CHK_RET(ValidateMemExportHandle(memHandle, allRegisteredBuffers_, localUbRmaBuffer));
96 :
97 : // 获取序列化信息
98 0 : CHK_RET(GetMemDesc(endpointDesc, localUbRmaBuffer));
99 :
100 0 : *memDescLen = static_cast<uint32_t>(localUbRmaBuffer->Desc.size());
101 0 : *memDesc = static_cast<void*>(localUbRmaBuffer->Desc.data());
102 :
103 0 : return HCCL_SUCCESS;
104 1 : }
105 :
106 0 : HcclResult UbRegedMemMgr::GetParamsFromMemDesc(
107 : const void* memDesc, uint32_t descLen, EndpointDesc& endpointDesc, Hccl::ExchangeUbBufferDto& dto)
108 : {
109 0 : const char* description = static_cast<const char*>(memDesc);
110 :
111 0 : CHK_PRT_RET(
112 : descLen < sizeof(EndpointDesc),
113 : HCCL_ERROR("[%s] descLen[%u] is too small, expected at least %zu", __func__, descLen, sizeof(EndpointDesc)),
114 : HCCL_E_PARA);
115 : // 从memDesc末尾提取EndpointDesc
116 0 : if (memcpy_s(
117 0 : &endpointDesc, sizeof(EndpointDesc), description + descLen - sizeof(EndpointDesc), sizeof(EndpointDesc))
118 0 : != EOK) {
119 0 : HCCL_ERROR(
120 : "[UbRegedMemMgr][GetParamsFromMemDesc] [%s] endpointDesc copy error. aim size:[%llu]", __func__,
121 : sizeof(EndpointDesc));
122 0 : return HCCL_E_INTERNAL;
123 : }
124 :
125 : // 反序列化
126 0 : std::vector<char> tempDesc{};
127 0 : tempDesc.resize(TRANSPORT_EMD_ESC_SIZE);
128 0 : tempDesc.assign(description, description + descLen - sizeof(EndpointDesc));
129 0 : Hccl::BinaryStream remoteUbRmaBufferStream(tempDesc);
130 0 : dto.Deserialize(remoteUbRmaBufferStream);
131 0 : return HCCL_SUCCESS;
132 0 : }
133 :
134 0 : HcclResult UbRegedMemMgr::MemoryImport(const void* memDesc, uint32_t descLen, HcommMem* outMem)
135 : {
136 0 : HCCL_INFO("[%s] Begin", __FUNCTION__);
137 0 : std::lock_guard<std::mutex> lock(memMtx_);
138 :
139 : EndpointDesc endpointDesc;
140 0 : Hccl::ExchangeUbBufferDto dto;
141 0 : CHK_RET(GetParamsFromMemDesc(memDesc, descLen, endpointDesc, dto));
142 :
143 : // 构造RemoteUbRmaBuffer
144 0 : std::shared_ptr<Hccl::RemoteUbRmaBuffer> remoteUbRmaBuffer;
145 0 : EXCEPTION_CATCH(remoteUbRmaBuffer = std::make_shared<Hccl::RemoteUbRmaBuffer>(rdmaHandle_, dto),
146 : return HCCL_E_PTR;);
147 0 : CHK_SMART_PTR_NULL(remoteUbRmaBuffer);
148 :
149 : // 放到RemoteUbRmaBufferMgr_
150 0 : hccl::BufferKey<uintptr_t, u64> tempKey(static_cast<uintptr_t>(dto.addr), dto.size);
151 0 : if (remoteUbRmaBufferMgrs_.find(endpointDesc) == remoteUbRmaBufferMgrs_.end()) {
152 0 : std::unique_ptr<RemoteUbRmaBufferMgr> remoteUbRmaBufferMgr;
153 0 : EXCEPTION_CATCH((remoteUbRmaBufferMgr = std::make_unique<RemoteUbRmaBufferMgr>()), return HCCL_E_PTR);
154 0 : CHK_SMART_PTR_NULL(remoteUbRmaBufferMgr);
155 0 : remoteUbRmaBufferMgrs_[endpointDesc] = std::move(remoteUbRmaBufferMgr);
156 0 : HCCL_INFO("remoteUbRmaBufferMgrs_ add remoteUbRmaBufferMgr successfully!");
157 0 : }
158 :
159 0 : auto resultPair = remoteUbRmaBufferMgrs_[endpointDesc]->Add(tempKey, remoteUbRmaBuffer);
160 0 : if (!resultPair.second) {
161 0 : HCCL_ERROR("[UbRegedMemMgr][MemoryImport] This memDesc has already been imported!");
162 0 : return HCCL_E_AGAIN;
163 : }
164 :
165 0 : outMem->addr = reinterpret_cast<void*>(remoteUbRmaBuffer->GetAddr());
166 0 : outMem->size = remoteUbRmaBuffer->GetSize();
167 :
168 0 : return HCCL_SUCCESS;
169 0 : }
170 :
171 0 : HcclResult UbRegedMemMgr::MemoryUnimport(const void* memDesc, uint32_t descLen)
172 : {
173 0 : HCCL_INFO("[%s] Begin", __FUNCTION__);
174 0 : std::lock_guard<std::mutex> lock(memMtx_);
175 :
176 : EndpointDesc endpointDesc;
177 0 : Hccl::ExchangeUbBufferDto dto;
178 0 : CHK_RET(GetParamsFromMemDesc(memDesc, descLen, endpointDesc, dto));
179 :
180 0 : if (remoteUbRmaBufferMgrs_.find(endpointDesc) == remoteUbRmaBufferMgrs_.end()) {
181 0 : HCCL_ERROR("[UrmaRegedMemMgr][MemoryUnimport] Remote buffer manager Not Found.");
182 0 : return HCCL_E_NOT_FOUND;
183 : }
184 :
185 : // 删除RemoteUbRmaBuffer
186 0 : HCCL_INFO("[MemoryUnimport][Ub] MemoryUnimport");
187 0 : hccl::BufferKey<uintptr_t, u64> tempKey(static_cast<uintptr_t>(dto.addr), dto.size);
188 :
189 0 : bool resultPair = false;
190 0 : EXCEPTION_CATCH(resultPair = remoteUbRmaBufferMgrs_[endpointDesc]->Del(tempKey), return HCCL_E_NOT_FOUND);
191 : // 计数器大于1时,返回false,说明框架层有其它设备在使用这段内存,返回HCCL_E_AGAIN
192 0 : if (!resultPair) {
193 0 : HCCL_INFO("[UrmaRegedMemMgr][[MemoryUnimport] Memory reference count is larger than 0"
194 : "(used by other RemoteRank).");
195 0 : return HCCL_E_AGAIN;
196 : }
197 0 : if (!remoteUbRmaBufferMgrs_[endpointDesc]->size()) {
198 0 : remoteUbRmaBufferMgrs_.erase(endpointDesc);
199 : }
200 0 : return HCCL_SUCCESS;
201 0 : }
202 :
203 5 : HcclResult UbRegedMemMgr::GetAllMemHandles(void** memHandles, uint32_t* memHandleNum)
204 : {
205 5 : HCCL_INFO("[%s] Begin", __FUNCTION__);
206 5 : std::lock_guard<std::mutex> lock(memMtx_);
207 5 : CHK_PTR_NULL(memHandles);
208 5 : CHK_PTR_NULL(memHandleNum);
209 5 : *memHandleNum = static_cast<uint32_t>(handlesRecords_.size());
210 5 : *memHandles = handlesRecords_.empty() ? nullptr : static_cast<void*>(handlesRecords_.data());
211 5 : HCCL_INFO("[UbRegedMemMgr][GetAllMemHandles] memHandleNum[%u]", *memHandleNum);
212 5 : return HCCL_SUCCESS;
213 5 : }
214 :
215 : } // namespace hcomm
|