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